添加获取历史对话数据的接口
This commit is contained in:
@@ -1,10 +1,8 @@
|
||||
from contextlib import asynccontextmanager
|
||||
from collections.abc import AsyncGenerator
|
||||
import json
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Request, status
|
||||
from openai import AsyncOpenAI
|
||||
from starlette.responses import JSONResponse, Response, StreamingResponse
|
||||
|
||||
from config import Settings
|
||||
from models import (
|
||||
@@ -12,23 +10,12 @@ from models import (
|
||||
CreateSessionResponse,
|
||||
SendMessageRequest,
|
||||
SendMessageResponse,
|
||||
SessionHistoryResponse,
|
||||
)
|
||||
from service import ChatProviderError, ChatService
|
||||
from storage import JsonSessionStorage, SessionNotFoundError, SessionStorageError
|
||||
|
||||
|
||||
async def encode_sse(
|
||||
events: AsyncGenerator[dict[str, object], None],
|
||||
) -> AsyncGenerator[str, None]:
|
||||
try:
|
||||
async for event in events:
|
||||
yield f"data: {json.dumps(event, ensure_ascii=False)}\n\n"
|
||||
except (ChatProviderError, SessionStorageError) as exc:
|
||||
error = {"type": "error", "detail": str(exc)}
|
||||
yield f"data: {json.dumps(error, ensure_ascii=False)}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
|
||||
def create_app(
|
||||
settings: Settings | None = None,
|
||||
client: AsyncOpenAI | None = None,
|
||||
@@ -83,30 +70,13 @@ def create_app(
|
||||
@app.post(
|
||||
"/sessions/{session_id}/messages",
|
||||
response_model=SendMessageResponse,
|
||||
responses={
|
||||
200: {
|
||||
"content": {"text/event-stream": {}},
|
||||
"description": "JSON response or SSE stream when stream=true",
|
||||
}
|
||||
},
|
||||
)
|
||||
async def send_message(
|
||||
session_id: str, body: SendMessageRequest, request: Request
|
||||
) -> Response:
|
||||
) -> SendMessageResponse:
|
||||
service: ChatService = request.app.state.chat_service
|
||||
try:
|
||||
if body.stream:
|
||||
await service.ensure_session_exists(session_id)
|
||||
return StreamingResponse(
|
||||
encode_sse(
|
||||
service.generate_response_stream(session_id, body.content)
|
||||
),
|
||||
media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
|
||||
)
|
||||
|
||||
response = await service.generate_response(session_id, body.content)
|
||||
return JSONResponse(content=response.model_dump(mode="json"))
|
||||
return await service.generate_response(session_id, body.content)
|
||||
except SessionNotFoundError as exc:
|
||||
raise HTTPException(status_code=404, detail="session not found") from exc
|
||||
except ChatProviderError as exc:
|
||||
@@ -114,6 +84,21 @@ def create_app(
|
||||
except SessionStorageError as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
@app.get(
|
||||
"/sessions/{session_id}/messages",
|
||||
response_model=SessionHistoryResponse,
|
||||
)
|
||||
async def get_session_history(
|
||||
session_id: str, request: Request
|
||||
) -> SessionHistoryResponse:
|
||||
service: ChatService = request.app.state.chat_service
|
||||
try:
|
||||
return await service.get_session_history(session_id)
|
||||
except SessionNotFoundError as exc:
|
||||
raise HTTPException(status_code=404, detail="session not found") from exc
|
||||
except SessionStorageError as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
return app
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user