Anti-Hallucination / agent /background.py
MHRDYN7's picture
Deploy Anti-Hallucination Chat to personal Space
f3793b5 verified
Raw
History Blame Contribute Delete
5.02 kB
"""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]