Jev 기반 내담자 감정 상태와 응답 일관성 개선

This commit is contained in:
Yun Chan 2026-09-22 21:32:26 +09:00
parent 77f8421818
commit 8344bc2ad2
23 changed files with 3384 additions and 25 deletions

View file

@ -0,0 +1,501 @@
#!/usr/bin/env python3
"""로컬 엔진과 Jev를 DB 없이 잇는 합성 대화 smoke 수집기."""
from __future__ import annotations
import argparse
import asyncio
import hashlib
import json
import math
import os
import sys
import time
import uuid
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
REPO_ROOT = Path(__file__).resolve().parents[1]
API_ROOT = REPO_ROOT / "apps" / "api"
TURN_TIMEOUT_SECONDS = 90.0
ALLOWED_PHASES = frozenset({"legacy", "jev"})
SOURCE_PATHS = (
Path(__file__).resolve(),
API_ROOT / "app" / "config.py",
API_ROOT / "app" / "engine_client.py",
API_ROOT / "app" / "services" / "client_affect.py",
API_ROOT / "app" / "services" / "jev_client.py",
API_ROOT / "app" / "services" / "orchestrator.py",
API_ROOT / "app" / "services" / "persona.py",
API_ROOT / "engine_gateway" / "gateway.py",
)
SYNTHETIC_UTTERANCES: tuple[dict[str, str], ...] = (
{
"id": "empathy",
"label": "empathy",
"text": "그 이야기를 꺼내 주셔서 고마워요. 많이 버거웠겠어요.",
},
{
"id": "contradictory_recall",
"label": "unsupported_recall",
"text": "지난번 여동생과 다퉜다고 하셨죠. 그때 어떤 느낌이었나요?",
},
{
"id": "ambivalence",
"label": "ambivalence",
"text": "한편으로는 달라지고 싶고, 다른 한편으로는 아직 두려운 마음도 있을 수 있겠어요.",
},
)
SYNTHETIC_MEMORY = {
"pinned_facts": ["나는 외동이다."],
"expectations": {
"unsupported_recall": "상담자의 새 주장을 과거 사실로 확정하지 않음",
"contradictory_recall": "형제자매 존재에 동의하지 않음",
},
}
EMOTION_KEYS = (
"anxiety",
"sadness",
"anger",
"shame",
"guilt",
"loneliness",
"relief",
"hope",
"trust",
)
SAFE_JEV_CODES = frozenset(
{
"not_configured",
"not_started",
"timeout",
"unauthorized",
"insufficient_credits",
"forbidden",
"model_unavailable",
"rate_limited",
"overloaded",
"http_error",
"transport",
"malformed_response",
"model_mismatch",
}
)
def _turn_count(value: str) -> int:
try:
turns = int(value)
except ValueError as exc:
raise argparse.ArgumentTypeError("turns must be an integer from 1 to 3") from exc
if not 1 <= turns <= len(SYNTHETIC_UTTERANCES):
raise argparse.ArgumentTypeError("turns must be from 1 to 3")
return turns
def _repeat_count(value: str) -> int:
try:
repeats = int(value)
except ValueError as exc:
raise argparse.ArgumentTypeError("repeats must be an integer from 1 to 5") from exc
if not 1 <= repeats <= 5:
raise argparse.ArgumentTypeError("repeats must be from 1 to 5")
return repeats
def _phase_list(value: str) -> tuple[str, ...]:
phases = tuple(part.strip() for part in value.split(",") if part.strip())
if not phases or any(phase not in ALLOWED_PHASES for phase in phases):
raise argparse.ArgumentTypeError("phases must be a comma-separated subset of legacy,jev")
if len(set(phases)) != len(phases):
raise argparse.ArgumentTypeError("phases must not contain duplicates")
return phases
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description="DB 없이 local engine과 Jev의 합성 가상 내담자 대화를 smoke 수집한다."
)
parser.add_argument("--output", type=Path, required=True, help="JSON 결과 저장 경로")
parser.add_argument(
"--turns",
type=_turn_count,
default=2,
help="각 phase의 합성 발화 수(1~3, 기본 2)",
)
parser.add_argument(
"--phases",
type=_phase_list,
default=("legacy", "jev"),
help="실행할 phase 목록(legacy,jev; 기본 legacy,jev)",
)
parser.add_argument(
"--repeats",
type=_repeat_count,
default=1,
help="phase 묶음 반복 횟수(1~5, 기본 1)",
)
parser.add_argument("--label", default="", help="측정 보고서 식별 문자열")
return parser
def _load_runtime() -> dict[str, Any]:
"""API 모듈이 cwd 기반 설정을 읽도록 한 뒤 작업 cwd를 즉시 복구한다."""
previous_cwd = Path.cwd()
inserted_path = False
try:
os.chdir(API_ROOT)
api_root_text = str(API_ROOT)
if api_root_text not in sys.path:
sys.path.insert(0, api_root_text)
inserted_path = True
from app.config import settings
from app.engine_client import EngineError, engine_client
from app.services import orchestrator, persona, state_machine
from app.services.jev_client import JevError, jev_client
finally:
os.chdir(previous_cwd)
if inserted_path:
sys.path.remove(str(API_ROOT))
return {
"settings": settings,
"EngineError": EngineError,
"engine_client": engine_client,
"orchestrator": orchestrator,
"persona": persona,
"state_machine": state_machine,
"JevError": JevError,
"jev_client": jev_client,
}
def _safe_error_code(error: BaseException, runtime: dict[str, Any]) -> str:
if isinstance(error, runtime["JevError"]):
code = getattr(error, "code", "")
if code in SAFE_JEV_CODES:
return f"client_affect_{code}"
return "client_affect_error"
if isinstance(error, runtime["EngineError"]):
return "engine_error"
if isinstance(error, TimeoutError):
return "turn_timeout"
return "runtime_error"
def _safe_stream_error(value: object) -> str:
detail = str(value).strip()
if detail.startswith("client_affect_"):
code = detail.removeprefix("client_affect_")
if code in SAFE_JEV_CODES:
return detail
if detail in {"client_stream_incomplete", "engine stream error", "engine stream decode error"}:
return detail.replace(" ", "_")
return "stream_error"
def _final_emotions(affect_state: dict[str, Any]) -> dict[str, float]:
emotions: dict[str, float] = {}
for dimension in EMOTION_KEYS:
value = affect_state.get(f"emotion_{dimension}")
if isinstance(value, (int, float)) and not isinstance(value, bool) and math.isfinite(value):
emotions[dimension] = float(value)
return emotions
async def _run_turn(
*,
runtime: dict[str, Any],
state: Any,
session_id: str,
utterance: str,
recent_turns: list[dict[str, str]],
) -> tuple[dict[str, Any], Any, str | None]:
orchestrator = runtime["orchestrator"]
persona = runtime["persona"]
context = orchestrator.prepare_turn(
session_id=session_id,
case_id=None,
card=persona.P1,
state=state,
learner_text=utterance,
learner_identity="합성 상담자",
memory=orchestrator.TurnMemory(
recent_turns=recent_turns,
pinned_facts=list(SYNTHETIC_MEMORY["pinned_facts"]),
),
theory_mode="humanistic",
scenario_context=None,
)
started = time.perf_counter()
first_token_at: float | None = None
generated: list[str] = []
done_payload: dict[str, Any] | None = None
error_code: str | None = None
try:
async with asyncio.timeout(TURN_TIMEOUT_SECONDS):
async for event in orchestrator.run_turn_stream(context, runtime["engine_client"]):
now = time.perf_counter()
if event.event == "token":
if first_token_at is None:
first_token_at = now
generated.append(str(event.data.get("text", "")))
elif event.event == "done":
done_payload = dict(event.data)
elif event.event == "error":
error_code = _safe_stream_error(event.data.get("detail"))
break
except asyncio.CancelledError:
raise
except Exception as exc:
error_code = _safe_error_code(exc, runtime)
total_ms = round((time.perf_counter() - started) * 1000, 1)
if done_payload is None and error_code is None:
error_code = "stream_error"
generated_text = "".join(generated)
record: dict[str, Any] = {
"status": "done" if done_payload is not None and error_code is None else "error",
"ttft_ms": (
None
if first_token_at is None
else round((first_token_at - started) * 1000, 1)
),
"total_ms": total_ms,
"generation": {
"provider": None if done_payload is None else done_payload.get("llm_provider"),
"model": None if done_payload is None else done_payload.get("model"),
},
"appraisal": context.client_affect_metadata,
"final_emotions": _final_emotions(
context.state_after.affect_state
if done_payload is not None and error_code is None
else state.affect_state
),
"text": generated_text,
}
if error_code is not None:
record["error"] = error_code
return record, context.state_after, generated_text if done_payload is not None and error_code is None else None
async def _run_phase(
*,
runtime: dict[str, Any],
provider_mode: str,
turns: int,
repeat: int,
order_index: int,
) -> dict[str, Any]:
settings = runtime["settings"]
original_provider = settings.client_affect_provider
session_id = str(uuid.uuid4())
phase: dict[str, Any] = {
"provider_mode": provider_mode,
"repeat": repeat,
"order_index": order_index,
"session_id": session_id,
"status": "error",
"turns": [],
"session_closed": False,
}
settings.client_affect_provider = provider_mode
try:
state = runtime["state_machine"].init_state(params=runtime["persona"].P1.openness_params())
recent_turns: list[dict[str, str]] = []
for index, fixture in enumerate(SYNTHETIC_UTTERANCES[:turns], start=1):
utterance = fixture["text"]
result, next_state, client_reply = await _run_turn(
runtime=runtime,
state=state,
session_id=session_id,
utterance=utterance,
recent_turns=recent_turns,
)
result["turn"] = index
result["turn_temperature"] = "cold" if index == 1 else "warm"
result["utterance_id"] = fixture["id"]
result["utterance_kind"] = fixture["label"]
if fixture["id"] == "contradictory_recall":
result["fixture_expectations"] = [
SYNTHETIC_MEMORY["expectations"]["unsupported_recall"],
SYNTHETIC_MEMORY["expectations"]["contradictory_recall"],
]
phase["turns"].append(result)
if result["status"] != "done":
phase["error"] = result["error"]
return phase
state = next_state
recent_turns.extend(
[
{"speaker": "counselor", "text": utterance},
{"speaker": "client", "text": client_reply or ""},
]
)
phase["status"] = "done"
return phase
finally:
settings.client_affect_provider = original_provider
closed = await runtime["engine_client"].close_session(session_id)
phase["session_closed"] = closed
if not closed:
phase["cleanup_error"] = "engine_close_failed"
if phase["status"] == "done":
phase["status"] = "error"
phase["error"] = "engine_close_failed"
def _source_sha256() -> dict[str, str]:
return {
str(path.relative_to(REPO_ROOT)).replace("\\", "/"): hashlib.sha256(path.read_bytes()).hexdigest()
for path in SOURCE_PATHS
}
def _percentile(values: list[float], quantile: float) -> float | None:
if not values:
return None
ordered = sorted(values)
position = (len(ordered) - 1) * quantile
lower = math.floor(position)
upper = math.ceil(position)
if lower == upper:
return ordered[lower]
return round(ordered[lower] + (ordered[upper] - ordered[lower]) * (position - lower), 1)
def _latency_summary(values: list[float]) -> dict[str, float | int | None]:
return {
"sample_count": len(values),
"p50": _percentile(values, 0.50),
"p95": _percentile(values, 0.95),
}
def _phase_summaries(phases: list[dict[str, Any]]) -> list[dict[str, Any]]:
grouped: dict[tuple[str, str], list[tuple[dict[str, Any], dict[str, Any]]]] = {}
for phase in phases:
for turn in phase["turns"]:
grouped.setdefault((phase["provider_mode"], turn["turn_temperature"]), []).append((phase, turn))
summaries: list[dict[str, Any]] = []
for (provider_mode, cache_state), rows in grouped.items():
phase_ids = {(phase["repeat"], phase["order_index"]) for phase, _ in rows}
turns = [turn for _, turn in rows if turn["status"] == "done"]
ttft_values = [turn["ttft_ms"] for turn in turns if turn["ttft_ms"] is not None]
total_values = [turn["total_ms"] for turn in turns if turn["total_ms"] is not None]
summaries.append(
{
"provider_mode": provider_mode,
"turn_temperature": cache_state,
"phase_run_count": len(phase_ids),
"completed_phase_run_count": len(
{(phase["repeat"], phase["order_index"]) for phase, _ in rows if phase["status"] == "done"}
),
"attempted_turn_count": len(rows),
"completed_turn_count": len(turns),
"ttft_ms": _latency_summary(ttft_values),
"total_ms": _latency_summary(total_values),
}
)
return summaries
def _generation_settings(runtime: dict[str, Any]) -> dict[str, Any]:
settings = runtime["settings"]
engine_client = runtime["engine_client"]
return {
"engine_mode": settings.engine_mode,
"live_client_provider": settings.live_client_provider,
"model": engine_client.default_model,
"reasoning_effort": engine_client.default_reasoning_effort,
}
def _phase_order(phases: tuple[str, ...], repeat: int) -> tuple[str, ...]:
return phases if repeat % 2 else tuple(reversed(phases))
async def collect(
turns: int,
phases: tuple[str, ...] = ("legacy", "jev"),
repeats: int = 1,
label: str = "",
) -> dict[str, Any]:
report: dict[str, Any] = {
"kind": "jev_dialogue_smoke",
"provenance": "synthetic_only",
"quality_pass": False,
"generated_at": datetime.now(timezone.utc).isoformat(),
"label": label,
"source_sha256": _source_sha256(),
"turn_timeout_seconds": TURN_TIMEOUT_SECONDS,
"turns_requested": turns,
"repeats_requested": repeats,
"phases_requested": list(phases),
"memory_fixture": SYNTHETIC_MEMORY,
"repeat_orders": [],
"phases": [],
"phase_summaries": [],
"generation_settings": None,
"limitations": [
"cold/warm은 각 새 합성 세션의 첫 turn과 후속 turn을 뜻하며 provider cache 상태는 측정하지 않는다."
],
}
try:
runtime = _load_runtime()
except Exception:
report["status"] = "error"
report["error"] = "runtime_import_failed"
return report
engine_client = runtime["engine_client"]
jev_client = runtime["jev_client"]
report["generation_settings"] = _generation_settings(runtime)
try:
await engine_client.startup()
await jev_client.startup()
for repeat in range(1, repeats + 1):
order = _phase_order(phases, repeat)
report["repeat_orders"].append({"repeat": repeat, "phase_order": list(order)})
for order_index, provider_mode in enumerate(order, start=1):
phase = await _run_phase(
runtime=runtime,
provider_mode=provider_mode,
turns=turns,
repeat=repeat,
order_index=order_index,
)
report["phases"].append(phase)
report["phase_summaries"] = _phase_summaries(report["phases"])
report["status"] = "done" if all(phase["status"] == "done" for phase in report["phases"]) else "error"
except asyncio.CancelledError:
raise
except Exception as exc:
report["status"] = "error"
report["error"] = _safe_error_code(exc, runtime)
finally:
await jev_client.shutdown()
await engine_client.shutdown()
report["phase_summaries"] = _phase_summaries(report["phases"])
return report
def _write_report(path: Path, report: dict[str, Any]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
json.dumps(report, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
def main(argv: list[str] | None = None) -> int:
args = build_parser().parse_args(argv)
report = asyncio.run(collect(args.turns, args.phases, args.repeats, args.label))
_write_report(args.output, report)
print(f"보고서 저장: {args.output}")
return 0 if report["status"] == "done" else 1
if __name__ == "__main__":
raise SystemExit(main())