"""Training session store. A session = one trainee's one-shot chat attempt against one persona. It records the full transcript + internal state + outcome + debrief. One user may have at most one session per persona (one-shot rule), enforced here. """ from __future__ import annotations import datetime from pathlib import Path from typing import Any from ..storage.store import JsonStore, new_id def _now() -> str: return datetime.datetime.now(datetime.timezone.utc).isoformat() class SessionStore: def __init__(self, data_dir: Path) -> None: self.sessions = JsonStore(data_dir / "sessions") def create( self, *, user_id: str, group_id: str, persona_id: str, persona_name: str, persona_meta: dict[str, Any] | None = None, ) -> dict[str, Any]: # One-shot: reject if the user already has a finished session on this persona existing = self.sessions.where( lambda r: r.get("user_id") == user_id and r.get("persona_id") == persona_id and r.get("outcome") in ("won", "lost") ) if existing: raise ValueError("you have already trained on this persona (one-shot)") sid = new_id("session") session = { "id": sid, "user_id": user_id, "group_id": group_id, "persona_id": persona_id, "persona_name": persona_name, "persona_meta": persona_meta or {}, "status": "active", # active | finished "outcome": None, # won | lost | abandoned "messages": [], # [{role, text, ts}] "internal": {"trust": 50, "pain_progress": {}, "buying_signals": [], "tier": None}, "debrief": None, "created_at": _now(), "updated_at": _now(), } return self.sessions.create(session, key=sid) def get(self, sid: str) -> dict[str, Any]: return self.sessions.get(sid) def get_or_none(self, sid: str) -> dict[str, Any] | None: return self.sessions.get_or_none(sid) def update(self, sid: str, **fields: Any) -> dict[str, Any]: fields.setdefault("updated_at", _now()) return self.sessions.update(sid, **fields) def active_for_persona(self, user_id: str, persona_id: str) -> dict[str, Any] | None: hits = self.sessions.where( lambda r: r.get("user_id") == user_id and r.get("persona_id") == persona_id and r.get("status") == "active" ) return hits[0] if hits else None def list_for_user(self, user_id: str) -> list[dict[str, Any]]: return sorted( self.sessions.where(lambda r: r.get("user_id") == user_id), key=lambda r: r.get("created_at", ""), )