130 lines
4.1 KiB
Python
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
|