121 lines
4.2 KiB
Python
121 lines
4.2 KiB
Python
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 (
|
|
CreateSessionRequest,
|
|
CreateSessionResponse,
|
|
SendMessageRequest,
|
|
SendMessageResponse,
|
|
)
|
|
from service import ChatService, DeepSeekUpstreamError
|
|
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 (DeepSeekUpstreamError, 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,
|
|
) -> FastAPI:
|
|
resolved_settings = settings or Settings.from_env()
|
|
storage = JsonSessionStorage(resolved_settings.data_dir)
|
|
|
|
@asynccontextmanager
|
|
async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
|
|
resolved_client = client
|
|
if resolved_client is None:
|
|
if not resolved_settings.api_key:
|
|
raise RuntimeError("DEEPSEEK_API_KEY is required")
|
|
resolved_client = AsyncOpenAI(
|
|
api_key=resolved_settings.api_key,
|
|
base_url=resolved_settings.base_url,
|
|
)
|
|
app.state.chat_service = ChatService(
|
|
storage=storage,
|
|
client=resolved_client,
|
|
model=resolved_settings.model,
|
|
)
|
|
yield
|
|
if client is None:
|
|
await resolved_client.close()
|
|
|
|
app = FastAPI(title="Simple Chat API", version="0.1.0", lifespan=lifespan)
|
|
|
|
@app.post(
|
|
"/sessions",
|
|
response_model=CreateSessionResponse,
|
|
status_code=status.HTTP_201_CREATED,
|
|
)
|
|
async def create_session(
|
|
body: CreateSessionRequest | None = None,
|
|
) -> CreateSessionResponse:
|
|
system_prompt = (
|
|
body.system_prompt
|
|
if body is not None and body.system_prompt is not None
|
|
else resolved_settings.default_system_prompt
|
|
)
|
|
try:
|
|
session = await storage.create(system_prompt)
|
|
except SessionStorageError as exc:
|
|
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
|
return CreateSessionResponse(
|
|
session_id=session.session_id,
|
|
system_prompt=session.system_prompt,
|
|
created_at=session.created_at,
|
|
)
|
|
|
|
@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:
|
|
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"))
|
|
except SessionNotFoundError as exc:
|
|
raise HTTPException(status_code=404, detail="session not found") from exc
|
|
except DeepSeekUpstreamError as exc:
|
|
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
|
except SessionStorageError as exc:
|
|
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
|
|
|
return app
|
|
|
|
|
|
app = create_app()
|