501 lines
17 KiB
Python
501 lines
17 KiB
Python
#!/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())
|