增加鉴权系统
This commit is contained in:
@@ -1,17 +1,19 @@
|
||||
from fastapi import APIRouter, HTTPException, Request, status
|
||||
|
||||
from auth import CurrentUser
|
||||
from schemas import (
|
||||
CreateSessionRequest,
|
||||
CreateSessionResponse,
|
||||
SendMessageRequest,
|
||||
SendMessageResponse,
|
||||
SessionHistoryResponse,
|
||||
UserUsageResponse,
|
||||
)
|
||||
from service import ChatProviderError, ChatService
|
||||
from storage import SessionNotFoundError, SessionStorageError
|
||||
|
||||
|
||||
router = APIRouter(prefix="/sessions", tags=["sessions"])
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def get_chat_service(request: Request) -> ChatService:
|
||||
@@ -19,32 +21,41 @@ def get_chat_service(request: Request) -> ChatService:
|
||||
|
||||
|
||||
@router.post(
|
||||
"",
|
||||
"/sessions",
|
||||
response_model=CreateSessionResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
tags=["sessions"],
|
||||
)
|
||||
async def create_session(
|
||||
request: Request,
|
||||
_user: CurrentUser,
|
||||
body: CreateSessionRequest | None = None,
|
||||
) -> CreateSessionResponse:
|
||||
try:
|
||||
return await get_chat_service(request).create_session(
|
||||
body.system_prompt if body is not None else None
|
||||
_user.user_id,
|
||||
body.system_prompt if body is not None else None,
|
||||
)
|
||||
except SessionStorageError as exc:
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/{session_id}/messages", response_model=SendMessageResponse)
|
||||
@router.post(
|
||||
"/sessions/{session_id}/messages",
|
||||
response_model=SendMessageResponse,
|
||||
tags=["sessions"],
|
||||
)
|
||||
async def send_message(
|
||||
session_id: str,
|
||||
body: SendMessageRequest,
|
||||
request: Request,
|
||||
user: CurrentUser,
|
||||
) -> SendMessageResponse:
|
||||
try:
|
||||
return await get_chat_service(request).generate_response(
|
||||
session_id, body.content
|
||||
response = await get_chat_service(request).generate_response(
|
||||
session_id, body.content, user.user_id
|
||||
)
|
||||
return response
|
||||
except SessionNotFoundError as exc:
|
||||
raise HTTPException(status_code=404, detail="session not found") from exc
|
||||
except ChatProviderError as exc:
|
||||
@@ -53,14 +64,29 @@ async def send_message(
|
||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.get("/{session_id}/messages", response_model=SessionHistoryResponse)
|
||||
@router.get(
|
||||
"/sessions/{session_id}/messages",
|
||||
response_model=SessionHistoryResponse,
|
||||
tags=["sessions"],
|
||||
)
|
||||
async def get_session_history(
|
||||
session_id: str,
|
||||
request: Request,
|
||||
_user: CurrentUser,
|
||||
) -> SessionHistoryResponse:
|
||||
try:
|
||||
return await get_chat_service(request).get_session_history(session_id)
|
||||
return await get_chat_service(request).get_session_history(
|
||||
session_id, _user.user_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
|
||||
|
||||
|
||||
@router.get("/usage", response_model=UserUsageResponse, tags=["usage"])
|
||||
async def get_current_user_usage(
|
||||
request: Request,
|
||||
user: CurrentUser,
|
||||
) -> UserUsageResponse:
|
||||
return await get_chat_service(request).get_user_usage(user.user_id)
|
||||
|
||||
Reference in New Issue
Block a user