런타임 계약과 학습자 흐름 보강
This commit is contained in:
parent
f456b8997a
commit
206018b088
56 changed files with 4306 additions and 1008 deletions
|
|
@ -124,6 +124,20 @@ class ManagedUserMemoryInput:
|
|||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ManagedUserUpsertInput:
|
||||
email: str
|
||||
display_name: str
|
||||
role: RoleName
|
||||
admin_access: bool | None = None
|
||||
cohort_ids: list[str] | None = None
|
||||
user_id: str | None = None
|
||||
external_id: str | None = None
|
||||
affiliation: str | None = None
|
||||
account_status: AccountStatus | None = None
|
||||
reactivate: bool = False
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ManagedUserPatch:
|
||||
display_name: str | None = None
|
||||
|
|
@ -1123,31 +1137,19 @@ def _memory_upsert_managed_user(data: ManagedUserMemoryInput) -> ManagedUser:
|
|||
return user
|
||||
|
||||
|
||||
async def upsert_managed_user(
|
||||
*,
|
||||
email: str,
|
||||
display_name: str,
|
||||
role: RoleName,
|
||||
admin_access: bool | None = None,
|
||||
cohort_ids: list[str] | None = None,
|
||||
user_id: str | None = None,
|
||||
external_id: str | None = None,
|
||||
affiliation: str | None = None,
|
||||
account_status: AccountStatus | None = None,
|
||||
reactivate: bool = False,
|
||||
) -> ManagedUser:
|
||||
normalized_email = _normalize_email(email)
|
||||
normalized_external_id = _normalize_external_id(external_id, normalized_email)
|
||||
async def upsert_managed_user(data: ManagedUserUpsertInput) -> ManagedUser:
|
||||
normalized_email = _normalize_email(data.email)
|
||||
normalized_external_id = _normalize_external_id(data.external_id, normalized_email)
|
||||
manual_external_id = f"email:{normalized_email}"
|
||||
desired_account_status = account_status or _initial_account_status(
|
||||
desired_account_status = data.account_status or _initial_account_status(
|
||||
email=normalized_email,
|
||||
external_id=normalized_external_id,
|
||||
reactivate=reactivate,
|
||||
reactivate=data.reactivate,
|
||||
)
|
||||
try:
|
||||
pool = get_pool()
|
||||
async with pool.acquire() as conn:
|
||||
if user_id is None and normalized_external_id != manual_external_id:
|
||||
if data.user_id is None and normalized_external_id != manual_external_id:
|
||||
row = await conn.fetchrow(
|
||||
"""
|
||||
UPDATE app.app_user SET
|
||||
|
|
@ -1198,9 +1200,9 @@ async def upsert_managed_user(
|
|||
""",
|
||||
normalized_external_id,
|
||||
normalized_email,
|
||||
(display_name.strip() if display_name else normalized_email),
|
||||
_cohort_value(cohort_ids),
|
||||
affiliation or DEFAULT_AFFILIATION,
|
||||
(data.display_name.strip() if data.display_name else normalized_email),
|
||||
_cohort_value(data.cohort_ids),
|
||||
data.affiliation or DEFAULT_AFFILIATION,
|
||||
manual_external_id,
|
||||
)
|
||||
if row is not None:
|
||||
|
|
@ -1270,13 +1272,13 @@ async def upsert_managed_user(
|
|||
""",
|
||||
normalized_external_id,
|
||||
normalized_email,
|
||||
(display_name.strip() if display_name else normalized_email),
|
||||
_db_role(role),
|
||||
_cohort_value(cohort_ids),
|
||||
affiliation or DEFAULT_AFFILIATION,
|
||||
reactivate,
|
||||
(data.display_name.strip() if data.display_name else normalized_email),
|
||||
_db_role(data.role),
|
||||
_cohort_value(data.cohort_ids),
|
||||
data.affiliation or DEFAULT_AFFILIATION,
|
||||
data.reactivate,
|
||||
desired_account_status,
|
||||
admin_access,
|
||||
data.admin_access,
|
||||
)
|
||||
if row is None:
|
||||
_inactive_emails.add(normalized_email)
|
||||
|
|
@ -1288,8 +1290,10 @@ async def upsert_managed_user(
|
|||
raise
|
||||
except Exception:
|
||||
require_runtime_fallback_allowed("managed user")
|
||||
current = _users.get(user_id or "") or _users.get(_email_index.get(normalized_email, ""))
|
||||
fallback_uid = user_id or (current.user_id if current is not None else user_id_from_external_id(normalized_external_id))
|
||||
current = _users.get(data.user_id or "") or _users.get(_email_index.get(normalized_email, ""))
|
||||
fallback_uid = data.user_id or (
|
||||
current.user_id if current is not None else user_id_from_external_id(normalized_external_id)
|
||||
)
|
||||
fallback_account_status = desired_account_status
|
||||
if current is not None:
|
||||
if current.account_status == "suspended":
|
||||
|
|
@ -1299,14 +1303,14 @@ async def upsert_managed_user(
|
|||
return _memory_upsert_managed_user(
|
||||
ManagedUserMemoryInput(
|
||||
email=normalized_email,
|
||||
display_name=display_name,
|
||||
role=role,
|
||||
admin_access=admin_access,
|
||||
display_name=data.display_name,
|
||||
role=data.role,
|
||||
admin_access=data.admin_access,
|
||||
account_status=fallback_account_status,
|
||||
cohort_ids=cohort_ids,
|
||||
cohort_ids=data.cohort_ids,
|
||||
user_id=fallback_uid,
|
||||
affiliation=affiliation,
|
||||
reactivate=reactivate,
|
||||
affiliation=data.affiliation,
|
||||
reactivate=data.reactivate,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -1799,13 +1803,15 @@ async def create_session(
|
|||
raw_sid = secrets.token_urlsafe(32)
|
||||
normalized_email = _normalize_email(email)
|
||||
managed = await upsert_managed_user(
|
||||
email=normalized_email,
|
||||
display_name=display_name,
|
||||
role=role,
|
||||
cohort_ids=cohort_ids,
|
||||
user_id=user_id,
|
||||
external_id=external_id,
|
||||
reactivate=False,
|
||||
ManagedUserUpsertInput(
|
||||
email=normalized_email,
|
||||
display_name=display_name,
|
||||
role=role,
|
||||
cohort_ids=cohort_ids,
|
||||
user_id=user_id,
|
||||
external_id=external_id,
|
||||
reactivate=False,
|
||||
)
|
||||
)
|
||||
expires_at = time.time() + settings.session_ttl_seconds
|
||||
user = SessionUser(
|
||||
|
|
|
|||
|
|
@ -104,6 +104,14 @@ class Settings(BaseSettings):
|
|||
default=256,
|
||||
validation_alias="EVALUATOR_SEMANTIC_CACHE_MAX_ENTRIES",
|
||||
)
|
||||
session_digest_worker_enabled: bool = Field(
|
||||
default=False,
|
||||
validation_alias="SESSION_DIGEST_WORKER_ENABLED",
|
||||
)
|
||||
session_digest_worker_model: str = Field(
|
||||
default="",
|
||||
validation_alias="SESSION_DIGEST_WORKER_MODEL",
|
||||
)
|
||||
|
||||
# ── 외부 LLM 키 (게이트웨이가 못 받을 때 직접 폴백, PII 마스킹 후만) ──
|
||||
anthropic_api_key: str = Field(default="", validation_alias="ANTHROPIC_API_KEY")
|
||||
|
|
|
|||
|
|
@ -24,6 +24,16 @@ ENGINE_GATEWAY_SSE_EVENTS: tuple[EngineGatewaySseEvent, ...] = (
|
|||
ENGINE_GATEWAY_SSE_DONE,
|
||||
ENGINE_GATEWAY_SSE_ERROR,
|
||||
)
|
||||
ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL = "gateway-default"
|
||||
|
||||
|
||||
def normalize_engine_gateway_model(model: Optional[str]) -> Optional[str]:
|
||||
"""Return an explicit model override, or None for gateway default routing."""
|
||||
|
||||
value = (model or "").strip()
|
||||
if not value or value == ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL:
|
||||
return None
|
||||
return value
|
||||
|
||||
|
||||
class EngineMessage(BaseModel):
|
||||
|
|
@ -58,6 +68,27 @@ class GenerateResponse(BaseModel):
|
|||
structured: Optional[dict[str, Any]] = None
|
||||
|
||||
|
||||
def structured_payload_from_response(resp: GenerateResponse) -> dict[str, Any] | None:
|
||||
"""Return structured output, or a JSON object embedded in legacy text."""
|
||||
|
||||
if isinstance(resp.structured, dict):
|
||||
return resp.structured
|
||||
|
||||
raw = (resp.text or "").strip()
|
||||
if not raw:
|
||||
return None
|
||||
raw = _strip_json_fence(raw)
|
||||
|
||||
parsed = _json_object_or_none(raw)
|
||||
if parsed is not None:
|
||||
return parsed
|
||||
|
||||
start, end = raw.find("{"), raw.rfind("}")
|
||||
if 0 <= start < end:
|
||||
return _json_object_or_none(raw[start : end + 1])
|
||||
return None
|
||||
|
||||
|
||||
class StreamTokenEvent(BaseModel):
|
||||
text: str
|
||||
|
||||
|
|
@ -188,6 +219,23 @@ def _to_stream_error_event(payload: Any) -> StreamErrorEvent:
|
|||
return StreamErrorEvent(detail="engine stream error")
|
||||
|
||||
|
||||
def _strip_json_fence(value: str) -> str:
|
||||
if not value.startswith("```"):
|
||||
return value
|
||||
text = value.split("```", 2)[1] if value.count("```") >= 2 else value.strip("`")
|
||||
if text.lstrip().lower().startswith("json"):
|
||||
text = text.lstrip()[4:]
|
||||
return text.strip()
|
||||
|
||||
|
||||
def _json_object_or_none(value: str) -> dict[str, Any] | None:
|
||||
try:
|
||||
parsed = json.loads(value)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return None
|
||||
return parsed if isinstance(parsed, dict) else None
|
||||
|
||||
|
||||
def _safe_int(value: Any) -> int:
|
||||
try:
|
||||
return int(value or 0)
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ from .contracts.engine_gateway import (
|
|||
GenerateRequest,
|
||||
GenerateResponse,
|
||||
StreamRequest,
|
||||
normalize_engine_gateway_model,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -66,13 +67,6 @@ class EngineClient:
|
|||
await self._client.aclose()
|
||||
self._client = None
|
||||
|
||||
@staticmethod
|
||||
def _model_override(model: Optional[str]) -> Optional[str]:
|
||||
value = (model or "").strip()
|
||||
if not value or value == "gateway-default":
|
||||
return None
|
||||
return value
|
||||
|
||||
async def configure(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -81,7 +75,7 @@ class EngineClient:
|
|||
default_model: Optional[str] = None,
|
||||
) -> None:
|
||||
next_url = base_url.rstrip("/")
|
||||
next_model = self._model_override(default_model)
|
||||
next_model = normalize_engine_gateway_model(default_model)
|
||||
async with self._lock:
|
||||
url_changed = next_url != self.base_url
|
||||
self.base_url = next_url
|
||||
|
|
@ -136,7 +130,7 @@ class EngineClient:
|
|||
}
|
||||
|
||||
async def generate(self, req: GenerateRequest) -> GenerateResponse:
|
||||
"""단발 생성. TODO: 게이트웨이 응답 스키마 확정 후 cost 텔레메트리 turns 적재."""
|
||||
"""단발 생성."""
|
||||
try:
|
||||
r = await self.client.post("/v1/generate", json=self._payload(req))
|
||||
r.raise_for_status()
|
||||
|
|
|
|||
166
apps/api/app/persona_generation_contract.py
Normal file
166
apps/api/app/persona_generation_contract.py
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
"""Persona draft generation contract helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from typing import Any
|
||||
|
||||
from .contracts.engine_gateway import GenerateResponse, structured_payload_from_response
|
||||
from .persona_read_model import PersonaDraftGenerateRequest, PersonaDraftPayload
|
||||
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_ID = "persona-draft-rag"
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_VERSION = "2026-06-28.1"
|
||||
PERSONA_DRAFT_SYSTEM_PROMPT = (
|
||||
"출력은 반드시 structured_schema를 따른다. code는 P숫자 형식을 선호하되 "
|
||||
"힌트가 없으면 빈 문자열 대신 임시값 P로 둔다. source_provenance에는 "
|
||||
"RAG source_id와 첨부 근거 기반 초안임을 남긴다. evidence chunk id를 "
|
||||
"임상 필드 본문에 그대로 노출하지 않는다."
|
||||
)
|
||||
PERSONA_DRAFT_USER_PROMPT_PREAMBLE = (
|
||||
"너는 Vignette 임상 콘텐츠 저작 보조자다. 아래 RAG 근거 청크만 바탕으로 교육용 "
|
||||
"가상내담자 페르소나 초안을 만든다. 첨부 원문은 KB 문서가 SSOT이며, 근거 밖 내용을 "
|
||||
"임의로 꾸며 핵심 임상 정보처럼 쓰지 않는다. 실제 개인정보는 이미 마스킹됐으며, "
|
||||
"원문 표현을 복사하지 말고 "
|
||||
"범주화·합성화된 임상 훈련용 설정으로 변환한다. CCD/DSM/역린은 런타임 내부 설정이므로 "
|
||||
"내담자 발화에 직접 노출되지 않는 형태로 작성한다."
|
||||
)
|
||||
|
||||
|
||||
def persona_draft_prompt_bundle() -> dict[str, str]:
|
||||
payload = "\n".join(
|
||||
[
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_ID,
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_VERSION,
|
||||
PERSONA_DRAFT_SYSTEM_PROMPT,
|
||||
PERSONA_DRAFT_USER_PROMPT_PREAMBLE,
|
||||
]
|
||||
)
|
||||
return {
|
||||
"id": PERSONA_DRAFT_PROMPT_BUNDLE_ID,
|
||||
"version": PERSONA_DRAFT_PROMPT_BUNDLE_VERSION,
|
||||
"hash": hashlib.sha256(payload.encode("utf-8")).hexdigest()[:12],
|
||||
}
|
||||
|
||||
|
||||
def persona_generation_schema() -> dict[str, Any]:
|
||||
return {
|
||||
"type": "object",
|
||||
"additionalProperties": False,
|
||||
"properties": {
|
||||
"draft": {
|
||||
"type": "object",
|
||||
"additionalProperties": False,
|
||||
"properties": {
|
||||
"code": {"type": "string"},
|
||||
"display_name": {"type": "string"},
|
||||
"difficulty": {"type": "string", "enum": ["easy", "moderate", "hard"]},
|
||||
"theory_target": {"type": "array", "items": {"type": "string"}},
|
||||
"demographics": {"type": "object"},
|
||||
"presenting": {"type": "object"},
|
||||
"history": {"type": "object"},
|
||||
"big5": {"type": "object"},
|
||||
"resistance": {"type": "object"},
|
||||
"speech_style": {"type": "object"},
|
||||
"affect_baseline": {"type": "object"},
|
||||
"ccd": {"type": "object"},
|
||||
"dsm5_dimensional": {"type": "object"},
|
||||
"triggers": {"type": "object"},
|
||||
"source_provenance": {"type": "string"},
|
||||
"is_synthetic": {"type": "boolean"},
|
||||
},
|
||||
"required": [
|
||||
"code",
|
||||
"display_name",
|
||||
"difficulty",
|
||||
"theory_target",
|
||||
"demographics",
|
||||
"presenting",
|
||||
"history",
|
||||
"big5",
|
||||
"resistance",
|
||||
"speech_style",
|
||||
"affect_baseline",
|
||||
"ccd",
|
||||
"dsm5_dimensional",
|
||||
"triggers",
|
||||
"source_provenance",
|
||||
"is_synthetic",
|
||||
],
|
||||
},
|
||||
"source_summary": {"type": "string"},
|
||||
"warnings": {"type": "array", "items": {"type": "string"}},
|
||||
},
|
||||
"required": ["draft", "source_summary", "warnings"],
|
||||
}
|
||||
|
||||
|
||||
def persona_generation_payload_from_response(response: GenerateResponse) -> dict[str, Any]:
|
||||
return structured_payload_from_response(response) or {}
|
||||
|
||||
|
||||
def coerce_persona_generated_draft(
|
||||
payload: dict[str, Any],
|
||||
request: PersonaDraftGenerateRequest,
|
||||
) -> PersonaDraftPayload:
|
||||
raw = payload.get("draft") if isinstance(payload.get("draft"), dict) else payload
|
||||
if not isinstance(raw, dict):
|
||||
raw = {}
|
||||
theory_target = raw.get("theory_target")
|
||||
theory_values = (
|
||||
[str(item).strip().lower() for item in theory_target if str(item).strip()]
|
||||
if isinstance(theory_target, list)
|
||||
else [value.strip().lower() for value in request.theory_target if value.strip()]
|
||||
)
|
||||
code = str(raw.get("code") or request.code_hint or "").strip().upper()
|
||||
display_name = str(raw.get("display_name") or request.display_name_hint or "자료 기반 새 페르소나").strip()
|
||||
difficulty = str(raw.get("difficulty") or request.difficulty)
|
||||
if difficulty not in {"easy", "moderate", "hard"}:
|
||||
difficulty = request.difficulty
|
||||
return PersonaDraftPayload(
|
||||
code=code or "P",
|
||||
display_name=display_name,
|
||||
difficulty=difficulty, # type: ignore[arg-type]
|
||||
theory_target=theory_values or ["humanistic"],
|
||||
demographics=_json_object(raw.get("demographics")),
|
||||
presenting=_json_object(raw.get("presenting")),
|
||||
history=_json_object(raw.get("history")),
|
||||
big5=_float_dict(raw.get("big5")) or {"O": 0.5, "C": 0.5, "E": 0.5, "A": 0.5, "N": 0.5},
|
||||
resistance=_float_dict(raw.get("resistance"))
|
||||
or {
|
||||
"base_resistance": 0.5,
|
||||
"unlock_rate": 0.1,
|
||||
"decay_floor": 0.05,
|
||||
"silence_prob": 0.15,
|
||||
"deflection_prob": 0.25,
|
||||
},
|
||||
speech_style=_json_object(raw.get("speech_style")),
|
||||
affect_baseline=_float_dict(raw.get("affect_baseline"))
|
||||
or {
|
||||
"negative_affect": 0.45,
|
||||
"hopelessness": 0.2,
|
||||
"anhedonia": 0.2,
|
||||
"sleep": 0.2,
|
||||
"anxiety": 0.35,
|
||||
"suicide_ideation_stage": 1,
|
||||
},
|
||||
ccd=_json_object(raw.get("ccd")),
|
||||
dsm5_dimensional=_json_object(raw.get("dsm5_dimensional")),
|
||||
triggers=_json_object(raw.get("triggers")),
|
||||
source_provenance=str(raw.get("source_provenance") or f"masked {request.source_kind}"),
|
||||
is_synthetic=bool(raw.get("is_synthetic", True)),
|
||||
submit_for_review=False,
|
||||
)
|
||||
|
||||
|
||||
def _float_dict(value: Any) -> dict[str, float]:
|
||||
if not isinstance(value, dict):
|
||||
return {}
|
||||
result: dict[str, float] = {}
|
||||
for key, item in value.items():
|
||||
if isinstance(item, (int, float)):
|
||||
result[str(key)] = float(item)
|
||||
return result
|
||||
|
||||
|
||||
def _json_object(value: Any) -> dict[str, Any]:
|
||||
return value if isinstance(value, dict) else {}
|
||||
|
|
@ -14,6 +14,7 @@ from pydantic import BaseModel, Field
|
|||
from ..auth_types import AccountStatus, RoleName
|
||||
from ..auth_sessions import (
|
||||
ManagedUserPatch,
|
||||
ManagedUserUpsertInput,
|
||||
active_session_count,
|
||||
deactivate_managed_user,
|
||||
get_managed_user,
|
||||
|
|
@ -24,6 +25,7 @@ from ..auth_sessions import (
|
|||
upsert_managed_user,
|
||||
)
|
||||
from ..config import settings
|
||||
from ..contracts.engine_gateway import ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL
|
||||
from ..db import acquire, get_pool, healthcheck
|
||||
from ..deps import Principal, require_admin_access
|
||||
from ..engine_client import engine_client
|
||||
|
|
@ -955,7 +957,7 @@ def _default_engine_config() -> AdminEngineConfigResponse:
|
|||
return AdminEngineConfigResponse(
|
||||
engine_mode=settings.engine_mode,
|
||||
engine_url=settings.engine_url,
|
||||
model="gateway-default",
|
||||
model=ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL,
|
||||
durable=False,
|
||||
source="runtime_default",
|
||||
)
|
||||
|
|
@ -1725,14 +1727,16 @@ async def create_user(
|
|||
if body.role == "admin" or body.admin_access:
|
||||
_assert_super_admin(principal)
|
||||
user = await upsert_managed_user(
|
||||
email=_normalize_email(body.email),
|
||||
display_name=body.display_name,
|
||||
role=body.role,
|
||||
admin_access=body.admin_access,
|
||||
account_status=body.account_status,
|
||||
affiliation=body.affiliation,
|
||||
cohort_ids=body.cohort_ids,
|
||||
reactivate=True,
|
||||
ManagedUserUpsertInput(
|
||||
email=_normalize_email(body.email),
|
||||
display_name=body.display_name,
|
||||
role=body.role,
|
||||
admin_access=body.admin_access,
|
||||
account_status=body.account_status,
|
||||
affiliation=body.affiliation,
|
||||
cohort_ids=body.cohort_ids,
|
||||
reactivate=True,
|
||||
)
|
||||
)
|
||||
users, durable = await list_managed_users()
|
||||
if not durable:
|
||||
|
|
|
|||
|
|
@ -137,14 +137,11 @@ async def reevaluate_session(
|
|||
raise HTTPException(status.HTTP_503_SERVICE_UNAVAILABLE, detail=result.error)
|
||||
|
||||
await session_persistence.save_session_evaluation(
|
||||
session_id=session_id,
|
||||
learner_id=sess.learner_id,
|
||||
status="error" if result.error else "ready",
|
||||
source="engine",
|
||||
scope=result.scope,
|
||||
stage=result.stage,
|
||||
payload=result.to_dict(),
|
||||
error=result.error,
|
||||
session_persistence.SessionEvaluationWrite.from_result(
|
||||
session_id=session_id,
|
||||
learner_id=sess.learner_id,
|
||||
result=result,
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
||||
|
|
@ -182,7 +179,7 @@ async def reevaluate_turn(
|
|||
client_reply = sess.turns[target_idx + 1].text_masked
|
||||
|
||||
# 평가용 경량 TurnContext 재구성(prepare_turn 의 결정론 산출과 동형). 엔진 호출 없음.
|
||||
from ..services.orchestrator import TurnContext # 지연 import(소유권 경계)
|
||||
from ..services.orchestrator import TurnContext, TurnMemory # 지연 import(소유권 경계)
|
||||
|
||||
recent = [
|
||||
{"speaker": tr.speaker, "text": tr.text_masked} for tr in sess.turns[max(0, target_idx - 4):target_idx]
|
||||
|
|
@ -195,7 +192,7 @@ async def reevaluate_turn(
|
|||
learner_text_raw=learner.text,
|
||||
learner_text_masked=learner.text_masked,
|
||||
state_after=sess.state, # 조회 시점 상태(정밀 재현은 DB 스냅샷 도입 시)
|
||||
recent_turns=recent,
|
||||
memory=TurnMemory(recent_turns=recent),
|
||||
)
|
||||
|
||||
result = await evaluator.evaluate_turn(
|
||||
|
|
|
|||
|
|
@ -13,6 +13,14 @@ from fastapi import APIRouter, Depends, HTTPException, Response, status
|
|||
from ..db import acquire
|
||||
from ..deps import CurrentPrincipal, Principal, Role, require_role
|
||||
from ..deps import AIView
|
||||
from ..persona_generation_contract import (
|
||||
PERSONA_DRAFT_SYSTEM_PROMPT,
|
||||
PERSONA_DRAFT_USER_PROMPT_PREAMBLE,
|
||||
coerce_persona_generated_draft,
|
||||
persona_draft_prompt_bundle,
|
||||
persona_generation_payload_from_response,
|
||||
persona_generation_schema,
|
||||
)
|
||||
from ..persona_repository import (
|
||||
archive_persona_family,
|
||||
create_persona_draft,
|
||||
|
|
@ -57,38 +65,6 @@ PERSONA_SOURCE_CITATION: dict[str, str] = {
|
|||
"textbook_guide": "교수자 첨부 교재/가이드 환언·발췌 근거 — 저작권 검수 필요",
|
||||
"mixed_notes": "교수자 첨부 혼합 메모 PII 마스킹 파생본",
|
||||
}
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_ID = "persona-draft-rag"
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_VERSION = "2026-06-28.1"
|
||||
PERSONA_DRAFT_SYSTEM_PROMPT = (
|
||||
"출력은 반드시 structured_schema를 따른다. code는 P숫자 형식을 선호하되 "
|
||||
"힌트가 없으면 빈 문자열 대신 임시값 P로 둔다. source_provenance에는 "
|
||||
"RAG source_id와 첨부 근거 기반 초안임을 남긴다. evidence chunk id를 "
|
||||
"임상 필드 본문에 그대로 노출하지 않는다."
|
||||
)
|
||||
PERSONA_DRAFT_USER_PROMPT_PREAMBLE = (
|
||||
"너는 Vignette 임상 콘텐츠 저작 보조자다. 아래 RAG 근거 청크만 바탕으로 교육용 "
|
||||
"가상내담자 페르소나 초안을 만든다. 첨부 원문은 KB 문서가 SSOT이며, 근거 밖 내용을 "
|
||||
"임의로 꾸며 핵심 임상 정보처럼 쓰지 않는다. 실제 개인정보는 이미 마스킹됐으며, "
|
||||
"원문 표현을 복사하지 말고 "
|
||||
"범주화·합성화된 임상 훈련용 설정으로 변환한다. CCD/DSM/역린은 런타임 내부 설정이므로 "
|
||||
"내담자 발화에 직접 노출되지 않는 형태로 작성한다."
|
||||
)
|
||||
|
||||
|
||||
def _persona_draft_prompt_bundle() -> dict[str, str]:
|
||||
payload = "\n".join(
|
||||
[
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_ID,
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_VERSION,
|
||||
PERSONA_DRAFT_SYSTEM_PROMPT,
|
||||
PERSONA_DRAFT_USER_PROMPT_PREAMBLE,
|
||||
]
|
||||
)
|
||||
return {
|
||||
"id": PERSONA_DRAFT_PROMPT_BUNDLE_ID,
|
||||
"version": PERSONA_DRAFT_PROMPT_BUNDLE_VERSION,
|
||||
"hash": hashlib.sha256(payload.encode("utf-8")).hexdigest()[:12],
|
||||
}
|
||||
|
||||
|
||||
def _card_from_draft_payload(request: PersonaDraftPayload):
|
||||
|
|
@ -454,142 +430,6 @@ def _format_generation_evidence(evidence: list[PersonaGenerationEvidence]) -> st
|
|||
return "\n\n".join(lines)
|
||||
|
||||
|
||||
def _persona_generation_schema() -> dict[str, Any]:
|
||||
return {
|
||||
"type": "object",
|
||||
"additionalProperties": False,
|
||||
"properties": {
|
||||
"draft": {
|
||||
"type": "object",
|
||||
"additionalProperties": False,
|
||||
"properties": {
|
||||
"code": {"type": "string"},
|
||||
"display_name": {"type": "string"},
|
||||
"difficulty": {"type": "string", "enum": ["easy", "moderate", "hard"]},
|
||||
"theory_target": {"type": "array", "items": {"type": "string"}},
|
||||
"demographics": {"type": "object"},
|
||||
"presenting": {"type": "object"},
|
||||
"history": {"type": "object"},
|
||||
"big5": {"type": "object"},
|
||||
"resistance": {"type": "object"},
|
||||
"speech_style": {"type": "object"},
|
||||
"affect_baseline": {"type": "object"},
|
||||
"ccd": {"type": "object"},
|
||||
"dsm5_dimensional": {"type": "object"},
|
||||
"triggers": {"type": "object"},
|
||||
"source_provenance": {"type": "string"},
|
||||
"is_synthetic": {"type": "boolean"},
|
||||
},
|
||||
"required": [
|
||||
"code",
|
||||
"display_name",
|
||||
"difficulty",
|
||||
"theory_target",
|
||||
"demographics",
|
||||
"presenting",
|
||||
"history",
|
||||
"big5",
|
||||
"resistance",
|
||||
"speech_style",
|
||||
"affect_baseline",
|
||||
"ccd",
|
||||
"dsm5_dimensional",
|
||||
"triggers",
|
||||
"source_provenance",
|
||||
"is_synthetic",
|
||||
],
|
||||
},
|
||||
"source_summary": {"type": "string"},
|
||||
"warnings": {"type": "array", "items": {"type": "string"}},
|
||||
},
|
||||
"required": ["draft", "source_summary", "warnings"],
|
||||
}
|
||||
|
||||
|
||||
def _json_payload_from_generation(text: str) -> dict[str, Any]:
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
except json.JSONDecodeError:
|
||||
start = text.find("{")
|
||||
end = text.rfind("}")
|
||||
if start >= 0 and end > start:
|
||||
try:
|
||||
parsed = json.loads(text[start : end + 1])
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
except json.JSONDecodeError:
|
||||
return {}
|
||||
return {}
|
||||
|
||||
|
||||
def _float_dict(value: Any) -> dict[str, float]:
|
||||
if not isinstance(value, dict):
|
||||
return {}
|
||||
result: dict[str, float] = {}
|
||||
for key, item in value.items():
|
||||
if isinstance(item, (int, float)):
|
||||
result[str(key)] = float(item)
|
||||
return result
|
||||
|
||||
|
||||
def _coerce_generated_draft(
|
||||
payload: dict[str, Any],
|
||||
request: PersonaDraftGenerateRequest,
|
||||
) -> PersonaDraftPayload:
|
||||
raw = payload.get("draft") if isinstance(payload.get("draft"), dict) else payload
|
||||
if not isinstance(raw, dict):
|
||||
raw = {}
|
||||
theory_target = raw.get("theory_target")
|
||||
theory_values = (
|
||||
[str(item).strip().lower() for item in theory_target if str(item).strip()]
|
||||
if isinstance(theory_target, list)
|
||||
else [value.strip().lower() for value in request.theory_target if value.strip()]
|
||||
)
|
||||
code = str(raw.get("code") or request.code_hint or "").strip().upper()
|
||||
display_name = str(raw.get("display_name") or request.display_name_hint or "자료 기반 새 페르소나").strip()
|
||||
difficulty = str(raw.get("difficulty") or request.difficulty)
|
||||
if difficulty not in {"easy", "moderate", "hard"}:
|
||||
difficulty = request.difficulty
|
||||
return PersonaDraftPayload(
|
||||
code=code or "P",
|
||||
display_name=display_name,
|
||||
difficulty=difficulty, # type: ignore[arg-type]
|
||||
theory_target=theory_values or ["humanistic"],
|
||||
demographics=_json_object(raw.get("demographics")),
|
||||
presenting=_json_object(raw.get("presenting")),
|
||||
history=_json_object(raw.get("history")),
|
||||
big5=_float_dict(raw.get("big5")) or {"O": 0.5, "C": 0.5, "E": 0.5, "A": 0.5, "N": 0.5},
|
||||
resistance=_float_dict(raw.get("resistance"))
|
||||
or {
|
||||
"base_resistance": 0.5,
|
||||
"unlock_rate": 0.1,
|
||||
"decay_floor": 0.05,
|
||||
"silence_prob": 0.15,
|
||||
"deflection_prob": 0.25,
|
||||
},
|
||||
speech_style=_json_object(raw.get("speech_style")),
|
||||
affect_baseline=_float_dict(raw.get("affect_baseline"))
|
||||
or {
|
||||
"negative_affect": 0.45,
|
||||
"hopelessness": 0.2,
|
||||
"anhedonia": 0.2,
|
||||
"sleep": 0.2,
|
||||
"anxiety": 0.35,
|
||||
"suicide_ideation_stage": 1,
|
||||
},
|
||||
ccd=_json_object(raw.get("ccd")),
|
||||
dsm5_dimensional=_json_object(raw.get("dsm5_dimensional")),
|
||||
triggers=_json_object(raw.get("triggers")),
|
||||
source_provenance=str(raw.get("source_provenance") or f"masked {request.source_kind}"),
|
||||
is_synthetic=bool(raw.get("is_synthetic", True)),
|
||||
submit_for_review=False,
|
||||
)
|
||||
|
||||
|
||||
def _json_object(value: Any) -> dict[str, Any]:
|
||||
return value if isinstance(value, dict) else {}
|
||||
|
||||
|
||||
def _ensure_teacher_or_admin(principal: Principal) -> None:
|
||||
if principal.role not in {Role.TEACHER, Role.ADMIN}:
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail="only teachers and admins can review personas")
|
||||
|
|
@ -775,7 +615,7 @@ async def generate_persona_draft_route(
|
|||
query=evidence_query or "페르소나 저작 근거",
|
||||
)
|
||||
evidence_text = _format_generation_evidence(evidence)
|
||||
prompt_bundle = _persona_draft_prompt_bundle()
|
||||
prompt_bundle = persona_draft_prompt_bundle()
|
||||
prompt = (
|
||||
f"{PERSONA_DRAFT_USER_PROMPT_PREAMBLE}\n\n"
|
||||
f"자료 종류: {request.source_kind}\n"
|
||||
|
|
@ -799,7 +639,7 @@ async def generate_persona_draft_route(
|
|||
],
|
||||
max_tokens=2200,
|
||||
temperature=0.2,
|
||||
structured_schema=_persona_generation_schema(),
|
||||
structured_schema=persona_generation_schema(),
|
||||
metadata={
|
||||
"feature": "persona_draft_generation",
|
||||
"prompt_bundle": prompt_bundle,
|
||||
|
|
@ -814,8 +654,8 @@ async def generate_persona_draft_route(
|
|||
status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail=f"persona draft generator unavailable: {exc}",
|
||||
) from exc
|
||||
payload = response.structured or _json_payload_from_generation(response.text)
|
||||
draft = _coerce_generated_draft(payload, request)
|
||||
payload = persona_generation_payload_from_response(response)
|
||||
draft = coerce_persona_generated_draft(payload, request)
|
||||
provenance = (
|
||||
f"prompt={prompt_bundle['id']}@{prompt_bundle['version']}#{prompt_bundle['hash']}; "
|
||||
f"RAG sources={','.join(source_ids)}; "
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from __future__ import annotations
|
|||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import secrets
|
||||
from datetime import datetime
|
||||
from typing import Literal, Optional
|
||||
|
|
@ -25,7 +26,16 @@ from ..deps import CurrentPrincipal, Principal, Role
|
|||
from ..engine_client import EngineError, engine_client
|
||||
from ..persona_repository import get_catalog_persona
|
||||
from ..runtime_policy import require_runtime_fallback_allowed
|
||||
from ..services import evaluator, guardrail, live_coach, memory, orchestrator, rag, state_machine
|
||||
from ..services import (
|
||||
evaluator,
|
||||
guardrail,
|
||||
live_coach,
|
||||
memory,
|
||||
orchestrator,
|
||||
rag,
|
||||
session_digest_worker,
|
||||
state_machine,
|
||||
)
|
||||
from ..session_read_model import (
|
||||
LearnerDashboardResponse,
|
||||
LearnerSessionsResponse,
|
||||
|
|
@ -58,6 +68,7 @@ from ..session_read_model import (
|
|||
from ..store import InProcSession, TurnRecord, store
|
||||
|
||||
router = APIRouter(prefix="/sessions", tags=["sessions"])
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
TheoryMode = Literal["humanistic", "cbt", "integrative"]
|
||||
EndStateValue = str | int | float | bool | None | dict[str, float]
|
||||
|
|
@ -394,10 +405,12 @@ async def _prepare_turn_context(
|
|||
card=sess.persona,
|
||||
state=sess.state,
|
||||
learner_text=learner_text,
|
||||
recall_summary=recall.recall_summary,
|
||||
pinned_facts=recall.pinned_facts,
|
||||
recent_turns=sess.recent_turns(visible_to="client"),
|
||||
kb_behavior_cues=kb_cues,
|
||||
memory=orchestrator.TurnMemory(
|
||||
recall_summary=recall.recall_summary,
|
||||
pinned_facts=recall.pinned_facts,
|
||||
recent_turns=sess.recent_turns(visible_to="client"),
|
||||
kb_behavior_cues=kb_cues,
|
||||
),
|
||||
theory_mode=sess.theory_mode,
|
||||
)
|
||||
assert ctx.state_after is not None
|
||||
|
|
@ -502,12 +515,52 @@ async def _end_persisted_session(sess: InProcSession, carry: memory.CarryOver) -
|
|||
sess.ended = True
|
||||
sess.ended_at = datetime.now().timestamp()
|
||||
store.put(sess)
|
||||
if _should_schedule_session_digest_worker(carry):
|
||||
asyncio.create_task(_run_session_digest_worker_for_session(sess.session_id))
|
||||
asyncio.create_task(_write_episodic_embeddings(sess))
|
||||
return
|
||||
require_runtime_fallback_allowed("session end")
|
||||
store.end(sess.session_id)
|
||||
|
||||
|
||||
def _should_schedule_session_digest_worker(carry: memory.CarryOver) -> bool:
|
||||
return bool(settings.session_digest_worker_enabled and carry.compression_job is not None)
|
||||
|
||||
|
||||
async def _run_session_digest_worker_for_session(session_id: str) -> None:
|
||||
"""Best-effort M2 LLM digest compressor.
|
||||
|
||||
The DB connection is held only for load/apply. Engine generation runs outside
|
||||
the transaction so a slow provider cannot pin the pool.
|
||||
"""
|
||||
|
||||
try:
|
||||
db.get_pool()
|
||||
async with db.acquire(role="admin") as conn:
|
||||
loaded = await session_digest_worker.load_session_digest_job(conn, session_id)
|
||||
if loaded is None:
|
||||
return
|
||||
model = settings.session_digest_worker_model.strip() or None
|
||||
worker = await session_digest_worker.run_session_digest_worker(
|
||||
loaded.job,
|
||||
engine_client,
|
||||
existing_case_digest=loaded.existing_case_digest,
|
||||
model=model,
|
||||
audit_hook=session_persistence.record_llm_call_audit,
|
||||
)
|
||||
if worker.apply_plan is None:
|
||||
return
|
||||
async with db.acquire(role="admin") as conn:
|
||||
await session_digest_worker.apply_session_digest_plan(
|
||||
conn,
|
||||
worker.apply_plan,
|
||||
learner_id=loaded.learner_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning("session digest worker failed for session_id=%s", session_id, exc_info=True)
|
||||
return
|
||||
|
||||
|
||||
async def _write_episodic_embeddings(sess: InProcSession) -> None:
|
||||
"""Best-effort M2 episodic writer.
|
||||
|
||||
|
|
@ -601,27 +654,22 @@ async def _generate_and_save_session_evaluation(sess: InProcSession) -> None:
|
|||
),
|
||||
timeout=min(float(settings.engine_timeout), 45.0),
|
||||
)
|
||||
status_value = "error" if result.error else "ready"
|
||||
await session_persistence.save_session_evaluation(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
status=status_value,
|
||||
source="engine",
|
||||
scope=result.scope,
|
||||
stage=result.stage,
|
||||
payload=result.to_dict(),
|
||||
error=result.error,
|
||||
session_persistence.SessionEvaluationWrite.from_result(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
result=result,
|
||||
)
|
||||
)
|
||||
except Exception as exc:
|
||||
await session_persistence.save_session_evaluation(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
status="error",
|
||||
source="engine",
|
||||
scope="session_end",
|
||||
stage=_stage_label(sess.state.stage),
|
||||
payload={},
|
||||
error=str(exc),
|
||||
session_persistence.SessionEvaluationWrite.from_error(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
scope="session_end",
|
||||
stage=_stage_label(sess.state.stage),
|
||||
error=str(exc),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -177,6 +177,10 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
receiving = False
|
||||
audio_started_at: float | None = None
|
||||
last_audio_end_at: float | None = None
|
||||
audio_format: str | None = None
|
||||
audio_sample_rate: int | None = None
|
||||
audio_channels: int | None = None
|
||||
audio_sample_width: int | None = None
|
||||
|
||||
try:
|
||||
while True:
|
||||
|
|
@ -217,6 +221,10 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
if ctype == "audio_start":
|
||||
receiving = True
|
||||
audio_started_at = time.monotonic()
|
||||
audio_format = _safe_str(ctrl.get("format"))
|
||||
audio_sample_rate = _safe_int(ctrl.get("sample_rate"))
|
||||
audio_channels = _safe_int(ctrl.get("channels"))
|
||||
audio_sample_width = _safe_int(ctrl.get("sample_width"))
|
||||
audio_buf.clear()
|
||||
await _safe_send_json(websocket, {"type": "state", "state": "listening"})
|
||||
|
||||
|
|
@ -226,13 +234,17 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
silence_ms = _safe_int(ctrl.get("silence_ms"))
|
||||
if silence_ms is None and last_audio_end_at is not None and audio_started_at is not None:
|
||||
silence_ms = max(0, int((audio_started_at - last_audio_end_at) * 1000))
|
||||
end_format = _safe_str(ctrl.get("format")) or audio_format
|
||||
await _handle_utterance(
|
||||
websocket,
|
||||
session_id=session_id,
|
||||
principal=principal,
|
||||
voice_preset=voice_preset,
|
||||
audio=bytes(audio_buf),
|
||||
fmt=ctrl.get("format"),
|
||||
fmt=end_format,
|
||||
sample_rate=_safe_int(ctrl.get("sample_rate")) or audio_sample_rate,
|
||||
channels=_safe_int(ctrl.get("channels")) or audio_channels,
|
||||
sample_width=_safe_int(ctrl.get("sample_width")) or audio_sample_width,
|
||||
audio_started_at=audio_started_at,
|
||||
audio_ended_at=audio_ended_at,
|
||||
silence_ms=silence_ms,
|
||||
|
|
@ -241,6 +253,10 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
)
|
||||
last_audio_end_at = audio_ended_at
|
||||
audio_started_at = None
|
||||
audio_format = None
|
||||
audio_sample_rate = None
|
||||
audio_channels = None
|
||||
audio_sample_width = None
|
||||
audio_buf.clear()
|
||||
|
||||
elif ctype == "text_turn":
|
||||
|
|
@ -257,6 +273,27 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
learner_text=learner_text,
|
||||
)
|
||||
|
||||
elif ctype == "stt_result":
|
||||
receiving = False
|
||||
audio_buf.clear()
|
||||
stt_received_at = time.monotonic()
|
||||
await _handle_stt_result_control(
|
||||
websocket,
|
||||
session_id=session_id,
|
||||
principal=principal,
|
||||
voice_preset=voice_preset,
|
||||
ctrl=ctrl,
|
||||
audio_started_at=audio_started_at,
|
||||
audio_ended_at=stt_received_at,
|
||||
last_audio_end_at=last_audio_end_at,
|
||||
)
|
||||
last_audio_end_at = stt_received_at
|
||||
audio_started_at = None
|
||||
audio_format = None
|
||||
audio_sample_rate = None
|
||||
audio_channels = None
|
||||
audio_sample_width = None
|
||||
|
||||
elif ctype == "ping":
|
||||
await _safe_send_json(websocket, {"type": "pong"})
|
||||
|
||||
|
|
@ -271,6 +308,60 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
await _safe_close(websocket)
|
||||
|
||||
|
||||
async def _handle_stt_result_control(
|
||||
websocket: WebSocket,
|
||||
*,
|
||||
session_id: str,
|
||||
principal: Principal,
|
||||
voice_preset: VoicePreset,
|
||||
ctrl: dict[str, object],
|
||||
audio_started_at: float | None = None,
|
||||
audio_ended_at: float | None = None,
|
||||
last_audio_end_at: float | None = None,
|
||||
) -> None:
|
||||
learner_text = str(ctrl.get("text") or "").strip()
|
||||
transcript_final = _safe_bool(ctrl.get("final"))
|
||||
silence_ms = _safe_int(ctrl.get("silence_ms"))
|
||||
if silence_ms is None and last_audio_end_at is not None and audio_started_at is not None:
|
||||
silence_ms = max(0, int((audio_started_at - last_audio_end_at) * 1000))
|
||||
provider_events = _safe_provider_events(ctrl.get("provider_events"))
|
||||
decision = voice_svc.assess_end_of_turn(
|
||||
transcript_text=learner_text,
|
||||
transcript_final=bool(transcript_final),
|
||||
silence_ms=silence_ms,
|
||||
)
|
||||
await _safe_send_json(
|
||||
websocket,
|
||||
{
|
||||
"type": "eot",
|
||||
"ready": decision.ready,
|
||||
"reason": decision.reason,
|
||||
"silence_ms": decision.silence_ms,
|
||||
"threshold_ms": decision.threshold_ms,
|
||||
},
|
||||
)
|
||||
if not decision.ready:
|
||||
await _safe_send_json(websocket, {"type": "state", "state": "listening"})
|
||||
return
|
||||
|
||||
await _safe_send_json(websocket, {"type": "state", "state": "thinking"})
|
||||
await _safe_send_json(
|
||||
websocket,
|
||||
{"type": "transcript", "text": learner_text, "final": True, "speaker": "counselor"},
|
||||
)
|
||||
await _run_turn_and_speak(
|
||||
websocket,
|
||||
session_id=session_id,
|
||||
principal=principal,
|
||||
voice_preset=voice_preset,
|
||||
learner_text=learner_text,
|
||||
duration_s=_elapsed_seconds(audio_started_at, audio_ended_at),
|
||||
silence_ms=decision.silence_ms,
|
||||
barge_in=_safe_bool(ctrl.get("barge_in")),
|
||||
provider_events=provider_events,
|
||||
)
|
||||
|
||||
|
||||
async def _handle_utterance(
|
||||
websocket: WebSocket,
|
||||
*,
|
||||
|
|
@ -279,6 +370,9 @@ async def _handle_utterance(
|
|||
voice_preset: VoicePreset,
|
||||
audio: bytes,
|
||||
fmt: Optional[str],
|
||||
sample_rate: int | None = None,
|
||||
channels: int | None = None,
|
||||
sample_width: int | None = None,
|
||||
audio_started_at: float | None = None,
|
||||
audio_ended_at: float | None = None,
|
||||
silence_ms: int | None = None,
|
||||
|
|
@ -293,10 +387,17 @@ async def _handle_utterance(
|
|||
|
||||
# STT begins after the learner stops speaking.
|
||||
await _safe_send_json(websocket, {"type": "state", "state": "thinking"})
|
||||
filename, content_type = _audio_meta(fmt)
|
||||
upload_audio, upload_fmt = _normalize_audio_upload(
|
||||
audio,
|
||||
fmt=fmt,
|
||||
sample_rate=sample_rate,
|
||||
channels=channels,
|
||||
sample_width=sample_width,
|
||||
)
|
||||
filename, content_type = _audio_meta(upload_fmt)
|
||||
try:
|
||||
stt = await voice_service.transcribe(
|
||||
audio, filename=filename, content_type=content_type
|
||||
upload_audio, filename=filename, content_type=content_type
|
||||
)
|
||||
except VoiceUnavailable as e:
|
||||
await _safe_send_json(websocket, {"type": "degraded", "reason": str(e)})
|
||||
|
|
@ -308,7 +409,7 @@ async def _handle_utterance(
|
|||
return
|
||||
|
||||
learner_text = stt.text
|
||||
audio_ref = _voice_audio_ref(audio, fmt)
|
||||
audio_ref = _voice_audio_ref(upload_audio, upload_fmt)
|
||||
duration_s = stt.duration or _elapsed_seconds(audio_started_at, audio_ended_at)
|
||||
speech_rate = _estimate_speech_rate(learner_text, duration_s)
|
||||
provider_events = _merge_provider_events(provider_events, getattr(stt, "provider_events", []))
|
||||
|
|
@ -342,12 +443,15 @@ async def _run_turn_and_speak(
|
|||
voice_preset: VoicePreset,
|
||||
learner_text: str,
|
||||
audio_ref: str | None = None,
|
||||
duration_s: float | None = None,
|
||||
silence_ms: int | None = None,
|
||||
speech_rate: float | None = None,
|
||||
barge_in: bool | None = None,
|
||||
provider_events: list[dict[str, object]] | None = None,
|
||||
) -> None:
|
||||
"""Run one counseling turn and stream synthesized client speech."""
|
||||
if speech_rate is None:
|
||||
speech_rate = _estimate_speech_rate(learner_text, duration_s)
|
||||
sess, err = await _load_voice_session(session_id, principal)
|
||||
if sess is None:
|
||||
await _safe_send_json(websocket, {"type": "error", "detail": err or "session not found or ended"})
|
||||
|
|
@ -364,10 +468,12 @@ async def _run_turn_and_speak(
|
|||
card=sess.persona,
|
||||
state=sess.state,
|
||||
learner_text=learner_text,
|
||||
recall_summary=recall.recall_summary,
|
||||
pinned_facts=recall.pinned_facts,
|
||||
recent_turns=sess.recent_turns(visible_to="client"),
|
||||
kb_behavior_cues=kb_cues,
|
||||
memory=orchestrator.TurnMemory(
|
||||
recall_summary=recall.recall_summary,
|
||||
pinned_facts=recall.pinned_facts,
|
||||
recent_turns=sess.recent_turns(visible_to="client"),
|
||||
kb_behavior_cues=kb_cues,
|
||||
),
|
||||
theory_mode=sess.theory_mode,
|
||||
)
|
||||
assert ctx.state_after is not None
|
||||
|
|
@ -656,6 +762,56 @@ def _audio_meta(fmt: Optional[str]) -> tuple[str, str]:
|
|||
return table.get(f, ("audio.webm", "audio/webm"))
|
||||
|
||||
|
||||
def _normalize_audio_upload(
|
||||
audio: bytes,
|
||||
*,
|
||||
fmt: Optional[str],
|
||||
sample_rate: int | None = None,
|
||||
channels: int | None = None,
|
||||
sample_width: int | None = None,
|
||||
) -> tuple[bytes, str]:
|
||||
f = (fmt or "webm").lower().lstrip(".") or "webm"
|
||||
if f != "pcm":
|
||||
return audio, f
|
||||
if sample_width not in (None, 2):
|
||||
raise ValueError("pcm sample_width must be 2 bytes")
|
||||
return _wav_from_pcm16(
|
||||
audio,
|
||||
sample_rate=_bounded_int(sample_rate, default=48000, minimum=8000, maximum=96000),
|
||||
channels=_bounded_int(channels, default=1, minimum=1, maximum=2),
|
||||
), "wav"
|
||||
|
||||
|
||||
def _bounded_int(value: int | None, *, default: int, minimum: int, maximum: int) -> int:
|
||||
if value is None:
|
||||
return default
|
||||
return min(maximum, max(minimum, value))
|
||||
|
||||
|
||||
def _wav_from_pcm16(pcm: bytes, *, sample_rate: int, channels: int) -> bytes:
|
||||
byte_rate = sample_rate * channels * 2
|
||||
block_align = channels * 2
|
||||
data_size = len(pcm)
|
||||
header = b"".join(
|
||||
[
|
||||
b"RIFF",
|
||||
(36 + data_size).to_bytes(4, "little"),
|
||||
b"WAVE",
|
||||
b"fmt ",
|
||||
(16).to_bytes(4, "little"),
|
||||
(1).to_bytes(2, "little"),
|
||||
channels.to_bytes(2, "little"),
|
||||
sample_rate.to_bytes(4, "little"),
|
||||
byte_rate.to_bytes(4, "little"),
|
||||
block_align.to_bytes(2, "little"),
|
||||
(16).to_bytes(2, "little"),
|
||||
b"data",
|
||||
data_size.to_bytes(4, "little"),
|
||||
]
|
||||
)
|
||||
return header + pcm
|
||||
|
||||
|
||||
def _voice_audio_ref(audio: bytes, fmt: Optional[str]) -> str | None:
|
||||
if not audio:
|
||||
return None
|
||||
|
|
@ -688,6 +844,13 @@ def _safe_int(value: object) -> int | None:
|
|||
return None
|
||||
|
||||
|
||||
def _safe_str(value: object) -> str | None:
|
||||
if isinstance(value, str):
|
||||
text = value.strip()
|
||||
return text or None
|
||||
return None
|
||||
|
||||
|
||||
def _safe_bool(value: object) -> bool | None:
|
||||
if value is None:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -34,12 +34,15 @@ from typing import TYPE_CHECKING, Any, Optional
|
|||
from pydantic import BaseModel, Field
|
||||
|
||||
from ..config import settings
|
||||
from ..contracts.engine_gateway import (
|
||||
ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL,
|
||||
structured_payload_from_response,
|
||||
)
|
||||
from ..engine_client import (
|
||||
EngineClient,
|
||||
EngineError,
|
||||
EngineMessage,
|
||||
GenerateRequest,
|
||||
GenerateResponse,
|
||||
)
|
||||
from ..taxonomy import (
|
||||
CLIENT_STATE_KO,
|
||||
|
|
@ -121,7 +124,7 @@ def _evaluator_cache_key(req: GenerateRequest) -> str:
|
|||
"version": _EVALUATOR_CACHE_VERSION,
|
||||
"ai_role": req.ai_role,
|
||||
"messages": [m.model_dump() for m in req.messages],
|
||||
"model": req.model or "gateway-default",
|
||||
"model": req.model or ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL,
|
||||
"max_tokens": req.max_tokens,
|
||||
"temperature": req.temperature,
|
||||
"structured_schema": req.structured_schema,
|
||||
|
|
@ -480,7 +483,7 @@ def build_fast_messages(ctx: "TurnContext", client_reply: str) -> list[EngineMes
|
|||
client_reply_masked = guardrail.mask_pii(client_reply).text_masked
|
||||
recent = "\n".join(
|
||||
f"{('상담자' if t.get('speaker') == 'counselor' else '내담자')}: {t.get('text', '')}"
|
||||
for t in (ctx.recent_turns or [])[-4:]
|
||||
for t in (ctx.memory.recent_turns or [])[-4:]
|
||||
) or "(직전 맥락 없음)"
|
||||
|
||||
crisis_note = ""
|
||||
|
|
@ -572,36 +575,6 @@ def build_deep_messages(
|
|||
]
|
||||
|
||||
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
# 4. 응답 파싱 — structured 우선, 없으면 text(JSON) 폴백, 실패는 빈 결과
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
def _structured_payload(resp: GenerateResponse) -> Optional[dict[str, Any]]:
|
||||
"""게이트웨이 structured 우선, 없으면 text 에서 JSON 추출(코드펜스/잡텍스트 관용)."""
|
||||
if resp.structured is not None and isinstance(resp.structured, dict):
|
||||
return resp.structured
|
||||
raw = (resp.text or "").strip()
|
||||
if not raw:
|
||||
return None
|
||||
# ```json ... ``` 펜스 제거
|
||||
if raw.startswith("```"):
|
||||
raw = raw.split("```", 2)[1] if raw.count("```") >= 2 else raw.strip("`")
|
||||
if raw.lstrip().lower().startswith("json"):
|
||||
raw = raw.lstrip()[4:]
|
||||
try:
|
||||
obj = json.loads(raw)
|
||||
return obj if isinstance(obj, dict) else None
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
# 본문 안에 묻힌 첫 객체만 시도
|
||||
start, end = raw.find("{"), raw.rfind("}")
|
||||
if 0 <= start < end:
|
||||
try:
|
||||
obj = json.loads(raw[start : end + 1])
|
||||
return obj if isinstance(obj, dict) else None
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _parse_intent_deviation(d: Any) -> Optional[IntentDeviation]:
|
||||
if not isinstance(d, dict):
|
||||
return None
|
||||
|
|
@ -770,7 +743,7 @@ async def evaluate_turn(
|
|||
base.error = f"eval_error: {e}"
|
||||
return base
|
||||
|
||||
payload = _structured_payload(resp)
|
||||
payload = structured_payload_from_response(resp)
|
||||
if payload is None:
|
||||
base.error = "no_structured_output"
|
||||
return base
|
||||
|
|
@ -853,7 +826,7 @@ async def evaluate_session(
|
|||
base.error = f"eval_error: {e}"
|
||||
return base
|
||||
|
||||
payload = _structured_payload(resp)
|
||||
payload = structured_payload_from_response(resp)
|
||||
if payload is None:
|
||||
base.error = "no_structured_output"
|
||||
return base
|
||||
|
|
|
|||
|
|
@ -21,7 +21,8 @@ from typing import TYPE_CHECKING, Any, Literal, Optional
|
|||
from pydantic import BaseModel, Field, field_validator
|
||||
|
||||
from ..config import settings
|
||||
from ..engine_client import EngineClient, EngineError, EngineMessage, GenerateRequest, GenerateResponse
|
||||
from ..contracts.engine_gateway import structured_payload_from_response
|
||||
from ..engine_client import EngineClient, EngineError, EngineMessage, GenerateRequest
|
||||
from ..paths import repo_root, repo_path
|
||||
from ..session_read_model import StageLabel, stage_label_or_none
|
||||
from . import guardrail
|
||||
|
|
@ -433,30 +434,6 @@ def _schema() -> dict[str, Any]:
|
|||
}
|
||||
|
||||
|
||||
def _structured_payload(resp: GenerateResponse) -> Optional[dict[str, Any]]:
|
||||
if isinstance(resp.structured, dict):
|
||||
return resp.structured
|
||||
raw = (resp.text or "").strip()
|
||||
if not raw:
|
||||
return None
|
||||
if raw.startswith("```"):
|
||||
raw = raw.split("```", 2)[1] if raw.count("```") >= 2 else raw.strip("`")
|
||||
if raw.lstrip().lower().startswith("json"):
|
||||
raw = raw.lstrip()[4:]
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
return data if isinstance(data, dict) else None
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
start, end = raw.find("{"), raw.rfind("}")
|
||||
if 0 <= start < end:
|
||||
try:
|
||||
data = json.loads(raw[start : end + 1])
|
||||
return data if isinstance(data, dict) else None
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _clip(value: Any, limit: int) -> Optional[str]:
|
||||
text = str(value or "").strip()
|
||||
if not text:
|
||||
|
|
@ -674,7 +651,7 @@ async def generate_live_coaching(
|
|||
inference_geo=resp.inference_geo,
|
||||
latency_ms=latency_ms,
|
||||
)
|
||||
payload = _structured_payload(resp)
|
||||
payload = structured_payload_from_response(resp)
|
||||
if payload is None:
|
||||
return _fallback_suggestion(
|
||||
item,
|
||||
|
|
|
|||
|
|
@ -37,6 +37,21 @@ _COUNSELING_AGREEMENT_WITHDRAWAL_RE = re.compile(
|
|||
r"더\s*이상.*(상담|회기).*(안\s*하|하지\s*않|못\s*하)"
|
||||
)
|
||||
_DIGEST_SPEAKERS = {"counselor", "client"}
|
||||
_LLM_DIGEST_MIN_CHARS = 60
|
||||
_LLM_DIGEST_INTERNAL_MARKERS = (
|
||||
"rapport_credit",
|
||||
"effective_openness",
|
||||
"end_state",
|
||||
"evaluation",
|
||||
"정답",
|
||||
"점수",
|
||||
"평가 payload",
|
||||
"평가점수",
|
||||
"평가 점수",
|
||||
"core_belief",
|
||||
"CCD",
|
||||
)
|
||||
_SESSION_DIGEST_PREFIX_RE = re.compile(r"^S(?P<session_no>\d+):")
|
||||
|
||||
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
|
|
@ -139,7 +154,7 @@ class SessionDigestInput:
|
|||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SessionDigestResult:
|
||||
"""Digest writer output contract shared by fallback and future LLM worker."""
|
||||
"""Digest writer output contract shared by fallback and LLM worker paths."""
|
||||
|
||||
session_id: str
|
||||
case_id: str | None
|
||||
|
|
@ -149,6 +164,35 @@ class SessionDigestResult:
|
|||
source: Literal["fallback", "llm"] = "fallback"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DigestQualityAssessment:
|
||||
"""Local quality gate result before an LLM digest can replace fallback."""
|
||||
|
||||
accepted: bool
|
||||
reason: Literal[
|
||||
"ok",
|
||||
"empty",
|
||||
"too_short",
|
||||
"forbidden_substring",
|
||||
"internal_marker",
|
||||
"wrong_session_prefix",
|
||||
]
|
||||
retryable: bool = False
|
||||
details: tuple[str, ...] = ()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SessionDigestWorkerOutcome:
|
||||
"""Accepted LLM digest result or the reason fallback must remain authoritative."""
|
||||
|
||||
result: SessionDigestResult | None
|
||||
quality: DigestQualityAssessment
|
||||
|
||||
@property
|
||||
def fallback_required(self) -> bool:
|
||||
return self.result is None
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class CompressionJob:
|
||||
"""회기종료 narrative 압축 작업(LLM, 비동기 비블로킹). 큐에 적재될 페이로드.
|
||||
|
|
@ -305,6 +349,101 @@ def build_compression_messages(job: CompressionJob) -> list[dict[str, str]]:
|
|||
return [{"role": "system", "content": system}, {"role": "user", "content": user}]
|
||||
|
||||
|
||||
def _normalize_digest_text(value: Any) -> str:
|
||||
return " ".join(str(value or "").split())
|
||||
|
||||
|
||||
def assess_llm_digest_quality(
|
||||
digest_input: SessionDigestInput,
|
||||
digest: Any,
|
||||
*,
|
||||
forbidden_substrings: tuple[str, ...] = (),
|
||||
min_chars: int = _LLM_DIGEST_MIN_CHARS,
|
||||
) -> DigestQualityAssessment:
|
||||
"""Validate an LLM digest candidate without using raw text or end-state data."""
|
||||
|
||||
text = _normalize_digest_text(digest)
|
||||
if not text:
|
||||
return DigestQualityAssessment(accepted=False, reason="empty", retryable=True)
|
||||
|
||||
prefix_match = _SESSION_DIGEST_PREFIX_RE.match(text)
|
||||
if prefix_match and int(prefix_match.group("session_no")) != digest_input.session_no:
|
||||
return DigestQualityAssessment(
|
||||
accepted=False,
|
||||
reason="wrong_session_prefix",
|
||||
retryable=True,
|
||||
details=(prefix_match.group(0),),
|
||||
)
|
||||
|
||||
forbidden_hits = tuple(
|
||||
marker
|
||||
for marker in (str(item).strip() for item in forbidden_substrings)
|
||||
if marker and marker in text
|
||||
)
|
||||
if forbidden_hits:
|
||||
return DigestQualityAssessment(
|
||||
accepted=False,
|
||||
reason="forbidden_substring",
|
||||
retryable=True,
|
||||
details=forbidden_hits[:5],
|
||||
)
|
||||
|
||||
lowered = text.lower()
|
||||
marker_hits = tuple(
|
||||
marker
|
||||
for marker in _LLM_DIGEST_INTERNAL_MARKERS
|
||||
if marker.lower() in lowered
|
||||
)
|
||||
if marker_hits:
|
||||
return DigestQualityAssessment(
|
||||
accepted=False,
|
||||
reason="internal_marker",
|
||||
retryable=True,
|
||||
details=marker_hits,
|
||||
)
|
||||
|
||||
if len(text) < max(1, int(min_chars)):
|
||||
return DigestQualityAssessment(accepted=False, reason="too_short", retryable=True)
|
||||
|
||||
return DigestQualityAssessment(accepted=True, reason="ok")
|
||||
|
||||
|
||||
def build_llm_digest_worker_outcome(
|
||||
digest_input: SessionDigestInput,
|
||||
digest: Any,
|
||||
*,
|
||||
forbidden_substrings: tuple[str, ...] = (),
|
||||
min_chars: int = _LLM_DIGEST_MIN_CHARS,
|
||||
) -> SessionDigestWorkerOutcome:
|
||||
"""Coerce a candidate LLM digest into the shared result contract if it passes."""
|
||||
|
||||
quality = assess_llm_digest_quality(
|
||||
digest_input,
|
||||
digest,
|
||||
forbidden_substrings=forbidden_substrings,
|
||||
min_chars=min_chars,
|
||||
)
|
||||
if not quality.accepted:
|
||||
return SessionDigestWorkerOutcome(result=None, quality=quality)
|
||||
|
||||
normalized = _normalize_digest_text(digest)
|
||||
prefix = f"S{digest_input.session_no}:"
|
||||
if not normalized.startswith(prefix):
|
||||
normalized = f"{prefix} {normalized}"
|
||||
|
||||
return SessionDigestWorkerOutcome(
|
||||
result=SessionDigestResult(
|
||||
session_id=digest_input.session_id,
|
||||
case_id=digest_input.case_id,
|
||||
session_no=digest_input.session_no,
|
||||
digest=normalized,
|
||||
open_threads=digest_input.open_threads,
|
||||
source="llm",
|
||||
),
|
||||
quality=quality,
|
||||
)
|
||||
|
||||
|
||||
def _compact(value: Any, *, limit: int = _SESSION_DIGEST_EXCERPT_CHARS) -> str:
|
||||
text = " ".join(str(value or "").split())
|
||||
if len(text) <= limit:
|
||||
|
|
@ -518,9 +657,13 @@ __all__ = [
|
|||
"MaskedDigestTurn",
|
||||
"SessionDigestInput",
|
||||
"SessionDigestResult",
|
||||
"DigestQualityAssessment",
|
||||
"SessionDigestWorkerOutcome",
|
||||
"make_carry_over",
|
||||
"build_session_digest_input",
|
||||
"build_compression_messages",
|
||||
"assess_llm_digest_quality",
|
||||
"build_llm_digest_worker_outcome",
|
||||
"build_fallback_digest_result",
|
||||
"build_fallback_session_digest",
|
||||
"PinnedFactCandidate",
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from ..engine_client import (
|
|||
StreamRequest,
|
||||
)
|
||||
from ..contracts.engine_gateway import (
|
||||
ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL,
|
||||
ENGINE_GATEWAY_SSE_DONE,
|
||||
ENGINE_GATEWAY_SSE_ERROR,
|
||||
EngineGatewaySseDecodeError,
|
||||
|
|
@ -40,7 +41,7 @@ from ..contracts.engine_gateway import (
|
|||
StreamTokenEvent,
|
||||
)
|
||||
from . import guardrail, persona, state_machine
|
||||
from .persona import PersonaCard, PersonaStateContext
|
||||
from .persona import PersonaCard, PersonaStateContext, TurnMemory
|
||||
from .state_machine import SessionState, Stage
|
||||
|
||||
# 평가 훅 타입: U_t(수련생 마스킹 발화) + 내담자응답 + 상태 → 평가 결과(dict)
|
||||
|
|
@ -63,10 +64,7 @@ class TurnContext:
|
|||
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)
|
||||
memory: TurnMemory = field(default_factory=TurnMemory)
|
||||
# 회기 이론모드(학습자 선택: humanistic|cbt|integrative). 평가 이론부합·생성 프레이밍에 사용.
|
||||
theory_mode: Optional[str] = None
|
||||
|
||||
|
|
@ -113,10 +111,7 @@ def prepare_turn(
|
|||
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,
|
||||
memory: Optional[TurnMemory] = None,
|
||||
theory_mode: Optional[str] = None,
|
||||
eval_rapport_signal: Optional[float] = None,
|
||||
) -> TurnContext:
|
||||
|
|
@ -125,16 +120,19 @@ def prepare_turn(
|
|||
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,
|
||||
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 []),
|
||||
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,
|
||||
)
|
||||
|
||||
|
|
@ -169,10 +167,7 @@ def prepare_turn(
|
|||
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,
|
||||
memory=ctx.memory,
|
||||
theory_mode=ctx.theory_mode,
|
||||
)
|
||||
return ctx
|
||||
|
|
@ -388,7 +383,11 @@ async def run_turn_stream(
|
|||
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"),
|
||||
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")),
|
||||
|
|
@ -405,7 +404,11 @@ async def run_turn_stream(
|
|||
"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"),
|
||||
"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")),
|
||||
|
|
@ -454,6 +457,7 @@ __all__ = [
|
|||
"EvalHook",
|
||||
"LlmAuditHook",
|
||||
"TurnContext",
|
||||
"TurnMemory",
|
||||
"TurnResult",
|
||||
"StreamEvent",
|
||||
"prepare_turn",
|
||||
|
|
|
|||
|
|
@ -90,6 +90,16 @@ class PersonaStateContext:
|
|||
affect_state: dict[str, float] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TurnMemory:
|
||||
"""L2/L4/L6 turn memory inputs assembled before persona prompt rendering."""
|
||||
|
||||
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)
|
||||
|
||||
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
# L0 — 역할 + 안전 가드레일 + 도식노출금지 (전 페르소나 공통, cache 대상)
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
|
|
@ -228,10 +238,7 @@ def build_turn_messages(
|
|||
state: PersonaStateContext,
|
||||
learner_text_masked: 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,
|
||||
memory: Optional[TurnMemory] = None,
|
||||
theory_mode: Optional[str] = None,
|
||||
) -> list[EngineMessage]:
|
||||
"""한 턴의 EngineMessage[] 조립 (L0~L6).
|
||||
|
|
@ -240,15 +247,14 @@ def build_turn_messages(
|
|||
card : 불변 페르소나(L0+L1)
|
||||
state : state_machine 산출 상태(L3)
|
||||
learner_text_masked : PII 마스킹된 수련생 발화(L5)
|
||||
recall_summary : 회기 시작 회상(L2-EP, 큰그림→세부 요약). CCD/정답 미포함.
|
||||
pinned_facts : 무손실 사실 hard-pin(L4). "자기 기억"으로만 표현.
|
||||
recent_turns : [{speaker, text}] 최근 K턴 버퍼(L6 직전 맥락)
|
||||
kb_behavior_cues : KB 증상 '행동단서'만(본문 비노출, sensitivity<=1)
|
||||
memory : 회상/고정 사실/최근 턴/KB 행동단서 묶음
|
||||
theory_mode : 회기 이론모드. 내담자 반응 프레이밍에만 사용.
|
||||
|
||||
반환 messages 순서: system(L0+L1, cache) → system(L2/L3/L4, cache 미설정) →
|
||||
assistant/user 히스토리 → user(이번 발화). 게이트웨이가 마지막 user 를 stdin 으로.
|
||||
assistant/user 최근 턴 기록 → user(이번 발화).
|
||||
현재 Python gateway split boundary 는 system 묶음과 마지막 user payload 만 소비한다.
|
||||
"""
|
||||
memory = memory or TurnMemory()
|
||||
messages: list[EngineMessage] = []
|
||||
|
||||
# L0+L1 — 정적, cache_control 대상
|
||||
|
|
@ -256,10 +262,10 @@ def build_turn_messages(
|
|||
|
||||
# L2 — 회상 + KB 행동단서 (회기 내 1회 로드, 캐시 친화)
|
||||
l2_parts: list[str] = []
|
||||
if recall_summary:
|
||||
l2_parts.append(f"[L2 회상 — 지난 맥락(큰그림→세부, 정답/평가 미포함)]\n{recall_summary}")
|
||||
if kb_behavior_cues:
|
||||
cues = "\n".join(f"- {c}" for c in kb_behavior_cues)
|
||||
if memory.recall_summary:
|
||||
l2_parts.append(f"[L2 회상 — 지난 맥락(큰그림→세부, 정답/평가 미포함)]\n{memory.recall_summary}")
|
||||
if memory.kb_behavior_cues:
|
||||
cues = "\n".join(f"- {c}" for c in memory.kb_behavior_cues)
|
||||
l2_parts.append(f"[L2 증상 행동단서(본문 비노출, 이렇게 '행동'으로만 드러난다)]\n{cues}")
|
||||
if l2_parts:
|
||||
messages.append(EngineMessage(role="system", content="\n\n".join(l2_parts), cache=True))
|
||||
|
|
@ -282,20 +288,20 @@ def build_turn_messages(
|
|||
messages.append(EngineMessage(role="system", content=theory_guidance, cache=False))
|
||||
|
||||
# L4 — pinned fact hard-pin (무손실, "자기 기억"으로만)
|
||||
if pinned_facts:
|
||||
pinned = "\n".join(f"- {f}" for f in pinned_facts)
|
||||
if memory.pinned_facts:
|
||||
pinned = "\n".join(f"- {f}" for f in memory.pinned_facts)
|
||||
messages.append(EngineMessage(
|
||||
role="system",
|
||||
content=("[L4 고정 사실 — 당신이 *이미 말했거나 사실인* 것. 모순되게 말하지 말 것]\n" + pinned),
|
||||
cache=False,
|
||||
))
|
||||
|
||||
# L6 — 직전 K턴 맥락 (히스토리). 게이트웨이가 단발이면 system 뒤 맥락으로 직렬화.
|
||||
if recent_turns:
|
||||
for t in recent_turns:
|
||||
# L6 — 직전 K턴 맥락. 현재 Python gateway 는 마지막 user payload 만 보내므로
|
||||
# non-system history records 는 요청 계약상 보존하고, 별도 prompt 동작 변경에서 소비한다.
|
||||
if memory.recent_turns:
|
||||
for t in memory.recent_turns:
|
||||
role = "assistant" if t.get("speaker") == "counselor" else "user"
|
||||
# 내담자(자기) 과거 발화는 assistant, 상담자 발화는 user 로 매핑하면
|
||||
# 게이트웨이가 [이전 상담자/내담자 발화]로 직렬화한다.
|
||||
# 내담자(자기) 과거 발화는 assistant, 상담자 발화는 user 로 매핑한다.
|
||||
messages.append(EngineMessage(role=role, content=t.get("text", ""), cache=False))
|
||||
|
||||
# L5 — 이번 수련생 발화 (마스킹 후)
|
||||
|
|
@ -476,6 +482,7 @@ def get_seed_persona(code: str) -> Optional[PersonaCard]:
|
|||
__all__ = [
|
||||
"PersonaCard",
|
||||
"PersonaStateContext",
|
||||
"TurnMemory",
|
||||
"L0_SAFETY",
|
||||
"build_persona_system_text",
|
||||
"build_turn_messages",
|
||||
|
|
|
|||
366
apps/api/app/services/session_digest_worker.py
Normal file
366
apps/api/app/services/session_digest_worker.py
Normal file
|
|
@ -0,0 +1,366 @@
|
|||
"""One-shot session digest worker boundary.
|
||||
|
||||
This module intentionally stops short of scheduling. It converts an existing
|
||||
CompressionJob into the shared engine gateway contract, applies the local digest
|
||||
quality gate, and updates persisted fallback rows only when the LLM candidate is
|
||||
accepted.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Awaitable, Callable, Protocol
|
||||
|
||||
from ..contracts.engine_gateway import (
|
||||
EngineMessage,
|
||||
GenerateRequest,
|
||||
GenerateResponse,
|
||||
normalize_engine_gateway_model,
|
||||
)
|
||||
from . import memory
|
||||
|
||||
LlmAuditHook = Callable[[dict[str, Any]], Awaitable[None]]
|
||||
|
||||
|
||||
class SessionDigestEngine(Protocol):
|
||||
async def generate(self, req: GenerateRequest) -> GenerateResponse: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LoadedSessionDigestJob:
|
||||
job: memory.CompressionJob
|
||||
existing_case_digest: str | None
|
||||
learner_id: str | None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SessionDigestApplyPlan:
|
||||
result: memory.SessionDigestResult
|
||||
case_digest: str | None
|
||||
compressed_by: str
|
||||
token_count: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SessionDigestWorkerRun:
|
||||
request: GenerateRequest
|
||||
response: GenerateResponse
|
||||
outcome: memory.SessionDigestWorkerOutcome
|
||||
apply_plan: SessionDigestApplyPlan | None
|
||||
|
||||
@property
|
||||
def fallback_required(self) -> bool:
|
||||
return self.apply_plan is None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SessionDigestOneShotResult:
|
||||
loaded: LoadedSessionDigestJob | None
|
||||
worker: SessionDigestWorkerRun | None
|
||||
applied: bool = False
|
||||
|
||||
@property
|
||||
def found(self) -> bool:
|
||||
return self.loaded is not None
|
||||
|
||||
|
||||
def build_session_digest_request(
|
||||
job: memory.CompressionJob,
|
||||
*,
|
||||
model: str | None = None,
|
||||
) -> GenerateRequest:
|
||||
"""Build the Node-compatible gateway request for narrative compression."""
|
||||
|
||||
return GenerateRequest(
|
||||
ai_role="evaluator",
|
||||
messages=[
|
||||
EngineMessage.model_validate(message)
|
||||
for message in memory.build_compression_messages(job)
|
||||
],
|
||||
model=normalize_engine_gateway_model(model),
|
||||
max_tokens=700,
|
||||
temperature=0.2,
|
||||
session_id=job.session_id,
|
||||
metadata={
|
||||
"loop": "session_digest",
|
||||
"case_id": job.case_id,
|
||||
"session_no": job.session_no,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def build_session_digest_apply_plan(
|
||||
digest_input: memory.SessionDigestInput,
|
||||
response: GenerateResponse,
|
||||
*,
|
||||
existing_case_digest: str | None = None,
|
||||
forbidden_substrings: tuple[str, ...] = (),
|
||||
) -> tuple[memory.SessionDigestWorkerOutcome, SessionDigestApplyPlan | None]:
|
||||
"""Validate an LLM response and prepare idempotent persistence arguments."""
|
||||
|
||||
outcome = memory.build_llm_digest_worker_outcome(
|
||||
digest_input,
|
||||
response.text,
|
||||
forbidden_substrings=forbidden_substrings,
|
||||
)
|
||||
if outcome.result is None:
|
||||
return outcome, None
|
||||
|
||||
case_digest = None
|
||||
if outcome.result.case_id is not None:
|
||||
case_digest = memory.merge_case_digest(
|
||||
existing_digest=existing_case_digest,
|
||||
session_no=outcome.result.session_no,
|
||||
session_digest=outcome.result.digest,
|
||||
)
|
||||
return outcome, SessionDigestApplyPlan(
|
||||
result=outcome.result,
|
||||
case_digest=case_digest,
|
||||
compressed_by=_compressed_by(response),
|
||||
token_count=_token_count(response),
|
||||
)
|
||||
|
||||
|
||||
async def run_session_digest_worker(
|
||||
job: memory.CompressionJob,
|
||||
engine: SessionDigestEngine,
|
||||
*,
|
||||
existing_case_digest: str | None = None,
|
||||
forbidden_substrings: tuple[str, ...] = (),
|
||||
model: str | None = None,
|
||||
audit_hook: LlmAuditHook | None = None,
|
||||
) -> SessionDigestWorkerRun:
|
||||
"""Run a single digest candidate through engine, audit, quality gate, plan."""
|
||||
|
||||
request = build_session_digest_request(job, model=model)
|
||||
started = time.perf_counter()
|
||||
response = await engine.generate(request)
|
||||
latency_ms = int((time.perf_counter() - started) * 1000)
|
||||
await _record_llm_audit(audit_hook, response, request.session_id, latency_ms)
|
||||
outcome, apply_plan = build_session_digest_apply_plan(
|
||||
job.digest_input,
|
||||
response,
|
||||
existing_case_digest=existing_case_digest,
|
||||
forbidden_substrings=forbidden_substrings,
|
||||
)
|
||||
return SessionDigestWorkerRun(
|
||||
request=request,
|
||||
response=response,
|
||||
outcome=outcome,
|
||||
apply_plan=apply_plan,
|
||||
)
|
||||
|
||||
|
||||
async def load_session_digest_job(conn: Any, session_id: str) -> LoadedSessionDigestJob | None:
|
||||
"""Load a persisted fallback summary plus masked client-visible transcript."""
|
||||
|
||||
summary = await conn.fetchrow(
|
||||
"""
|
||||
SELECT
|
||||
ss.session_id, ss.case_id, ss.session_no, ss.open_threads,
|
||||
cp.case_digest, s.learner_id
|
||||
FROM app.session_summary ss
|
||||
JOIN app.sessions s ON s.id = ss.session_id
|
||||
LEFT JOIN app.case_profile cp
|
||||
ON cp.case_id = ss.case_id
|
||||
AND cp.learner_id = s.learner_id
|
||||
WHERE ss.session_id = $1::uuid
|
||||
AND ss.compressed_by IS NULL
|
||||
""",
|
||||
session_id,
|
||||
)
|
||||
if summary is None:
|
||||
return None
|
||||
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT id, speaker, text_masked, visible_to
|
||||
FROM app.turns
|
||||
WHERE session_id = $1::uuid
|
||||
AND speaker = ANY($2::text[])
|
||||
ORDER BY seq ASC
|
||||
""",
|
||||
session_id,
|
||||
["counselor", "client"],
|
||||
)
|
||||
digest_input = memory.build_session_digest_input(
|
||||
session_id=str(_row_get(summary, "session_id", session_id)),
|
||||
case_id=_optional_str(_row_get(summary, "case_id")),
|
||||
session_no=int(_row_get(summary, "session_no", 0) or 0),
|
||||
masked_turns=[_turn_from_row(row) for row in rows],
|
||||
open_threads=_open_threads(_row_get(summary, "open_threads")),
|
||||
)
|
||||
return LoadedSessionDigestJob(
|
||||
job=memory.CompressionJob(digest_input=digest_input),
|
||||
existing_case_digest=_optional_str(_row_get(summary, "case_digest")),
|
||||
learner_id=_optional_str(_row_get(summary, "learner_id")),
|
||||
)
|
||||
|
||||
|
||||
async def apply_session_digest_plan(
|
||||
conn: Any,
|
||||
plan: SessionDigestApplyPlan,
|
||||
*,
|
||||
learner_id: str | None,
|
||||
) -> bool:
|
||||
"""Replace fallback digest rows after quality acceptance only."""
|
||||
|
||||
applied = _update_applied(await conn.execute(
|
||||
"""
|
||||
UPDATE app.session_summary
|
||||
SET digest = $2,
|
||||
open_threads = $3::jsonb,
|
||||
compressed_by = $4,
|
||||
token_count = $5
|
||||
WHERE session_id = $1::uuid
|
||||
AND compressed_by IS NULL
|
||||
""",
|
||||
plan.result.session_id,
|
||||
plan.result.digest,
|
||||
list(plan.result.open_threads),
|
||||
plan.compressed_by,
|
||||
plan.token_count,
|
||||
))
|
||||
if not applied:
|
||||
return False
|
||||
if plan.result.case_id is None or learner_id is None or plan.case_digest is None:
|
||||
return True
|
||||
await conn.execute(
|
||||
"""
|
||||
UPDATE app.case_profile
|
||||
SET case_digest = $3,
|
||||
updated_at = now()
|
||||
WHERE case_id = $1::uuid
|
||||
AND learner_id = $2::uuid
|
||||
""",
|
||||
plan.result.case_id,
|
||||
learner_id,
|
||||
plan.case_digest,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
async def run_session_digest_once(
|
||||
conn: Any,
|
||||
*,
|
||||
session_id: str,
|
||||
engine: SessionDigestEngine,
|
||||
forbidden_substrings: tuple[str, ...] = (),
|
||||
model: str | None = None,
|
||||
audit_hook: LlmAuditHook | None = None,
|
||||
persist_accepted: bool = True,
|
||||
) -> SessionDigestOneShotResult:
|
||||
"""One-shot DB loader/worker/apply helper for a single ended session.
|
||||
|
||||
This is convenient for tests and dry-run CLIs. A production scheduler should
|
||||
load the job, release the DB connection, call the engine, then briefly
|
||||
reacquire a connection for apply_session_digest_plan().
|
||||
"""
|
||||
|
||||
loaded = await load_session_digest_job(conn, session_id)
|
||||
if loaded is None:
|
||||
return SessionDigestOneShotResult(loaded=None, worker=None)
|
||||
|
||||
worker = await run_session_digest_worker(
|
||||
loaded.job,
|
||||
engine,
|
||||
existing_case_digest=loaded.existing_case_digest,
|
||||
forbidden_substrings=forbidden_substrings,
|
||||
model=model,
|
||||
audit_hook=audit_hook,
|
||||
)
|
||||
applied = False
|
||||
if persist_accepted and worker.apply_plan is not None:
|
||||
applied = await apply_session_digest_plan(
|
||||
conn,
|
||||
worker.apply_plan,
|
||||
learner_id=loaded.learner_id,
|
||||
)
|
||||
return SessionDigestOneShotResult(loaded=loaded, worker=worker, applied=applied)
|
||||
|
||||
|
||||
async def _record_llm_audit(
|
||||
audit_hook: LlmAuditHook | None,
|
||||
response: GenerateResponse,
|
||||
session_id: str | None,
|
||||
latency_ms: int,
|
||||
) -> None:
|
||||
if audit_hook is None:
|
||||
return
|
||||
try:
|
||||
await audit_hook(
|
||||
{
|
||||
"session_id": session_id,
|
||||
"provider": response.provider,
|
||||
"model": response.model,
|
||||
"tokens_in": response.tokens_in,
|
||||
"tokens_out": response.tokens_out,
|
||||
"cost_usd": response.cost_usd,
|
||||
"inference_geo": response.inference_geo,
|
||||
"latency_ms": latency_ms,
|
||||
}
|
||||
)
|
||||
except Exception:
|
||||
return
|
||||
|
||||
|
||||
def _compressed_by(response: GenerateResponse) -> str:
|
||||
provider = (response.provider or "unknown").strip() or "unknown"
|
||||
model = (response.model or "unknown").strip() or "unknown"
|
||||
return f"llm:{provider}/{model}"
|
||||
|
||||
|
||||
def _token_count(response: GenerateResponse) -> int:
|
||||
return max(0, int(response.tokens_in or 0)) + max(0, int(response.tokens_out or 0))
|
||||
|
||||
|
||||
def _update_applied(status: Any) -> bool:
|
||||
return str(status).upper().strip().endswith(" 1")
|
||||
|
||||
|
||||
def _row_get(row: Any, key: str, default: Any = None) -> Any:
|
||||
if isinstance(row, dict):
|
||||
return row.get(key, default)
|
||||
try:
|
||||
return row[key]
|
||||
except (KeyError, IndexError, TypeError):
|
||||
return default
|
||||
|
||||
|
||||
def _optional_str(value: Any) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
text = str(value).strip()
|
||||
return text or None
|
||||
|
||||
|
||||
def _open_threads(value: Any) -> list[str]:
|
||||
if not isinstance(value, list):
|
||||
return []
|
||||
return [str(item).strip() for item in value if str(item).strip()]
|
||||
|
||||
|
||||
def _turn_from_row(row: Any) -> dict[str, Any]:
|
||||
return {
|
||||
"turn_id": _optional_str(_row_get(row, "id")),
|
||||
"speaker": _optional_str(_row_get(row, "speaker")) or "",
|
||||
"text_masked": _optional_str(_row_get(row, "text_masked")) or "",
|
||||
"text": "",
|
||||
"visible_to": _row_get(row, "visible_to"),
|
||||
}
|
||||
|
||||
|
||||
__all__ = [
|
||||
"LoadedSessionDigestJob",
|
||||
"SessionDigestApplyPlan",
|
||||
"SessionDigestEngine",
|
||||
"SessionDigestOneShotResult",
|
||||
"SessionDigestWorkerRun",
|
||||
"apply_session_digest_plan",
|
||||
"build_session_digest_apply_plan",
|
||||
"build_session_digest_request",
|
||||
"load_session_digest_job",
|
||||
"run_session_digest_once",
|
||||
"run_session_digest_worker",
|
||||
]
|
||||
|
|
@ -88,18 +88,32 @@ def turn_rapport(ev: dict[str, Any]) -> float | None:
|
|||
return max(-1.0, min(1.0, value))
|
||||
|
||||
|
||||
def turn_technique_label(item: object) -> str | None:
|
||||
if isinstance(item, dict):
|
||||
label = (
|
||||
item.get("label_ko")
|
||||
or item.get("label")
|
||||
or item.get("name")
|
||||
or item.get("id")
|
||||
or item.get("code")
|
||||
)
|
||||
else:
|
||||
label = item
|
||||
if label is None:
|
||||
return None
|
||||
text = str(label).strip()
|
||||
return text or None
|
||||
|
||||
|
||||
def turn_techniques(ev: dict[str, Any]) -> list[str]:
|
||||
raw = ev.get("techniques")
|
||||
if not isinstance(raw, list):
|
||||
return []
|
||||
labels: list[str] = []
|
||||
for item in raw:
|
||||
if isinstance(item, dict):
|
||||
label = item.get("label") or item.get("name") or item.get("id") or item.get("code")
|
||||
else:
|
||||
label = item
|
||||
label = turn_technique_label(item)
|
||||
if label:
|
||||
labels.append(str(label))
|
||||
labels.append(label)
|
||||
return labels
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import uuid
|
|||
import hashlib
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Iterable
|
||||
from typing import Any, Iterable, Protocol
|
||||
|
||||
from .db import acquire, get_pool
|
||||
from .deps import Principal
|
||||
|
|
@ -56,6 +56,78 @@ class SessionSummaryWrite:
|
|||
open_threads: list[str]
|
||||
|
||||
|
||||
class _SessionEvaluationResult(Protocol):
|
||||
scope: str
|
||||
stage: str
|
||||
error: str | None
|
||||
|
||||
def to_dict(self) -> dict[str, Any]: ...
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SessionEvaluationWrite:
|
||||
session_id: str
|
||||
learner_id: str
|
||||
status: str
|
||||
source: str
|
||||
scope: str
|
||||
stage: str
|
||||
payload: dict[str, Any]
|
||||
error: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_result(
|
||||
cls,
|
||||
*,
|
||||
session_id: str,
|
||||
learner_id: str,
|
||||
result: _SessionEvaluationResult,
|
||||
source: str = "engine",
|
||||
) -> "SessionEvaluationWrite":
|
||||
return cls(
|
||||
session_id=session_id,
|
||||
learner_id=learner_id,
|
||||
status="error" if result.error else "ready",
|
||||
source=source,
|
||||
scope=result.scope,
|
||||
stage=result.stage,
|
||||
payload=result.to_dict(),
|
||||
error=result.error,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_error(
|
||||
cls,
|
||||
*,
|
||||
session_id: str,
|
||||
learner_id: str,
|
||||
scope: str,
|
||||
stage: str,
|
||||
error: BaseException | str,
|
||||
source: str = "engine",
|
||||
) -> "SessionEvaluationWrite":
|
||||
return cls(
|
||||
session_id=session_id,
|
||||
learner_id=learner_id,
|
||||
status="error",
|
||||
source=source,
|
||||
scope=scope,
|
||||
stage=stage,
|
||||
payload={},
|
||||
error=str(error),
|
||||
)
|
||||
|
||||
def cache_record(self) -> dict[str, Any]:
|
||||
return {
|
||||
"status": self.status,
|
||||
"source": self.source,
|
||||
"scope": self.scope,
|
||||
"stage": self.stage,
|
||||
"payload": self.payload,
|
||||
"error": self.error,
|
||||
}
|
||||
|
||||
|
||||
_JOINED_CARD_COLUMNS = (
|
||||
"card_persona_id",
|
||||
"card_code",
|
||||
|
|
@ -1125,30 +1197,13 @@ async def ensure_review_tables() -> None:
|
|||
return
|
||||
|
||||
|
||||
async def save_session_evaluation(
|
||||
*,
|
||||
session_id: str,
|
||||
learner_id: str,
|
||||
status: str,
|
||||
source: str,
|
||||
scope: str,
|
||||
stage: str,
|
||||
payload: dict[str, Any],
|
||||
error: str | None = None,
|
||||
) -> bool:
|
||||
record = {
|
||||
"status": status,
|
||||
"source": source,
|
||||
"scope": scope,
|
||||
"stage": stage,
|
||||
"payload": payload,
|
||||
"error": error,
|
||||
}
|
||||
async def save_session_evaluation(write: SessionEvaluationWrite) -> bool:
|
||||
record = write.cache_record()
|
||||
if runtime_fallback_allowed():
|
||||
_EVALUATION_CACHE[session_id] = record
|
||||
_EVALUATION_CACHE[write.session_id] = record
|
||||
try:
|
||||
get_pool()
|
||||
async with acquire(role="learner", user_id=learner_id) as conn:
|
||||
async with acquire(role="learner", user_id=write.learner_id) as conn:
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO app.session_evaluation (
|
||||
|
|
@ -1165,13 +1220,13 @@ async def save_session_evaluation(
|
|||
error = EXCLUDED.error,
|
||||
updated_at = now()
|
||||
""",
|
||||
session_id,
|
||||
status,
|
||||
source,
|
||||
scope,
|
||||
stage,
|
||||
payload,
|
||||
error,
|
||||
write.session_id,
|
||||
write.status,
|
||||
write.source,
|
||||
write.scope,
|
||||
write.stage,
|
||||
write.payload,
|
||||
write.error,
|
||||
)
|
||||
return True
|
||||
except Exception:
|
||||
|
|
@ -2301,7 +2356,9 @@ async def end_session(sess: InProcSession, carry: memory.CarryOver) -> bool:
|
|||
end_state = EXCLUDED.end_state,
|
||||
rapport_delta = EXCLUDED.rapport_delta,
|
||||
digest = EXCLUDED.digest,
|
||||
open_threads = EXCLUDED.open_threads
|
||||
open_threads = EXCLUDED.open_threads,
|
||||
compressed_by = NULL,
|
||||
token_count = NULL
|
||||
""",
|
||||
summary_write.session_id,
|
||||
summary_write.case_id,
|
||||
|
|
|
|||
|
|
@ -385,10 +385,12 @@ class AuthProviderScaffoldTest(unittest.IsolatedAsyncioTestCase):
|
|||
patch.object(auth_sessions, "get_pool", side_effect=RuntimeError("no db")),
|
||||
):
|
||||
target = await auth_sessions.upsert_managed_user(
|
||||
email="admin-grant-target@hs.ac.kr",
|
||||
display_name="Grant Target",
|
||||
role="learner",
|
||||
external_id="dev:admin-grant-target@hs.ac.kr",
|
||||
auth_sessions.ManagedUserUpsertInput(
|
||||
email="admin-grant-target@hs.ac.kr",
|
||||
display_name="Grant Target",
|
||||
role="learner",
|
||||
external_id="dev:admin-grant-target@hs.ac.kr",
|
||||
)
|
||||
)
|
||||
operator = Principal(
|
||||
user_id="00000000-0000-0000-0000-000000000602",
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from unittest.mock import patch
|
|||
from .deps import Principal, Role
|
||||
from . import session_persistence
|
||||
from .routes import sessions
|
||||
from .services import evaluator
|
||||
from .services import evaluator, session_metrics
|
||||
from .services import persona as persona_service
|
||||
from .services import state_machine
|
||||
from .store import InProcSession
|
||||
|
|
@ -45,6 +45,16 @@ class FakeAcquire:
|
|||
|
||||
|
||||
class EvaluationPersistenceMappingTest(unittest.TestCase):
|
||||
def test_session_metrics_prefers_rehydrated_technique_label_ko(self) -> None:
|
||||
ev = {
|
||||
"techniques": [
|
||||
{"code": "empathy", "label_ko": "공감", "label": "legacy empathy"},
|
||||
{"code": "open_question"},
|
||||
]
|
||||
}
|
||||
|
||||
self.assertEqual(session_metrics.turn_techniques(ev), ["공감", "open_question"])
|
||||
|
||||
def test_fast_evaluator_masks_client_reply_before_prompting(self) -> None:
|
||||
card = persona_service.P1
|
||||
state = state_machine.init_state(params=card.openness_params())
|
||||
|
|
@ -106,6 +116,65 @@ class EvaluationPersistenceMappingTest(unittest.TestCase):
|
|||
self.assertEqual(rows["technique:empathy"]["rationale"], "감정을 명시적으로 반영했다.")
|
||||
self.assertEqual(rows["client_state:affect_contact"]["rationale"], "내담자가 감정을 언급했다.")
|
||||
|
||||
def test_alternative_rows_accept_string_and_dict_shapes(self) -> None:
|
||||
rows = session_persistence._evaluation_alternative_rows(
|
||||
{
|
||||
"alternative_utterances": [
|
||||
"감정을 먼저 반영해 보세요.",
|
||||
{"text": "조언 전에 의미를 확인해 보세요.", "rationale": "성급한 해결 방지"},
|
||||
{"suggestion": "침묵을 허용해 보세요."},
|
||||
{"rationale": "빈 제안은 저장하지 않음"},
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
rows,
|
||||
[
|
||||
{"suggestion": "감정을 먼저 반영해 보세요.", "rationale": None},
|
||||
{"suggestion": "조언 전에 의미를 확인해 보세요.", "rationale": "성급한 해결 방지"},
|
||||
{"suggestion": "침묵을 허용해 보세요.", "rationale": None},
|
||||
],
|
||||
)
|
||||
|
||||
def test_session_evaluation_write_from_result_preserves_payload_shape(self) -> None:
|
||||
result = evaluator.SessionEvaluation(
|
||||
session_id="session-1",
|
||||
stage="explore",
|
||||
scope="session_end",
|
||||
turns_evaluated=2,
|
||||
)
|
||||
|
||||
write = session_persistence.SessionEvaluationWrite.from_result(
|
||||
session_id="session-1",
|
||||
learner_id="learner-1",
|
||||
result=result,
|
||||
)
|
||||
|
||||
self.assertEqual(write.status, "ready")
|
||||
self.assertEqual(write.source, "engine")
|
||||
self.assertEqual(write.scope, "session_end")
|
||||
self.assertEqual(write.stage, "explore")
|
||||
self.assertEqual(write.payload, result.to_dict())
|
||||
self.assertNotIn("error", write.payload)
|
||||
self.assertIsNone(write.error)
|
||||
|
||||
def test_session_evaluation_write_from_error_preserves_fallback_shape(self) -> None:
|
||||
write = session_persistence.SessionEvaluationWrite.from_error(
|
||||
session_id="session-1",
|
||||
learner_id="learner-1",
|
||||
scope="session_end",
|
||||
stage="explore",
|
||||
error=RuntimeError("engine timeout"),
|
||||
)
|
||||
|
||||
self.assertEqual(write.status, "error")
|
||||
self.assertEqual(write.source, "engine")
|
||||
self.assertEqual(write.scope, "session_end")
|
||||
self.assertEqual(write.stage, "explore")
|
||||
self.assertEqual(write.payload, {})
|
||||
self.assertEqual(write.error, "engine timeout")
|
||||
|
||||
def test_rebuild_turn_evaluation_restores_review_shape(self) -> None:
|
||||
rebuilt = session_persistence._rebuild_turn_evaluations(
|
||||
[("11111111-1111-1111-1111-111111111111", 2, "탐색")],
|
||||
|
|
|
|||
|
|
@ -51,12 +51,14 @@ def _prepare_context() -> orchestrator.TurnContext:
|
|||
card=persona.P1,
|
||||
state=_initial_state(),
|
||||
learner_text=RAW_TEXT,
|
||||
recent_turns=[
|
||||
{
|
||||
"speaker": "counselor",
|
||||
"text": "Previous learner contact was already masked: [PHONE] [EMAIL] [RRN].",
|
||||
}
|
||||
],
|
||||
memory=orchestrator.TurnMemory(
|
||||
recent_turns=[
|
||||
{
|
||||
"speaker": "counselor",
|
||||
"text": "Previous learner contact was already masked: [PHONE] [EMAIL] [RRN].",
|
||||
}
|
||||
],
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -163,11 +165,13 @@ class OrchestratorMaskingGateTest(unittest.IsolatedAsyncioTestCase):
|
|||
card=persona.P1,
|
||||
state=_initial_state(),
|
||||
learner_text="Current text has no identifiers.",
|
||||
recall_summary=f"Recall mentioned {RAW_PHONE}.",
|
||||
pinned_facts=[f"Pinned email {RAW_EMAIL}."],
|
||||
recent_turns=[
|
||||
{"speaker": "counselor", "text": f"Previous raw RRN {RAW_RRN}."},
|
||||
],
|
||||
memory=orchestrator.TurnMemory(
|
||||
recall_summary=f"Recall mentioned {RAW_PHONE}.",
|
||||
pinned_facts=[f"Pinned email {RAW_EMAIL}."],
|
||||
recent_turns=[
|
||||
{"speaker": "counselor", "text": f"Previous raw RRN {RAW_RRN}."},
|
||||
],
|
||||
),
|
||||
)
|
||||
|
||||
blob = _message_blob(ctx.messages)
|
||||
|
|
@ -203,11 +207,13 @@ class OrchestratorMaskingGateTest(unittest.IsolatedAsyncioTestCase):
|
|||
card=persona.P1,
|
||||
state=_initial_state(),
|
||||
learner_text=RAW_KO_TEXT,
|
||||
recall_summary=f"지난 회기 요약에 {RAW_KO_NAME}과 {RAW_KO_ORG}가 남아 있었다.",
|
||||
pinned_facts=[f"소속 {RAW_KO_DEPT}"],
|
||||
recent_turns=[
|
||||
{"speaker": "counselor", "text": f"{RAW_KO_NAME} 씨가 상담실에 왔다."},
|
||||
],
|
||||
memory=orchestrator.TurnMemory(
|
||||
recall_summary=f"지난 회기 요약에 {RAW_KO_NAME}과 {RAW_KO_ORG}가 남아 있었다.",
|
||||
pinned_facts=[f"소속 {RAW_KO_DEPT}"],
|
||||
recent_turns=[
|
||||
{"speaker": "counselor", "text": f"{RAW_KO_NAME} 씨가 상담실에 왔다."},
|
||||
],
|
||||
),
|
||||
)
|
||||
|
||||
blob = _message_blob(ctx.messages)
|
||||
|
|
|
|||
110
apps/api/app/test_persona_generation_contract.py
Normal file
110
apps/api/app/test_persona_generation_contract.py
Normal file
|
|
@ -0,0 +1,110 @@
|
|||
"""Regression tests for persona draft generation contract helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from .contracts.engine_gateway import GenerateResponse
|
||||
from . import persona_generation_contract as contract
|
||||
from .persona_read_model import PersonaDraftGenerateRequest
|
||||
|
||||
|
||||
class PersonaGenerationContractTest(unittest.TestCase):
|
||||
def test_persona_generation_schema_pins_required_fields(self) -> None:
|
||||
schema = contract.persona_generation_schema()
|
||||
draft_schema = schema["properties"]["draft"]
|
||||
|
||||
self.assertFalse(schema["additionalProperties"])
|
||||
self.assertEqual(schema["required"], ["draft", "source_summary", "warnings"])
|
||||
self.assertFalse(draft_schema["additionalProperties"])
|
||||
self.assertIn("source_provenance", draft_schema["required"])
|
||||
self.assertEqual(draft_schema["properties"]["difficulty"]["enum"], ["easy", "moderate", "hard"])
|
||||
|
||||
def test_persona_draft_prompt_bundle_is_stable(self) -> None:
|
||||
bundle = contract.persona_draft_prompt_bundle()
|
||||
|
||||
self.assertEqual(bundle["id"], "persona-draft-rag")
|
||||
self.assertEqual(bundle["version"], "2026-06-28.1")
|
||||
self.assertRegex(bundle["hash"], r"^[0-9a-f]{12}$")
|
||||
|
||||
def test_persona_generation_payload_from_response_uses_legacy_json_fallback(self) -> None:
|
||||
response = GenerateResponse(
|
||||
text='```json\n{"source_summary":"요약","warnings":["검수 필요"]}\n```',
|
||||
model="test-model",
|
||||
provider="test-provider",
|
||||
structured=None,
|
||||
)
|
||||
|
||||
payload = contract.persona_generation_payload_from_response(response)
|
||||
|
||||
self.assertEqual(payload["source_summary"], "요약")
|
||||
self.assertEqual(payload["warnings"], ["검수 필요"])
|
||||
|
||||
def test_persona_generation_payload_from_response_returns_empty_dict_for_invalid_output(self) -> None:
|
||||
response = GenerateResponse(text="[1, 2, 3]", model="test-model", provider="test-provider")
|
||||
|
||||
self.assertEqual(contract.persona_generation_payload_from_response(response), {})
|
||||
|
||||
def test_coerce_persona_generated_draft_normalizes_nested_draft(self) -> None:
|
||||
request = PersonaDraftGenerateRequest(
|
||||
source_ids=["persona_authoring_test"],
|
||||
source_kind="client_record",
|
||||
code_hint="p9",
|
||||
display_name_hint="힌트 페르소나",
|
||||
difficulty="hard",
|
||||
theory_target=["CBT"],
|
||||
)
|
||||
|
||||
draft = contract.coerce_persona_generated_draft(
|
||||
{
|
||||
"draft": {
|
||||
"difficulty": "unsupported",
|
||||
"theory_target": [" Humanistic ", ""],
|
||||
"big5": {"O": 0.8, "C": "ignored"},
|
||||
"resistance": {},
|
||||
"affect_baseline": {},
|
||||
"is_synthetic": False,
|
||||
}
|
||||
},
|
||||
request,
|
||||
)
|
||||
|
||||
self.assertEqual(draft.code, "P9")
|
||||
self.assertEqual(draft.display_name, "힌트 페르소나")
|
||||
self.assertEqual(draft.difficulty, "hard")
|
||||
self.assertEqual(draft.theory_target, ["humanistic"])
|
||||
self.assertEqual(draft.big5, {"O": 0.8})
|
||||
self.assertEqual(draft.resistance["base_resistance"], 0.5)
|
||||
self.assertEqual(draft.source_provenance, "masked client_record")
|
||||
self.assertFalse(draft.is_synthetic)
|
||||
self.assertFalse(draft.submit_for_review)
|
||||
|
||||
def test_coerce_persona_generated_draft_accepts_top_level_payload_without_draft(self) -> None:
|
||||
request = PersonaDraftGenerateRequest(
|
||||
source_ids=["persona_authoring_test"],
|
||||
source_kind="mixed_notes",
|
||||
difficulty="moderate",
|
||||
)
|
||||
|
||||
draft = contract.coerce_persona_generated_draft(
|
||||
{
|
||||
"code": "p10",
|
||||
"display_name": "상위 페르소나",
|
||||
"difficulty": "easy",
|
||||
"theory_target": ["CBT"],
|
||||
"demographics": "not-a-dict",
|
||||
"presenting": {"summary": "불안"},
|
||||
},
|
||||
request,
|
||||
)
|
||||
|
||||
self.assertEqual(draft.code, "P10")
|
||||
self.assertEqual(draft.display_name, "상위 페르소나")
|
||||
self.assertEqual(draft.difficulty, "easy")
|
||||
self.assertEqual(draft.theory_target, ["cbt"])
|
||||
self.assertEqual(draft.demographics, {})
|
||||
self.assertEqual(draft.presenting, {"summary": "불안"})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
@ -11,7 +11,7 @@ from unittest.mock import AsyncMock, patch
|
|||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from . import persona_repository, session_persistence
|
||||
from . import persona_generation_contract, persona_repository, session_persistence
|
||||
from .deps import Principal, Role
|
||||
from .engine_client import GenerateResponse
|
||||
from .persona_repository import PersonaDraftRecord, PersonaReviewItem
|
||||
|
|
@ -804,6 +804,7 @@ class PersonaReviewQueueTest(unittest.IsolatedAsyncioTestCase):
|
|||
self.assertEqual(captured[0].metadata["prompt_bundle"]["id"], "persona-draft-rag")
|
||||
self.assertEqual(captured[0].metadata["prompt_bundle"]["version"], "2026-06-28.1")
|
||||
self.assertRegex(captured[0].metadata["prompt_bundle"]["hash"], r"^[0-9a-f]{12}$")
|
||||
self.assertEqual(captured[0].structured_schema, persona_generation_contract.persona_generation_schema())
|
||||
sent_text = captured[0].messages[-1].content
|
||||
self.assertNotIn("010-1234-5678", sent_text)
|
||||
self.assertIn("[PHONE]", sent_text)
|
||||
|
|
|
|||
|
|
@ -343,7 +343,7 @@ class LearnerSessionIdorTest(unittest.IsolatedAsyncioTestCase):
|
|||
|
||||
async def successful_turn(ctx, engine, **kwargs):
|
||||
nonlocal captured_recent_turns
|
||||
captured_recent_turns = list(ctx.recent_turns)
|
||||
captured_recent_turns = list(ctx.memory.recent_turns)
|
||||
assert ctx.state_after is not None
|
||||
return sessions.orchestrator.TurnResult(
|
||||
turn_seq=ctx.state_after.turn_seq,
|
||||
|
|
|
|||
312
apps/api/app/test_session_digest_worker.py
Normal file
312
apps/api/app/test_session_digest_worker.py
Normal file
|
|
@ -0,0 +1,312 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from .contracts.engine_gateway import GenerateResponse
|
||||
from .services import memory, session_digest_worker
|
||||
|
||||
|
||||
SESSION_ID = "00000000-0000-0000-0000-00000000feed"
|
||||
CASE_ID = "00000000-0000-0000-0000-00000000ca5e"
|
||||
LEARNER_ID = "00000000-0000-0000-0000-000000000101"
|
||||
|
||||
|
||||
def _digest_input(session_no: int = 2) -> memory.SessionDigestInput:
|
||||
return memory.build_session_digest_input(
|
||||
session_id=SESSION_ID,
|
||||
case_id=CASE_ID,
|
||||
session_no=session_no,
|
||||
masked_turns=[
|
||||
{
|
||||
"speaker": "client",
|
||||
"text": "저는 [NAME]이고 가족 갈등을 조심스럽게 설명했습니다.",
|
||||
"visible_to": ["client", "evaluator"],
|
||||
},
|
||||
{
|
||||
"speaker": "counselor",
|
||||
"text": "그 이야기는 다음 회기에서 천천히 이어가겠습니다.",
|
||||
"visible_to": ["client"],
|
||||
},
|
||||
{
|
||||
"speaker": "client",
|
||||
"text": "평가자 전용 발화",
|
||||
"visible_to": ["evaluator"],
|
||||
},
|
||||
],
|
||||
open_threads=["가족 갈등을 다음 회기에서 이어가기"],
|
||||
)
|
||||
|
||||
|
||||
def _accepted_text(session_no: int = 2) -> str:
|
||||
return (
|
||||
f"S{session_no}: 내담자는 [NAME]으로 지칭되며 가족 갈등을 조심스럽게 꺼냈다. "
|
||||
"상담자는 감정을 서두르지 않고 확인했고 다음 회기에서 같은 주제를 이어가기로 했다. "
|
||||
"내담자는 관계 이야기를 계속 다루는 데 약간의 부담과 기대를 함께 보였다."
|
||||
)
|
||||
|
||||
|
||||
class FakeEngine:
|
||||
def __init__(self, text: str) -> None:
|
||||
self.text = text
|
||||
self.requests = []
|
||||
|
||||
async def generate(self, req):
|
||||
self.requests.append(req)
|
||||
return GenerateResponse(
|
||||
text=self.text,
|
||||
provider="openai",
|
||||
model="gpt-4.1-mini",
|
||||
tokens_in=20,
|
||||
tokens_out=13,
|
||||
cost_usd=0.001,
|
||||
inference_geo="us",
|
||||
)
|
||||
|
||||
|
||||
class SessionDigestWorkerPureTest(unittest.IsolatedAsyncioTestCase):
|
||||
def test_request_uses_gateway_contract_without_internal_state(self) -> None:
|
||||
job = memory.CompressionJob(digest_input=_digest_input())
|
||||
|
||||
request = session_digest_worker.build_session_digest_request(job)
|
||||
prompt = "\n".join(message.content for message in request.messages)
|
||||
|
||||
self.assertEqual(request.ai_role, "evaluator")
|
||||
self.assertEqual(request.session_id, SESSION_ID)
|
||||
self.assertEqual(request.metadata["loop"], "session_digest")
|
||||
self.assertEqual(request.metadata["case_id"], CASE_ID)
|
||||
self.assertEqual(request.metadata["session_no"], 2)
|
||||
self.assertEqual([message.role for message in request.messages], ["system", "user"])
|
||||
self.assertIn("[NAME]", prompt)
|
||||
self.assertIn("가족 갈등을 다음 회기에서 이어가기", prompt)
|
||||
self.assertNotIn("김서연", prompt)
|
||||
self.assertNotIn("end_state", prompt)
|
||||
self.assertNotIn("rapport_credit", prompt)
|
||||
self.assertNotIn("CCD", prompt)
|
||||
self.assertNotIn("평가자 전용", prompt)
|
||||
|
||||
async def test_run_builds_apply_plan_and_audits_accepted_digest(self) -> None:
|
||||
job = memory.CompressionJob(digest_input=_digest_input())
|
||||
engine = FakeEngine(_accepted_text())
|
||||
audits = []
|
||||
|
||||
async def audit_hook(payload):
|
||||
audits.append(payload)
|
||||
|
||||
run = await session_digest_worker.run_session_digest_worker(
|
||||
job,
|
||||
engine,
|
||||
existing_case_digest="S1: 이전 회기\nS2: 오래된 요약",
|
||||
audit_hook=audit_hook,
|
||||
)
|
||||
|
||||
self.assertFalse(run.fallback_required)
|
||||
self.assertIsNotNone(run.apply_plan)
|
||||
assert run.apply_plan is not None
|
||||
self.assertEqual(run.apply_plan.compressed_by, "llm:openai/gpt-4.1-mini")
|
||||
self.assertEqual(run.apply_plan.token_count, 33)
|
||||
self.assertIn("S1: 이전 회기", run.apply_plan.case_digest or "")
|
||||
self.assertIn(_accepted_text(), run.apply_plan.case_digest or "")
|
||||
self.assertNotIn("오래된 요약", run.apply_plan.case_digest or "")
|
||||
self.assertEqual(audits[0]["session_id"], SESSION_ID)
|
||||
self.assertEqual(audits[0]["provider"], "openai")
|
||||
self.assertEqual(audits[0]["tokens_in"], 20)
|
||||
self.assertEqual(len(engine.requests), 1)
|
||||
|
||||
async def test_rejected_digest_does_not_create_apply_plan(self) -> None:
|
||||
job = memory.CompressionJob(digest_input=_digest_input(session_no=3))
|
||||
engine = FakeEngine("S2: 내담자 김서연은 잘못된 회기 prefix와 raw 이름을 포함했다.")
|
||||
|
||||
run = await session_digest_worker.run_session_digest_worker(
|
||||
job,
|
||||
engine,
|
||||
forbidden_substrings=("김서연",),
|
||||
)
|
||||
|
||||
self.assertTrue(run.fallback_required)
|
||||
self.assertIsNone(run.apply_plan)
|
||||
self.assertTrue(run.outcome.fallback_required)
|
||||
self.assertIn(run.outcome.quality.reason, {"wrong_session_prefix", "forbidden_substring"})
|
||||
|
||||
|
||||
class SessionDigestWorkerPersistenceTest(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_load_session_digest_job_rebuilds_masked_client_visible_input(self) -> None:
|
||||
class FakeConn:
|
||||
def __init__(self) -> None:
|
||||
self.fetchrow_query = ""
|
||||
|
||||
async def fetchrow(self, query: str, *args):
|
||||
self.fetchrow_query = query
|
||||
return {
|
||||
"session_id": SESSION_ID,
|
||||
"case_id": CASE_ID,
|
||||
"session_no": 4,
|
||||
"open_threads": ["다음 회기 주제"],
|
||||
"case_digest": "S3: 이전 회기",
|
||||
"learner_id": LEARNER_ID,
|
||||
}
|
||||
|
||||
async def fetch(self, query: str, *args):
|
||||
return [
|
||||
{
|
||||
"id": "00000000-0000-0000-0000-000000000201",
|
||||
"speaker": "client",
|
||||
"text": "raw 김서연",
|
||||
"text_masked": "저는 [NAME]입니다.",
|
||||
"visible_to": ["client", "evaluator"],
|
||||
},
|
||||
{
|
||||
"id": "00000000-0000-0000-0000-000000000202",
|
||||
"speaker": "client",
|
||||
"text": "평가자 전용 raw",
|
||||
"text_masked": "평가자 전용 masked",
|
||||
"visible_to": ["evaluator"],
|
||||
},
|
||||
{
|
||||
"id": "00000000-0000-0000-0000-000000000203",
|
||||
"speaker": "client",
|
||||
"text": "raw 김서연",
|
||||
"text_masked": "",
|
||||
"visible_to": ["client"],
|
||||
},
|
||||
{
|
||||
"id": "00000000-0000-0000-0000-000000000204",
|
||||
"speaker": "system",
|
||||
"text": "시스템 메모",
|
||||
"text_masked": "시스템 메모",
|
||||
"visible_to": ["client"],
|
||||
},
|
||||
]
|
||||
|
||||
conn = FakeConn()
|
||||
loaded = await session_digest_worker.load_session_digest_job(conn, SESSION_ID)
|
||||
|
||||
self.assertIsNotNone(loaded)
|
||||
assert loaded is not None
|
||||
self.assertIn("ss.compressed_by IS NULL", conn.fetchrow_query)
|
||||
digest_input = loaded.job.digest_input
|
||||
self.assertEqual(digest_input.session_no, 4)
|
||||
self.assertEqual(digest_input.open_threads, ("다음 회기 주제",))
|
||||
self.assertEqual(len(digest_input.masked_turns), 1)
|
||||
self.assertEqual(digest_input.masked_turns[0].text, "저는 [NAME]입니다.")
|
||||
self.assertNotIn("김서연", " ".join(turn.text for turn in digest_input.masked_turns))
|
||||
self.assertEqual(loaded.existing_case_digest, "S3: 이전 회기")
|
||||
self.assertEqual(loaded.learner_id, LEARNER_ID)
|
||||
|
||||
async def test_load_session_digest_job_skips_already_compressed_rows(self) -> None:
|
||||
class FakeConn:
|
||||
def __init__(self) -> None:
|
||||
self.fetchrow_query = ""
|
||||
self.fetch_called = False
|
||||
|
||||
async def fetchrow(self, query: str, *args):
|
||||
self.fetchrow_query = query
|
||||
return None
|
||||
|
||||
async def fetch(self, query: str, *args):
|
||||
self.fetch_called = True
|
||||
return []
|
||||
|
||||
conn = FakeConn()
|
||||
loaded = await session_digest_worker.load_session_digest_job(conn, SESSION_ID)
|
||||
|
||||
self.assertIsNone(loaded)
|
||||
self.assertIn("ss.compressed_by IS NULL", conn.fetchrow_query)
|
||||
self.assertFalse(conn.fetch_called)
|
||||
|
||||
async def test_run_session_digest_once_applies_only_accepted_plan(self) -> None:
|
||||
class FakeConn:
|
||||
def __init__(self) -> None:
|
||||
self.executed = []
|
||||
|
||||
async def fetchrow(self, query: str, *args):
|
||||
return {
|
||||
"session_id": SESSION_ID,
|
||||
"case_id": CASE_ID,
|
||||
"session_no": 5,
|
||||
"open_threads": ["가족 갈등"],
|
||||
"case_digest": "S4: 이전 회기\nS5: 오래된 요약",
|
||||
"learner_id": LEARNER_ID,
|
||||
}
|
||||
|
||||
async def fetch(self, query: str, *args):
|
||||
return [
|
||||
{
|
||||
"id": "00000000-0000-0000-0000-000000000201",
|
||||
"speaker": "client",
|
||||
"text": "raw 김서연",
|
||||
"text_masked": "저는 [NAME]이고 가족 갈등이 부담됩니다.",
|
||||
"visible_to": ["client"],
|
||||
},
|
||||
{
|
||||
"id": "00000000-0000-0000-0000-000000000202",
|
||||
"speaker": "counselor",
|
||||
"text": "다음 회기에 이어가겠습니다.",
|
||||
"text_masked": "다음 회기에 이어가겠습니다.",
|
||||
"visible_to": ["client"],
|
||||
},
|
||||
]
|
||||
|
||||
async def execute(self, query: str, *args):
|
||||
self.executed.append((query, args))
|
||||
return "UPDATE 1"
|
||||
|
||||
conn = FakeConn()
|
||||
result = await session_digest_worker.run_session_digest_once(
|
||||
conn,
|
||||
session_id=SESSION_ID,
|
||||
engine=FakeEngine(_accepted_text(session_no=5)),
|
||||
)
|
||||
|
||||
self.assertTrue(result.found)
|
||||
self.assertTrue(result.applied)
|
||||
self.assertEqual(len(conn.executed), 2)
|
||||
summary_query, summary_args = conn.executed[0]
|
||||
self.assertIn("UPDATE app.session_summary", summary_query)
|
||||
self.assertEqual(summary_args[0], SESSION_ID)
|
||||
self.assertIn("[NAME]", summary_args[1])
|
||||
self.assertEqual(summary_args[2], ["가족 갈등"])
|
||||
self.assertEqual(summary_args[3], "llm:openai/gpt-4.1-mini")
|
||||
self.assertEqual(summary_args[4], 33)
|
||||
case_query, case_args = conn.executed[1]
|
||||
self.assertIn("UPDATE app.case_profile", case_query)
|
||||
self.assertEqual(case_args[0], CASE_ID)
|
||||
self.assertEqual(case_args[1], LEARNER_ID)
|
||||
self.assertIn("S4: 이전 회기", case_args[2])
|
||||
self.assertIn(_accepted_text(session_no=5), case_args[2])
|
||||
self.assertNotIn("오래된 요약", case_args[2])
|
||||
|
||||
async def test_apply_session_digest_plan_uses_cas_before_case_update(self) -> None:
|
||||
class FakeConn:
|
||||
def __init__(self) -> None:
|
||||
self.executed = []
|
||||
|
||||
async def execute(self, query: str, *args):
|
||||
self.executed.append((query, args))
|
||||
return "UPDATE 0"
|
||||
|
||||
plan = session_digest_worker.SessionDigestApplyPlan(
|
||||
result=memory.SessionDigestResult(
|
||||
session_id=SESSION_ID,
|
||||
case_id=CASE_ID,
|
||||
session_no=5,
|
||||
digest=_accepted_text(session_no=5),
|
||||
open_threads=("가족 갈등",),
|
||||
source="llm",
|
||||
),
|
||||
case_digest="S5: accepted",
|
||||
compressed_by="llm:openai/gpt-4.1-mini",
|
||||
token_count=33,
|
||||
)
|
||||
conn = FakeConn()
|
||||
|
||||
applied = await session_digest_worker.apply_session_digest_plan(
|
||||
conn,
|
||||
plan,
|
||||
learner_id=LEARNER_ID,
|
||||
)
|
||||
|
||||
self.assertFalse(applied)
|
||||
self.assertEqual(len(conn.executed), 1)
|
||||
summary_query, summary_args = conn.executed[0]
|
||||
self.assertIn("AND compressed_by IS NULL", summary_query)
|
||||
self.assertEqual(summary_args[0], SESSION_ID)
|
||||
|
|
@ -1,6 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
from contextlib import asynccontextmanager
|
||||
from unittest.mock import patch
|
||||
|
||||
from . import session_persistence
|
||||
|
|
@ -129,6 +130,85 @@ class SessionMemoryPureTest(unittest.TestCase):
|
|||
self.assertIn("[NAME]", result.digest)
|
||||
self.assertNotIn("김서연", result.digest)
|
||||
|
||||
def test_llm_digest_worker_outcome_accepts_masked_contract_result(self) -> None:
|
||||
digest_input = memory.build_session_digest_input(
|
||||
session_id="00000000-0000-0000-0000-00000000feed",
|
||||
case_id="00000000-0000-0000-0000-00000000ca5e",
|
||||
session_no=7,
|
||||
masked_turns=[
|
||||
{"speaker": "client", "text": "저는 [NAME]이고 가족 이야기가 어렵습니다.", "visible_to": ["client"]},
|
||||
{"speaker": "counselor", "text": "그 주제를 다음 회기에 이어가겠습니다.", "visible_to": ["client"]},
|
||||
],
|
||||
open_threads=["가족 갈등을 다음 회기에 이어가기"],
|
||||
)
|
||||
|
||||
outcome = memory.build_llm_digest_worker_outcome(
|
||||
digest_input,
|
||||
(
|
||||
"내담자는 [NAME]으로 지칭되며 가족 갈등을 조심스럽게 설명했다. "
|
||||
"상담자는 감정 확인과 다음 회기에서 이어갈 주제를 함께 정리했다."
|
||||
),
|
||||
forbidden_substrings=("김서연",),
|
||||
)
|
||||
|
||||
self.assertFalse(outcome.fallback_required)
|
||||
self.assertTrue(outcome.quality.accepted)
|
||||
self.assertEqual(outcome.quality.reason, "ok")
|
||||
self.assertIsNotNone(outcome.result)
|
||||
assert outcome.result is not None
|
||||
self.assertEqual(outcome.result.source, "llm")
|
||||
self.assertEqual(outcome.result.session_no, 7)
|
||||
self.assertEqual(outcome.result.open_threads, ("가족 갈등을 다음 회기에 이어가기",))
|
||||
self.assertTrue(outcome.result.digest.startswith("S7:"))
|
||||
self.assertIn("[NAME]", outcome.result.digest)
|
||||
self.assertNotIn("김서연", outcome.result.digest)
|
||||
|
||||
def test_llm_digest_quality_rejects_raw_forbidden_substring(self) -> None:
|
||||
digest_input = memory.build_session_digest_input(
|
||||
session_id="00000000-0000-0000-0000-00000000feed",
|
||||
case_id="00000000-0000-0000-0000-00000000ca5e",
|
||||
session_no=8,
|
||||
masked_turns=[
|
||||
{"speaker": "client", "text": "저는 [NAME]입니다.", "visible_to": ["client"]},
|
||||
],
|
||||
)
|
||||
|
||||
outcome = memory.build_llm_digest_worker_outcome(
|
||||
digest_input,
|
||||
"S8: 내담자 김서연은 가족 갈등을 설명했고 상담자는 다음 회기에서 이어갈 주제를 정리했다.",
|
||||
forbidden_substrings=("김서연",),
|
||||
)
|
||||
|
||||
self.assertTrue(outcome.fallback_required)
|
||||
self.assertIsNone(outcome.result)
|
||||
self.assertFalse(outcome.quality.accepted)
|
||||
self.assertEqual(outcome.quality.reason, "forbidden_substring")
|
||||
self.assertEqual(outcome.quality.details, ("김서연",))
|
||||
|
||||
def test_llm_digest_quality_rejects_internal_markers_and_wrong_session(self) -> None:
|
||||
digest_input = memory.build_session_digest_input(
|
||||
session_id="00000000-0000-0000-0000-00000000feed",
|
||||
case_id="00000000-0000-0000-0000-00000000ca5e",
|
||||
session_no=9,
|
||||
masked_turns=[
|
||||
{"speaker": "client", "text": "저는 [NAME]입니다.", "visible_to": ["client"]},
|
||||
],
|
||||
)
|
||||
|
||||
marker_outcome = memory.build_llm_digest_worker_outcome(
|
||||
digest_input,
|
||||
"S9: rapport_credit 수치와 evaluation payload를 근거로 요약을 작성했다. 다음 회기 주제를 유지한다.",
|
||||
)
|
||||
wrong_session_outcome = memory.build_llm_digest_worker_outcome(
|
||||
digest_input,
|
||||
"S8: 내담자는 [NAME]으로 지칭되며 가족 갈등을 설명했다. 상담자는 다음 회기 주제를 정리했다.",
|
||||
)
|
||||
|
||||
self.assertTrue(marker_outcome.fallback_required)
|
||||
self.assertEqual(marker_outcome.quality.reason, "internal_marker")
|
||||
self.assertTrue(wrong_session_outcome.fallback_required)
|
||||
self.assertEqual(wrong_session_outcome.quality.reason, "wrong_session_prefix")
|
||||
|
||||
def test_extract_pinned_fact_candidates_is_conservative_and_masked(self) -> None:
|
||||
facts = memory.extract_pinned_fact_candidates(
|
||||
[
|
||||
|
|
@ -304,6 +384,190 @@ class SessionMemoryPersistenceTest(unittest.IsolatedAsyncioTestCase):
|
|||
self.assertTrue(sess.ended)
|
||||
self.assertEqual(len(scheduled), 1)
|
||||
|
||||
async def test_end_persisted_session_schedules_digest_worker_only_when_enabled(self) -> None:
|
||||
scheduled: list[str] = []
|
||||
|
||||
def fake_create_task(coro):
|
||||
scheduled.append(coro.cr_code.co_name)
|
||||
coro.close()
|
||||
return None
|
||||
|
||||
sess = InProcSession(
|
||||
session_id="00000000-0000-0000-0000-00000000feed",
|
||||
case_id="00000000-0000-0000-0000-00000000ca5e",
|
||||
learner_id="00000000-0000-0000-0000-000000000101",
|
||||
persona_code=P1.code,
|
||||
theory_mode="humanistic",
|
||||
persona=P1,
|
||||
state=state_machine.SessionState(),
|
||||
session_no=1,
|
||||
)
|
||||
digest_input = memory.build_session_digest_input(
|
||||
session_id=sess.session_id,
|
||||
case_id=sess.case_id,
|
||||
session_no=sess.session_no,
|
||||
masked_turns=[
|
||||
{
|
||||
"speaker": "client",
|
||||
"text_masked": "다음 회기에 가족 이야기를 이어가고 싶어요.",
|
||||
"visible_to": ["client"],
|
||||
}
|
||||
],
|
||||
open_threads=["가족 이야기"],
|
||||
)
|
||||
carry = memory.CarryOver(
|
||||
end_state={},
|
||||
rapport_delta=0.0,
|
||||
compression_job=memory.CompressionJob(digest_input=digest_input),
|
||||
)
|
||||
|
||||
with patch.object(session_persistence, "end_session", return_value=True), patch.object(
|
||||
sessions.settings,
|
||||
"session_digest_worker_enabled",
|
||||
True,
|
||||
), patch.object(
|
||||
sessions.asyncio,
|
||||
"create_task",
|
||||
fake_create_task,
|
||||
):
|
||||
await sessions._end_persisted_session(sess, carry)
|
||||
|
||||
self.assertEqual(
|
||||
scheduled,
|
||||
["_run_session_digest_worker_for_session", "_write_episodic_embeddings"],
|
||||
)
|
||||
|
||||
async def test_end_persisted_session_keeps_digest_worker_default_off(self) -> None:
|
||||
scheduled: list[str] = []
|
||||
|
||||
def fake_create_task(coro):
|
||||
scheduled.append(coro.cr_code.co_name)
|
||||
coro.close()
|
||||
return None
|
||||
|
||||
sess = InProcSession(
|
||||
session_id="00000000-0000-0000-0000-00000000feed",
|
||||
case_id="00000000-0000-0000-0000-00000000ca5e",
|
||||
learner_id="00000000-0000-0000-0000-000000000101",
|
||||
persona_code=P1.code,
|
||||
theory_mode="humanistic",
|
||||
persona=P1,
|
||||
state=state_machine.SessionState(),
|
||||
session_no=1,
|
||||
)
|
||||
digest_input = memory.build_session_digest_input(
|
||||
session_id=sess.session_id,
|
||||
case_id=sess.case_id,
|
||||
session_no=sess.session_no,
|
||||
masked_turns=[
|
||||
{
|
||||
"speaker": "client",
|
||||
"text_masked": "다음 회기에 가족 이야기를 이어가고 싶어요.",
|
||||
"visible_to": ["client"],
|
||||
}
|
||||
],
|
||||
open_threads=["가족 이야기"],
|
||||
)
|
||||
carry = memory.CarryOver(
|
||||
end_state={},
|
||||
rapport_delta=0.0,
|
||||
compression_job=memory.CompressionJob(digest_input=digest_input),
|
||||
)
|
||||
|
||||
with patch.object(session_persistence, "end_session", return_value=True), patch.object(
|
||||
sessions.settings,
|
||||
"session_digest_worker_enabled",
|
||||
False,
|
||||
), patch.object(
|
||||
sessions.asyncio,
|
||||
"create_task",
|
||||
fake_create_task,
|
||||
):
|
||||
await sessions._end_persisted_session(sess, carry)
|
||||
|
||||
self.assertEqual(scheduled, ["_write_episodic_embeddings"])
|
||||
|
||||
async def test_session_digest_worker_releases_db_connection_during_engine_call(self) -> None:
|
||||
order: list[str] = []
|
||||
digest_input = memory.build_session_digest_input(
|
||||
session_id="00000000-0000-0000-0000-00000000feed",
|
||||
case_id="00000000-0000-0000-0000-00000000ca5e",
|
||||
session_no=1,
|
||||
masked_turns=[
|
||||
{
|
||||
"speaker": "client",
|
||||
"text_masked": "다음 회기에 가족 이야기를 이어가고 싶어요.",
|
||||
"visible_to": ["client"],
|
||||
}
|
||||
],
|
||||
open_threads=["가족 이야기"],
|
||||
)
|
||||
loaded = sessions.session_digest_worker.LoadedSessionDigestJob(
|
||||
job=memory.CompressionJob(digest_input=digest_input),
|
||||
existing_case_digest="S0: 이전",
|
||||
learner_id="00000000-0000-0000-0000-000000000101",
|
||||
)
|
||||
|
||||
class Worker:
|
||||
apply_plan = object()
|
||||
|
||||
acquire_count = 0
|
||||
|
||||
@asynccontextmanager
|
||||
async def fake_acquire(**kwargs):
|
||||
nonlocal acquire_count
|
||||
acquire_count += 1
|
||||
label = "load" if acquire_count == 1 else "apply"
|
||||
order.append(f"enter-{label}")
|
||||
try:
|
||||
yield object()
|
||||
finally:
|
||||
order.append(f"exit-{label}")
|
||||
|
||||
async def fake_load(conn, session_id: str):
|
||||
order.append("load")
|
||||
return loaded
|
||||
|
||||
async def fake_run(job, engine, *, existing_case_digest=None, model=None, audit_hook=None):
|
||||
order.append("engine")
|
||||
self.assertIs(engine, sessions.engine_client)
|
||||
self.assertEqual(existing_case_digest, "S0: 이전")
|
||||
self.assertIs(audit_hook, session_persistence.record_llm_call_audit)
|
||||
return Worker()
|
||||
|
||||
async def fake_apply(conn, apply_plan, *, learner_id=None):
|
||||
order.append("apply")
|
||||
self.assertEqual(learner_id, loaded.learner_id)
|
||||
return True
|
||||
|
||||
with patch.object(sessions.db, "get_pool", return_value=object()), patch.object(
|
||||
sessions.db,
|
||||
"acquire",
|
||||
fake_acquire,
|
||||
), patch.object(
|
||||
sessions.session_digest_worker,
|
||||
"load_session_digest_job",
|
||||
fake_load,
|
||||
), patch.object(
|
||||
sessions.session_digest_worker,
|
||||
"run_session_digest_worker",
|
||||
fake_run,
|
||||
), patch.object(
|
||||
sessions.session_digest_worker,
|
||||
"apply_session_digest_plan",
|
||||
fake_apply,
|
||||
), patch.object(
|
||||
sessions.settings,
|
||||
"session_digest_worker_model",
|
||||
"",
|
||||
):
|
||||
await sessions._run_session_digest_worker_for_session(loaded.job.session_id)
|
||||
|
||||
self.assertEqual(
|
||||
order,
|
||||
["enter-load", "load", "exit-load", "engine", "enter-apply", "apply", "exit-apply"],
|
||||
)
|
||||
|
||||
async def test_seed_recall_loads_case_digest_and_client_visible_pinned_facts(self) -> None:
|
||||
test_case = self
|
||||
|
||||
|
|
|
|||
|
|
@ -168,9 +168,9 @@ class SessionTurnPersistenceTest(unittest.IsolatedAsyncioTestCase):
|
|||
|
||||
async def successful_turn(ctx, engine, **kwargs):
|
||||
assert ctx.state_after is not None
|
||||
self.assertIn("[NAME]", ctx.recall_summary or "")
|
||||
self.assertNotIn("김서연", ctx.recall_summary or "")
|
||||
self.assertEqual(ctx.pinned_facts, ["[NAME]와 주 1회 상담 약속"])
|
||||
self.assertIn("[NAME]", ctx.memory.recall_summary or "")
|
||||
self.assertNotIn("김서연", ctx.memory.recall_summary or "")
|
||||
self.assertEqual(ctx.memory.pinned_facts, ["[NAME]와 주 1회 상담 약속"])
|
||||
return orchestrator.TurnResult(
|
||||
turn_seq=ctx.state_after.turn_seq,
|
||||
stage=ctx.state_after.stage.value,
|
||||
|
|
@ -750,7 +750,7 @@ class SessionTurnPersistenceTest(unittest.IsolatedAsyncioTestCase):
|
|||
card=sess.persona,
|
||||
state=sess.state,
|
||||
learner_text="게이트웨이 스트림 테스트",
|
||||
recent_turns=[],
|
||||
memory=orchestrator.TurnMemory(recent_turns=[]),
|
||||
)
|
||||
|
||||
events = [
|
||||
|
|
@ -789,7 +789,7 @@ class SessionTurnPersistenceTest(unittest.IsolatedAsyncioTestCase):
|
|||
card=sess.persona,
|
||||
state=sess.state,
|
||||
learner_text="게이트웨이 오류 테스트",
|
||||
recent_turns=[],
|
||||
memory=orchestrator.TurnMemory(recent_turns=[]),
|
||||
)
|
||||
|
||||
events = [
|
||||
|
|
|
|||
|
|
@ -190,6 +190,9 @@ class VoiceWebSocketContractTest(unittest.IsolatedAsyncioTestCase):
|
|||
self.assertEqual(kwargs["voice_preset"], VOICE_PRESET)
|
||||
self.assertEqual(kwargs["audio"], b"chunk-onechunk-two")
|
||||
self.assertEqual(kwargs["fmt"], "webm")
|
||||
self.assertIsNone(kwargs["sample_rate"])
|
||||
self.assertIsNone(kwargs["channels"])
|
||||
self.assertIsNone(kwargs["sample_width"])
|
||||
self.assertEqual(kwargs["audio_started_at"], 10.0)
|
||||
self.assertEqual(kwargs["audio_ended_at"], 12.0)
|
||||
self.assertEqual(kwargs["silence_ms"], 450)
|
||||
|
|
@ -212,6 +215,62 @@ class VoiceWebSocketContractTest(unittest.IsolatedAsyncioTestCase):
|
|||
],
|
||||
)
|
||||
|
||||
def test_pcm_upload_is_wrapped_as_wav_before_stt(self) -> None:
|
||||
audio, fmt = voice_routes._normalize_audio_upload(
|
||||
b"\x00\x00\xff\x7f",
|
||||
fmt="pcm",
|
||||
sample_rate=16000,
|
||||
channels=1,
|
||||
sample_width=2,
|
||||
)
|
||||
|
||||
self.assertEqual(fmt, "wav")
|
||||
self.assertTrue(audio.startswith(b"RIFF"))
|
||||
self.assertEqual(audio[8:12], b"WAVE")
|
||||
self.assertEqual(audio[12:16], b"fmt ")
|
||||
self.assertEqual(int.from_bytes(audio[24:28], "little"), 16000)
|
||||
self.assertEqual(int.from_bytes(audio[22:24], "little"), 1)
|
||||
self.assertEqual(audio[36:40], b"data")
|
||||
self.assertEqual(int.from_bytes(audio[40:44], "little"), 4)
|
||||
self.assertEqual(audio[44:], b"\x00\x00\xff\x7f")
|
||||
|
||||
async def test_pcm_control_metadata_flows_to_handle_utterance(self) -> None:
|
||||
websocket = FakeWebSocket(
|
||||
[
|
||||
_control({"type": "audio_start", "format": "pcm", "sample_rate": 16000, "channels": 1, "sample_width": 2}),
|
||||
_binary(b"\x00\x00\xff\x7f"),
|
||||
_control({"type": "audio_end", "format": "pcm"}),
|
||||
_control({"type": "close"}),
|
||||
]
|
||||
)
|
||||
handle_utterance = AsyncMock()
|
||||
|
||||
with patch.object(
|
||||
voice_routes,
|
||||
"_principal_from_websocket",
|
||||
AsyncMock(return_value=_principal()),
|
||||
), patch.object(
|
||||
voice_routes,
|
||||
"_bind_session",
|
||||
AsyncMock(return_value=self._bind_result()),
|
||||
), patch.object(
|
||||
voice_routes.voice_service,
|
||||
"is_available",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
voice_routes,
|
||||
"_handle_utterance",
|
||||
handle_utterance,
|
||||
):
|
||||
await voice_routes.voice_ws(websocket) # type: ignore[arg-type]
|
||||
|
||||
handle_utterance.assert_awaited_once()
|
||||
kwargs = handle_utterance.await_args.kwargs
|
||||
self.assertEqual(kwargs["fmt"], "pcm")
|
||||
self.assertEqual(kwargs["sample_rate"], 16000)
|
||||
self.assertEqual(kwargs["channels"], 1)
|
||||
self.assertEqual(kwargs["sample_width"], 2)
|
||||
|
||||
async def test_text_turn_strips_text_runs_turn_and_ping_close_still_work(self) -> None:
|
||||
websocket = FakeWebSocket(
|
||||
[
|
||||
|
|
@ -259,6 +318,152 @@ class VoiceWebSocketContractTest(unittest.IsolatedAsyncioTestCase):
|
|||
self.assertEqual(kwargs["voice_preset"], VOICE_PRESET)
|
||||
self.assertEqual(kwargs["learner_text"], "I need help practicing.")
|
||||
|
||||
async def test_stt_result_waits_for_final_transcript_before_running_turn(self) -> None:
|
||||
websocket = FakeWebSocket(
|
||||
[
|
||||
_control(
|
||||
{
|
||||
"type": "stt_result",
|
||||
"text": "I am still talking",
|
||||
"final": False,
|
||||
"silence_ms": 2500,
|
||||
}
|
||||
),
|
||||
_control({"type": "close"}),
|
||||
]
|
||||
)
|
||||
run_turn = AsyncMock()
|
||||
handle_utterance = AsyncMock()
|
||||
|
||||
with patch.object(
|
||||
voice_routes,
|
||||
"_principal_from_websocket",
|
||||
AsyncMock(return_value=_principal()),
|
||||
), patch.object(
|
||||
voice_routes,
|
||||
"_bind_session",
|
||||
AsyncMock(return_value=self._bind_result()),
|
||||
), patch.object(
|
||||
voice_routes.voice_service,
|
||||
"is_available",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
voice_routes,
|
||||
"_handle_utterance",
|
||||
handle_utterance,
|
||||
), patch.object(
|
||||
voice_routes,
|
||||
"_run_turn_and_speak",
|
||||
run_turn,
|
||||
):
|
||||
await voice_routes.voice_ws(websocket) # type: ignore[arg-type]
|
||||
|
||||
handle_utterance.assert_not_awaited()
|
||||
run_turn.assert_not_awaited()
|
||||
self.assertIn(
|
||||
{
|
||||
"type": "eot",
|
||||
"ready": False,
|
||||
"reason": "final_transcript_pending",
|
||||
"silence_ms": 2500,
|
||||
"threshold_ms": 1200,
|
||||
},
|
||||
websocket.sent_json,
|
||||
)
|
||||
self.assertEqual(websocket.sent_json[-1], {"type": "state", "state": "listening"})
|
||||
|
||||
async def test_stt_result_runs_turn_only_after_eot_ready(self) -> None:
|
||||
websocket = FakeWebSocket(
|
||||
[
|
||||
_control({"type": "audio_start", "format": "pcm"}),
|
||||
_control(
|
||||
{
|
||||
"type": "stt_result",
|
||||
"text": " I am done now. ",
|
||||
"final": True,
|
||||
"silence_ms": 1300,
|
||||
"barge_in": False,
|
||||
"provider_events": [
|
||||
{
|
||||
"type": "speech_final",
|
||||
"confidence": 0.91,
|
||||
"text": "raw transcript must not persist",
|
||||
}
|
||||
],
|
||||
}
|
||||
),
|
||||
_control({"type": "close"}),
|
||||
]
|
||||
)
|
||||
run_turn = AsyncMock()
|
||||
handle_utterance = AsyncMock()
|
||||
|
||||
with patch.object(
|
||||
voice_routes,
|
||||
"_principal_from_websocket",
|
||||
AsyncMock(return_value=_principal()),
|
||||
), patch.object(
|
||||
voice_routes,
|
||||
"_bind_session",
|
||||
AsyncMock(return_value=self._bind_result()),
|
||||
), patch.object(
|
||||
voice_routes.voice_service,
|
||||
"is_available",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
voice_routes,
|
||||
"_handle_utterance",
|
||||
handle_utterance,
|
||||
), patch.object(
|
||||
voice_routes,
|
||||
"_run_turn_and_speak",
|
||||
run_turn,
|
||||
), patch.object(
|
||||
voice_routes.time,
|
||||
"monotonic",
|
||||
side_effect=[10.0, 12.0],
|
||||
):
|
||||
await voice_routes.voice_ws(websocket) # type: ignore[arg-type]
|
||||
|
||||
handle_utterance.assert_not_awaited()
|
||||
run_turn.assert_awaited_once()
|
||||
self.assertIn(
|
||||
{
|
||||
"type": "eot",
|
||||
"ready": True,
|
||||
"reason": "ready",
|
||||
"silence_ms": 1300,
|
||||
"threshold_ms": 1200,
|
||||
},
|
||||
websocket.sent_json,
|
||||
)
|
||||
self.assertIn(
|
||||
{
|
||||
"type": "transcript",
|
||||
"text": "I am done now.",
|
||||
"final": True,
|
||||
"speaker": "counselor",
|
||||
},
|
||||
websocket.sent_json,
|
||||
)
|
||||
self.assertIn({"type": "state", "state": "thinking"}, websocket.sent_json)
|
||||
kwargs = run_turn.await_args.kwargs
|
||||
self.assertEqual(kwargs["learner_text"], "I am done now.")
|
||||
self.assertEqual(kwargs["duration_s"], 2.0)
|
||||
self.assertEqual(kwargs["silence_ms"], 1300)
|
||||
self.assertIs(kwargs["barge_in"], False)
|
||||
self.assertEqual(
|
||||
kwargs["provider_events"],
|
||||
[
|
||||
{
|
||||
"type": "speech_final",
|
||||
"confidence": 0.91,
|
||||
"event_type": "speech_final",
|
||||
"category": "speech_activity",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
async def test_oversize_binary_audio_reports_error_and_drops_utterance(self) -> None:
|
||||
websocket = FakeWebSocket(
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue