Files
simple-chat-api/service.py
T
2026-07-03 20:49:36 +08:00

216 lines
8.4 KiB
Python

import asyncio
from datetime import UTC, datetime
from typing import Any
from openai import APIError, AsyncOpenAI
from domain import (
AssistantMessage,
Message,
Session,
ToolCall,
ToolFunctionCall,
ToolMessage,
UserMessage,
)
from schemas import (
CreateSessionResponse,
FinalAssistantMessage,
SendMessageResponse,
SessionHistoryResponse,
TokenUsage,
)
from storage import JsonSessionStorage
from tools import ToolRegistry
class ChatProviderError(Exception):
pass
class ChatService:
def __init__(
self,
storage: JsonSessionStorage,
client: AsyncOpenAI,
model: str,
default_system_prompt: str = "You are a helpful assistant.",
tool_registry: ToolRegistry | None = None,
max_tool_rounds: int = 5,
max_tool_calls_per_turn: int = 10,
) -> None:
if max_tool_rounds <= 0 or max_tool_calls_per_turn <= 0:
raise ValueError("tool limits must be greater than zero")
self.storage = storage
self.client = client
self.model = model
self.default_system_prompt = default_system_prompt
self.tool_registry = tool_registry or ToolRegistry()
self.max_tool_rounds = max_tool_rounds
self.max_tool_calls_per_turn = max_tool_calls_per_turn
self._locks: dict[str, asyncio.Lock] = {}
async def create_session(
self, system_prompt: str | None = None
) -> CreateSessionResponse:
session = await self.storage.create(
system_prompt or self.default_system_prompt
)
return CreateSessionResponse(
session_id=session.session_id,
system_prompt=session.system_prompt,
created_at=session.created_at,
)
async def generate_response(
self, session_id: str, content: str
) -> 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)
api_messages = self._build_api_messages(session)
api_messages.append(self._message_to_api(user_message))
pending_messages: list[Message] = [user_message]
tools_use: list[str] = []
usage = TokenUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0)
has_usage = False
tool_call_count = 0
for _ in range(self.max_tool_rounds):
completion = await self._request_completion(api_messages)
if completion.usage is not None:
self._add_usage(usage, completion.usage)
has_usage = True
provider_message = completion.choices[0].message
provider_tool_calls = getattr(provider_message, "tool_calls", None) or []
if not provider_tool_calls:
assistant_content = provider_message.content
if not assistant_content:
raise ChatProviderError(
"Upstream chat provider returned an empty response"
)
final_message = AssistantMessage(
content=assistant_content,
created_at=datetime.now(UTC),
)
pending_messages.append(final_message)
session.messages.extend(pending_messages)
session.updated_at = final_message.created_at
await self.storage.write(session)
return SendMessageResponse(
session_id=session_id,
message=FinalAssistantMessage(
content=assistant_content,
created_at=final_message.created_at,
),
tools_use=tools_use,
usage=usage if has_usage else None,
)
tool_call_count += len(provider_tool_calls)
if tool_call_count > self.max_tool_calls_per_turn:
raise ChatProviderError("Tool call limit exceeded")
tool_calls = [
ToolCall(
id=tool_call.id,
function=ToolFunctionCall(
name=tool_call.function.name,
arguments=tool_call.function.arguments,
),
)
for tool_call in provider_tool_calls
]
assistant_tool_message = AssistantMessage(
content=provider_message.content,
tool_calls=tool_calls,
created_at=datetime.now(UTC),
)
pending_messages.append(assistant_tool_message)
api_messages.append(self._message_to_api(assistant_tool_message))
for tool_call in tool_calls:
if tool_call.function.name not in tools_use:
tools_use.append(tool_call.function.name)
result = await self.tool_registry.execute(
tool_call.function.name,
tool_call.function.arguments,
)
tool_message = ToolMessage(
content=result.content,
tool_call_id=tool_call.id,
name=tool_call.function.name,
is_error=result.is_error,
created_at=datetime.now(UTC),
)
pending_messages.append(tool_message)
api_messages.append(self._message_to_api(tool_message))
raise ChatProviderError("Tool round limit exceeded")
async def get_session_history(self, session_id: str) -> SessionHistoryResponse:
lock = self._locks.setdefault(session_id, asyncio.Lock())
async with lock:
session = await self.storage.read(session_id)
return SessionHistoryResponse(
session_id=session.session_id,
system_prompt=session.system_prompt,
created_at=session.created_at,
updated_at=session.updated_at,
messages=session.messages,
)
async def _request_completion(
self, messages: list[dict[str, Any]]
) -> Any:
try:
completion = await self.client.chat.completions.create(
model=self.model,
messages=messages, # type: ignore[arg-type]
tools=self.tool_registry.definitions(), # type: ignore[arg-type]
extra_body={"thinking": {"type": "disabled"}},
)
if not completion.choices:
raise ChatProviderError("Upstream chat provider returned no choices")
return completion
except ChatProviderError:
raise
except (APIError, IndexError, AttributeError) as exc:
raise ChatProviderError("Upstream chat request failed") from exc
@classmethod
def _build_api_messages(cls, session: Session) -> list[dict[str, Any]]:
messages: list[dict[str, Any]] = [
{"role": "system", "content": session.system_prompt}
]
messages.extend(cls._message_to_api(message) for message in session.messages)
return messages
@staticmethod
def _message_to_api(message: Message) -> dict[str, Any]:
if isinstance(message, UserMessage):
return {"role": "user", "content": message.content}
if isinstance(message, AssistantMessage):
result: dict[str, Any] = {
"role": "assistant",
"content": message.content,
}
if message.tool_calls:
result["tool_calls"] = [
tool_call.model_dump() for tool_call in message.tool_calls
]
return result
return {
"role": "tool",
"content": message.content,
"tool_call_id": message.tool_call_id,
}
@staticmethod
def _add_usage(total: TokenUsage, usage: object) -> None:
total.prompt_tokens += usage.prompt_tokens # type: ignore[attr-defined]
total.completion_tokens += usage.completion_tokens # type: ignore[attr-defined]
total.total_tokens += usage.total_tokens # type: ignore[attr-defined]