366 lines
11 KiB
Python
366 lines
11 KiB
Python
"""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",
|
|
]
|