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

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

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