Files

130 lines
4.1 KiB
Python

"""SQLAlchemy adapter for the session/message aggregate."""
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any
from sqlalchemy import select
from sqlalchemy.orm import Session
from ..models import Message, TrainingSession
from .errors import EntityNotFoundError, InvalidRepositoryField
class SqlAlchemySessionRepository:
"""Tenant-scoped training-session and message persistence adapter."""
_MUTABLE_SESSION_FIELDS = frozenset(
{
"status",
"outcome",
"scenario_json",
"internal_json",
"debrief_json",
}
)
def __init__(self, session: Session):
self._session = session
def get(self, session_id: str, *, org_id: str) -> TrainingSession | None:
statement = select(TrainingSession).where(
TrainingSession.id == session_id,
TrainingSession.org_id == org_id,
)
return self._session.scalars(statement).first()
def list_for_user(self, org_id: str, user_id: str) -> list[TrainingSession]:
statement = (
select(TrainingSession)
.where(
TrainingSession.org_id == org_id,
TrainingSession.user_id == user_id,
)
.order_by(TrainingSession.id)
)
return list(self._session.scalars(statement).all())
def create(
self,
*,
session_id: str,
org_id: str,
user_id: str,
group_id: str,
persona_id: str,
mode: str = "trainee",
status: str = "active",
outcome: str | None = None,
scenario_json: dict[str, Any] | None = None,
internal_json: dict[str, Any] | None = None,
debrief_json: dict[str, Any] | None = None,
) -> TrainingSession:
training_session = TrainingSession(
id=session_id,
org_id=org_id,
user_id=user_id,
group_id=group_id,
persona_id=persona_id,
mode=mode,
status=status,
outcome=outcome,
scenario_json=scenario_json,
internal_json=internal_json,
debrief_json=debrief_json,
)
self._session.add(training_session)
self._session.flush()
return training_session
def update(self, session_id: str, *, org_id: str, **fields: object) -> TrainingSession:
training_session = self.get(session_id, org_id=org_id)
if training_session is None:
raise EntityNotFoundError(f"session not found: {session_id}")
invalid = set(fields) - self._MUTABLE_SESSION_FIELDS
if invalid:
names = ", ".join(sorted(invalid))
raise InvalidRepositoryField(f"immutable or unknown repository fields: {names}")
for field, value in fields.items():
setattr(training_session, field, value)
training_session.updated_at = datetime.now(timezone.utc)
self._session.flush()
return training_session
def list_messages(self, session_id: str, *, org_id: str) -> list[Message]:
self._require_session(session_id, org_id=org_id)
statement = (
select(Message)
.where(Message.session_id == session_id)
.order_by(Message.sequence)
)
return list(self._session.scalars(statement).all())
def append_message(
self,
*,
session_id: str,
org_id: str,
message_id: str,
sequence: int,
role: str,
text: str,
) -> Message:
self._require_session(session_id, org_id=org_id)
message = Message(
id=message_id,
session_id=session_id,
sequence=sequence,
role=role,
text=text,
)
self._session.add(message)
self._session.flush()
return message
def _require_session(self, session_id: str, *, org_id: str) -> TrainingSession:
training_session = self.get(session_id, org_id=org_id)
if training_session is None:
raise EntityNotFoundError(f"session not found: {session_id}")
return training_session