增加鉴权系统

This commit is contained in:
2026-07-03 21:31:25 +08:00
parent 8073dd63a7
commit 7e65fadd83
12 changed files with 301 additions and 20 deletions
+46 -7
View File
@@ -19,8 +19,9 @@ from schemas import (
SendMessageResponse,
SessionHistoryResponse,
TokenUsage,
UserUsageResponse,
)
from storage import JsonSessionStorage
from storage import JsonSessionStorage, SessionNotFoundError
from tools import ToolRegistry
@@ -51,10 +52,11 @@ class ChatService:
self._locks: dict[str, asyncio.Lock] = {}
async def create_session(
self, system_prompt: str | None = None
self, user_id: str, system_prompt: str | None = None
) -> CreateSessionResponse:
session = await self.storage.create(
system_prompt or self.default_system_prompt
system_prompt or self.default_system_prompt,
user_id,
)
return CreateSessionResponse(
session_id=session.session_id,
@@ -63,12 +65,12 @@ class ChatService:
)
async def generate_response(
self, session_id: str, content: str
self, session_id: str, content: str, user_id: str | None = None
) -> SendMessageResponse:
lock = self._locks.setdefault(session_id, asyncio.Lock())
async with lock:
user_message = UserMessage(content=content, created_at=datetime.now(UTC))
session = await self.storage.read(session_id)
session = await self._read_owned_session(session_id, user_id)
api_messages = self._build_api_messages(session)
api_messages.append(self._message_to_api(user_message))
pending_messages: list[Message] = [user_message]
@@ -98,6 +100,11 @@ class ChatService:
pending_messages.append(final_message)
session.messages.extend(pending_messages)
session.updated_at = final_message.created_at
session.api_calls += 1
if has_usage:
session.prompt_tokens += usage.prompt_tokens
session.completion_tokens += usage.completion_tokens
session.total_tokens += usage.total_tokens
await self.storage.write(session)
return SendMessageResponse(
session_id=session_id,
@@ -150,10 +157,12 @@ class ChatService:
raise ChatProviderError("Tool round limit exceeded")
async def get_session_history(self, session_id: str) -> SessionHistoryResponse:
async def get_session_history(
self, session_id: str, user_id: str | None = None
) -> SessionHistoryResponse:
lock = self._locks.setdefault(session_id, asyncio.Lock())
async with lock:
session = await self.storage.read(session_id)
session = await self._read_owned_session(session_id, user_id)
return SessionHistoryResponse(
session_id=session.session_id,
system_prompt=session.system_prompt,
@@ -162,6 +171,36 @@ class ChatService:
messages=session.messages,
)
async def get_user_usage(self, user_id: str) -> UserUsageResponse:
sessions = await self.storage.list_by_user(user_id)
updated_at = max(
(session.updated_at for session in sessions),
default=datetime.now(UTC),
)
return UserUsageResponse(
user_id=user_id,
api_calls=sum(session.api_calls for session in sessions),
prompt_tokens=sum(session.prompt_tokens for session in sessions),
completion_tokens=sum(
session.completion_tokens for session in sessions
),
total_tokens=sum(session.total_tokens for session in sessions),
updated_at=updated_at,
)
async def _read_owned_session(
self, session_id: str, user_id: str | None
) -> Session:
session = await self.storage.read(session_id)
if user_id is None:
return session
if session.user_id is None:
session.user_id = user_id
await self.storage.write(session)
elif session.user_id != user_id:
raise SessionNotFoundError(session_id)
return session
async def _request_completion(
self, messages: list[dict[str, Any]]
) -> Any: