현재 작업 전체 반영
This commit is contained in:
parent
5560638e54
commit
c0dddab594
85 changed files with 11322 additions and 539 deletions
|
|
@ -19,6 +19,7 @@ MASTERPLAN §2.2 / MEMORY_DESIGN §2-B 턴 사이클:
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, AsyncIterator, Awaitable, Callable, Optional
|
||||
|
||||
|
|
@ -37,6 +38,7 @@ 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)
|
||||
|
|
@ -84,6 +86,8 @@ class TurnResult:
|
|||
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
|
||||
|
|
@ -161,6 +165,7 @@ def prepare_turn(
|
|||
pinned_facts=ctx.pinned_facts,
|
||||
recent_turns=ctx.recent_turns,
|
||||
kb_behavior_cues=ctx.kb_behavior_cues,
|
||||
theory_mode=ctx.theory_mode,
|
||||
)
|
||||
return ctx
|
||||
|
||||
|
|
@ -192,6 +197,7 @@ async def run_turn_generate(
|
|||
engine: EngineClient,
|
||||
*,
|
||||
eval_hook: Optional[EvalHook] = None,
|
||||
audit_hook: Optional[LlmAuditHook] = None,
|
||||
) -> TurnResult:
|
||||
"""동기 턴 실행(4~8). 내담자 응답을 한 번에 받아 가드레일·평가 순차 적용.
|
||||
|
||||
|
|
@ -200,6 +206,9 @@ async def run_turn_generate(
|
|||
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",
|
||||
|
|
@ -207,7 +216,20 @@ async def run_turn_generate(
|
|||
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 상한
|
||||
|
|
@ -242,6 +264,24 @@ async def run_turn_generate(
|
|||
)
|
||||
|
||||
|
||||
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)
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
|
|
@ -256,6 +296,8 @@ class StreamEvent:
|
|||
async def run_turn_stream(
|
||||
ctx: TurnContext,
|
||||
engine: EngineClient,
|
||||
*,
|
||||
audit_hook: Optional[LlmAuditHook] = None,
|
||||
) -> AsyncIterator[StreamEvent]:
|
||||
"""스트리밍 턴 실행(4~8). 게이트웨이 SSE 를 받아 token/done/safety/error 로 재방출.
|
||||
|
||||
|
|
@ -277,10 +319,34 @@ async def run_turn_stream(
|
|||
stream_meta: dict[str, Any] = {}
|
||||
if ctx.crisis is not None and ctx.crisis.escalate:
|
||||
flagged = True
|
||||
yield StreamEvent("safety", {"reason": "learner_real_crisis", "level": ctx.crisis.risk_level})
|
||||
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:
|
||||
current_event = "message"
|
||||
started = time.perf_counter()
|
||||
async for raw in engine.stream(req):
|
||||
# engine_client.stream 은 게이트웨이 SSE 의 *원시 라인*을 그대로 yield 한다.
|
||||
# 게이트웨이 프레이밍: "event: token|done|error" + "data: {...}".
|
||||
|
|
@ -317,6 +383,19 @@ async def run_turn_stream(
|
|||
|
||||
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",
|
||||
{
|
||||
|
|
@ -336,6 +415,18 @@ async def run_turn_stream(
|
|||
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 _extract_sse_payload(raw_line: str) -> Any:
|
||||
"""게이트웨이 SSE data 라인의 JSON payload를 추출.
|
||||
|
||||
|
|
@ -364,6 +455,13 @@ def _payload_text(payload: Any) -> Optional[str]:
|
|||
return None
|
||||
|
||||
|
||||
def _optional_str(value: Any) -> Optional[str]:
|
||||
if value is None:
|
||||
return None
|
||||
text = str(value).strip()
|
||||
return text or None
|
||||
|
||||
|
||||
def _payload_detail(payload: Any, fallback: str) -> str:
|
||||
if isinstance(payload, dict) and payload.get("detail"):
|
||||
return str(payload["detail"])
|
||||
|
|
@ -388,6 +486,7 @@ def _safe_float(value: Any) -> float:
|
|||
|
||||
__all__ = [
|
||||
"EvalHook",
|
||||
"LlmAuditHook",
|
||||
"TurnContext",
|
||||
"TurnResult",
|
||||
"StreamEvent",
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue