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]