"""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", ]