64 lines
2.2 KiB
Python
64 lines
2.2 KiB
Python
"""Chat orchestration without interface or provider-specific details."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
|
|
from javis.memory.sqlite_store import ChatSession, SQLiteSessionStore
|
|
from javis.providers.base import ChatMessage, LocalModelProvider
|
|
|
|
|
|
class SessionProviderMismatchError(RuntimeError):
|
|
"""The active provider cannot safely continue the stored session."""
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class LoadedSession:
|
|
session: ChatSession
|
|
messages: list[ChatMessage]
|
|
|
|
|
|
class ChatService:
|
|
def __init__(
|
|
self,
|
|
store: SQLiteSessionStore,
|
|
provider: LocalModelProvider,
|
|
) -> None:
|
|
self.store = store
|
|
self.provider = provider
|
|
|
|
def new_session(self) -> ChatSession:
|
|
return self.store.create_session(self.provider.name, self.provider.model)
|
|
|
|
def list_sessions(self) -> list[ChatSession]:
|
|
return self.store.list_sessions()
|
|
|
|
def load_session(self, session_id: str) -> LoadedSession:
|
|
session = self.store.get_session(session_id)
|
|
supports_session = getattr(self.provider, "supports_session", None)
|
|
compatible = (
|
|
supports_session(session.provider, session.model)
|
|
if callable(supports_session)
|
|
else session.provider == self.provider.name and session.model == self.provider.model
|
|
)
|
|
if not compatible:
|
|
raise SessionProviderMismatchError(
|
|
"Die Sitzung verwendet "
|
|
f"{session.provider}/{session.model}, aktiv ist "
|
|
f"{self.provider.name}/{self.provider.model}."
|
|
)
|
|
return LoadedSession(session, self.store.get_messages(session_id))
|
|
|
|
def send(self, session_id: str, text: str) -> str:
|
|
normalized = text.strip()
|
|
if not normalized:
|
|
raise ValueError("Eine leere Nachricht wird nicht gesendet.")
|
|
loaded = self.load_session(session_id)
|
|
messages = [*loaded.messages, ChatMessage("user", normalized)]
|
|
response = self.provider.chat(messages)
|
|
self.store.append_exchange(session_id, normalized, response)
|
|
return response
|
|
|
|
def clear_session(self, session_id: str) -> None:
|
|
self.store.clear_messages(session_id)
|