"""Background per-message audit workers (threaded, serialized).""" from __future__ import annotations import queue import threading from typing import Literal from agent.nodes.audit import audit_one_message from agent.state import ChatTurn, MessageAudit from config import AUDIT_USE_WEB_SEARCH _lock = threading.Lock() _STORE: dict[str, dict[str, MessageAudit]] = {} _RUNNING: dict[str, set[str]] = {} _QUEUES: dict[str, queue.Queue] = {} _WORKERS: dict[str, threading.Thread] = {} # One Gemma audit at a time — avoids rate-limit hangs from parallel calls. _AUDIT_SEMAPHORE = threading.Semaphore(1) def _placeholder( turn: ChatTurn, status: Literal["pending", "running", "error"] = "running", phase: Literal["queued", "extracting", "searching", "judging", "done", "error"] = "queued", ) -> MessageAudit: return { "message_id": turn["id"], "role": turn["role"], "content": turn["content"], "rewritten": "", "paragraphs": [], "claims": [], "status": status, "phase": phase, "evidence_text": "", } def _publish(session_id: str, mid: str, audit: MessageAudit) -> None: with _lock: bucket = _STORE.setdefault(session_id, {}) bucket[mid] = audit def _get_audit(session_id: str, message_id: str) -> MessageAudit | None: with _lock: return _STORE.get(session_id, {}).get(message_id) def _prior_audits(session_id: str, turns: list[ChatTurn], before_id: str) -> list[MessageAudit]: """Only use completed audits as context — never block waiting.""" prior: list[MessageAudit] = [] for turn in turns: if turn["id"] == before_id: break audit = _get_audit(session_id, turn["id"]) if audit and audit.get("status") == "done": prior.append(audit) return prior def _run_audit(session_id: str, turn: ChatTurn, turns: list[ChatTurn]) -> None: mid = turn["id"] def on_partial(partial: MessageAudit) -> None: _publish(session_id, mid, partial) try: with _AUDIT_SEMAPHORE: prior = _prior_audits(session_id, turns, mid) result = audit_one_message( turn, prior, use_web_search=AUDIT_USE_WEB_SEARCH, on_partial=on_partial, ) _publish(session_id, mid, result) except Exception as exc: # noqa: BLE001 err = _placeholder(turn, "error", "error") err["rewritten"] = str(exc) _publish(session_id, mid, err) finally: with _lock: _RUNNING.get(session_id, set()).discard(mid) def _session_worker(session_id: str) -> None: q = _QUEUES[session_id] while True: item = q.get() try: if item is None: return turn, turns = item _run_audit(session_id, turn, turns) finally: q.task_done() def _ensure_worker(session_id: str) -> None: with _lock: if session_id in _WORKERS and _WORKERS[session_id].is_alive(): return _QUEUES.setdefault(session_id, queue.Queue()) thread = threading.Thread( target=_session_worker, args=(session_id,), daemon=True, name=f"audit-worker-{session_id[:8]}", ) _WORKERS[session_id] = thread thread.start() def schedule_audit(session_id: str, turn: ChatTurn, turns: list[ChatTurn]) -> None: """Queue one message for background audit (non-blocking).""" mid = turn["id"] with _lock: bucket = _STORE.setdefault(session_id, {}) existing = bucket.get(mid) if ( existing and existing.get("status") == "done" and existing.get("content") == turn["content"] ): return if existing and existing.get("status") == "running": return bucket[mid] = _placeholder(turn, "running", "queued") _RUNNING.setdefault(session_id, set()).add(mid) _ensure_worker(session_id) _QUEUES[session_id].put((turn, turns)) def schedule_missing_audits(session_id: str, turns: list[ChatTurn]) -> None: """Safety net: queue turns that somehow never started an audit.""" for turn in turns: audit = _get_audit(session_id, turn["id"]) if not audit or audit.get("content") != turn["content"]: schedule_audit(session_id, turn, turns) def sync_audits(session_id: str, target: dict[str, MessageAudit]) -> bool: with _lock: bucket = dict(_STORE.get(session_id, {})) running = bool(_RUNNING.get(session_id)) any_active = running for mid, audit in bucket.items(): target[mid] = audit if audit.get("status") in ("pending", "running"): any_active = True return any_active def audits_in_order( turns: list[ChatTurn], by_id: dict[str, MessageAudit] ) -> list[MessageAudit]: return [by_id[t["id"]] for t in turns if t["id"] in by_id]