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

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

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