优化目录结构
This commit is contained in:
@@ -1,19 +1,13 @@
|
||||
from contextlib import asynccontextmanager
|
||||
from collections.abc import AsyncGenerator
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Request, status
|
||||
from fastapi import FastAPI
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
from config import Settings
|
||||
from models import (
|
||||
CreateSessionRequest,
|
||||
CreateSessionResponse,
|
||||
SendMessageRequest,
|
||||
SendMessageResponse,
|
||||
SessionHistoryResponse,
|
||||
)
|
||||
from service import ChatProviderError, ChatService
|
||||
from storage import JsonSessionStorage, SessionNotFoundError, SessionStorageError
|
||||
from routes import router
|
||||
from service import ChatService
|
||||
from storage import JsonSessionStorage
|
||||
from tools import ToolRegistry
|
||||
|
||||
|
||||
@@ -22,7 +16,6 @@ def create_app(
|
||||
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]:
|
||||
@@ -34,75 +27,23 @@ def create_app(
|
||||
api_key=resolved_settings.api_key,
|
||||
base_url=resolved_settings.base_url,
|
||||
)
|
||||
|
||||
app.state.chat_service = ChatService(
|
||||
storage=storage,
|
||||
storage=JsonSessionStorage(resolved_settings.data_dir),
|
||||
client=resolved_client,
|
||||
model=resolved_settings.model,
|
||||
default_system_prompt=resolved_settings.default_system_prompt,
|
||||
tool_registry=ToolRegistry(resolved_settings.tool_timeout_seconds),
|
||||
max_tool_rounds=resolved_settings.max_tool_rounds,
|
||||
max_tool_calls_per_turn=resolved_settings.max_tool_calls_per_turn,
|
||||
)
|
||||
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,
|
||||
)
|
||||
async def send_message(
|
||||
session_id: str, body: SendMessageRequest, request: Request
|
||||
) -> SendMessageResponse:
|
||||
service: ChatService = request.app.state.chat_service
|
||||
try:
|
||||
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:
|
||||
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
||||
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
|
||||
|
||||
app.include_router(router)
|
||||
return app
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user