"""REST/WS 공용 턴 런타임 헬퍼. 세션 로드, 오너십 검증, 완료 턴 영속화, 상태 갱신은 REST 세션 라우트와 음성 WebSocket 라우트가 같은 규칙을 공유해야 한다. """ from __future__ import annotations from enum import Enum import logging from typing import Optional from . import db, session_persistence from .deps import Principal from .runtime_policy import require_runtime_fallback_allowed, runtime_fallback_allowed from .services import guardrail, orchestrator, state_machine from .stage_contract import stage_label from .store import InProcSession, TurnRecord, store logger = logging.getLogger(__name__) _LIVE_COACH_RECHARGE_MIN_RAPPORT = 0.35 _LIVE_COACH_RECHARGE_MIN_OPENNESS_GAIN = 0.02 class SessionAccessError(str, Enum): NOT_FOUND = "not_found" FORBIDDEN = "forbidden" ENDED = "ended" async def load_owned_session( session_id: str, principal: Principal, *, allow_ended: bool = False, include_turn_evaluation: bool = False, ) -> tuple[InProcSession | None, Optional[SessionAccessError]]: """DB 우선으로 학습자 소유 세션을 로드하고 접근 오류를 코드로 반환한다.""" sess = await session_persistence.load_session( session_id, principal, allow_ended=True, include_turn_evaluation=include_turn_evaluation, ) if sess is not None: store.put(sess) elif runtime_fallback_allowed(): sess = store.get(session_id) if sess is None: return None, SessionAccessError.NOT_FOUND if sess.learner_id != principal.user_id: return None, SessionAccessError.FORBIDDEN if sess.ended and not allow_ended: return None, SessionAccessError.ENDED return sess, None async def append_completed_turn( sess: InProcSession, turn: TurnRecord, *, context: str, ) -> None: """완료된 턴을 DB와 in-process 미러에 기록한다.""" 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(context) store.append_turn(sess.session_id, turn) async def update_session_state( sess: InProcSession, state: state_machine.SessionState, *, context: str, ) -> None: """working state를 DB와 in-process 미러에 반영한다.""" 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(context) store.update_state(sess.session_id, state) async def record_completed_turn( sess: InProcSession, ctx: orchestrator.TurnContext, result: orchestrator.TurnResult, *, context_prefix: str, counselor_turn: TurnRecord | None = None, ) -> None: """상담자 발화와 내담자 응답을 한 번에 기록하고 상태를 갱신한다.""" assert ctx.state_after is not None learner_turn = counselor_turn or TurnRecord( turn_seq=ctx.state_after.turn_seq, speaker="counselor", stage=stage_label(ctx.state_after.stage), text=ctx.learner_text_raw, text_masked=ctx.learner_text_masked, evaluation=result.evaluation, ) await append_completed_turn( sess, learner_turn, context=f"{context_prefix} turn append", ) if result.client_reply: client_mask = guardrail.mask_pii(result.client_reply) await append_completed_turn( sess, TurnRecord( turn_seq=result.turn_seq, speaker="client", stage=stage_label(result.state_after.stage), text=result.client_reply, text_masked=client_mask.text_masked, llm_provider=result.llm_provider, model=result.model, tokens_in=result.tokens_in, tokens_out=result.tokens_out, cost_usd=result.cost_usd, ), context=f"{context_prefix} turn append", ) await update_session_state( sess, result.state_after, context=f"{context_prefix} state update", ) async def record_safety_event( sess: InProcSession, ctx: orchestrator.TurnContext, result: orchestrator.TurnResult, ) -> None: """위기 escalate 시 app.safety_events에 교수자 확인용 알림 레코드를 남긴다.""" crisis = getattr(ctx, "crisis", None) if crisis is None or not getattr(crisis, "escalate", False): return kind = getattr(crisis.kind, "value", None) or str(getattr(crisis, "kind", "crisis")) detail = { "matched": list(getattr(crisis, "matched", []) or []), "stage": getattr(result, "stage", None), "turn_seq": getattr(result, "turn_seq", None), "conversation_stopped": getattr(result, "conversation_stopped", False), "crisis_resource": getattr(result, "crisis_resource", None), "alert_status": "teacher_dashboard", } try: async with db.acquire(ai_context=True) as conn: await conn.execute( """ INSERT INTO app.safety_events (session_id, trigger_type, ko_risk_level, escalated, detail) VALUES ($1::uuid, $2, $3, TRUE, $4::jsonb) """, sess.session_id, kind, int(getattr(crisis, "risk_level", 0) or 0), detail, ) except Exception: logger.exception( "safety event persistence failed: session_id=%s trigger_type=%s", sess.session_id, kind, ) require_runtime_fallback_allowed("safety event") def should_recharge_live_coach_credit( evaluation: dict | None, before: state_machine.SessionState, after: state_machine.SessionState, ) -> tuple[bool, str]: """Good-score recharge gate based on stored evaluator/state-machine evidence.""" if not isinstance(evaluation, dict): return False, "" if evaluation.get("appropriateness") != "pos": return False, "" try: rapport = float(evaluation.get("rapport_signal") or 0) except (TypeError, ValueError): rapport = 0.0 if rapport < _LIVE_COACH_RECHARGE_MIN_RAPPORT: return False, "" openness_gain = float(after.effective_openness or 0) - float(before.effective_openness or 0) stage_changed = after.stage != before.stage if not stage_changed and openness_gain < _LIVE_COACH_RECHARGE_MIN_OPENNESS_GAIN: return False, "" if stage_changed: return True, "좋은 발화로 내담자 단계가 열려 코칭 기회 1개를 충전했습니다." return True, "좋은 발화 뒤 내담자 개방도가 올라 코칭 기회 1개를 충전했습니다." async def maybe_recharge_live_coach_credit( sess: InProcSession, ctx: orchestrator.TurnContext, result: orchestrator.TurnResult, ) -> None: """Record one live-coach recharge when a strong learner turn changes client state.""" assert ctx.state_after is not None should_recharge, reason = should_recharge_live_coach_credit( result.evaluation, ctx.state_before, result.state_after, ) if not should_recharge: return try: await session_persistence.record_live_coach_recharge( session_id=sess.session_id, learner_id=sess.learner_id, turn_seq=result.turn_seq, stage=stage_label(result.state_after.stage), reason=reason, ) except Exception: logger.exception( "live coach recharge persistence failed: session_id=%s turn_seq=%s", sess.session_id, result.turn_seq, ) require_runtime_fallback_allowed("live coach recharge") async def finalize_completed_turn( sess: InProcSession, ctx: orchestrator.TurnContext, result: orchestrator.TurnResult, *, context_prefix: str, counselor_turn: TurnRecord | None = None, ) -> None: """Persist a completed turn and emit any derived safety alert in route-safe order.""" await record_completed_turn( sess, ctx, result, context_prefix=context_prefix, counselor_turn=counselor_turn, ) await maybe_recharge_live_coach_credit(sess, ctx, result) await record_safety_event(sess, ctx, result) __all__ = [ "SessionAccessError", "append_completed_turn", "finalize_completed_turn", "load_owned_session", "maybe_recharge_live_coach_credit", "record_safety_event", "record_completed_turn", "should_recharge_live_coach_credit", "stage_label", "update_session_state", ]