462 lines
18 KiB
Python
462 lines
18 KiB
Python
"""턴 오케스트레이터 — 상담 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_SSE_DONE,
|
|
ENGINE_GATEWAY_SSE_ERROR,
|
|
EngineGatewaySseDecodeError,
|
|
StreamDoneEvent,
|
|
StreamErrorEvent,
|
|
StreamTokenEvent,
|
|
)
|
|
from . import guardrail, persona, state_machine
|
|
from .persona import PersonaCard, PersonaStateContext
|
|
from .state_machine import SessionState, Stage
|
|
|
|
# 평가 훅 타입: U_t(수련생 마스킹 발화) + 내담자응답 + 상태 → 평가 결과(dict)
|
|
# Features evaluator 가 이 시그니처에 맞춰 함수를 주입한다(여기선 호출만).
|
|
EvalHook = Callable[["TurnContext", str], Awaitable[Optional[dict]]]
|
|
LlmAuditHook = Callable[[dict[str, Any]], Awaitable[None]]
|
|
|
|
|
|
@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 에서 옴)
|
|
recall_summary: Optional[str] = None
|
|
pinned_facts: list[str] = field(default_factory=list)
|
|
recent_turns: list[dict[str, str]] = field(default_factory=list)
|
|
kb_behavior_cues: list[str] = field(default_factory=list)
|
|
# 회기 이론모드(학습자 선택: 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
|
|
|
|
|
|
# ════════════════════════════════════════════════════════════════════════════
|
|
# 1~3단계 — 입력 가드레일 + 상태머신 + 페르소나 컨텍스트 (엔진 호출 전 결정론)
|
|
# ════════════════════════════════════════════════════════════════════════════
|
|
def prepare_turn(
|
|
*,
|
|
session_id: str,
|
|
case_id: Optional[str],
|
|
card: PersonaCard,
|
|
state: SessionState,
|
|
learner_text: str,
|
|
recall_summary: Optional[str] = None,
|
|
pinned_facts: Optional[list[str]] = None,
|
|
recent_turns: Optional[list[dict[str, str]]] = None,
|
|
kb_behavior_cues: Optional[list[str]] = 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 의 경량 휴리스틱으로 라포 신호를 추정한다.
|
|
"""
|
|
ctx = TurnContext(
|
|
session_id=session_id,
|
|
case_id=case_id,
|
|
persona=card,
|
|
state_before=state,
|
|
learner_text_raw=learner_text,
|
|
recall_summary=_mask_optional_text(recall_summary),
|
|
pinned_facts=_mask_text_list(pinned_facts),
|
|
recent_turns=_mask_recent_turns(recent_turns),
|
|
kb_behavior_cues=list(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,
|
|
recall_summary=ctx.recall_summary,
|
|
pinned_facts=ctx.pinned_facts,
|
|
recent_turns=ctx.recent_turns,
|
|
kb_behavior_cues=ctx.kb_behavior_cues,
|
|
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
|
|
|
|
|
|
# ════════════════════════════════════════════════════════════════════════════
|
|
# 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},
|
|
)
|
|
started = time.perf_counter()
|
|
resp: GenerateResponse = await engine.generate(req)
|
|
latency_ms = int((time.perf_counter() - started) * 1000)
|
|
await _record_llm_audit(
|
|
audit_hook,
|
|
session_id=ctx.session_id,
|
|
provider=resp.provider,
|
|
model=resp.model,
|
|
tokens_in=resp.tokens_in,
|
|
tokens_out=resp.tokens_out,
|
|
cost_usd=resp.cost_usd,
|
|
inference_geo=resp.inference_geo,
|
|
latency_ms=latency_ms,
|
|
)
|
|
reply = resp.text
|
|
|
|
# 5) 출력 가드레일 — 수단 차단 + ideation 상한
|
|
guard = guardrail.sanitize_client_reply(reply, ideation_stage=st.ideation_stage)
|
|
safety_flagged = guard.blocked or (ctx.crisis is not None and ctx.crisis.escalate)
|
|
if guard.needs_regeneration:
|
|
# 수단정보 누출 → 안전 대체 응답으로 치환(1차). 재생성 루프는 후속.
|
|
reply = "…(말을 잇지 못하고 잠시 침묵한다)"
|
|
|
|
# 6) 평가 훅(주입형) — 평가 AI 4차원 태깅 (Features 소유)
|
|
evaluation: Optional[dict] = None
|
|
if eval_hook is not None:
|
|
try:
|
|
evaluation = await eval_hook(ctx, reply)
|
|
except Exception:
|
|
evaluation = None # 평가 실패가 상담 루프를 막지 않게(비치명적)
|
|
|
|
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
|
|
stream_meta: dict[str, Any] = {}
|
|
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)
|
|
if guard.needs_regeneration and not flagged:
|
|
flagged = True
|
|
yield StreamEvent("safety", {"reason": "means_info_blocked"})
|
|
# 토큰은 더 내보내지 않고 안전 대체로 종결
|
|
accumulated = "…(말을 잇지 못하고 잠시 침묵한다)"
|
|
break
|
|
|
|
yield StreamEvent("token", {"text": text_piece})
|
|
|
|
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 "gateway-default"),
|
|
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 "gateway-default"),
|
|
"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")),
|
|
},
|
|
)
|
|
except EngineGatewaySseDecodeError as e:
|
|
yield StreamEvent("error", {"detail": str(e)})
|
|
except EngineError as e:
|
|
yield StreamEvent("error", {"detail": str(e)})
|
|
|
|
|
|
async def _record_llm_audit(
|
|
audit_hook: Optional[LlmAuditHook],
|
|
**payload: Any,
|
|
) -> None:
|
|
if audit_hook is None:
|
|
return
|
|
try:
|
|
await audit_hook(payload)
|
|
except Exception:
|
|
return
|
|
|
|
|
|
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",
|
|
"TurnResult",
|
|
"StreamEvent",
|
|
"prepare_turn",
|
|
"run_turn_generate",
|
|
"run_turn_stream",
|
|
]
|