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