"""턴 오케스트레이터 — 상담 1턴 파이프라인 1~8단계 조립. MASTERPLAN §2.2 / MEMORY_DESIGN §2-B 턴 사이클: 1. [입력 가드레일] PII 마스킹 + 위기분류(실제위기 vs 연기) (guardrail) 2. [상태머신] effective_openness 결정론 계산 + 단계전이 (state_machine) 3. [페르소나 컨텍스트] L0~L6 messages 조립 (persona) 4. [내담자 AI] engine_client.stream/generate (CCD 비노출) (engine_client) 5. [출력 가드레일] 자살수단 차단, ideation 상한 (guardrail) 6. [평가 훅] 주입형 — 평가 함수는 *인자로 받는다*(Features 소유) (hook) 7. [상태 갱신] working state 반영(체크포인트는 호출부가 DB/store UPSERT) 8. [로깅 훅] 주입형 — turns insert/임베딩은 호출부가 주입 설계 원칙: - 평가(evaluator) 함수와 로깅 함수는 *주입*받는다(이 모듈은 evaluator.py 를 import 하지 않음). - DB 는 인터페이스로 추상화하되 asyncpg conn 도 받을 수 있게 했다(현재는 hook 으로만 사용). - generate(동기, 폴백/테스트) + stream(SSE 토큰) 두 경로 모두 제공. - 엔진 장애는 EngineError 로 전파 → 라우트가 503/SSE error 프레임으로 변환. """ from __future__ import annotations import time from dataclasses import dataclass, field from typing import Any, AsyncIterator, Awaitable, Callable, Optional from ..engine_client import ( EngineClient, EngineError, EngineMessage, GenerateRequest, GenerateResponse, StreamRequest, ) from ..contracts.engine_gateway import ( ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL, ENGINE_GATEWAY_SSE_DONE, ENGINE_GATEWAY_SSE_ERROR, EngineGatewaySseDecodeError, StreamDoneEvent, StreamErrorEvent, StreamTokenEvent, ) from . import guardrail, persona, state_machine from .llm_audit import LlmAuditHook, generate_with_audit, record_llm_audit from .persona import PersonaCard, PersonaStateContext, TurnMemory from .state_machine import SessionState # 평가 훅 타입: U_t(수련생 마스킹 발화) + 내담자응답 + 상태 → 평가 결과(dict) # Features evaluator 가 이 시그니처에 맞춰 함수를 주입한다(여기선 호출만). EvalHook = Callable[["TurnContext", str], Awaitable[Optional[dict]]] def turn_evaluation_error_payload(ctx: "TurnContext", error: BaseException | str) -> dict[str, Any]: """Represent a non-fatal fast-loop evaluator failure without hiding it.""" st = ctx.state_after or ctx.state_before if isinstance(error, BaseException): # Exception messages can contain raw learner/provider text. Keep review-facing # evidence to the exception type; detailed trace stays in server logs. detail = type(error).__name__ else: detail = str(error).strip() or "unknown turn evaluation error" masked = guardrail.mask_pii(detail).text_masked.strip() or "unknown turn evaluation error" return { "loop": "fast", "turn_seq": st.turn_seq, "stage": st.stage.value, "appropriateness": "neutral", "error": masked, } @dataclass(slots=True) class TurnContext: """한 턴 파이프라인을 관통하는 컨텍스트(가드레일/상태/페르소나 산출 집약).""" session_id: str case_id: Optional[str] persona: PersonaCard state_before: SessionState learner_text_raw: str learner_text_masked: str = "" crisis: Optional[guardrail.CrisisResult] = None state_after: Optional[SessionState] = None messages: list[EngineMessage] = field(default_factory=list) # 회상/메모리 주입(memory.RecallContext 에서 옴) memory: TurnMemory = field(default_factory=TurnMemory) # 회기 이론모드(학습자 선택: humanistic|cbt|integrative). 평가 이론부합·생성 프레이밍에 사용. theory_mode: Optional[str] = None def to_state_context(self) -> PersonaStateContext: st = self.state_after or self.state_before return PersonaStateContext( stage=st.stage.value, effective_openness=st.effective_openness, resistance=st.resistance, rapport_credit=st.rapport_credit, ideation_stage=st.ideation_stage, affect_state=st.affect_state, ) @dataclass(slots=True) class TurnResult: """동기(generate) 턴 결과.""" turn_seq: int stage: str effective_openness: float client_reply: Optional[str] safety_flagged: bool state_after: SessionState evaluation: Optional[dict] = None crisis_kind: str = "none" crisis_resource: Optional[dict[str, str]] = None conversation_stopped: bool = False llm_provider: Optional[str] = None model: Optional[str] = None tokens_in: int = 0 tokens_out: int = 0 cost_usd: float = 0.0 output_error: Optional[str] = None # ════════════════════════════════════════════════════════════════════════════ # 1~3단계 — 입력 가드레일 + 상태머신 + 페르소나 컨텍스트 (엔진 호출 전 결정론) # ════════════════════════════════════════════════════════════════════════════ def prepare_turn( *, session_id: str, case_id: Optional[str], card: PersonaCard, state: SessionState, learner_text: str, memory: Optional[TurnMemory] = None, theory_mode: Optional[str] = None, eval_rapport_signal: Optional[float] = None, ) -> TurnContext: """엔진 호출 전 결정론 전처리(1~3단계). 순수 — IO/LLM 없음. eval_rapport_signal 이 주어지면(평가 AI fast-loop 신호) 그걸 쓰고, 없으면 state_machine 의 경량 휴리스틱으로 라포 신호를 추정한다. """ turn_memory = memory or TurnMemory() ctx = TurnContext( session_id=session_id, case_id=case_id, persona=card, state_before=state, learner_text_raw=learner_text, memory=TurnMemory( recall_summary=_mask_optional_text(turn_memory.recall_summary), pinned_facts=_mask_text_list(turn_memory.pinned_facts), recent_turns=_mask_recent_turns(turn_memory.recent_turns), kb_behavior_cues=list(turn_memory.kb_behavior_cues or []), ), theory_mode=theory_mode, ) # 1) 입력 가드레일 — PII 마스킹 + 위기분류 mask = guardrail.mask_pii(learner_text) ctx.learner_text_masked = mask.text_masked ctx.crisis = guardrail.classify_crisis(learner_text, speaker_is_persona_context=True) # 2) 상태머신 — 라포 신호 → 결정론 전이 signal = ( eval_rapport_signal if eval_rapport_signal is not None else state_machine.estimate_rapport_signal(ctx.learner_text_masked) ) # 위기분류가 관측한 risk_level(>0)을 상태머신에 ideation_observed 로 전달 → # ideation_stage 보수적 상향(절대 하향 안 함, 안전 R5). C2 위기 관측 반영. crisis_ideation = ( ctx.crisis.risk_level if ctx.crisis is not None and ctx.crisis.risk_level > 0 else None ) ctx.state_after = state_machine.evolve( state, rapport_signal=signal, unlock_rate=card.unlock_rate(), decay_floor=card.decay_floor(), ideation_observed=crisis_ideation, ) # 3) 페르소나 컨텍스트 — L0~L6 messages 조립 (CCD 는 행동으로만, L0 가 강제) ctx.messages = persona.build_turn_messages( card, ctx.to_state_context(), ctx.learner_text_masked, memory=ctx.memory, theory_mode=ctx.theory_mode, ) return ctx def _mask_optional_text(text: Optional[str]) -> Optional[str]: if text is None: return None return guardrail.mask_pii(text).text_masked def _mask_text_list(values: Optional[list[str]]) -> list[str]: return [guardrail.mask_pii(value).text_masked for value in (values or [])] def _mask_recent_turns(turns: Optional[list[dict[str, str]]]) -> list[dict[str, str]]: masked: list[dict[str, str]] = [] for turn in turns or []: item = dict(turn) item["text"] = guardrail.mask_pii(str(item.get("text", ""))).text_masked masked.append(item) return masked def _latest_client_reply(turns: list[dict[str, str]]) -> Optional[str]: for turn in reversed(turns): if turn.get("speaker") == "client": text = str(turn.get("text", "")).strip() if text: return text return None # ════════════════════════════════════════════════════════════════════════════ # 4~8단계 — 동기 생성 경로 (폴백/테스트) # ════════════════════════════════════════════════════════════════════════════ async def run_turn_generate( ctx: TurnContext, engine: EngineClient, *, eval_hook: Optional[EvalHook] = None, audit_hook: Optional[LlmAuditHook] = None, ) -> TurnResult: """동기 턴 실행(4~8). 내담자 응답을 한 번에 받아 가드레일·평가 순차 적용. eval_hook 은 Features 가 주입(없으면 생략). 엔진 장애는 EngineError 전파. """ assert ctx.state_after is not None st = ctx.state_after if ctx.crisis is not None and ctx.crisis.escalate: return _crisis_gate_result(ctx) # 4) 내담자 AI 생성 req = GenerateRequest( ai_role="client", messages=ctx.messages, session_id=ctx.session_id, metadata={"stage": st.stage.value}, ) previous_client_reply = _latest_client_reply(ctx.memory.recent_turns) resp: GenerateResponse | None = None reply = "" safety_flagged = ctx.crisis is not None and ctx.crisis.escalate for attempt in range(2): resp = await generate_with_audit(engine, req, audit_hook) # 5) 출력 가드레일 — 수단 차단 + persona 품질 재생성 guard = guardrail.sanitize_client_reply( resp.text, ideation_stage=st.ideation_stage, turn_seq=st.turn_seq, previous_client_reply=previous_client_reply, ) if guard.needs_regeneration: if attempt == 0: continue return TurnResult( turn_seq=st.turn_seq, stage=st.stage.value, effective_openness=st.effective_openness, client_reply=None, safety_flagged=True, state_after=st, evaluation=None, crisis_kind=ctx.crisis.kind.value if ctx.crisis else "none", llm_provider=resp.provider, model=resp.model, tokens_in=resp.tokens_in, tokens_out=resp.tokens_out, cost_usd=resp.cost_usd, output_error="client_reply_quality_retryable", ) safety_flagged = safety_flagged or guard.blocked reply = guard.text break assert resp is not None # 6) 평가 훅(주입형) — 평가 AI 4차원 태깅 (Features 소유) evaluation: Optional[dict] = None if eval_hook is not None: try: evaluation = await eval_hook(ctx, reply) except Exception as exc: # 평가 실패는 상담 루프를 막지 않되, 리뷰/대시보드에서 조용히 사라지지 않게 남긴다. evaluation = turn_evaluation_error_payload(ctx, exc) return TurnResult( turn_seq=st.turn_seq, stage=st.stage.value, effective_openness=st.effective_openness, client_reply=reply, safety_flagged=safety_flagged, state_after=st, evaluation=evaluation, crisis_kind=ctx.crisis.kind.value if ctx.crisis else "none", llm_provider=resp.provider, model=resp.model, tokens_in=resp.tokens_in, tokens_out=resp.tokens_out, cost_usd=resp.cost_usd, ) def _crisis_gate_result(ctx: TurnContext) -> TurnResult: """실제 위기 신호는 LLM 호출 전에 중단하고 109 리소스를 반환한다.""" assert ctx.state_after is not None crisis = ctx.crisis return TurnResult( turn_seq=ctx.state_after.turn_seq, stage=ctx.state_after.stage.value, effective_openness=ctx.state_after.effective_openness, client_reply=None, safety_flagged=True, state_after=ctx.state_after, evaluation=None, crisis_kind=crisis.kind.value if crisis else "learner_real", crisis_resource=guardrail.crisis_resource(), conversation_stopped=True, ) # ════════════════════════════════════════════════════════════════════════════ # 4~8단계 — SSE 스트림 경로 (기본 UX) # ════════════════════════════════════════════════════════════════════════════ @dataclass(slots=True) class StreamEvent: """SSE 재방출용 이벤트. 라우트가 sse_starlette 형식으로 변환.""" event: str # 'token' | 'done' | 'safety' | 'error' data: dict[str, Any] async def run_turn_stream( ctx: TurnContext, engine: EngineClient, *, audit_hook: Optional[LlmAuditHook] = None, ) -> AsyncIterator[StreamEvent]: """스트리밍 턴 실행(4~8). 게이트웨이 SSE 를 받아 token/done/safety/error 로 재방출. 출력 가드레일은 *누적 텍스트* 기준으로 수단정보를 감지(스트림 중 발견 시 safety 이벤트 + 재생성 신호). 토큰 단위 완벽 차단은 후속(현재는 누적 스캔). """ assert ctx.state_after is not None st = ctx.state_after req = StreamRequest( ai_role="client", messages=ctx.messages, session_id=ctx.session_id, metadata={"stage": st.stage.value}, ) accumulated = "" flagged = False output_error: str | None = None stream_meta: dict[str, Any] = {} previous_client_reply = _latest_client_reply(ctx.memory.recent_turns) if ctx.crisis is not None and ctx.crisis.escalate: flagged = True resource = guardrail.crisis_resource() yield StreamEvent( "safety", { "reason": "learner_real_crisis", "level": ctx.crisis.risk_level, "crisis_resource": resource, "conversation_stopped": True, }, ) yield StreamEvent( "done", { "session_id": ctx.session_id, "stage": st.stage.value, "effective_openness": round(st.effective_openness, 4), "turn_seq": st.turn_seq, "safety_flagged": True, "crisis_kind": ctx.crisis.kind.value, "crisis_resource": resource, "conversation_stopped": True, }, ) return try: started = time.perf_counter() async for packet in engine.stream_packets(req): if packet.event == ENGINE_GATEWAY_SSE_ERROR: payload = packet.payload detail = payload.detail if isinstance(payload, StreamErrorEvent) else "engine stream error" yield StreamEvent("error", {"detail": detail}) return if packet.event == ENGINE_GATEWAY_SSE_DONE: payload = packet.payload if isinstance(payload, StreamDoneEvent): stream_meta = payload.model_dump() break payload = packet.payload if not isinstance(payload, StreamTokenEvent): continue text_piece = payload.text accumulated += text_piece guard = guardrail.sanitize_client_reply( accumulated, ideation_stage=st.ideation_stage, turn_seq=st.turn_seq, previous_client_reply=previous_client_reply, ) if guard.needs_regeneration: flagged = True output_error = "client_reply_quality_retryable" yield StreamEvent( "safety", { "reason": output_error, "reasons": guard.reasons, }, ) accumulated = "" else: flagged = flagged or guard.blocked if guard.text: yield StreamEvent("token", {"text": guard.text}) latency_ms = int((time.perf_counter() - started) * 1000) await record_llm_audit( audit_hook, session_id=ctx.session_id, provider=str(stream_meta.get("provider") or engine.engine_mode), model=str( stream_meta.get("model") or engine.default_model or ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL ), tokens_in=_safe_int(stream_meta.get("tokens_in")), tokens_out=_safe_int(stream_meta.get("tokens_out")), cost_usd=_safe_float(stream_meta.get("cost_usd")), inference_geo=_optional_str(stream_meta.get("inference_geo")), latency_ms=latency_ms, ) yield StreamEvent( "done", { "session_id": ctx.session_id, "stage": st.stage.value, "effective_openness": round(st.effective_openness, 4), "turn_seq": st.turn_seq, "safety_flagged": flagged, "llm_provider": str(stream_meta.get("provider") or engine.engine_mode), "model": str( stream_meta.get("model") or engine.default_model or ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL ), "tokens_in": _safe_int(stream_meta.get("tokens_in")), "tokens_out": _safe_int(stream_meta.get("tokens_out")), "cost_usd": _safe_float(stream_meta.get("cost_usd")), "output_error": output_error, }, ) except EngineGatewaySseDecodeError as e: yield StreamEvent("error", {"detail": str(e)}) except EngineError as e: yield StreamEvent("error", {"detail": str(e)}) def _optional_str(value: Any) -> Optional[str]: if value is None: return None text = str(value).strip() return text or None def _safe_int(value: Any) -> int: try: return int(value or 0) except (TypeError, ValueError): return 0 def _safe_float(value: Any) -> float: try: return float(value or 0.0) except (TypeError, ValueError): return 0.0 __all__ = [ "EvalHook", "LlmAuditHook", "TurnContext", "TurnMemory", "TurnResult", "StreamEvent", "prepare_turn", "run_turn_generate", "run_turn_stream", ]