vignette/apps/api/app/routes/sessions.py
2026-06-26 14:47:00 +09:00

990 lines
32 KiB
Python

"""Counseling session routes.
The DB-backed source of truth is still pending, so this route uses the existing
in-process session store when DB is degraded. Unlike the previous dev fallback,
all browser calls now require a verified server-side auth session and every
session operation checks learner ownership.
"""
from __future__ import annotations
import asyncio
import json
from collections import Counter
from datetime import datetime
from typing import Literal, Optional
from fastapi import APIRouter, HTTPException, status
from pydantic import BaseModel, Field
from sse_starlette.sse import EventSourceResponse
from .. import session_persistence
from ..config import settings
from ..deps import CurrentPrincipal, Principal, Role
from ..engine_client import EngineError, engine_client
from ..persona_repository import get_catalog_persona
from ..runtime_policy import require_runtime_fallback_allowed, runtime_fallback_allowed
from ..services import evaluator, memory, orchestrator, state_machine
from ..store import InProcSession, TurnRecord, store
router = APIRouter(prefix="/sessions", tags=["sessions"])
TheoryMode = Literal["humanistic", "cbt", "integrative"]
class SessionStartRequest(BaseModel):
persona_code: str = Field(..., examples=["P1"])
theory_mode: TheoryMode = "humanistic"
class SessionStartResponse(BaseModel):
session_id: str
case_id: str
session_no: int
stage: str
effective_openness: float
recall_summary: Optional[str] = None
degraded: bool = False
class TurnRequest(BaseModel):
text: str = Field(..., min_length=1)
class TurnResponse(BaseModel):
turn_seq: int
stage: str
effective_openness: float
client_reply: Optional[str] = None
safety_flagged: bool = False
crisis_kind: str = "none"
class SessionEndResponse(BaseModel):
session_id: str
session_no: int
digest_pending: bool
end_state: dict
class LearnerSessionSummary(BaseModel):
session_id: str
persona_code: str
persona_name: str
session_no: int
status: Literal["active", "ended"]
stage: str
turn_count: int
learner_turn_count: int
client_turn_count: int
started_at: str
ended_at: str | None = None
review_ready: bool = False
class LearnerSessionsResponse(BaseModel):
source: str = "runtime"
sessions: list[LearnerSessionSummary] = Field(default_factory=list)
class SessionDetailTurn(BaseModel):
turn_seq: int
speaker: Literal["learner", "client"]
stage: str
text: str
created_at: str
class SessionDetailResponse(BaseModel):
session_id: str
case_id: str
persona_code: str
persona_name: str
theory_mode: str
status: Literal["active", "ended"]
stage: str
effective_openness: float
started_at: str
ended_at: str | None = None
turns: list[SessionDetailTurn] = Field(default_factory=list)
review_ready: bool = False
class ReviewClient(BaseModel):
name: str
initial: str
persona: str
class ReviewTechnique(BaseModel):
kind: str
label: str
class ReviewNote(BaseModel):
author: str
tone: str
title: str
body: str
quote: Optional[str] = None
class ReviewTurn(BaseModel):
id: str
ts: str
speaker: Literal["learner", "client"]
who: str
text: str
techniques: list[ReviewTechnique] = Field(default_factory=list)
note: Optional[ReviewNote] = None
class ReviewPhaseSegment(BaseModel):
key: str
label: str
weight: float
class ReviewValencePoint(BaseModel):
t: float
v: float
class ReviewRubricRow(BaseModel):
name: str
cluster: str
ratio: float
quality: Literal["good", "watch"]
freq: str
class ReviewPoint(BaseModel):
title: str
body: str
jumpTo: Optional[str] = None
class SessionReviewResponse(BaseModel):
session_id: str
client: ReviewClient
date: str
durationLabel: str
durationSeconds: int
reachedPhase: str
sessionSignal: str
supervisorState: str
supervisorName: str
summary: str
phases: list[ReviewPhaseSegment] = Field(default_factory=list)
phaseAxis: list[str] = Field(default_factory=list)
valenceAxis: list[str] = Field(default_factory=list)
clientValence: list[ReviewValencePoint] = Field(default_factory=list)
counselorBaseline: list[ReviewValencePoint] = Field(default_factory=list)
turns: list[ReviewTurn] = Field(default_factory=list)
rubric: list[ReviewRubricRow] = Field(default_factory=list)
goodMoments: list[ReviewPoint] = Field(default_factory=list)
growthPoints: list[ReviewPoint] = Field(default_factory=list)
nextLine: Optional[str] = None
clientFeedback: Optional[str] = None
audioUrl: Optional[str] = None
pdfExportUrl: Optional[str] = None
degraded: bool = True
reviewReady: bool = False
_RECALL_CACHE: dict[str, memory.RecallContext] = {}
_PHASE_KEY_BY_LABEL = {
"라포": "rapport",
"탐색": "explore",
"개입": "intervene",
"정리": "closing",
}
def _stage_label(stage: object) -> str:
name = getattr(stage, "name", "")
return {
"RAPPORT": "라포",
"EXPLORE": "탐색",
"INTERVENE": "개입",
"CLOSE": "정리",
}.get(name, str(getattr(stage, "value", stage)))
def _ensure_learner(principal: Principal) -> None:
if principal.role != Role.LEARNER:
raise HTTPException(status.HTTP_403_FORBIDDEN, detail="only learners can use sessions")
async def _load_session_or_404(
session_id: str,
principal: Principal,
*,
allow_ended: bool = False,
) -> InProcSession:
sess = await session_persistence.load_session(session_id, principal, allow_ended=True)
if sess is not None:
store.put(sess)
elif runtime_fallback_allowed():
sess = store.get(session_id)
if sess is None:
raise HTTPException(status.HTTP_404_NOT_FOUND, detail="session not found")
if sess.learner_id != principal.user_id:
raise HTTPException(status.HTTP_403_FORBIDDEN, detail="session does not belong to user")
if sess.ended and not allow_ended:
raise HTTPException(status.HTTP_409_CONFLICT, detail="session already ended")
return sess
async def _append_session_turn(sess: InProcSession, turn: TurnRecord) -> None:
if await session_persistence.append_turn(
session_id=sess.session_id,
learner_id=sess.learner_id,
turn=turn,
):
sess.turns.append(turn)
store.put(sess)
return
require_runtime_fallback_allowed("session turn append")
store.append_turn(sess.session_id, turn)
async def _update_session_state(
sess: InProcSession,
state: state_machine.SessionState,
) -> None:
if await session_persistence.update_state(
session_id=sess.session_id,
learner_id=sess.learner_id,
state=state,
):
sess.state = state
store.put(sess)
return
require_runtime_fallback_allowed("session state update")
store.update_state(sess.session_id, state)
async def _end_persisted_session(sess: InProcSession, carry: memory.CarryOver) -> None:
if await session_persistence.end_session(sess, carry):
sess.ended = True
sess.ended_at = datetime.now().timestamp()
store.put(sess)
return
require_runtime_fallback_allowed("session end")
store.end(sess.session_id)
def _offset_label(seconds: float) -> str:
whole = max(0, int(round(seconds)))
minutes, sec = divmod(whole, 60)
return f"{minutes}:{sec:02d}"
def _iso(ts: float | None) -> str | None:
if ts is None:
return None
return datetime.fromtimestamp(ts).isoformat(timespec="seconds")
def _duration_label(seconds: int) -> str:
if seconds < 60:
return f"{seconds}"
minutes, sec = divmod(seconds, 60)
return f"{minutes}{sec}"
def _client_name(raw: str) -> str:
name = raw.split("(", 1)[0].strip()
return name or raw.strip() or "내담자"
def _review_summary(*, client_name: str, reached_phase: str, turns: list[ReviewTurn]) -> str:
if not turns:
return (
"아직 실제 발화가 없어 리뷰를 만들 수 없습니다. 회기를 진행한 뒤 종료하면 "
"저장된 축어록을 기준으로 리뷰가 표시됩니다."
)
learner_count = sum(1 for turn in turns if turn.speaker == "learner")
client_count = sum(1 for turn in turns if turn.speaker == "client")
return (
f"이 리뷰는 현재 세션에 저장된 실제 축어록 {len(turns)}개를 기반으로 합니다. "
f"{client_name}와의 회기는 {reached_phase} 단계까지 진행되었고, "
f"학습자 발화 {learner_count}개와 내담자 응답 {client_count}개가 기록되었습니다. "
"평가 AI 또는 교수자 코멘트가 아직 생성되지 않은 항목은 빈 상태로 남겨 둡니다."
)
def _phase_segments(stage_labels: list[str]) -> list[ReviewPhaseSegment]:
counts = Counter(stage_labels)
return [
ReviewPhaseSegment(
key=_PHASE_KEY_BY_LABEL.get(label, label),
label=label,
weight=float(count),
)
for label, count in counts.items()
if count > 0
]
def _clamp_ratio(value: float) -> float:
return round(max(0.0, min(1.0, value)), 3)
def _compact_text(text: str) -> str:
return " ".join(text.split())
def _clip_text(text: str, limit: int = 180) -> str:
compact = _compact_text(text)
if len(compact) <= limit:
return compact
return f"{compact[: max(0, limit - 1)].rstrip()}..."
def _point_title(text: str, fallback: str) -> str:
compact = _clip_text(text, 72)
for sep in (".", "", "!", "?", "\n"):
if sep in compact:
first = compact.split(sep, 1)[0].strip()
if first:
return _clip_text(first, 44)
return _clip_text(compact, 44) or fallback
def _ai_review_points(values: object, *, fallback_prefix: str) -> list[ReviewPoint]:
if not isinstance(values, list):
return []
points: list[ReviewPoint] = []
for index, value in enumerate(values, start=1):
body = _compact_text(str(value or ""))
if not body:
continue
points.append(
ReviewPoint(
title=_point_title(body, f"{fallback_prefix} {index}"),
body=body,
jumpTo=None,
)
)
return points[:3]
def _intent_deviation_points(values: object) -> list[ReviewPoint]:
if not isinstance(values, list):
return []
points: list[ReviewPoint] = []
for index, value in enumerate(values, start=1):
if not isinstance(value, dict):
continue
dimension = _compact_text(str(value.get("dimension") or f"의도 이탈 {index}"))
expected = _compact_text(str(value.get("expected") or ""))
actual = _compact_text(str(value.get("actual") or ""))
severity = _compact_text(str(value.get("severity") or "minor"))
body_parts = []
if expected:
body_parts.append(f"기대: {expected}")
if actual:
body_parts.append(f"실제: {actual}")
if severity:
body_parts.append(f"심각도: {severity}")
if body_parts:
points.append(
ReviewPoint(
title=dimension,
body=" · ".join(body_parts),
jumpTo=None,
)
)
return points[:3]
def _rubric_from_evaluation(payload: dict[str, object]) -> list[ReviewRubricRow]:
distribution = payload.get("distribution")
if not isinstance(distribution, dict):
return []
by_category = distribution.get("by_category")
if not isinstance(by_category, dict):
return []
total = int(distribution.get("total") or 0)
if total <= 0:
return []
overused = {str(item) for item in distribution.get("overused") or []}
underused = {str(item) for item in distribution.get("underused") or []}
rows: list[ReviewRubricRow] = []
for category, raw_count in sorted(by_category.items(), key=lambda item: str(item[0])):
try:
count = int(raw_count)
except (TypeError, ValueError):
continue
code = str(category)
watch = code in overused or code in underused
rows.append(
ReviewRubricRow(
name=code.replace("_", " ").title(),
cluster="평가 AI 기법 분포",
ratio=_clamp_ratio(count / max(1, total)),
quality="watch" if watch else "good",
freq=f"{count}/{total} labels",
)
)
return rows
def _review_summary_from_evaluation(
*,
fallback: str,
evaluation_record: dict[str, object] | None,
payload: dict[str, object],
) -> str:
if not evaluation_record:
return fallback
status = str(evaluation_record.get("status") or "")
if status != "ready":
error = _compact_text(str(evaluation_record.get("error") or payload.get("error") or ""))
return (
"저장된 축어록은 확인했지만 평가 AI 산출물이 아직 준비되지 않았습니다. "
+ (f"사유: {error}" if error else "평가가 완료되면 코칭 항목이 갱신됩니다.")
)
rationale = _compact_text(str(payload.get("supervisor_rationale") or ""))
critique = _compact_text(str(payload.get("supervisor_critique") or ""))
evaluated = payload.get("turns_evaluated")
prefix = f"평가 AI가 학습자 발화 {evaluated}개를 deep-loop로 분석했습니다. "
details = " ".join(part for part in [rationale, critique] if part)
return prefix + (details if details else "아래 코칭 항목은 저장된 축어록과 평가 AI 결과를 기준으로 합니다.")
def _next_line_from_evaluation(payload: dict[str, object]) -> str | None:
alternatives = payload.get("alternative_utterances")
if not isinstance(alternatives, list):
return None
for value in alternatives:
line = _compact_text(str(value or ""))
if line:
return line
return None
def _latest_client_feedback(turns: list[ReviewTurn]) -> str | None:
for turn in reversed(turns):
if turn.speaker == "client":
return _clip_text(turn.text)
return None
def _evaluation_payload(record: dict[str, object] | None) -> dict[str, object]:
if not record:
return {}
payload = record.get("payload")
return payload if isinstance(payload, dict) else {}
async def _generate_and_save_session_evaluation(sess: InProcSession) -> None:
if not sess.turns:
return
enriched: list[dict[str, object]] = []
for index, turn in enumerate(sess.masked_turns(), start=1):
item: dict[str, object] = dict(turn)
item["seq"] = index
enriched.append(item)
try:
result = await asyncio.wait_for(
evaluator.evaluate_session(
session_id=sess.session_id,
stage=_stage_label(sess.state.stage),
masked_turns=enriched,
engine=engine_client,
technique_codes=[],
theory_mode=sess.theory_mode,
scope="session_end",
),
timeout=min(float(settings.engine_timeout), 45.0),
)
status_value = "error" if result.error else "ready"
await session_persistence.save_session_evaluation(
session_id=sess.session_id,
learner_id=sess.learner_id,
status=status_value,
source="engine",
scope=result.scope,
stage=result.stage,
payload=result.to_dict(),
error=result.error,
)
except Exception as exc:
await session_persistence.save_session_evaluation(
session_id=sess.session_id,
learner_id=sess.learner_id,
status="error",
source="engine",
scope="session_end",
stage=_stage_label(sess.state.stage),
payload={},
error=str(exc),
)
def _schedule_session_evaluation(sess: InProcSession) -> None:
if not sess.turns:
return
asyncio.create_task(_generate_and_save_session_evaluation(sess))
def _learner_summary(sess: InProcSession, *, review_ready: bool = False) -> LearnerSessionSummary:
learner_turns = sum(1 for turn in sess.turns if turn.speaker == "counselor")
client_turns = sum(1 for turn in sess.turns if turn.speaker == "client")
return LearnerSessionSummary(
session_id=sess.session_id,
persona_code=sess.persona_code,
persona_name=sess.persona.display_name,
session_no=sess.session_no,
status="ended" if sess.ended else "active",
stage=_stage_label(sess.state.stage),
turn_count=len(sess.turns),
learner_turn_count=learner_turns,
client_turn_count=client_turns,
started_at=_iso(sess.created_at) or "",
ended_at=_iso(sess.ended_at),
review_ready=review_ready,
)
async def _review_ready(sess: InProcSession, principal: Principal) -> bool:
if not sess.ended or not sess.turns:
return False
evaluation_record, _ = await session_persistence.load_session_evaluation(
sess.session_id,
principal,
)
return bool(evaluation_record and evaluation_record.get("status") == "ready")
def _session_detail(
sess: InProcSession,
*,
review_ready: bool = False,
) -> SessionDetailResponse:
return SessionDetailResponse(
session_id=sess.session_id,
case_id=sess.case_id,
persona_code=sess.persona_code,
persona_name=sess.persona.display_name,
theory_mode=sess.theory_mode,
status="ended" if sess.ended else "active",
stage=_stage_label(sess.state.stage),
effective_openness=round(sess.state.effective_openness, 4),
started_at=_iso(sess.created_at) or "",
ended_at=_iso(sess.ended_at),
turns=[
SessionDetailTurn(
turn_seq=turn.turn_seq,
speaker="learner" if turn.speaker == "counselor" else "client",
stage=turn.stage,
text=turn.text_masked,
created_at=_iso(turn.created_at) or "",
)
for turn in sess.turns
],
review_ready=review_ready,
)
@router.get("", response_model=LearnerSessionsResponse)
async def list_learner_sessions(principal: CurrentPrincipal) -> LearnerSessionsResponse:
"""Return the current learner's real practice sessions."""
_ensure_learner(principal)
sessions, durable = await session_persistence.list_sessions(principal)
if not durable:
require_runtime_fallback_allowed("session list")
sessions = [
sess
for sess in store.list()
if sess.learner_id == principal.user_id
]
sessions.sort(key=lambda sess: sess.created_at, reverse=True)
summaries: list[LearnerSessionSummary] = []
for sess in sessions[:20]:
summaries.append(_learner_summary(sess, review_ready=await _review_ready(sess, principal)))
return LearnerSessionsResponse(
source="database" if durable else "runtime",
sessions=summaries,
)
@router.get("/{session_id}", response_model=SessionDetailResponse)
async def get_session_detail(
session_id: str,
principal: CurrentPrincipal,
) -> SessionDetailResponse:
"""Return a learner-owned session with transcript for resume/history."""
_ensure_learner(principal)
sess = await _load_session_or_404(session_id, principal, allow_ended=True)
return _session_detail(sess, review_ready=await _review_ready(sess, principal))
@router.post("", response_model=SessionStartResponse, status_code=status.HTTP_201_CREATED)
async def start_session(
body: SessionStartRequest,
principal: CurrentPrincipal,
) -> SessionStartResponse:
"""Start a learner-owned practice session."""
_ensure_learner(principal)
try:
catalog_persona = await get_catalog_persona(body.persona_code)
except Exception as exc:
raise HTTPException(
status.HTTP_503_SERVICE_UNAVAILABLE,
detail="persona catalog database unavailable",
) from exc
if catalog_persona is None:
raise HTTPException(status.HTTP_404_NOT_FOUND, detail=f"unknown persona {body.persona_code}")
card = catalog_persona.card
recall = memory.build_recall_context()
st = state_machine.init_state(
base_resistance=card.base_resistance(),
unlock_rate=card.unlock_rate(),
decay_floor=card.decay_floor(),
ideation_baseline=card.ideation_baseline(),
carry=recall.carry,
)
carry_rapport = st.rapport_credit
sess = await session_persistence.create_session(
learner_id=principal.user_id,
card=card,
theory_mode=body.theory_mode,
state=st,
session_no=1,
carry_rapport=carry_rapport,
persona_id=catalog_persona.persona_id,
persona_version=catalog_persona.version,
)
degraded = catalog_persona.degraded or sess is None
if sess is None:
require_runtime_fallback_allowed("session creation")
sess = store.create(
learner_id=principal.user_id,
persona=card,
theory_mode=body.theory_mode,
state=st,
session_no=1,
carry_rapport=carry_rapport,
)
else:
store.put(sess)
_RECALL_CACHE[sess.session_id] = recall
return SessionStartResponse(
session_id=sess.session_id,
case_id=sess.case_id,
session_no=sess.session_no,
stage=_stage_label(st.stage),
effective_openness=round(st.effective_openness, 4),
recall_summary=recall.recall_summary,
degraded=degraded,
)
@router.get("/{session_id}/review", response_model=SessionReviewResponse)
async def get_session_review(
session_id: str,
principal: CurrentPrincipal,
) -> SessionReviewResponse:
"""Return a learner-safe review built only from the stored session transcript."""
_ensure_learner(principal)
sess = await _load_session_or_404(session_id, principal, allow_ended=True)
end_ts = sess.ended_at or datetime.now().timestamp()
duration_seconds = max(0, int(round(end_ts - sess.created_at)))
client_name = _client_name(sess.persona.display_name)
client_initial = client_name[:1] or ""
reached_phase = _stage_label(sess.state.stage)
stage_labels = [turn.stage for turn in sess.turns] or [reached_phase]
axis = ["0:00"]
if duration_seconds > 0:
axis.append(_offset_label(duration_seconds))
evaluation_record, evaluation_durable = await session_persistence.load_session_evaluation(
session_id,
principal,
)
evaluation_payload = _evaluation_payload(evaluation_record)
evaluation_status = str(evaluation_record.get("status") or "") if evaluation_record else ""
evaluation_ready = evaluation_status == "ready"
first_turn_ts = sess.turns[0].created_at if sess.turns else sess.created_at
turns: list[ReviewTurn] = []
for index, turn in enumerate(sess.turns):
speaker: Literal["learner", "client"] = (
"learner" if turn.speaker == "counselor" else "client"
)
turns.append(
ReviewTurn(
id=f"t{index + 1}",
ts=_offset_label(turn.created_at - first_turn_ts),
speaker=speaker,
who="학습자" if speaker == "learner" else client_name,
text=turn.text_masked,
techniques=[],
note=None,
)
)
if not turns:
session_signal = "기록 없음"
elif sess.ended:
session_signal = "종료됨"
else:
session_signal = "진행 중"
transcript_summary = _review_summary(
client_name=client_name,
reached_phase=reached_phase,
turns=turns,
)
rubric: list[ReviewRubricRow] = []
good_moments: list[ReviewPoint] = []
growth_points: list[ReviewPoint] = []
next_line: str | None = None
if evaluation_ready:
rubric = _rubric_from_evaluation(evaluation_payload)
good_moments = _ai_review_points(
evaluation_payload.get("strengths"),
fallback_prefix="강점",
)
growth_points = _ai_review_points(
evaluation_payload.get("improvements"),
fallback_prefix="개선점",
)
if not growth_points:
growth_points = _intent_deviation_points(evaluation_payload.get("intent_deviations"))
next_line = _next_line_from_evaluation(evaluation_payload)
client_feedback = _latest_client_feedback(turns)
review_degraded = bool(turns) and not evaluation_ready
if evaluation_ready:
supervisor_state = "평가 완료"
elif evaluation_status == "error":
supervisor_state = "평가 실패"
elif turns:
supervisor_state = "평가 대기"
else:
supervisor_state = "기록 대기"
summary = _review_summary_from_evaluation(
fallback=transcript_summary,
evaluation_record=evaluation_record,
payload=evaluation_payload,
)
if evaluation_record and not evaluation_durable:
summary += " 현재 평가는 런타임 캐시에서 복원되었습니다."
return SessionReviewResponse(
session_id=session_id,
client=ReviewClient(
name=client_name,
initial=client_initial,
persona=f"{sess.persona_code} · {sess.persona.difficulty}",
),
date=datetime.fromtimestamp(sess.created_at).strftime("%Y-%m-%d"),
durationLabel=_duration_label(duration_seconds),
durationSeconds=duration_seconds,
reachedPhase=reached_phase,
sessionSignal=session_signal,
supervisorState=supervisor_state,
supervisorName="AI",
summary=summary,
phases=_phase_segments(stage_labels),
phaseAxis=axis,
valenceAxis=axis,
clientValence=[],
counselorBaseline=[],
turns=turns,
rubric=rubric,
goodMoments=good_moments,
growthPoints=growth_points,
nextLine=next_line,
clientFeedback=client_feedback,
audioUrl=None,
pdfExportUrl=None,
degraded=review_degraded,
reviewReady=evaluation_ready,
)
@router.post("/{session_id}/turn", response_model=TurnResponse)
async def submit_turn(
session_id: str,
body: TurnRequest,
principal: CurrentPrincipal,
) -> TurnResponse:
"""Submit one trainee utterance and return the generated client reply."""
_ensure_learner(principal)
sess = await _load_session_or_404(session_id, principal)
recall = _RECALL_CACHE.get(session_id) or memory.RecallContext()
ctx = orchestrator.prepare_turn(
session_id=session_id,
case_id=sess.case_id,
card=sess.persona,
state=sess.state,
learner_text=body.text,
recall_summary=recall.recall_summary,
pinned_facts=recall.pinned_facts,
recent_turns=sess.recent_turns(),
)
assert ctx.state_after is not None
try:
result = await orchestrator.run_turn_generate(ctx, engine_client)
except EngineError as exc:
raise HTTPException(
status.HTTP_503_SERVICE_UNAVAILABLE,
detail=f"engine unavailable: {exc}",
) from exc
await _append_session_turn(
sess,
TurnRecord(
turn_seq=ctx.state_after.turn_seq,
speaker="counselor",
stage=_stage_label(ctx.state_after.stage),
text=body.text,
text_masked=ctx.learner_text_masked,
),
)
if result.client_reply:
await _append_session_turn(
sess,
TurnRecord(
turn_seq=result.turn_seq,
speaker="client",
stage=_stage_label(result.state_after.stage),
text=result.client_reply,
text_masked=result.client_reply,
),
)
await _update_session_state(sess, result.state_after)
return TurnResponse(
turn_seq=result.turn_seq,
stage=_stage_label(result.state_after.stage),
effective_openness=round(result.effective_openness, 4),
client_reply=result.client_reply,
safety_flagged=result.safety_flagged,
crisis_kind=result.crisis_kind,
)
@router.post("/{session_id}/stream")
async def stream_turn(
session_id: str,
body: TurnRequest,
principal: CurrentPrincipal,
):
"""Stream a generated client reply for one trainee utterance."""
_ensure_learner(principal)
sess = await _load_session_or_404(session_id, principal)
recall = _RECALL_CACHE.get(session_id) or memory.RecallContext()
ctx = orchestrator.prepare_turn(
session_id=session_id,
case_id=sess.case_id,
card=sess.persona,
state=sess.state,
learner_text=body.text,
recall_summary=recall.recall_summary,
pinned_facts=recall.pinned_facts,
recent_turns=sess.recent_turns(),
)
assert ctx.state_after is not None
async def event_generator():
last_beat = asyncio.get_running_loop().time()
final_reply = ""
try:
async for ev in orchestrator.run_turn_stream(ctx, engine_client):
if ev.event == "token":
text = str(ev.data.get("text", ""))
final_reply += text
yield {"event": "token", "data": text}
elif ev.event == "done":
data = {**ev.data, "stage": _stage_label(ctx.state_after.stage)}
await _append_session_turn(
sess,
TurnRecord(
turn_seq=ctx.state_after.turn_seq,
speaker="counselor",
stage=_stage_label(ctx.state_after.stage),
text=body.text,
text_masked=ctx.learner_text_masked,
),
)
await _update_session_state(sess, ctx.state_after)
if final_reply:
await _append_session_turn(
sess,
TurnRecord(
turn_seq=ctx.state_after.turn_seq,
speaker="client",
stage=_stage_label(ctx.state_after.stage),
text=final_reply,
text_masked=final_reply,
),
)
yield {"event": "done", "data": json.dumps(data, ensure_ascii=False)}
else:
yield {"event": ev.event, "data": json.dumps(ev.data, ensure_ascii=False)}
now = asyncio.get_running_loop().time()
if now - last_beat >= settings.sse_heartbeat_seconds:
yield {"event": "ping", "data": "{}"}
last_beat = now
except Exception as exc:
yield {"event": "error", "data": json.dumps({"detail": str(exc)}, ensure_ascii=False)}
return
return EventSourceResponse(event_generator())
@router.post("/{session_id}/end", response_model=SessionEndResponse)
async def end_session(
session_id: str,
principal: CurrentPrincipal,
) -> SessionEndResponse:
"""End a learner-owned session and prepare carry-over state."""
_ensure_learner(principal)
sess = await _load_session_or_404(session_id, principal, allow_ended=True)
recall = _RECALL_CACHE.get(session_id) or memory.RecallContext()
carry = memory.make_carry_over(
state=sess.state,
session_id=session_id,
case_id=sess.case_id,
session_no=sess.session_no,
masked_turns=sess.masked_turns(),
prev_rapport_credit=sess.prev_rapport_credit,
open_threads=recall.open_threads,
)
await _end_persisted_session(sess, carry)
_RECALL_CACHE.pop(session_id, None)
_schedule_session_evaluation(sess)
return SessionEndResponse(
session_id=session_id,
session_no=sess.session_no,
digest_pending=carry.compression_job is not None,
end_state=carry.end_state,
)