159 lines
5.7 KiB
Python
159 lines
5.7 KiB
Python
import asyncio
|
|
from collections.abc import AsyncGenerator
|
|
from datetime import UTC, datetime
|
|
from openai import APIError, AsyncOpenAI
|
|
|
|
from models import Message, SendMessageResponse, Session, TokenUsage
|
|
from storage import JsonSessionStorage
|
|
|
|
|
|
class ChatProviderError(Exception):
|
|
pass
|
|
|
|
|
|
class ChatService:
|
|
def __init__(
|
|
self,
|
|
storage: JsonSessionStorage,
|
|
client: AsyncOpenAI,
|
|
model: str,
|
|
) -> None:
|
|
self.storage = storage
|
|
self.client = client
|
|
self.model = model
|
|
self._locks: dict[str, asyncio.Lock] = {}
|
|
|
|
async def generate_response(
|
|
self, session_id: str, content: str
|
|
) -> SendMessageResponse:
|
|
lock = self._locks.setdefault(session_id, asyncio.Lock())
|
|
async with lock:
|
|
user_created_at = datetime.now(UTC)
|
|
session = await self.storage.read(session_id)
|
|
api_messages = self._build_api_messages(session, content)
|
|
|
|
try:
|
|
completion = await self.client.chat.completions.create(
|
|
model=self.model,
|
|
messages=api_messages, # type: ignore[arg-type]
|
|
stream=False,
|
|
extra_body={"thinking": {"type": "disabled"}},
|
|
)
|
|
assistant_content = completion.choices[0].message.content
|
|
except (APIError, IndexError, AttributeError) as exc:
|
|
raise ChatProviderError("Upstream chat request failed") from exc
|
|
|
|
if assistant_content is None:
|
|
raise ChatProviderError("Upstream chat provider returned an empty response")
|
|
|
|
assistant_message = await self._save_exchange(
|
|
session, content, user_created_at, assistant_content
|
|
)
|
|
usage = self._parse_usage(completion.usage)
|
|
|
|
return SendMessageResponse(
|
|
session_id=session_id,
|
|
message=assistant_message,
|
|
usage=usage,
|
|
)
|
|
|
|
async def ensure_session_exists(self, session_id: str) -> None:
|
|
await self.storage.read(session_id)
|
|
|
|
async def generate_response_stream(
|
|
self, session_id: str, content: str
|
|
) -> AsyncGenerator[dict[str, object], None]:
|
|
lock = self._locks.setdefault(session_id, asyncio.Lock())
|
|
async with lock:
|
|
user_created_at = datetime.now(UTC)
|
|
session = await self.storage.read(session_id)
|
|
api_messages = self._build_api_messages(session, content)
|
|
|
|
try:
|
|
stream = await self.client.chat.completions.create(
|
|
model=self.model,
|
|
messages=api_messages, # type: ignore[arg-type]
|
|
stream=True,
|
|
stream_options={"include_usage": True},
|
|
extra_body={"thinking": {"type": "disabled"}},
|
|
)
|
|
parts: list[str] = []
|
|
usage: TokenUsage | None = None
|
|
|
|
try:
|
|
async for chunk in stream:
|
|
chunk_usage = getattr(chunk, "usage", None)
|
|
if chunk_usage is not None:
|
|
usage = self._parse_usage(chunk_usage)
|
|
|
|
for choice in chunk.choices:
|
|
delta = choice.delta.content
|
|
if delta:
|
|
parts.append(delta)
|
|
yield {"type": "delta", "content": delta}
|
|
finally:
|
|
await stream.close()
|
|
except (APIError, AttributeError, TypeError) as exc:
|
|
raise ChatProviderError("Upstream chat stream failed") from exc
|
|
|
|
assistant_content = "".join(parts)
|
|
if not assistant_content:
|
|
raise ChatProviderError("Upstream chat provider returned an empty response")
|
|
|
|
assistant_message = await self._save_exchange(
|
|
session, content, user_created_at, assistant_content
|
|
)
|
|
yield {
|
|
"type": "done",
|
|
"message": assistant_message.model_dump(mode="json"),
|
|
"usage": usage.model_dump() if usage is not None else None,
|
|
}
|
|
|
|
@staticmethod
|
|
def _build_api_messages(
|
|
session: Session, content: str
|
|
) -> list[dict[str, str]]:
|
|
messages = [{"role": "system", "content": session.system_prompt}]
|
|
messages.extend(
|
|
{"role": message.role, "content": message.content}
|
|
for message in session.messages
|
|
)
|
|
messages.append({"role": "user", "content": content})
|
|
return messages
|
|
|
|
async def _save_exchange(
|
|
self,
|
|
session: Session,
|
|
user_content: str,
|
|
user_created_at: datetime,
|
|
assistant_content: str,
|
|
) -> Message:
|
|
assistant_message = Message(
|
|
role="assistant",
|
|
content=assistant_content,
|
|
created_at=datetime.now(UTC),
|
|
)
|
|
session.messages.extend(
|
|
[
|
|
Message(
|
|
role="user",
|
|
content=user_content,
|
|
created_at=user_created_at,
|
|
),
|
|
assistant_message,
|
|
]
|
|
)
|
|
session.updated_at = assistant_message.created_at
|
|
await self.storage.write(session)
|
|
return assistant_message
|
|
|
|
@staticmethod
|
|
def _parse_usage(usage: object | None) -> TokenUsage | None:
|
|
if usage is None:
|
|
return None
|
|
return TokenUsage(
|
|
prompt_tokens=usage.prompt_tokens, # type: ignore[attr-defined]
|
|
completion_tokens=usage.completion_tokens, # type: ignore[attr-defined]
|
|
total_tokens=usage.total_tokens, # type: ignore[attr-defined]
|
|
)
|