LLM-translator/backend/sessions.py

111 lines
3.7 KiB
Python

"""Persistent translation session store with transient in-memory progress."""
import threading
from typing import Any
from backend.database import (
Database,
SessionArchivedError,
SessionConflictError,
database,
)
from backend.models import PhaseProgress, TranslationSession
class SessionStore:
def __init__(self, db: Database | None = None):
self.db = db or database
self._progress: dict[tuple[str, str], dict[str, Any]] = {}
self._progress_lock = threading.RLock()
def create(self, session_data: TranslationSession, owner_id: str) -> str:
return self.db.create_session(session_data, owner_id)
def get(self, session_id: str, owner_id: str) -> TranslationSession | None:
result = self.get_with_version(session_id, owner_id)
return result[0] if result else None
def get_with_version(
self, session_id: str, owner_id: str
) -> tuple[TranslationSession, int] | None:
result = self.db.get_session(session_id, owner_id)
if not result:
return None
session, version, _ = result
with self._progress_lock:
progress = self._progress.get((owner_id, session_id))
if progress:
session = session.model_copy(update={"progress": PhaseProgress(**progress)})
return session, version
def update(
self,
session_id: str,
data: dict[str, Any],
owner_id: str,
expected_version: int | None = None,
*,
snapshot_reason: str | None = None,
current_phase: int | None = None,
) -> TranslationSession | None:
result = self.db.update_session(
session_id,
owner_id,
data,
expected_version=expected_version,
snapshot_reason=snapshot_reason,
current_phase=current_phase,
)
if not result:
return None
self.set_progress(session_id, None, owner_id)
return result[0]
def delete(self, session_id: str, owner_id: str) -> bool:
self.set_progress(session_id, None, owner_id)
return self.db.delete_session(session_id, owner_id)
def set_progress(
self, session_id: str, progress: dict[str, Any] | None, owner_id: str
) -> bool:
key = (owner_id, session_id)
with self._progress_lock:
if progress is None:
self._progress.pop(key, None)
else:
self._progress[key] = progress
return True
def list(self, owner_id: str, include_archived: bool = False, limit: int = 50):
return self.db.list_sessions(
owner_id, include_archived=include_archived, limit=limit
)
def metadata(self, session_id: str, owner_id: str) -> dict[str, Any] | None:
result = self.db.get_session(session_id, owner_id)
return result[2] if result else None
def rename(self, session_id: str, owner_id: str, title: str) -> bool:
return self.db.rename_session(session_id, owner_id, title)
def archive(self, session_id: str, owner_id: str, archived: bool = True) -> bool:
return self.db.set_archived(session_id, owner_id, archived)
def clone(self, session_id: str, owner_id: str) -> str | None:
return self.db.clone_session(session_id, owner_id)
def get_revision(self, session_id: str, owner_id: str):
return self.db.get_revision(session_id, owner_id)
def restore_revision(self, session_id: str, owner_id: str) -> bool:
return self.db.restore_revision(session_id, owner_id)
session_store = SessionStore()
__all__ = [
"SessionArchivedError",
"SessionConflictError",
"SessionStore",
"session_store",
]