Spaces:
Running
Running
| """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] | |