런타임 계약과 학습자 흐름 보강

This commit is contained in:
Yun Chan 2026-06-29 08:12:14 +09:00
parent f456b8997a
commit 206018b088
56 changed files with 4306 additions and 1008 deletions

View file

@ -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(

View file

@ -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")

View file

@ -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)

View file

@ -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()

View 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 {}

View file

@ -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:

View file

@ -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(

View file

@ -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)}; "

View file

@ -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),
)
)

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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",

View file

@ -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",

View file

@ -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",

View 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",
]

View file

@ -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

View file

@ -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,

View file

@ -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",

View file

@ -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, "탐색")],

View file

@ -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)

View 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()

View file

@ -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)

View file

@ -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,

View 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)

View file

@ -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

View file

@ -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 = [

View file

@ -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(
[