增加鉴权系统
This commit is contained in:
+46
-7
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user