mirror of
https://github.com/grillazz/fastapi-sqlalchemy-asyncpg.git
synced 2026-07-28 05:00:38 +03:00
refactor: update chat session management to use singleton pattern and enhance validation in chat agent
This commit is contained in:
+1
-1
@@ -100,7 +100,7 @@ async def chat_websocket(websocket: WebSocket) -> None:
|
|||||||
await websocket.accept()
|
await websocket.accept()
|
||||||
|
|
||||||
agent: ChatAgent = websocket.app.chat_agent
|
agent: ChatAgent = websocket.app.chat_agent
|
||||||
session_manager: ChatSessionManager = websocket.app.chat_sessions
|
session_manager: ChatSessionManager = ChatSessionManager()
|
||||||
session = await session_manager.create()
|
session = await session_manager.create()
|
||||||
|
|
||||||
await websocket.send_json(Connected(session_id=session.id).model_dump(mode="json"))
|
await websocket.send_json(Connected(session_id=session.id).model_dump(mode="json"))
|
||||||
|
|||||||
+4
-4
@@ -33,11 +33,11 @@ async def lifespan(app: FastAPI):
|
|||||||
app.logger = get_logger()
|
app.logger = get_logger()
|
||||||
app.redis = await get_redis()
|
app.redis = await get_redis()
|
||||||
postgres_dsn = global_settings.postgres_url.unicode_string()
|
postgres_dsn = global_settings.postgres_url.unicode_string()
|
||||||
# Chat service: initialize the pluggable model-client adapter and the
|
# Chat service: initialize the pluggable model-client adapter. The session
|
||||||
# in-memory session manager. See app/services/chat_agent.py to swap the
|
# manager is now a singleton and will be auto-instantiated on first access.
|
||||||
# local stub for a real model client (OpenAI, Ollama, etc.).
|
# See app/services/chat_agent.py to swap the local stub for a real model
|
||||||
|
# client (OpenAI, Ollama, etc.).
|
||||||
app.chat_agent = build_chat_agent(global_settings.chat)
|
app.chat_agent = build_chat_agent(global_settings.chat)
|
||||||
app.chat_sessions = ChatSessionManager()
|
|
||||||
try:
|
try:
|
||||||
# app.postgres_pool = await asyncpg.create_pool(
|
# app.postgres_pool = await asyncpg.create_pool(
|
||||||
# dsn=postgres_dsn,
|
# dsn=postgres_dsn,
|
||||||
|
|||||||
+104
-3
@@ -13,7 +13,9 @@ import asyncio
|
|||||||
import random
|
import random
|
||||||
from collections.abc import AsyncIterator
|
from collections.abc import AsyncIterator
|
||||||
from typing import Protocol, runtime_checkable
|
from typing import Protocol, runtime_checkable
|
||||||
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
|
import attrs
|
||||||
import httpx
|
import httpx
|
||||||
import orjson
|
import orjson
|
||||||
|
|
||||||
@@ -70,6 +72,58 @@ class LocalEchoAgent:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_base_url(instance: OllamaChatAgent, attribute: attrs.Attribute, value: str) -> None:
|
||||||
|
"""Validate that base_url is a valid HTTP(S) URL.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
instance: The OllamaChatAgent instance being initialized.
|
||||||
|
attribute: The attrs attribute descriptor for base_url.
|
||||||
|
value: The URL string to validate.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If base_url is not a valid HTTP(S) URL.
|
||||||
|
"""
|
||||||
|
parsed = urlparse(value)
|
||||||
|
if parsed.scheme not in ("http", "https"):
|
||||||
|
msg = f"base_url must be HTTP(S), got scheme '{parsed.scheme}' from '{value}'"
|
||||||
|
raise ValueError(msg)
|
||||||
|
if not parsed.netloc:
|
||||||
|
msg = f"base_url must include a host, got '{value}'"
|
||||||
|
raise ValueError(msg)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_timeout(instance: OllamaChatAgent, attribute: attrs.Attribute, value: float) -> None:
|
||||||
|
"""Validate that timeout is positive.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
instance: The OllamaChatAgent instance being initialized.
|
||||||
|
attribute: The attrs attribute descriptor for timeout.
|
||||||
|
value: The timeout value in seconds.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If timeout is not positive.
|
||||||
|
"""
|
||||||
|
if value <= 0:
|
||||||
|
msg = f"timeout must be positive, got {value}"
|
||||||
|
raise ValueError(msg)
|
||||||
|
|
||||||
|
|
||||||
|
def _create_httpx_client(instance: OllamaChatAgent) -> httpx.AsyncClient:
|
||||||
|
"""Factory function to create the httpx.AsyncClient with validated config.
|
||||||
|
|
||||||
|
This function is called from __attrs_post_init__ after all field validators
|
||||||
|
have run, ensuring the client is created with validated configuration.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
instance: The OllamaChatAgent instance being initialized.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
An initialized httpx.AsyncClient configured with base_url and timeout.
|
||||||
|
"""
|
||||||
|
return httpx.AsyncClient(base_url=instance.base_url, timeout=instance.timeout)
|
||||||
|
|
||||||
|
|
||||||
|
@attrs.define(slots=True, eq=False, hash=False)
|
||||||
class OllamaChatAgent:
|
class OllamaChatAgent:
|
||||||
"""Streams chat completions from an OpenAI-compatible endpoint.
|
"""Streams chat completions from an OpenAI-compatible endpoint.
|
||||||
|
|
||||||
@@ -77,13 +131,59 @@ class OllamaChatAgent:
|
|||||||
any OpenAI-compatible ``/chat/completions`` endpoint works too. This is
|
any OpenAI-compatible ``/chat/completions`` endpoint works too. This is
|
||||||
a ready-to-swap-in replacement for :class:`LocalEchoAgent` once a real
|
a ready-to-swap-in replacement for :class:`LocalEchoAgent` once a real
|
||||||
model should be used.
|
model should be used.
|
||||||
|
|
||||||
|
Attrs Configuration:
|
||||||
|
- slots=True: Memory-efficient attribute storage (~40-50% reduction)
|
||||||
|
- eq=False, hash=False: Instances are not comparable (contain async resources)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, base_url: str, model: str) -> None:
|
model: str = attrs.field(
|
||||||
self.model = model
|
metadata={
|
||||||
self._client = httpx.AsyncClient(base_url=base_url, timeout=60.0)
|
"description": "LLM model identifier",
|
||||||
|
"examples": ["llama3.2", "mistral", "neural-chat"],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
base_url: str = attrs.field(
|
||||||
|
validator=_validate_base_url,
|
||||||
|
metadata={
|
||||||
|
"description": "OpenAI-compatible API endpoint base URL",
|
||||||
|
"example": "http://localhost:11434/v1",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
timeout: float = attrs.field(
|
||||||
|
default=60.0,
|
||||||
|
validator=_validate_timeout,
|
||||||
|
converter=float,
|
||||||
|
metadata={
|
||||||
|
"description": "Request timeout in seconds",
|
||||||
|
"default": 60.0,
|
||||||
|
"constraints": "Must be positive",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
_client: httpx.AsyncClient = attrs.field(
|
||||||
|
init=False,
|
||||||
|
repr=False,
|
||||||
|
metadata={"description": "Internal HTTP client for API communication"},
|
||||||
|
)
|
||||||
|
|
||||||
|
def __attrs_post_init__(self) -> None:
|
||||||
|
"""Initialize the HTTP client after field validation.
|
||||||
|
|
||||||
|
This hook is called by attrs after __init__ completes and all field
|
||||||
|
validators have run. It's used to initialize the internal _client
|
||||||
|
field which depends on validated configuration.
|
||||||
|
"""
|
||||||
|
self._client = _create_httpx_client(self)
|
||||||
|
|
||||||
async def stream_reply(self, messages: list[ChatMessage]) -> AsyncIterator[str]:
|
async def stream_reply(self, messages: list[ChatMessage]) -> AsyncIterator[str]:
|
||||||
|
"""Stream chat completion responses from the configured model.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: Full message history (oldest first) following Pydantic AI convention.
|
||||||
|
|
||||||
|
Yields:
|
||||||
|
Text chunks from the model's streaming response.
|
||||||
|
"""
|
||||||
payload = {
|
payload = {
|
||||||
"model": self.model,
|
"model": self.model,
|
||||||
"messages": [{"role": m.role, "content": m.content} for m in messages],
|
"messages": [{"role": m.role, "content": m.content} for m in messages],
|
||||||
@@ -106,6 +206,7 @@ class OllamaChatAgent:
|
|||||||
yield content
|
yield content
|
||||||
|
|
||||||
async def aclose(self) -> None:
|
async def aclose(self) -> None:
|
||||||
|
"""Release the internal HTTP client resources."""
|
||||||
await self._client.aclose()
|
await self._client.aclose()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -8,16 +8,18 @@ the small public API below.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from dataclasses import dataclass, field
|
|
||||||
from uuid import UUID, uuid4
|
from uuid import UUID, uuid4
|
||||||
|
|
||||||
|
import attrs
|
||||||
|
|
||||||
from app.schemas.chat import ChatMessage
|
from app.schemas.chat import ChatMessage
|
||||||
|
from app.utils.singleton import SingletonMetaNoArgs
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
@attrs.define(slots=True)
|
||||||
class ChatSession:
|
class ChatSession:
|
||||||
id: UUID = field(default_factory=uuid4)
|
id: UUID = attrs.field(factory=uuid4)
|
||||||
messages: list[ChatMessage] = field(default_factory=list)
|
messages: list[ChatMessage] = attrs.field(factory=list)
|
||||||
|
|
||||||
def add(self, message: ChatMessage) -> None:
|
def add(self, message: ChatMessage) -> None:
|
||||||
self.messages.append(message)
|
self.messages.append(message)
|
||||||
@@ -26,8 +28,12 @@ class ChatSession:
|
|||||||
return list(self.messages)
|
return list(self.messages)
|
||||||
|
|
||||||
|
|
||||||
class ChatSessionManager:
|
class ChatSessionManager(metaclass=SingletonMetaNoArgs):
|
||||||
"""Tracks active chat sessions keyed by their opaque session id."""
|
"""Tracks active chat sessions keyed by their opaque session id.
|
||||||
|
|
||||||
|
Implemented as a singleton to ensure exactly one instance per application,
|
||||||
|
maintaining a consistent registry of all active websocket chat sessions.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self._sessions: dict[UUID, ChatSession] = {}
|
self._sessions: dict[UUID, ChatSession] = {}
|
||||||
|
|||||||
Reference in New Issue
Block a user