111 lines
3.7 KiB
Python
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",
|
|
]
|