Livrare LOT 1 - Didi

This commit is contained in:
Dezvoltari Evotech 2026-06-25 14:13:25 -07:00
commit 5380c3fc63
990 changed files with 133308 additions and 0 deletions

View file

@ -0,0 +1,532 @@
"""Admin operations for Brain — used by AI platform dashboard.
Provides:
- paginated atom + verification cache browsers with filters
- manual expire (atom) / hard delete (verification cache)
- taxonomy snapshot + reload trigger
Read paths return lightweight rows (no result_raw, no full evidence_urls
arrays) to keep DataGrid responses small. Detail endpoints expose the full
JSON payload.
"""
from __future__ import annotations
import json
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any, Literal
from brain_api.db import db
from shared.logging import get_logger
log = get_logger(__name__)
# ============================================================================
# Atom admin
# ============================================================================
@dataclass(slots=True)
class AtomListRow:
atom_id: int
content_hash: str
content_preview: str | None
component: str
tier: str
cache_tier: str
prompt_hash: str
framework_version: str | None
model_used: str | None
llm_confidence: float | None
human_validated: bool
hit_count: int
last_hit_at: datetime | None
created_at: datetime
updated_at: datetime
expires_at: datetime | None
@dataclass(slots=True)
class AtomDetailRow(AtomListRow):
result_processed: dict
result_raw: dict | None
human_corrections: dict | None
validator_user_id: str | None
validated_at: datetime | None
@dataclass(slots=True)
class AtomListPage:
items: list[AtomListRow]
total: int
page: int
page_size: int
_FRESHNESS_OPTIONS = ("fresh", "expiring", "expired", "all")
def _atom_list_filters(
component: str | None,
tier: str | None,
freshness: str | None,
q: str | None,
) -> tuple[str, list]:
"""Build WHERE clause + params (1-indexed)."""
where = ["1=1"]
params: list[Any] = []
idx = 1
if component and component != "all":
where.append(f"component = ${idx}")
params.append(component)
idx += 1
if tier and tier != "all":
# 'tier' here = cache_tier (gold/silver/bronze) — that's what users want to filter by
where.append(f"cache_tier = ${idx}")
params.append(tier)
idx += 1
# Freshness on expires_at:
# fresh = not yet expired (expires_at IS NULL or > now())
# expiring = expires within 7 days
# expired = expires_at <= now()
if freshness == "fresh":
where.append("(expires_at IS NULL OR expires_at > now())")
elif freshness == "expiring":
where.append("expires_at IS NOT NULL AND expires_at > now() AND expires_at <= now() + interval '7 days'")
elif freshness == "expired":
where.append("expires_at IS NOT NULL AND expires_at <= now()")
# else "all" or None → no filter
if q:
where.append(f"(content_preview ILIKE ${idx} OR content_hash ILIKE ${idx})")
params.append(f"%{q}%")
idx += 1
return " AND ".join(where), params
async def list_atoms(
*,
component: str | None = None,
tier: str | None = None,
freshness: str | None = None,
q: str | None = None,
page: int = 1,
page_size: int = 25,
) -> AtomListPage:
page = max(1, page)
page_size = min(max(1, page_size), 100)
offset = (page - 1) * page_size
where_sql, params = _atom_list_filters(component, tier, freshness, q)
count_sql = f"SELECT COUNT(*) AS n FROM brain_analysis_atom WHERE {where_sql}"
list_sql = f"""
SELECT atom_id, content_hash, content_preview, component, tier, cache_tier,
prompt_hash, framework_version, model_used, llm_confidence,
human_validated, hit_count, last_hit_at,
created_at, updated_at, expires_at
FROM brain_analysis_atom
WHERE {where_sql}
ORDER BY (cache_tier = 'gold') DESC, updated_at DESC
LIMIT ${len(params)+1} OFFSET ${len(params)+2}
"""
async with db.pool.acquire() as conn:
count_row = await conn.fetchrow(count_sql, *params)
rows = await conn.fetch(list_sql, *params, page_size, offset)
items = [
AtomListRow(
atom_id=r["atom_id"],
content_hash=r["content_hash"],
content_preview=r["content_preview"],
component=r["component"],
tier=r["tier"],
cache_tier=r["cache_tier"],
prompt_hash=r["prompt_hash"],
framework_version=r["framework_version"],
model_used=r["model_used"],
llm_confidence=float(r["llm_confidence"]) if r["llm_confidence"] is not None else None,
human_validated=r["human_validated"],
hit_count=r["hit_count"] or 0,
last_hit_at=r["last_hit_at"],
created_at=r["created_at"],
updated_at=r["updated_at"],
expires_at=r["expires_at"],
)
for r in rows
]
return AtomListPage(
items=items,
total=count_row["n"] or 0,
page=page,
page_size=page_size,
)
async def get_atom(atom_id: int) -> AtomDetailRow | None:
sql = """
SELECT atom_id, content_hash, content_preview, component, tier, cache_tier,
prompt_hash, framework_version, model_used, llm_confidence,
human_validated, hit_count, last_hit_at,
created_at, updated_at, expires_at,
result_processed, result_raw, human_corrections,
validator_user_id, validated_at
FROM brain_analysis_atom
WHERE atom_id = $1
"""
async with db.pool.acquire() as conn:
r = await conn.fetchrow(sql, atom_id)
if not r:
return None
rp = r["result_processed"]
rr = r["result_raw"]
hc = r["human_corrections"]
if isinstance(rp, str):
rp = json.loads(rp)
if isinstance(rr, str):
rr = json.loads(rr)
if isinstance(hc, str):
hc = json.loads(hc)
return AtomDetailRow(
atom_id=r["atom_id"],
content_hash=r["content_hash"],
content_preview=r["content_preview"],
component=r["component"],
tier=r["tier"],
cache_tier=r["cache_tier"],
prompt_hash=r["prompt_hash"],
framework_version=r["framework_version"],
model_used=r["model_used"],
llm_confidence=float(r["llm_confidence"]) if r["llm_confidence"] is not None else None,
human_validated=r["human_validated"],
hit_count=r["hit_count"] or 0,
last_hit_at=r["last_hit_at"],
created_at=r["created_at"],
updated_at=r["updated_at"],
expires_at=r["expires_at"],
result_processed=rp or {},
result_raw=rr,
human_corrections=hc,
validator_user_id=r["validator_user_id"],
validated_at=r["validated_at"],
)
async def expire_atom(atom_id: int) -> bool:
"""Mark atom as expired immediately (soft delete — keeps row for audit).
Returns True if a row was updated, False if not found.
"""
sql = """
UPDATE brain_analysis_atom
SET expires_at = now(), updated_at = now()
WHERE atom_id = $1
"""
async with db.pool.acquire() as conn:
result = await conn.execute(sql, atom_id)
# asyncpg execute() returns string like "UPDATE 1"
return result.endswith("1")
# ============================================================================
# Verification cache admin
# ============================================================================
@dataclass(slots=True)
class VerificationListRow:
claim_hash: str
tier: str
model: str | None
prompt_hash: str
framework_version: str | None
schema_name: str
evidence_url_count: int
created_at: datetime
updated_at: datetime
expires_at: datetime
status: str | None # extracted from verification_processed.status if present
volatility: str | None
topic_codes: list[str]
@dataclass(slots=True)
class VerificationDetailRow(VerificationListRow):
evidence_urls: list[str]
verification_processed: dict
verification_raw: dict | None
@dataclass(slots=True)
class VerificationListPage:
items: list[VerificationListRow]
total: int
page: int
page_size: int
def _verif_filters(
tier: str | None,
q: str | None,
) -> tuple[str, list]:
where = ["1=1"]
params: list[Any] = []
idx = 1
if tier and tier != "all":
where.append(f"tier = ${idx}")
params.append(tier)
idx += 1
if q:
# Search on claim_hash prefix or model name
where.append(f"(claim_hash ILIKE ${idx} OR model ILIKE ${idx})")
params.append(f"%{q}%")
idx += 1
return " AND ".join(where), params
async def list_verifications(
*,
tier: str | None = None,
q: str | None = None,
page: int = 1,
page_size: int = 25,
) -> VerificationListPage:
page = max(1, page)
page_size = min(max(1, page_size), 100)
offset = (page - 1) * page_size
where_sql, params = _verif_filters(tier, q)
count_sql = f"SELECT COUNT(*) AS n FROM brain_verification_cache WHERE {where_sql}"
list_sql = f"""
SELECT claim_hash, tier, model, prompt_hash, framework_version, schema_name,
jsonb_array_length(evidence_urls) AS evidence_url_count,
verification_processed,
volatility, topic_codes,
created_at, updated_at, expires_at
FROM brain_verification_cache
WHERE {where_sql}
ORDER BY updated_at DESC
LIMIT ${len(params)+1} OFFSET ${len(params)+2}
"""
async with db.pool.acquire() as conn:
count_row = await conn.fetchrow(count_sql, *params)
rows = await conn.fetch(list_sql, *params, page_size, offset)
items: list[VerificationListRow] = []
for r in rows:
vp = r["verification_processed"]
if isinstance(vp, str):
try:
vp = json.loads(vp)
except Exception:
vp = {}
status = vp.get("status") if isinstance(vp, dict) else None
items.append(
VerificationListRow(
claim_hash=r["claim_hash"],
tier=r["tier"],
model=r["model"],
prompt_hash=r["prompt_hash"],
framework_version=r["framework_version"],
schema_name=r["schema_name"],
evidence_url_count=r["evidence_url_count"] or 0,
created_at=r["created_at"],
updated_at=r["updated_at"],
expires_at=r["expires_at"],
status=status,
volatility=r["volatility"],
topic_codes=list(r["topic_codes"]) if r["topic_codes"] else [],
)
)
return VerificationListPage(
items=items,
total=count_row["n"] or 0,
page=page,
page_size=page_size,
)
async def get_verification(
claim_hash: str, tier: Literal["free", "premium"]
) -> VerificationDetailRow | None:
sql = """
SELECT claim_hash, tier, evidence_hash, evidence_urls, model, prompt_hash,
framework_version, schema_name, verification_processed, verification_raw,
volatility, topic_codes,
created_at, updated_at, expires_at
FROM brain_verification_cache
WHERE claim_hash = $1 AND tier = $2
"""
async with db.pool.acquire() as conn:
r = await conn.fetchrow(sql, claim_hash, tier)
if not r:
return None
ev = r["evidence_urls"]
vp = r["verification_processed"]
vr = r["verification_raw"]
if isinstance(ev, str):
ev = json.loads(ev)
if isinstance(vp, str):
vp = json.loads(vp)
if isinstance(vr, str):
vr = json.loads(vr)
status = vp.get("status") if isinstance(vp, dict) else None
return VerificationDetailRow(
claim_hash=r["claim_hash"],
tier=r["tier"],
model=r["model"],
prompt_hash=r["prompt_hash"],
framework_version=r["framework_version"],
schema_name=r["schema_name"],
evidence_url_count=len(ev) if ev else 0,
created_at=r["created_at"],
updated_at=r["updated_at"],
expires_at=r["expires_at"],
status=status,
volatility=r["volatility"],
topic_codes=list(r["topic_codes"]) if r["topic_codes"] else [],
evidence_urls=list(ev) if ev else [],
verification_processed=vp or {},
verification_raw=vr,
)
async def delete_verification(
claim_hash: str, tier: Literal["free", "premium"]
) -> bool:
sql = "DELETE FROM brain_verification_cache WHERE claim_hash = $1 AND tier = $2"
async with db.pool.acquire() as conn:
result = await conn.execute(sql, claim_hash, tier)
return result.endswith("1")
# ============================================================================
# Taxonomy admin
# ============================================================================
async def get_taxonomy_info(resolver) -> dict:
"""Snapshot of currently loaded taxonomy.
`resolver.all` is a dict {path: tag_id}, so iterating it yields the path
strings directly (e.g. "Country", "Country/France", "Topics/Health/COVID").
"""
all_paths: list[str] = list(resolver.all) if hasattr(resolver, "all") else []
by_namespace: dict[str, int] = {}
for path in all_paths:
ns = path.split("/", 1)[0] if "/" in path else (path or "(root)")
by_namespace[ns] = by_namespace.get(ns, 0) + 1
return {
"total_tags": len(all_paths),
"by_namespace": by_namespace,
"namespaces": sorted(by_namespace.keys()),
}
async def reload_taxonomy(resolver, atomic) -> dict:
"""Re-fetch tags from atomic and replace the in-process resolver state.
Returns a small status dict.
"""
from shared.taxonomy import build_path_map_from_tags
try:
live_tags = await atomic.list_tags()
path_map = build_path_map_from_tags(live_tags)
before = len(resolver.all) if hasattr(resolver, "all") else 0
if path_map:
resolver.load_from_mapping(path_map)
after = len(resolver.all) if hasattr(resolver, "all") else 0
return {
"ok": True,
"before": before,
"after": after,
"fetched": len(path_map),
}
except Exception as e: # noqa: BLE001
log.exception("taxonomy_reload_failed")
return {"ok": False, "error": f"{type(e).__name__}: {e}"}
# ============================================================================
# Extended stats (existing /v1/analysis_atom/stats kept; this returns more)
# ============================================================================
async def get_stats_extended() -> dict:
"""Extended stats: per-tier hit rates, top components, recent activity."""
sql = """
WITH base AS (
SELECT
cache_tier,
component,
hit_count,
last_hit_at,
created_at,
validated_at,
human_validated
FROM brain_analysis_atom
)
SELECT
COUNT(*) AS total,
COUNT(*) FILTER (WHERE cache_tier='gold') AS gold,
COUNT(*) FILTER (WHERE cache_tier='silver') AS silver,
COUNT(*) FILTER (WHERE cache_tier='bronze') AS bronze,
COUNT(*) FILTER (WHERE component='techniques') AS c_tech,
COUNT(*) FILTER (WHERE component='ai_tampered') AS c_ai,
COUNT(*) FILTER (WHERE component='claims') AS c_claims,
COUNT(*) FILTER (WHERE created_at > now() - interval '24 hours') AS writes_24h,
SUM(hit_count) FILTER (WHERE last_hit_at > now() - interval '24 hours') AS hits_24h,
SUM(hit_count) FILTER (WHERE cache_tier='gold' AND last_hit_at > now() - interval '24 hours') AS hits_24h_gold,
SUM(hit_count) FILTER (WHERE cache_tier='silver' AND last_hit_at > now() - interval '24 hours') AS hits_24h_silver,
-- Count promotions by when the moderator validated, not when the atom was first written.
COUNT(*) FILTER (WHERE human_validated = true AND validated_at > now() - interval '24 hours') AS gold_promotions_24h
FROM base
"""
async with db.pool.acquire() as conn:
row = await conn.fetchrow(sql)
total = row["total"] or 0
hits_24h = int(row["hits_24h"] or 0)
writes_24h = int(row["writes_24h"] or 0)
hits_24h_gold = int(row["hits_24h_gold"] or 0)
hits_24h_silver = int(row["hits_24h_silver"] or 0)
denom = hits_24h + writes_24h
hit_rate_24h = (hits_24h / max(denom, 1)) if denom > 0 else None
return {
"total_atoms": total,
"by_tier": {
"gold": row["gold"] or 0,
"silver": row["silver"] or 0,
"bronze": row["bronze"] or 0,
},
"by_component": {
"techniques": row["c_tech"] or 0,
"ai_tampered": row["c_ai"] or 0,
"claims": row["c_claims"] or 0,
},
"hit_rate_24h": hit_rate_24h,
"hits_24h_gold": hits_24h_gold,
"hits_24h_silver": hits_24h_silver,
"writes_24h": writes_24h,
"gold_promotions_24h": int(row["gold_promotions_24h"] or 0),
}

View file

@ -0,0 +1,761 @@
"""Analysis atom cache — store full-component LLM results for techniques/ai_tampered/claims.
Contract with didi-backend agent-v3:
- Before LLM run, agent-v3 calls POST /v1/analysis_atom/lookup with
(content_hash, component, prompt_hash). On gold or silver+fresh hit,
backend skips LLM and uses cached result.
- After LLM run (only if tier='premium'), agent-v3 fires POST /v1/analysis_atom
to cache the result. cache_tier is silver by default; bronze if llm_confidence
is below the configured threshold.
- After moderator resolves a session with corrections, agent-v3 calls PATCH
/v1/analysis_atom/{atom_id} with human_validated=true to promote silvergold.
Storage rules:
- Lookup key: (content_hash, component, prompt_hash) tier excluded so
free users benefit from premium cached entries.
- Write: rejected if tier='free' (only premium runs ingest).
- Bronze atoms (low confidence) are stored for audit but NEVER served on
lookup. They can be promoted to silver if a future run produces higher
confidence on the same content.
- Gold atoms have expires_at=NULL (forever). Silver/bronze get TTL.
"""
from __future__ import annotations
import asyncio
import hashlib
import json
from dataclasses import dataclass, field
from datetime import datetime, timedelta, timezone
from typing import Any, Literal
from brain_api.db import db
from brain_api.services.classifier import (
ClaimVolatility,
classify_claim_volatility,
)
from shared.llm_client import LlmClient
from shared.logging import get_logger
log = get_logger(__name__)
# ----- legacy fallback TTLs when no classification is available -----
# These are now the *startup defaults* — at runtime they can be overridden
# live by the AI platform dashboard via RuntimeConfigClient (keys
# `brain.atom.silver_ttl_days`, `brain.atom.bronze_ttl_days`,
# `brain.atom.confidence_silver_threshold`). The helpers below read live
# values with these as fallback. Direct constant access is kept for tests
# and code paths that don't yet plumb through Settings.
SILVER_TTL_HOURS = 90 * 24 # 2160h, ~3 months
BRONZE_TTL_HOURS = 30 * 24 # 720h, ~1 month
DEFAULT_CONFIDENCE_THRESHOLD = 60.0 # below → bronze, at/above → silver
def _live_silver_ttl_hours() -> int:
"""Read silver TTL from runtime config (or fall back to settings/constant).
Priority: dashboard live value > Settings default > constant. Settings is
cached so this is essentially free; runtime_config falls back gracefully
if the dashboard is unreachable.
"""
from brain_api.runtime_config import get_global_client
from shared.config import get_settings
rc = get_global_client()
fallback_days = get_settings().atom_silver_ttl_days
if rc is not None and rc.enabled:
return rc.get_int("brain.atom.silver_ttl_days", fallback_days) * 24
return fallback_days * 24
def _live_bronze_ttl_hours() -> int:
from brain_api.runtime_config import get_global_client
from shared.config import get_settings
rc = get_global_client()
fallback_days = get_settings().atom_bronze_ttl_days
if rc is not None and rc.enabled:
return rc.get_int("brain.atom.bronze_ttl_days", fallback_days) * 24
return fallback_days * 24
def _live_confidence_threshold() -> float:
from brain_api.runtime_config import get_global_client
from shared.config import get_settings
rc = get_global_client()
fallback = get_settings().atom_confidence_silver_threshold
if rc is not None and rc.enabled:
return rc.get_float("brain.atom.confidence_silver_threshold", fallback)
return fallback
AtomComponent = Literal["techniques", "ai_tampered", "claims"]
AtomTier = Literal["free", "premium"]
AtomCacheTier = Literal["gold", "silver", "bronze"]
StalenessStatus = Literal["fresh", "stale_framework", "stale_prompt", "miss"]
# --------------------------------------------------------------------------- DTO
@dataclass(slots=True)
class AtomEntry:
atom_id: int
content_hash: str
component: str
tier: str
prompt_hash: str
framework_version: str | None
model_used: str | None
cache_tier: str
human_validated: bool
human_corrections: dict | None
validator_user_id: str | None
validated_at: datetime | None
result_processed: dict
result_raw: dict | None
llm_confidence: float | None
hit_count: int
last_hit_at: datetime | None
created_at: datetime
updated_at: datetime
expires_at: datetime | None
content_preview: str | None = None
# --------------------------------------------------------------------- helpers
def normalize_content_hash(content_hash: str) -> str:
"""Backend computes content_hash already; this is just a passthrough/sanity check.
Convention: backend sends sha256(text)[:16] or sha256(text). We store as-is.
"""
return content_hash.strip().lower()
def _decide_cache_tier(llm_confidence: float | None, override: str | None) -> str:
if override in ("gold", "silver", "bronze"):
return override
if llm_confidence is None:
# No confidence info → assume good enough (silver)
return "silver"
threshold = _live_confidence_threshold()
return "silver" if llm_confidence >= threshold else "bronze"
def _resolve_ttl_hours(
*,
cache_tier: str,
classification: ClaimVolatility | None,
) -> int | None:
"""Pick the effective TTL in hours, combining cache_tier + classification.
Rules:
- gold None (forever, regardless of classification)
- silver/bronze with classification use classifier estimate (already
capped per volatility tier in classifier.py)
- silver/bronze without classification fall back to legacy fixed TTLs
The classification's estimate is *already* clamped to per-tier hard caps
(volatile48h, evolving720h, stable26280h) by the classifier, so we
just trust it here.
"""
if cache_tier == "gold":
return None
if classification is not None and not classification.degraded:
return classification.estimated_validity_hours
return _live_silver_ttl_hours() if cache_tier == "silver" else _live_bronze_ttl_hours()
def _expires_at_from_hours(ttl_hours: int | None) -> datetime | None:
"""Convert TTL hours → expires_at timestamptz. None → no expiry (gold)."""
if ttl_hours is None:
return None
return datetime.now(tz=timezone.utc) + timedelta(hours=ttl_hours)
async def _get_or_compute_classification(
*,
classification: ClaimVolatility | None,
llm: LlmClient | None,
content_preview: str | None,
) -> ClaimVolatility | None:
"""Use caller-provided classification, else compute via LLM if possible.
Returns None if neither path is available caller falls back to legacy
behavior (no volatility, fixed TTL).
"""
if classification is not None:
return classification
if llm is None or not content_preview:
return None
try:
return await classify_claim_volatility(llm, claim=content_preview)
except Exception as e: # noqa: BLE001
log.warning(
"atom_upsert_classifier_failed",
error=f"{type(e).__name__}:{e}",
)
return None
async def _register_facts_async(
classification: ClaimVolatility | None,
*,
source_atom_id: int,
) -> None:
"""Best-effort fact registration after a successful upsert.
Imports lazily to avoid a circular import (fact_status imports classifier
types). Errors are swallowed fact registration is enrichment, not core.
"""
if classification is None or not classification.entity_bindings:
return
try:
# Local import: services.fact_status imports classifier types, so
# importing it at module load would create a cycle.
from brain_api.services.fact_status import register_facts_from_bindings
await register_facts_from_bindings(
classification.entity_bindings,
volatility=classification.volatility,
topic_codes=classification.topic_codes,
source_atom_id=str(source_atom_id),
)
except Exception as e: # noqa: BLE001
log.warning(
"fact_registration_failed",
atom_id=source_atom_id,
error=f"{type(e).__name__}:{e}",
)
def _row_to_entry(row) -> AtomEntry:
rp = row["result_processed"]
rr = row["result_raw"]
hc = row["human_corrections"]
if isinstance(rp, str):
rp = json.loads(rp)
if isinstance(rr, str):
rr = json.loads(rr)
if isinstance(hc, str):
hc = json.loads(hc)
return AtomEntry(
atom_id=row["atom_id"],
content_hash=row["content_hash"],
component=row["component"],
tier=row["tier"],
prompt_hash=row["prompt_hash"],
framework_version=row["framework_version"],
model_used=row["model_used"],
cache_tier=row["cache_tier"],
human_validated=row["human_validated"],
human_corrections=hc,
validator_user_id=row["validator_user_id"],
validated_at=row["validated_at"],
result_processed=rp or {},
result_raw=rr,
llm_confidence=float(row["llm_confidence"]) if row["llm_confidence"] is not None else None,
hit_count=row["hit_count"],
last_hit_at=row["last_hit_at"],
created_at=row["created_at"],
updated_at=row["updated_at"],
expires_at=row["expires_at"],
content_preview=row.get("content_preview") if hasattr(row, "get") else None,
)
def decide_freshness(
entry: AtomEntry | None,
current_prompt_hash: str | None,
current_framework_version: str | None,
) -> StalenessStatus:
"""fresh | stale_prompt | stale_framework | miss.
Gold atoms are ALWAYS fresh human-validated answers don't depend on prompt.
"""
if entry is None:
return "miss"
if entry.cache_tier == "gold":
return "fresh"
if current_prompt_hash and entry.prompt_hash != current_prompt_hash:
return "stale_prompt"
if (
current_framework_version
and entry.framework_version
and entry.framework_version != current_framework_version
):
return "stale_framework"
return "fresh"
# ------------------------------------------------------------------------- IO
async def lookup(
*,
content_hash: str,
component: AtomComponent,
prompt_hash: str,
framework_version: str | None = None,
) -> tuple[AtomEntry | None, StalenessStatus]:
"""Lookup an atom by (content_hash, component) — tier-agnostic.
Returns (entry_or_None, staleness). Bronze atoms are filtered out (NEVER
served). For multiple matches with different prompt_hash, prefers gold.
"""
ch = normalize_content_hash(content_hash)
sql = """
SELECT atom_id, content_hash, component, tier, prompt_hash, framework_version,
model_used, cache_tier, human_validated, human_corrections, validator_user_id,
validated_at, result_processed, result_raw, llm_confidence, hit_count,
last_hit_at, created_at, updated_at, expires_at, content_preview
FROM brain_analysis_atom
WHERE content_hash = $1
AND component = $2
AND cache_tier IN ('gold', 'silver')
AND (expires_at IS NULL OR expires_at > now())
ORDER BY (cache_tier = 'gold') DESC, updated_at DESC
LIMIT 1
"""
async with db.pool.acquire() as conn:
row = await conn.fetchrow(sql, ch, component)
if not row:
return None, "miss"
entry = _row_to_entry(row)
staleness = decide_freshness(entry, prompt_hash, framework_version)
# Increment hit_count (best-effort)
if staleness == "fresh":
try:
async with db.pool.acquire() as conn:
await conn.execute(
"UPDATE brain_analysis_atom SET hit_count = hit_count + 1, last_hit_at = now() WHERE atom_id = $1",
entry.atom_id,
)
except Exception: # noqa: BLE001
log.debug("hit_count_update_skipped", atom_id=entry.atom_id)
return entry, staleness
async def upsert(
*,
content_hash: str,
content_preview: str | None,
component: AtomComponent,
tier: AtomTier,
prompt_hash: str,
framework_version: str | None,
model_used: str | None,
result_processed: dict,
result_raw: dict | None,
llm_confidence: float | None,
cache_tier_override: str | None = None,
classification: ClaimVolatility | None = None,
llm: LlmClient | None = None,
) -> tuple[AtomEntry | None, str | None]:
"""Insert or update an atom row, with volatility classification.
Pipeline:
1. Reject tier='free' silently (premium-only ingest).
2. If no ``classification`` provided and an ``llm`` client is, run the
volatility classifier on ``content_preview`` to derive volatility,
topic_codes, entity_bindings, and the recommended TTL.
3. UPSERT the row with the new metadata columns. Gold rows are
preserved on every soft field (truth-preserving).
4. After successful write, schedule fact_status registration as a
background task (best-effort, errors swallowed).
Args:
content_hash, content_preview, component, tier, prompt_hash,
framework_version, model_used, result_processed, result_raw,
llm_confidence, cache_tier_override: same as before.
classification: Pre-computed ClaimVolatility from caller. If None,
attempts to compute via ``llm``.
llm: LLM client for classifier. Optional passing None disables
classification (legacy fixed-TTL behavior).
Returns:
``(entry, skip_reason)``. Skip cases (entry=None):
- tier='free' reject silently (premium-only ingest)
- SQL error propagated to caller
"""
if tier == "free":
return None, "tier=free (premium-only ingest)"
ch = normalize_content_hash(content_hash)
cache_tier = _decide_cache_tier(llm_confidence, cache_tier_override)
# Volatility classification — caller-provided or LLM-derived.
classification = await _get_or_compute_classification(
classification=classification,
llm=llm,
content_preview=content_preview,
)
ttl_hours = _resolve_ttl_hours(
cache_tier=cache_tier, classification=classification
)
expires_at = _expires_at_from_hours(ttl_hours)
volatility = classification.volatility if classification else None
topic_codes = classification.topic_codes if classification else []
entity_bindings_json = (
json.dumps(classification.entity_bindings_jsonb())
if classification
else "[]"
)
rp_json = json.dumps(result_processed)
rr_json = json.dumps(result_raw) if result_raw is not None else None
sql = """
INSERT INTO brain_analysis_atom (
content_hash, content_preview, component, tier, prompt_hash, framework_version,
model_used, result_processed, result_raw, llm_confidence,
cache_tier, expires_at,
volatility, topic_codes, entity_bindings, ttl_hours_used
)
VALUES (
$1, $2, $3, $4, $5, $6, $7, $8::jsonb, $9::jsonb, $10, $11, $12,
$13, $14, $15::jsonb, $16
)
ON CONFLICT (content_hash, component, prompt_hash) DO UPDATE SET
-- DO NOT downgrade gold to silver preserve human validation
result_processed = CASE
WHEN brain_analysis_atom.cache_tier = 'gold' THEN brain_analysis_atom.result_processed
ELSE EXCLUDED.result_processed
END,
result_raw = CASE
WHEN brain_analysis_atom.cache_tier = 'gold' THEN brain_analysis_atom.result_raw
ELSE EXCLUDED.result_raw
END,
cache_tier = CASE
WHEN brain_analysis_atom.cache_tier = 'gold' THEN brain_analysis_atom.cache_tier
ELSE EXCLUDED.cache_tier
END,
llm_confidence = CASE
WHEN brain_analysis_atom.cache_tier = 'gold' THEN brain_analysis_atom.llm_confidence
ELSE EXCLUDED.llm_confidence
END,
model_used = COALESCE(EXCLUDED.model_used, brain_analysis_atom.model_used),
framework_version = COALESCE(EXCLUDED.framework_version, brain_analysis_atom.framework_version),
content_preview = COALESCE(EXCLUDED.content_preview, brain_analysis_atom.content_preview),
tier = EXCLUDED.tier,
updated_at = now(),
expires_at = CASE
WHEN brain_analysis_atom.cache_tier = 'gold' THEN brain_analysis_atom.expires_at
ELSE EXCLUDED.expires_at
END,
-- Volatility metadata: prefer fresh values when present (a re-write
-- with classifier may have better data than the original write).
volatility = COALESCE(EXCLUDED.volatility, brain_analysis_atom.volatility),
topic_codes = CASE
WHEN array_length(EXCLUDED.topic_codes, 1) > 0
THEN EXCLUDED.topic_codes
ELSE brain_analysis_atom.topic_codes
END,
entity_bindings = CASE
WHEN jsonb_array_length(EXCLUDED.entity_bindings) > 0
THEN EXCLUDED.entity_bindings
ELSE brain_analysis_atom.entity_bindings
END,
ttl_hours_used = COALESCE(EXCLUDED.ttl_hours_used, brain_analysis_atom.ttl_hours_used)
RETURNING atom_id, content_hash, component, tier, prompt_hash, framework_version,
model_used, cache_tier, human_validated, human_corrections, validator_user_id,
validated_at, result_processed, result_raw, llm_confidence, hit_count,
last_hit_at, created_at, updated_at, expires_at, content_preview
"""
async with db.pool.acquire() as conn:
row = await conn.fetchrow(
sql,
ch, content_preview, component, tier, prompt_hash, framework_version,
model_used, rp_json, rr_json, llm_confidence, cache_tier, expires_at,
volatility, topic_codes, entity_bindings_json, ttl_hours,
)
assert row is not None
entry = _row_to_entry(row)
# Fire-and-forget fact registration (best-effort, never blocks the write).
if classification and classification.entity_bindings:
asyncio.create_task(
_register_facts_async(classification, source_atom_id=entry.atom_id)
)
log.info(
"atom_upsert_ok",
atom_id=entry.atom_id,
component=component,
tier=tier,
cache_tier=entry.cache_tier,
volatility=volatility,
ttl_hours=ttl_hours,
topic_codes=topic_codes,
binding_count=len(classification.entity_bindings) if classification else 0,
)
return entry, None
async def patch_to_gold(
*,
atom_id: int,
human_validated: bool = True,
human_corrections: dict | None = None,
validator_user_id: str | None = None,
result_processed: dict | None = None,
) -> AtomEntry | None:
"""Promote an atom to gold after human review.
If result_processed is provided (corrections applied), it replaces the LLM
result. Otherwise the existing result is kept (e.g. moderator approved as is).
"""
sets = [
"human_validated = $2",
"validator_user_id = $3",
"validated_at = now()",
"cache_tier = 'gold'",
"expires_at = NULL",
"updated_at = now()",
]
params: list[Any] = [atom_id, human_validated, validator_user_id]
next_idx = 4
if human_corrections is not None:
sets.append(f"human_corrections = ${next_idx}::jsonb")
params.append(json.dumps(human_corrections))
next_idx += 1
else:
sets.append("human_corrections = NULL")
if result_processed is not None:
sets.append(f"result_processed = ${next_idx}::jsonb")
params.append(json.dumps(result_processed))
next_idx += 1
sql = f"""
UPDATE brain_analysis_atom
SET {", ".join(sets)}
WHERE atom_id = $1
RETURNING atom_id, content_hash, component, tier, prompt_hash, framework_version,
model_used, cache_tier, human_validated, human_corrections, validator_user_id,
validated_at, result_processed, result_raw, llm_confidence, hit_count,
last_hit_at, created_at, updated_at, expires_at, content_preview
"""
async with db.pool.acquire() as conn:
row = await conn.fetchrow(sql, *params)
if not row:
return None
return _row_to_entry(row)
async def get_stats() -> dict: # noqa: PLR0915 (kept compact; flake later)
return await _get_stats_impl()
# ============================================================================
# Phase B2: confidence decay + judge integration
# ============================================================================
# These helpers run on cache HITS to decide whether the cached verdict is
# still trustworthy. They do not own the lookup query itself — callers
# (gather.py, the lookup endpoint, the daily auditor) call ``lookup()`` first,
# then optionally call ``judge_and_update`` if they have fresh evidence.
# Decay half-lives (hours) per volatility tier. Beyond half-life, the cached
# confidence is halved; at 2× half-life, quartered; etc. Stable claims do
# not decay.
DECAY_HALF_LIVES_HOURS: dict[str, float] = {
"volatile": 24.0, # ~half confidence after 1 day
"evolving": 168.0, # ~half after 1 week
"stable": float("inf"),
}
# Cap on audit_history length kept on each row — older entries are trimmed.
AUDIT_HISTORY_MAX = 50
# Minimum hours between audit-pass increments triggered by lookup judges.
# Without this, popular content hits the auditor 100×/day and consecutive_-
# audit_passes rockets, defeating the purpose. The daily auditor cron is
# the authoritative source of audit passes; lookups only nudge.
AUDIT_PASS_MIN_INTERVAL_HOURS = 6.0
# Confidence floor below which a "fresh" cache entry is treated as miss.
EFFECTIVE_CONFIDENCE_FLOOR = 60.0
def compute_effective_confidence(
*,
base_confidence: float | None,
volatility: str | None,
age_hours: float,
consecutive_audit_passes: int = 0,
) -> float | None:
"""Decay base confidence by age, modulated by volatility and audit history.
Stable rows do not decay. Volatile/evolving rows lose confidence with
exponential half-life. Atoms that survived many audits get a multiplier
boost (max +30% over base).
Returns:
Decayed confidence value, or None if base was None.
"""
if base_confidence is None:
return None
half = DECAY_HALF_LIVES_HOURS.get(volatility or "evolving", 168.0)
if half == float("inf") or age_hours <= 0:
decay = 1.0
else:
# Exponential decay: each half-life halves the confidence.
decay = 0.5 ** (age_hours / half)
audit_boost = min(0.3, 0.03 * max(0, consecutive_audit_passes))
return float(base_confidence) * decay * (1.0 + audit_boost)
def is_effectively_fresh(
*,
base_confidence: float | None,
volatility: str | None,
age_hours: float,
consecutive_audit_passes: int = 0,
floor: float = EFFECTIVE_CONFIDENCE_FLOOR,
) -> bool:
"""True if the decayed confidence is above the freshness floor.
Callers can use this *in addition* to ``decide_freshness`` to drop
entries that are technically not stale but have decayed below usable
confidence.
"""
eff = compute_effective_confidence(
base_confidence=base_confidence,
volatility=volatility,
age_hours=age_hours,
consecutive_audit_passes=consecutive_audit_passes,
)
if eff is None:
# No base confidence stored → trust the freshness flag from the SQL
# path; we have no other signal.
return True
return eff >= floor
async def apply_judge_verdict(
atom_id: int,
verdict: object, # JudgeVerdict — typed loosely to avoid circular import
) -> None:
"""Persist a JudgeVerdict to brain_analysis_atom.
Updates audit_history (append, cap at AUDIT_HISTORY_MAX), last_audited_at,
consecutive_audit_passes (incremented only when KEEP_CACHE and last
increment was >AUDIT_PASS_MIN_INTERVAL_HOURS ago), and expires_at on
INVALIDATE (sets to now() so the row is treated as expired).
Also writes a brain_audit_log row for global telemetry.
"""
if not db.pool:
raise RuntimeError("brain_db not connected")
# Lazy import to avoid circular: cache_judge ← nli ← (transitively) us.
from brain_api.services.cache_judge import JudgeVerdict
if not isinstance(verdict, JudgeVerdict):
raise TypeError(
f"apply_judge_verdict: expected JudgeVerdict, got {type(verdict).__name__}"
)
audit_entry = verdict.to_audit_entry()
audit_json = json.dumps(audit_entry)
sql = """
UPDATE brain_analysis_atom
SET
audit_history = (
-- Append new entry, then keep only the last AUDIT_HISTORY_MAX.
SELECT jsonb_agg(elem)
FROM (
SELECT elem
FROM jsonb_array_elements(
COALESCE(audit_history, '[]'::jsonb) || $2::jsonb
) WITH ORDINALITY AS t(elem, ord)
ORDER BY ord DESC
LIMIT $3
) recent
),
last_audited_at = now(),
consecutive_audit_passes = CASE
WHEN $4 = 'KEEP_CACHE' AND (
last_audited_at IS NULL
OR last_audited_at < now() - ($5 || ' hours')::interval
)
THEN consecutive_audit_passes + 1
WHEN $4 = 'INVALIDATE' THEN 0
ELSE consecutive_audit_passes
END,
expires_at = CASE
WHEN $4 = 'INVALIDATE' AND cache_tier <> 'gold' THEN now()
ELSE expires_at
END,
updated_at = now()
WHERE atom_id = $1
"""
async with db.pool.acquire() as conn:
await conn.execute(
sql,
atom_id,
json.dumps([audit_entry]), # wrap as JSONB array for concat
AUDIT_HISTORY_MAX,
verdict.decision,
str(int(AUDIT_PASS_MIN_INTERVAL_HOURS)),
)
# Audit log entry for cross-table telemetry / dashboards.
await conn.execute(
"""
INSERT INTO brain_audit_log (action, target_table, target_id, actor, payload)
VALUES ($1, 'brain_analysis_atom', $2, 'cache_judge', $3::jsonb)
""",
f"judge_{verdict.decision.lower()}",
str(atom_id),
audit_json,
)
async def _get_stats_impl() -> dict:
"""Return aggregated counts for monitoring."""
sql = """
SELECT
COUNT(*) AS total_atoms,
COUNT(*) FILTER (WHERE cache_tier = 'gold') AS gold,
COUNT(*) FILTER (WHERE cache_tier = 'silver') AS silver,
COUNT(*) FILTER (WHERE cache_tier = 'bronze') AS bronze,
COUNT(*) FILTER (WHERE component = 'techniques') AS c_techniques,
COUNT(*) FILTER (WHERE component = 'ai_tampered') AS c_ai_tampered,
COUNT(*) FILTER (WHERE component = 'claims') AS c_claims,
COUNT(*) FILTER (WHERE created_at > now() - interval '24 hours') AS writes_24h,
SUM(hit_count) FILTER (WHERE last_hit_at > now() - interval '24 hours') AS hits_24h
FROM brain_analysis_atom
"""
async with db.pool.acquire() as conn:
row = await conn.fetchrow(sql)
total = row["total_atoms"] or 0
hits_24h = int(row["hits_24h"] or 0)
writes_24h = int(row["writes_24h"] or 0)
hit_rate = None
if hits_24h + writes_24h > 0:
hit_rate = hits_24h / max(hits_24h + writes_24h, 1)
return {
"total_atoms": total,
"by_tier": {
"gold": row["gold"] or 0,
"silver": row["silver"] or 0,
"bronze": row["bronze"] or 0,
},
"by_component": {
"techniques": row["c_techniques"] or 0,
"ai_tampered": row["c_ai_tampered"] or 0,
"claims": row["c_claims"] or 0,
},
"hit_rate_24h": hit_rate,
"writes_24h": writes_24h,
}

View file

@ -0,0 +1,315 @@
"""Cache Judge — Pilon 2 of the cache freshness defense.
On every cache hit (verification_cache or analysis_atom), the judge runs NLI
between the cached truth direction and current top-K fresh evidence. If fresh
sources contradict the cached verdict, the cache is invalidated and the
caller is forced to recompute with current data.
Decision rules:
- **KEEP_CACHE** fresh evidence supports the cached verdict (or volatility
is stable and age is below threshold, where we skip NLI entirely).
- **INVALIDATE** fresh evidence contradicts the cached verdict above the
threshold; caller must treat this as a miss and recompute.
- **NEEDS_FULL_RECHECK** evidence is mostly neutral or split; caller may
still serve the cache but should mark it as low-confidence.
Cheap path: stable claims younger than ``STABLE_NLI_SKIP_HOURS`` skip NLI
entirely (no LLM call) pure cache hit.
Confidence boost: rows that passed many consecutive audits get a higher
contradiction threshold (we trust them more). A row that survived 10 daily
audits requires more contradicting evidence to invalidate than a fresh write.
Audit logging: callers should append the JudgeVerdict to brain_analysis_atom.
audit_history (or brain_verification_cache.audit_history) so we have a
running record of why a cache was kept or invalidated.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Literal
from brain_api.services.nli import NliResult, classify_batch
from shared.llm_client import LlmClient
from shared.logging import get_logger
log = get_logger(__name__)
Decision = Literal["KEEP_CACHE", "INVALIDATE", "NEEDS_FULL_RECHECK"]
CachedTruth = Literal["TRUE", "FALSE", "MIXED", "UNVERIFIED"]
# Skip NLI entirely for stable claims younger than this. Pure cache hit, zero
# LLM cost. Stable + old → still run NLI (rare facts can change).
STABLE_NLI_SKIP_HOURS = 720.0 # 30 days
# Truncate evidence text before sending to NLI — matches nli.MAX_EVIDENCE_CHARS.
MAX_EVIDENCE_CHARS = 1500
MAX_EVIDENCE_PIECES = 3 # only judge against top-3 fresh sources
# Decision thresholds. These are the fractions of NLI calls that determine
# the verdict.
#
# When cached truth is TRUE:
# contradicts_fraction >= INVALIDATE_THRESHOLD → INVALIDATE
# supports_fraction >= KEEP_THRESHOLD → KEEP_CACHE
# else → NEEDS_FULL_RECHECK
#
# When cached truth is FALSE: roles flip — supports invalidates, contradicts keeps.
INVALIDATE_THRESHOLD = 0.50 # 50% disagreeing → invalidate
KEEP_THRESHOLD = 0.50 # 50% agreeing → keep
MIN_CONFIDENCE = 0.50 # NLI calls below this confidence are ignored
# Confidence boost from consecutive audit passes — each pass nudges the
# invalidate threshold up by this much (so trusted atoms are harder to flip).
AUDIT_PASS_BONUS = 0.03 # +3% per pass
MAX_AUDIT_BONUS = 0.30 # cap at +30%
@dataclass(slots=True, frozen=True)
class EvidenceSnippet:
"""A single piece of fresh evidence to judge cached verdicts against.
Attributes:
url: Canonical source URL (used in audit log).
text: The text excerpt from the source. Will be truncated to
``MAX_EVIDENCE_CHARS`` before NLI.
published_at: When the source was published (ISO string or None).
Used by callers to filter out stale evidence before passing here.
"""
url: str
text: str
published_at: str | None = None
@dataclass(slots=True)
class JudgeVerdict:
"""Output of the cache judge — drives KEEP/INVALIDATE/RECHECK decision.
Attributes:
decision: One of "KEEP_CACHE", "INVALIDATE", "NEEDS_FULL_RECHECK".
nli_results: Per-evidence NLI labels (for audit_history).
nli_skipped: True if we took the cheap path and skipped LLM entirely.
supports_fraction: Fraction of high-confidence NLI calls labelled SUPPORTS.
contradicts_fraction: Fraction labelled CONTRADICTS.
neutral_fraction: Fraction labelled NEUTRAL (or low-confidence).
effective_invalidate_threshold: Threshold actually applied (after audit bonus).
reasoning: Short human-readable explanation, suitable for audit_history.
evaluated_at: When the judgment was made (UTC ISO).
"""
decision: Decision
nli_results: list[NliResult] = field(default_factory=list)
nli_skipped: bool = False
supports_fraction: float = 0.0
contradicts_fraction: float = 0.0
neutral_fraction: float = 0.0
effective_invalidate_threshold: float = INVALIDATE_THRESHOLD
reasoning: str = ""
evaluated_at: str = ""
def to_audit_entry(self) -> dict[str, object]:
"""Serialize for appending to brain_analysis_atom.audit_history."""
return {
"evaluated_at": self.evaluated_at,
"decision": self.decision,
"nli_skipped": self.nli_skipped,
"supports": round(self.supports_fraction, 3),
"contradicts": round(self.contradicts_fraction, 3),
"neutral": round(self.neutral_fraction, 3),
"threshold": round(self.effective_invalidate_threshold, 3),
"reasoning": self.reasoning[:200],
"evidence_count": len(self.nli_results),
}
def _compute_invalidate_threshold(consecutive_audit_passes: int) -> float:
"""Audit-pass bonus: trusted atoms are harder to invalidate.
Each consecutive daily audit that judged KEEP_CACHE increments the
threshold so a single contradicting source can't overturn an atom that's
been stable for weeks.
"""
bonus = min(
MAX_AUDIT_BONUS,
AUDIT_PASS_BONUS * max(0, consecutive_audit_passes),
)
return min(0.95, INVALIDATE_THRESHOLD + bonus)
def _aggregate_nli(results: list[NliResult]) -> tuple[float, float, float]:
"""Compute (supports, contradicts, neutral) fractions over high-confidence calls.
Low-confidence (< MIN_CONFIDENCE) and errored calls count as NEUTRAL we
don't want noisy signals to invalidate cache.
"""
if not results:
return 0.0, 0.0, 1.0
supports = 0
contradicts = 0
neutral = 0
for r in results:
if r.error or r.confidence < MIN_CONFIDENCE:
neutral += 1
elif r.label == "SUPPORTS":
supports += 1
elif r.label == "CONTRADICTS":
contradicts += 1
else:
neutral += 1
total = float(len(results))
return supports / total, contradicts / total, neutral / total
def _decide_for_truth_direction(
*,
cached_truth: CachedTruth,
supports_fraction: float,
contradicts_fraction: float,
invalidate_threshold: float,
) -> tuple[Decision, str]:
"""Map NLI aggregates to KEEP/INVALIDATE/RECHECK based on cached truth direction."""
if cached_truth == "TRUE":
# We expect SUPPORTS. CONTRADICTS is the danger signal.
if contradicts_fraction >= invalidate_threshold:
return "INVALIDATE", (
f"cached=TRUE but {contradicts_fraction:.0%} of fresh evidence "
f"contradicts (threshold {invalidate_threshold:.0%})"
)
if supports_fraction >= KEEP_THRESHOLD:
return "KEEP_CACHE", (
f"cached=TRUE confirmed by {supports_fraction:.0%} fresh evidence"
)
return "NEEDS_FULL_RECHECK", (
f"cached=TRUE but evidence is split: "
f"{supports_fraction:.0%}/{contradicts_fraction:.0%}"
)
if cached_truth == "FALSE":
# We expect CONTRADICTS. SUPPORTS is the danger signal (claim now true).
if supports_fraction >= invalidate_threshold:
return "INVALIDATE", (
f"cached=FALSE but {supports_fraction:.0%} of fresh evidence "
f"supports (threshold {invalidate_threshold:.0%})"
)
if contradicts_fraction >= KEEP_THRESHOLD:
return "KEEP_CACHE", (
f"cached=FALSE confirmed by {contradicts_fraction:.0%} fresh evidence"
)
return "NEEDS_FULL_RECHECK", (
f"cached=FALSE but evidence is split: "
f"{supports_fraction:.0%}/{contradicts_fraction:.0%}"
)
# MIXED / UNVERIFIED — caller couldn't decide originally either; if fresh
# evidence is now decisive in either direction, force a full recheck so
# a stronger verdict can be issued.
if supports_fraction >= KEEP_THRESHOLD or contradicts_fraction >= KEEP_THRESHOLD:
return "NEEDS_FULL_RECHECK", (
f"cached={cached_truth} but fresh evidence has shifted "
f"({supports_fraction:.0%}/{contradicts_fraction:.0%})"
)
return "KEEP_CACHE", (
f"cached={cached_truth}, fresh evidence still inconclusive"
)
async def judge_cache_validity(
llm: LlmClient,
*,
claim: str,
cached_truth: CachedTruth,
current_evidence: list[EvidenceSnippet],
volatility: str,
age_hours: float,
consecutive_audit_passes: int = 0,
) -> JudgeVerdict:
"""Decide whether a cached verdict still holds against current evidence.
Args:
llm: LLM client (used by NLI).
claim: The original claim text what the cache was written for.
cached_truth: The truth direction the cache claims (TRUE/FALSE/MIXED/UNVERIFIED).
current_evidence: Top-K fresh evidence snippets from /v1/gather. Caller
should already have applied recency filtering for volatile topics.
volatility: One of "volatile", "evolving", "stable" controls the
cheap-path skip and influences logging.
age_hours: How old the cache row is (for cheap-path eligibility).
consecutive_audit_passes: How many prior audits the cache survived.
Increases invalidation resistance.
Returns:
JudgeVerdict never raises. On NLI failure, individual evidence calls
return NEUTRAL with error set; aggregation handles it gracefully.
"""
now_iso = datetime.now(tz=timezone.utc).isoformat()
# Cheap path: stable + young → trust the cache without LLM.
if volatility == "stable" and age_hours < STABLE_NLI_SKIP_HOURS:
return JudgeVerdict(
decision="KEEP_CACHE",
nli_skipped=True,
reasoning=(
f"stable + age {age_hours:.0f}h < {STABLE_NLI_SKIP_HOURS:.0f}h "
f"(cheap path)"
),
evaluated_at=now_iso,
)
# No fresh evidence to check against → can't make a decision; let the
# caller treat as a recheck so they go and gather some.
if not current_evidence:
return JudgeVerdict(
decision="NEEDS_FULL_RECHECK",
reasoning="no fresh evidence available to judge against",
evaluated_at=now_iso,
)
# Truncate + cap evidence count.
snippets = current_evidence[:MAX_EVIDENCE_PIECES]
evidence_texts = [s.text[:MAX_EVIDENCE_CHARS] for s in snippets]
nli_results = await classify_batch(
llm,
claim=claim,
evidence_texts=evidence_texts,
)
supports, contradicts, neutral = _aggregate_nli(nli_results)
threshold = _compute_invalidate_threshold(consecutive_audit_passes)
decision, reasoning = _decide_for_truth_direction(
cached_truth=cached_truth,
supports_fraction=supports,
contradicts_fraction=contradicts,
invalidate_threshold=threshold,
)
log.info(
"cache_judge_done",
decision=decision,
cached_truth=cached_truth,
volatility=volatility,
age_hours=round(age_hours, 1),
supports=round(supports, 2),
contradicts=round(contradicts, 2),
neutral=round(neutral, 2),
threshold=round(threshold, 2),
audit_passes=consecutive_audit_passes,
)
return JudgeVerdict(
decision=decision,
nli_results=nli_results,
nli_skipped=False,
supports_fraction=supports,
contradicts_fraction=contradicts,
neutral_fraction=neutral,
effective_invalidate_threshold=threshold,
reasoning=reasoning,
evaluated_at=now_iso,
)

View file

@ -0,0 +1,189 @@
"""Temporal canonicalizer — Pilon 7 of the cache freshness defense.
Resolves relative time markers ("azi", "today", "săptămâna asta") and
underspecified entities ("alegerile") in a claim against the current date,
so the same surface text asked at different times produces different cache
keys. This prevents the most insidious form of cache staleness: a claim
phrased identically in 2024 and 2026 silently serving the 2024 verdict.
Pipeline position:
1. agent-v3 receives a user claim
2. agent-v3 calls /v1/canonicalize with claim + current_date
3. agent-v3 hashes the *canonical* form (not the original) for cache lookups
4. brain receives the same canonical form on subsequent identical-text
requests, but only if the same time horizon yields the same canonical
Failure mode: if LLM fails or returns invalid JSON, we return the original
claim verbatim with ``changed=false``. The cache then behaves as today
(no temporal disambiguation, but no regression either).
"""
from __future__ import annotations
import asyncio
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from shared.config import LlmRole
from shared.llm_client import LlmClient, LlmError
from shared.logging import get_logger
log = get_logger(__name__)
PROMPT_VERSION = "v1"
_PROMPT_PATH = (
Path(__file__).resolve().parent.parent
/ "prompts"
/ f"canonicalize_{PROMPT_VERSION}.md"
)
PER_CALL_TIMEOUT_S = 15.0
MAX_CLAIM_CHARS = 2000
_PROMPT_TEMPLATE: str | None = None
def _load_prompt() -> str:
global _PROMPT_TEMPLATE
if _PROMPT_TEMPLATE is None:
_PROMPT_TEMPLATE = _PROMPT_PATH.read_text(encoding="utf-8")
return _PROMPT_TEMPLATE
@dataclass(slots=True, frozen=True)
class Canonicalization:
"""Result of canonicalize_claim_temporal.
Attributes:
canonical: The rewritten claim with anchors applied.
original: The input claim, verbatim (for audit).
changed: Whether canonical differs meaningfully from original.
anchors_added: Short labels of what was disambiguated.
reasoning: One-line LLM justification.
error: Populated if LLM failed and we fell back to original.
"""
canonical: str
original: str
changed: bool
anchors_added: list[str]
reasoning: str
error: str | None = None
@property
def degraded(self) -> bool:
return self.error is not None
def _passthrough(claim: str, error: str | None = None) -> Canonicalization:
"""Build a no-op canonicalization (claim unchanged)."""
return Canonicalization(
canonical=claim,
original=claim,
changed=False,
anchors_added=[],
reasoning="passthrough" if error is None else "fallback_passthrough",
error=error,
)
def _parse_response(data: object, original: str) -> Canonicalization | None:
"""Validate the LLM JSON response. Returns None on bad shape."""
if not isinstance(data, dict):
return None
canonical = str(data.get("canonical", "")).strip()
if not canonical:
return None
changed = bool(data.get("changed", False))
anchors_raw = data.get("anchors_added") or []
if not isinstance(anchors_raw, list):
return None
anchors = [str(a).strip() for a in anchors_raw if str(a).strip()]
reasoning = str(data.get("reasoning", "")).strip()[:200]
return Canonicalization(
canonical=canonical,
original=original,
changed=changed,
anchors_added=anchors,
reasoning=reasoning,
)
async def canonicalize_claim_temporal(
llm: LlmClient,
*,
claim: str,
current_date: datetime | None = None,
) -> Canonicalization:
"""Resolve temporal markers and ambiguous entities in a claim.
On any failure, returns a passthrough Canonicalization (original=canonical,
changed=False) with ``error`` populated. The caller can log telemetry but
the cache lookup proceeds with the original text no regression.
Args:
llm: LLM client (uses REASONING role).
claim: User claim, possibly containing relative time markers.
current_date: Reference "now". Defaults to UTC now.
Returns:
Canonicalization with the rewritten claim or a passthrough on failure.
"""
if not claim or not claim.strip():
return _passthrough(claim, error="empty_claim")
truncated = claim.strip()[:MAX_CLAIM_CHARS]
today = (current_date or datetime.now(tz=timezone.utc)).date().isoformat()
prompt = (
_load_prompt()
.replace("{current_date}", today)
.replace("{claim}", truncated)
)
try:
result, _usage = await asyncio.wait_for(
llm.chat_json(
role=LlmRole.REASONING,
system=(
"You are a temporal disambiguator. Respond with strictly "
"valid JSON only, no commentary."
),
user=prompt,
max_tokens=400,
temperature=0.0,
),
timeout=PER_CALL_TIMEOUT_S,
)
except asyncio.TimeoutError:
log.warning("canonicalize_timeout", claim_preview=truncated[:80])
return _passthrough(truncated, error="timeout")
except LlmError as e:
log.warning("canonicalize_llm_error", error=str(e)[:200])
return _passthrough(truncated, error=f"llm:{e}")
except Exception as e: # noqa: BLE001
log.warning(
"canonicalize_unexpected_error",
error=f"{type(e).__name__}:{e}",
)
return _passthrough(truncated, error=f"{type(e).__name__}:{e}")
parsed = _parse_response(result, truncated)
if parsed is None:
log.warning("canonicalize_bad_response_shape", got=type(result).__name__)
return _passthrough(truncated, error="bad_response_shape")
if parsed.changed:
log.info(
"canonicalize_anchored",
anchors=parsed.anchors_added,
preview_in=truncated[:80],
preview_out=parsed.canonical[:80],
)
return parsed

View file

@ -0,0 +1,350 @@
"""Volatility classifier — Pilon 1 of the cache freshness defense.
Single LLM call returns the temporal characteristics of a claim:
- how fast it can become outdated (volatility: volatile|evolving|stable)
- which topics it touches (topic_codes)
- which entity-predicate-object triples it binds to (entity_bindings)
- how many hours from now its verification can be trusted
This is invoked BEFORE writing to brain_analysis_atom or brain_verification_cache,
so the resulting metadata becomes part of the cache row and drives:
- TTL (expires_at = now() + estimated_validity_hours, capped per tier)
- audit scheduling (volatile rows get audited daily by didibrain-auditor)
- mass invalidation by topic (didibrain-breaking-watcher)
- fact-status registration (entity_bindings brain_fact_status, Pilon 11)
Failure mode: if the LLM call fails or returns invalid JSON, we degrade
gracefully to a conservative fallback (volatility="evolving",
estimated_validity_hours=168) with `error` populated. The caller still gets a
usable classification and the cache write proceeds. The auditor picks up
non-stable rows on its next sweep and corrects misclassifications over time.
"""
from __future__ import annotations
import asyncio
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Literal
from shared.config import LlmRole
from shared.llm_client import LlmClient, LlmError
from shared.logging import get_logger
log = get_logger(__name__)
PROMPT_VERSION = "v1"
_PROMPT_PATH = (
Path(__file__).resolve().parent.parent
/ "prompts"
/ f"classifier_{PROMPT_VERSION}.md"
)
ALLOWED_VOLATILITY: set[str] = {"volatile", "evolving", "stable"}
# Hard caps (hours) per volatility level — LLM estimate is clamped to these
# upper bounds. Even if the LLM says "this is stable for 5 years", we cap.
HARD_CAPS_HOURS: dict[str, int] = {
"volatile": 48, # max 2 days
"evolving": 720, # max 30 days
"stable": 26280, # max ~3 years
}
# Sensible floors — clamp from below so we never get a 0-hour TTL.
HARD_FLOORS_HOURS: dict[str, int] = {
"volatile": 1,
"evolving": 24,
"stable": 720,
}
# Conservative defaults applied when classification fails.
DEFAULT_VOLATILITY: Literal["volatile", "evolving", "stable"] = "evolving"
DEFAULT_VALIDITY_HOURS = 168 # 7 days
PER_CALL_TIMEOUT_S = 20.0
MAX_CLAIM_CHARS = 2000
_PROMPT_TEMPLATE: str | None = None
def _load_prompt() -> str:
"""Load and cache the classifier prompt template."""
global _PROMPT_TEMPLATE
if _PROMPT_TEMPLATE is None:
_PROMPT_TEMPLATE = _PROMPT_PATH.read_text(encoding="utf-8")
return _PROMPT_TEMPLATE
@dataclass(slots=True, frozen=True)
class EntityBinding:
"""A single (subject, predicate, object) triple extracted from a claim.
Attributes:
subject: Canonical name of the entity (e.g., "Vladimir Putin").
predicate: Short relation name (e.g., "is_president_of").
obj: Target value of the relation. Named ``obj`` instead of ``object``
to avoid shadowing the Python builtin in callers.
confidence: LLM's certainty about this extraction, 0.0-1.0.
"""
subject: str
predicate: str
obj: str
confidence: float
def to_dict(self) -> dict[str, object]:
"""Serialize for JSONB storage in brain_fact_status."""
return {
"subject": self.subject,
"predicate": self.predicate,
"object": self.obj,
"confidence": self.confidence,
}
@dataclass(slots=True, frozen=True)
class ClaimVolatility:
"""Output of the volatility classifier — drives all caching decisions.
Attributes:
volatility: One of "volatile", "evolving", "stable".
topic_codes: List of topic identifiers this claim touches.
entity_bindings: Subject-predicate-object triples to track in
brain_fact_status for temporal versioning.
estimated_validity_hours: How many hours the cached verdict can be
trusted (already clamped by per-tier caps and floors).
time_sensitive: True if the claim has relative time markers.
reasoning: One-line LLM explanation of the volatility decision.
error: Populated only when classification fell back to defaults.
"""
volatility: Literal["volatile", "evolving", "stable"]
topic_codes: list[str]
entity_bindings: list[EntityBinding]
estimated_validity_hours: int
time_sensitive: bool
reasoning: str
error: str | None = None
@property
def degraded(self) -> bool:
"""True if classification fell back to defaults (LLM failed)."""
return self.error is not None
def entity_bindings_jsonb(self) -> list[dict[str, object]]:
"""Serialize entity_bindings for JSONB storage."""
return [b.to_dict() for b in self.entity_bindings]
def _conservative_default(error_msg: str) -> ClaimVolatility:
"""Build a safe-default classification on LLM failure.
The auditor will pick this up on its next sweep (since volatility is
"evolving", not "stable") and may correct it.
"""
return ClaimVolatility(
volatility=DEFAULT_VOLATILITY,
topic_codes=[],
entity_bindings=[],
estimated_validity_hours=DEFAULT_VALIDITY_HOURS,
time_sensitive=False,
reasoning="classifier_fallback",
error=error_msg,
)
def _clamp_validity_hours(raw: int, volatility: str) -> int:
"""Apply hard caps + floors per volatility tier."""
cap = HARD_CAPS_HOURS.get(volatility, DEFAULT_VALIDITY_HOURS)
floor = HARD_FLOORS_HOURS.get(volatility, 1)
try:
value = int(raw)
except (TypeError, ValueError):
value = DEFAULT_VALIDITY_HOURS
return max(floor, min(value, cap))
def _parse_response(data: object) -> ClaimVolatility | None:
"""Validate the LLM JSON response. Returns None on bad shape.
Strictly checks volatility label, list types, and binding structure.
Silently drops malformed entity_bindings rather than rejecting the whole
response.
"""
if not isinstance(data, dict):
return None
raw_vol = (data.get("volatility") or "").strip().lower()
if raw_vol not in ALLOWED_VOLATILITY:
return None
topic_codes_raw = data.get("topic_codes") or []
if not isinstance(topic_codes_raw, list):
return None
topic_codes = [
str(t).strip() for t in topic_codes_raw if str(t).strip()
]
bindings_raw = data.get("entity_bindings") or []
if not isinstance(bindings_raw, list):
return None
bindings: list[EntityBinding] = []
for item in bindings_raw:
if not isinstance(item, dict):
continue
subj = str(item.get("subject", "")).strip()
pred = str(item.get("predicate", "")).strip()
obj_ = str(item.get("object", "")).strip()
try:
conf = float(item.get("confidence", 0.5))
except (TypeError, ValueError):
conf = 0.5
if not (subj and pred and obj_):
continue
bindings.append(
EntityBinding(
subject=subj,
predicate=pred,
obj=obj_,
confidence=max(0.0, min(1.0, conf)),
)
)
validity = _clamp_validity_hours(
data.get("estimated_validity_hours", DEFAULT_VALIDITY_HOURS),
raw_vol,
)
time_sensitive = bool(data.get("time_sensitive", False))
reasoning = str(data.get("reasoning", "")).strip()[:200]
return ClaimVolatility(
volatility=raw_vol, # type: ignore[arg-type]
topic_codes=topic_codes,
entity_bindings=bindings,
estimated_validity_hours=validity,
time_sensitive=time_sensitive,
reasoning=reasoning,
)
async def classify_claim_volatility(
llm: LlmClient,
*,
claim: str,
current_date: datetime | None = None,
) -> ClaimVolatility:
"""Classify a claim's temporal characteristics with one LLM call.
Always returns a ClaimVolatility. On any failure (timeout, LLM error,
bad JSON), returns a conservative default with ``error`` populated so the
caller can log telemetry but still proceed with the cache write.
Args:
llm: Configured LLM client. Uses LlmRole.REASONING internally.
claim: The claim text to classify. Truncated at ``MAX_CLAIM_CHARS``.
current_date: The "now" reference for the classifier (used for
relative time resolution). Defaults to UTC now.
Returns:
ClaimVolatility with the parsed classification, or a conservative
fallback (volatility="evolving", validity=168h) on failure.
"""
if not claim or not claim.strip():
return _conservative_default("empty_claim")
truncated = claim.strip()[:MAX_CLAIM_CHARS]
today = (current_date or datetime.now(tz=timezone.utc)).date().isoformat()
prompt = (
_load_prompt()
.replace("{current_date}", today)
.replace("{claim}", truncated)
)
try:
result, _usage = await asyncio.wait_for(
llm.chat_json(
role=LlmRole.REASONING,
system=(
"You are a temporal volatility classifier. Respond with "
"strictly valid JSON only, no commentary."
),
user=prompt,
max_tokens=600,
temperature=0.0,
),
timeout=PER_CALL_TIMEOUT_S,
)
except asyncio.TimeoutError:
log.warning("classifier_timeout", claim_preview=truncated[:80])
return _conservative_default("timeout")
except LlmError as e:
log.warning("classifier_llm_error", error=str(e)[:200])
return _conservative_default(f"llm:{e}")
except Exception as e: # noqa: BLE001
log.warning(
"classifier_unexpected_error",
error=f"{type(e).__name__}:{e}",
)
return _conservative_default(f"{type(e).__name__}:{e}")
parsed = _parse_response(result)
if parsed is None:
log.warning(
"classifier_bad_response_shape",
got=type(result).__name__,
)
return _conservative_default("bad_response_shape")
# D1 — apply admin-configured topic overrides on top of LLM judgment.
# Lazy import to avoid a circular dependency if topic_volatility ever
# grows to import classifier types.
try:
from brain_api.services.topic_volatility import (
get_topic_overrides,
reconcile_with_classifier,
)
overrides = await get_topic_overrides()
eff_vol, eff_ttl = reconcile_with_classifier(
classifier_volatility=parsed.volatility,
classifier_validity_hours=parsed.estimated_validity_hours,
classifier_topics=parsed.topic_codes,
overrides=overrides,
)
if eff_vol != parsed.volatility or eff_ttl != parsed.estimated_validity_hours:
log.info(
"classifier_admin_override",
llm_volatility=parsed.volatility,
llm_ttl=parsed.estimated_validity_hours,
final_volatility=eff_vol,
final_ttl=eff_ttl,
topics=parsed.topic_codes,
)
# Re-clamp the final TTL against the per-tier hard caps.
eff_ttl = _clamp_validity_hours(eff_ttl, eff_vol)
parsed = ClaimVolatility(
volatility=eff_vol, # type: ignore[arg-type]
topic_codes=parsed.topic_codes,
entity_bindings=parsed.entity_bindings,
estimated_validity_hours=eff_ttl,
time_sensitive=parsed.time_sensitive,
reasoning=parsed.reasoning,
)
except Exception as e: # noqa: BLE001
# Override layer is best-effort — never block on it.
log.debug(
"classifier_override_skipped",
error=f"{type(e).__name__}:{e}",
)
log.info(
"classifier_ok",
volatility=parsed.volatility,
topics=parsed.topic_codes,
validity_hours=parsed.estimated_validity_hours,
bindings=len(parsed.entity_bindings),
)
return parsed

View file

@ -0,0 +1,758 @@
"""Fact Status — Pilon 11 of the cache freshness defense.
Versioned knowledge layer for entity-predicate-object triples extracted from
claims. Each fact has:
- a current truth value (TRUE / FALSE / NULL=unknown), in brain_fact_status
- a chronological history of (truth, valid_from, valid_to) windows in
brain_fact_version
When the world changes (a president loses an election, an official dies, a
ceasefire is signed), the active fact_version gets ``valid_to = now()`` and a
new version opens with the new truth value. The cache invalidation pipeline
queries this table at lookup time to decide whether any cached verdict
depends on a fact whose current truth no longer matches what the cache
assumed.
Population:
- extractor pipeline (services/ingest.py background task) upsert at
ingestion of new claim atoms, with truth=NULL until verified
- classifier (services/classifier.py) returns entity_bindings for every
classified claim these are upserted lazily on first cache write
- moderator override (admin endpoint) ``moderator_locked=true`` prevents
the auditor from reverting the moderator's decision
- breaking news watcher close current version + open new one with the
fresh truth value derived from the breaking story
Read by:
- cache lookups (services/cache_judge.py) if any binding is known-FALSE,
treat cache as INVALIDATE without running NLI
- admin dashboard (Phase D2) fact browser with timeline
This module is pure persistence + canonicalization. Truth detection (the
"is X currently true?" decision) lives in services/cache_judge.py + the
auditor cron they call upsert_fact_truth here when they have an answer.
"""
from __future__ import annotations
import hashlib
import json
import unicodedata
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any, Literal
from brain_api.db import db
from brain_api.services.classifier import EntityBinding
from shared.logging import get_logger
log = get_logger(__name__)
# Default re-check intervals per volatility tier (hours). Auditor uses these
# to schedule next_check_at when no per-fact override is set.
DEFAULT_CHECK_INTERVAL_HOURS: dict[str, int] = {
"volatile": 6,
"evolving": 168, # 7 days
"stable": 2160, # 90 days
}
CreatedBy = Literal[
"auto",
"moderator",
"breaking_news_watcher",
"auditor",
"extractor",
]
# --------------------------------------------------------------------- DTOs
@dataclass(slots=True)
class FactRecord:
"""Current state of a fact (one row in brain_fact_status)."""
fact_id: int
subject: str
predicate: str
obj: str
canonical_form: str
canonical_form_hash: str
current_truth: bool | None
current_version_id: int | None
current_confidence: float | None
last_verified_at: datetime | None
last_evidence_urls: list[str]
volatility: str | None
topic_codes: list[str]
next_check_at: datetime
check_interval_hours: int
moderator_locked: bool
moderator_user_id: str | None
moderator_notes: str | None
created_at: datetime
updated_at: datetime
@dataclass(slots=True)
class FactVersion:
"""One historical version of a fact (one row in brain_fact_version)."""
version_id: int
fact_id: int
truth_value: bool
confidence: float | None
valid_from: datetime
valid_to: datetime | None
source_atom_ids: list[str]
evidence_urls: list[str]
llm_reasoning: str | None
created_by: str
moderator_user_id: str | None
notes: str | None
created_at: datetime
# --------------------------------------------------------------- canonicalization
def _normalize_token(s: str) -> str:
"""Strip diacritics + lowercase + collapse whitespace.
Matches verification_cache.normalize_claim style so subject "România" and
"Romania" hash to the same fact.
"""
s = unicodedata.normalize("NFKD", s)
s = "".join(c for c in s if not unicodedata.combining(c))
s = s.lower().strip()
return " ".join(s.split())
def canonicalize_triple(subject: str, predicate: str, obj: str) -> str:
"""Build a canonical "subject predicate object" string.
Predicate is normalized to ``snake_case`` (already conventional in the
classifier prompt). Subject and object are lowercased, diacritic-stripped,
whitespace-collapsed.
"""
s = _normalize_token(subject)
p = _normalize_token(predicate).replace(" ", "_")
o = _normalize_token(obj)
return f"{s} {p} {o}"
def hash_canonical(canonical_form: str) -> str:
"""sha256[:32] of the canonical form. Matches the UNIQUE constraint width."""
return hashlib.sha256(canonical_form.encode("utf-8")).hexdigest()[:32]
# --------------------------------------------------------------------- writes
async def register_facts_from_bindings(
bindings: list[EntityBinding],
*,
volatility: str | None = None,
topic_codes: list[str] | None = None,
source_atom_id: str | None = None,
) -> list[int]:
"""Upsert fact_status rows from classifier-extracted bindings.
No truth value is asserted here bindings are recorded with
``current_truth=NULL`` until something verifies them (the auditor, a
breaking-news event, or a moderator). This is the lazy-registration
path called from the cache write hooks.
Returns:
List of fact_ids touched (one per binding, in the same order).
"""
if not bindings:
return []
if not db.pool:
raise RuntimeError("brain_db not connected")
interval_hours = DEFAULT_CHECK_INTERVAL_HOURS.get(volatility or "evolving", 168)
next_check = datetime.now(tz=timezone.utc) + timedelta(hours=interval_hours)
topics = topic_codes or []
sql = """
INSERT INTO brain_fact_status (
subject, predicate, object, canonical_form, canonical_form_hash,
volatility, topic_codes, next_check_at, check_interval_hours
)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)
ON CONFLICT (canonical_form_hash) DO UPDATE SET
topic_codes = (
SELECT ARRAY(
SELECT DISTINCT t FROM unnest(
brain_fact_status.topic_codes || EXCLUDED.topic_codes
) AS t
)
),
volatility = COALESCE(EXCLUDED.volatility, brain_fact_status.volatility),
check_interval_hours = LEAST(
brain_fact_status.check_interval_hours,
EXCLUDED.check_interval_hours
),
updated_at = now()
RETURNING fact_id
"""
fact_ids: list[int] = []
async with db.pool.acquire() as conn:
async with conn.transaction():
for b in bindings:
canonical = canonicalize_triple(b.subject, b.predicate, b.obj)
ch = hash_canonical(canonical)
row = await conn.fetchrow(
sql,
b.subject,
b.predicate,
b.obj,
canonical,
ch,
volatility,
topics,
next_check,
interval_hours,
)
if row:
fact_ids.append(row["fact_id"])
if source_atom_id and fact_ids:
log.debug(
"fact_status_registered",
atom_id=source_atom_id,
fact_count=len(fact_ids),
)
return fact_ids
async def assert_fact_truth(
*,
canonical_form_hash: str,
truth_value: bool,
confidence: float | None,
evidence_urls: list[str],
source_atom_ids: list[str] | None = None,
llm_reasoning: str | None = None,
created_by: CreatedBy = "auto",
moderator_user_id: str | None = None,
notes: str | None = None,
) -> tuple[FactRecord, bool] | None:
"""Set a fact's current truth value, opening a new version if it changed.
If the new truth_value matches the current truth (same boolean), the
existing version is touched (its evidence list is augmented) but no new
version is opened. If the value differs (or the fact had no truth yet),
the active version gets ``valid_to=now()`` and a new version opens.
Skipped silently if ``moderator_locked`` is set on the row moderator
overrides win until explicitly unlocked.
Args:
canonical_form_hash: The hash from ``hash_canonical``.
truth_value: TRUE or FALSE; pass through ``assert_fact_unknown`` if
you want to clear back to NULL.
confidence: 0-100 LLM confidence (or moderator confidence).
evidence_urls: Sources backing this assertion.
source_atom_ids: Atomic atom IDs that contributed to this assertion.
llm_reasoning: One-line LLM justification.
created_by: Provenance tag for the version.
moderator_user_id: Required if ``created_by='moderator'``.
notes: Free-form notes (especially useful for moderator overrides).
Returns:
``(FactRecord, version_changed)`` version_changed=True when a new
version was opened. Returns None if the fact is moderator_locked
and the caller is not a moderator.
"""
if not db.pool:
raise RuntimeError("brain_db not connected")
now = datetime.now(tz=timezone.utc)
async with db.pool.acquire() as conn:
async with conn.transaction():
# 1. Lock-and-load the fact row.
row = await conn.fetchrow(
"""
SELECT * FROM brain_fact_status
WHERE canonical_form_hash = $1
FOR UPDATE
""",
canonical_form_hash,
)
if row is None:
log.warning(
"assert_fact_truth_unknown_fact",
canonical_form_hash=canonical_form_hash,
)
return None
if row["moderator_locked"] and created_by != "moderator":
log.info(
"assert_fact_truth_locked",
fact_id=row["fact_id"],
canonical=row["canonical_form"][:80],
)
return _row_to_fact(row), False
old_truth = row["current_truth"]
value_changed = old_truth is None or bool(old_truth) != truth_value
new_version_id: int | None = None
if value_changed:
# Close the active version (if any).
await conn.execute(
"""
UPDATE brain_fact_version
SET valid_to = $1
WHERE fact_id = $2 AND valid_to IS NULL
""",
now,
row["fact_id"],
)
# Open a new version.
inserted = await conn.fetchrow(
"""
INSERT INTO brain_fact_version (
fact_id, truth_value, confidence,
valid_from, valid_to,
source_atom_ids, evidence_urls, llm_reasoning,
created_by, moderator_user_id, notes
)
VALUES ($1, $2, $3, $4, NULL, $5, $6::jsonb, $7, $8, $9, $10)
RETURNING version_id
""",
row["fact_id"],
truth_value,
confidence,
now,
source_atom_ids or [],
json.dumps(list(evidence_urls)),
llm_reasoning,
created_by,
moderator_user_id,
notes,
)
new_version_id = inserted["version_id"] if inserted else None
# 2. Update brain_fact_status (always — even on no-change we
# bump last_verified_at and merge evidence URLs).
interval_h = row["check_interval_hours"]
next_check = now + timedelta(hours=interval_h)
updated = await conn.fetchrow(
"""
UPDATE brain_fact_status
SET current_truth = $2,
current_version_id = COALESCE($3, current_version_id),
current_confidence = $4,
last_verified_at = $5,
last_evidence_urls = $6::jsonb,
next_check_at = $7,
moderator_user_id = COALESCE($8, moderator_user_id),
moderator_notes = COALESCE($9, moderator_notes),
updated_at = now()
WHERE fact_id = $1
RETURNING *
""",
row["fact_id"],
truth_value,
new_version_id,
confidence,
now,
json.dumps(list(evidence_urls)),
next_check,
moderator_user_id,
notes,
)
# 3. Audit log.
await conn.execute(
"""
INSERT INTO brain_audit_log (action, target_table, target_id, actor, payload)
VALUES ($1, 'brain_fact_status', $2, $3, $4::jsonb)
""",
"fact_truth_set" if not value_changed else "fact_truth_changed",
str(row["fact_id"]),
created_by if created_by != "moderator" else (moderator_user_id or "moderator"),
json.dumps({
"old_truth": old_truth,
"new_truth": truth_value,
"confidence": confidence,
"evidence_count": len(evidence_urls),
}),
)
return _row_to_fact(updated), value_changed # type: ignore[arg-type]
async def lock_fact(
*,
canonical_form_hash: str,
moderator_user_id: str,
moderator_notes: str | None = None,
) -> FactRecord | None:
"""Moderator override: prevent auditor from changing this fact.
Use when a human has decided the truth and machine judgment is unreliable
for the topic.
"""
if not db.pool:
raise RuntimeError("brain_db not connected")
async with db.pool.acquire() as conn:
row = await conn.fetchrow(
"""
UPDATE brain_fact_status
SET moderator_locked = true,
moderator_user_id = $2,
moderator_notes = COALESCE($3, moderator_notes),
updated_at = now()
WHERE canonical_form_hash = $1
RETURNING *
""",
canonical_form_hash,
moderator_user_id,
moderator_notes,
)
return _row_to_fact(row) if row else None
async def unlock_fact(
*, canonical_form_hash: str, moderator_user_id: str
) -> FactRecord | None:
"""Re-enable auditor updates on a previously locked fact."""
if not db.pool:
raise RuntimeError("brain_db not connected")
async with db.pool.acquire() as conn:
row = await conn.fetchrow(
"""
UPDATE brain_fact_status
SET moderator_locked = false,
moderator_user_id = $2,
updated_at = now()
WHERE canonical_form_hash = $1
RETURNING *
""",
canonical_form_hash,
moderator_user_id,
)
return _row_to_fact(row) if row else None
# --------------------------------------------------------------------- reads
async def check_fact_validity(
bindings: list[EntityBinding],
) -> dict[str, bool | None]:
"""For each binding, return current_truth from brain_fact_status.
Returns:
Dict keyed by canonical_form_hash. Values:
- True fact is currently TRUE (cache-aligned if cache assumed TRUE)
- False fact is currently FALSE (cache-aligned if cache assumed FALSE)
- None unknown / not registered yet
Caller (typically gather.py) compares the cache's assumption against this
dict and invalidates if any binding flipped against the cache.
"""
if not bindings:
return {}
if not db.pool:
raise RuntimeError("brain_db not connected")
hashes = [hash_canonical(canonicalize_triple(b.subject, b.predicate, b.obj)) for b in bindings]
sql = """
SELECT canonical_form_hash, current_truth
FROM brain_fact_status
WHERE canonical_form_hash = ANY($1::text[])
"""
out: dict[str, bool | None] = {h: None for h in hashes}
async with db.pool.acquire() as conn:
rows = await conn.fetch(sql, hashes)
for r in rows:
out[r["canonical_form_hash"]] = r["current_truth"]
return out
async def get_fact(canonical_form_hash: str) -> FactRecord | None:
"""Fetch one fact_status row by hash."""
if not db.pool:
raise RuntimeError("brain_db not connected")
async with db.pool.acquire() as conn:
row = await conn.fetchrow(
"SELECT * FROM brain_fact_status WHERE canonical_form_hash = $1",
canonical_form_hash,
)
return _row_to_fact(row) if row else None
async def get_fact_by_id(fact_id: int) -> FactRecord | None:
if not db.pool:
raise RuntimeError("brain_db not connected")
async with db.pool.acquire() as conn:
row = await conn.fetchrow(
"SELECT * FROM brain_fact_status WHERE fact_id = $1",
fact_id,
)
return _row_to_fact(row) if row else None
async def list_versions(fact_id: int) -> list[FactVersion]:
"""All versions for one fact, newest first."""
if not db.pool:
raise RuntimeError("brain_db not connected")
async with db.pool.acquire() as conn:
rows = await conn.fetch(
"""
SELECT * FROM brain_fact_version
WHERE fact_id = $1
ORDER BY valid_from DESC
""",
fact_id,
)
return [_row_to_version(r) for r in rows]
async def list_facts_due_for_check(
*, limit: int = 100, volatility: str | None = None
) -> list[FactRecord]:
"""Auditor entry point: facts whose ``next_check_at`` has passed.
Excludes moderator-locked facts (those are managed by humans).
"""
if not db.pool:
raise RuntimeError("brain_db not connected")
sql = """
SELECT * FROM brain_fact_status
WHERE moderator_locked = false
AND next_check_at <= now()
AND ($1::text IS NULL OR volatility = $1)
ORDER BY next_check_at ASC
LIMIT $2
"""
async with db.pool.acquire() as conn:
rows = await conn.fetch(sql, volatility, limit)
return [_row_to_fact(r) for r in rows]
# --------------------------------------------------------------------- admin
@dataclass(slots=True)
class FactPage:
"""Paginated fact_status listing."""
items: list[FactRecord]
total: int
page: int
page_size: int
async def list_facts_admin(
*,
entity: str | None = None,
predicate: str | None = None,
current_truth: bool | None = None,
locked_only: bool = False,
topic: str | None = None,
page: int = 1,
page_size: int = 25,
) -> FactPage:
"""Admin browser query — paginated, ILIKE search on subject/object.
Args:
entity: ILIKE search on subject OR object (e.g., "putin").
predicate: Exact match on predicate (e.g., "is_president_of").
current_truth: Filter to TRUE / FALSE only when set.
locked_only: Show only moderator-locked facts.
topic: Filter to facts whose topic_codes contain this code.
page, page_size: Pagination (1-based).
"""
if not db.pool:
raise RuntimeError("brain_db not connected")
page = max(1, page)
page_size = max(1, min(100, page_size))
where: list[str] = ["1=1"]
params: list[Any] = []
if entity:
params.append(f"%{entity}%")
where.append(f"(subject ILIKE ${len(params)} OR object ILIKE ${len(params)})")
if predicate:
params.append(predicate)
where.append(f"predicate = ${len(params)}")
if current_truth is not None:
params.append(current_truth)
where.append(f"current_truth = ${len(params)}")
if locked_only:
where.append("moderator_locked = true")
if topic:
params.append([topic])
where.append(f"topic_codes && ${len(params)}::text[]")
where_sql = " AND ".join(where)
async with db.pool.acquire() as conn:
total_row = await conn.fetchrow(
f"SELECT COUNT(*) AS c FROM brain_fact_status WHERE {where_sql}",
*params,
)
total = int(total_row["c"]) if total_row else 0
params.append(page_size)
params.append((page - 1) * page_size)
rows = await conn.fetch(
f"""SELECT * FROM brain_fact_status
WHERE {where_sql}
ORDER BY updated_at DESC
LIMIT ${len(params) - 1} OFFSET ${len(params)}""",
*params,
)
return FactPage(
items=[_row_to_fact(r) for r in rows],
total=total,
page=page,
page_size=page_size,
)
async def list_audit_log(
*,
action: str | None = None,
target_table: str | None = None,
actor: str | None = None,
since: datetime | None = None,
page: int = 1,
page_size: int = 50,
) -> tuple[list[dict[str, Any]], int]:
"""Browse brain_audit_log entries (paginated, newest first).
Returns ``(items, total)``.
"""
if not db.pool:
raise RuntimeError("brain_db not connected")
page = max(1, page)
page_size = max(1, min(200, page_size))
where: list[str] = ["1=1"]
params: list[Any] = []
if action:
params.append(f"{action}%")
where.append(f"action ILIKE ${len(params)}")
if target_table:
params.append(target_table)
where.append(f"target_table = ${len(params)}")
if actor:
params.append(f"%{actor}%")
where.append(f"actor ILIKE ${len(params)}")
if since is not None:
params.append(since)
where.append(f"created_at >= ${len(params)}")
where_sql = " AND ".join(where)
async with db.pool.acquire() as conn:
total_row = await conn.fetchrow(
f"SELECT COUNT(*) AS c FROM brain_audit_log WHERE {where_sql}",
*params,
)
total = int(total_row["c"]) if total_row else 0
params.append(page_size)
params.append((page - 1) * page_size)
rows = await conn.fetch(
f"""SELECT log_id, action, target_table, target_id, actor,
payload, created_at
FROM brain_audit_log
WHERE {where_sql}
ORDER BY created_at DESC
LIMIT ${len(params) - 1} OFFSET ${len(params)}""",
*params,
)
items: list[dict[str, Any]] = []
for r in rows:
payload = r["payload"]
if isinstance(payload, str):
try:
payload = json.loads(payload)
except (TypeError, ValueError):
payload = {}
items.append({
"log_id": r["log_id"],
"action": r["action"],
"target_table": r["target_table"],
"target_id": r["target_id"],
"actor": r["actor"],
"payload": payload,
"created_at": r["created_at"],
})
return items, total
# --------------------------------------------------------------------- helpers
def _row_to_fact(row: Any) -> FactRecord:
last_evidence = row["last_evidence_urls"]
if isinstance(last_evidence, str):
last_evidence = json.loads(last_evidence)
return FactRecord(
fact_id=row["fact_id"],
subject=row["subject"],
predicate=row["predicate"],
obj=row["object"],
canonical_form=row["canonical_form"],
canonical_form_hash=row["canonical_form_hash"],
current_truth=row["current_truth"],
current_version_id=row["current_version_id"],
current_confidence=(
float(row["current_confidence"])
if row["current_confidence"] is not None
else None
),
last_verified_at=row["last_verified_at"],
last_evidence_urls=list(last_evidence) if last_evidence else [],
volatility=row["volatility"],
topic_codes=list(row["topic_codes"]) if row["topic_codes"] else [],
next_check_at=row["next_check_at"],
check_interval_hours=row["check_interval_hours"],
moderator_locked=row["moderator_locked"],
moderator_user_id=row["moderator_user_id"],
moderator_notes=row["moderator_notes"],
created_at=row["created_at"],
updated_at=row["updated_at"],
)
def _row_to_version(row: Any) -> FactVersion:
ev_urls = row["evidence_urls"]
if isinstance(ev_urls, str):
ev_urls = json.loads(ev_urls)
return FactVersion(
version_id=row["version_id"],
fact_id=row["fact_id"],
truth_value=row["truth_value"],
confidence=(
float(row["confidence"]) if row["confidence"] is not None else None
),
valid_from=row["valid_from"],
valid_to=row["valid_to"],
source_atom_ids=(
list(row["source_atom_ids"]) if row["source_atom_ids"] else []
),
evidence_urls=list(ev_urls) if ev_urls else [],
llm_reasoning=row["llm_reasoning"],
created_by=row["created_by"],
moderator_user_id=row["moderator_user_id"],
notes=row["notes"],
created_at=row["created_at"],
)

View file

@ -0,0 +1,62 @@
"""POST /v1/fetch — look up stored atoms by URL and return their content.
Unlike the web module's /v1/fetch which fetches URLs live, ours just checks
if we already have an atom for each URL. URLs we don't have go into
`failed_urls` with reason "not_in_brain" Didi's backend can then fall back
to the live web module for those.
"""
from __future__ import annotations
import asyncio
import time
import uuid
from brain_api.schemas import (
BrainMeta,
FailedUrl,
FetchedPage,
FetchRequest,
FetchResponse,
)
from brain_api.services.mapping import doc_to_fetched_page
from shared.atomic_api import AtomicClient
from shared.logging import get_logger
log = get_logger(__name__)
async def fetch(req: FetchRequest, *, atomic: AtomicClient) -> FetchResponse:
t0 = time.perf_counter()
request_id = str(uuid.uuid4())
tasks = [atomic.get_atom_by_source_url(u) for u in req.urls]
results = await asyncio.gather(*tasks, return_exceptions=True)
pages: list[FetchedPage] = []
failed: list[FailedUrl] = []
for url, result in zip(req.urls, results, strict=True):
if isinstance(result, Exception):
failed.append(FailedUrl(url=url, error=f"brain_lookup_error: {result}"))
continue
if result is None:
failed.append(FailedUrl(url=url, error="not_in_brain"))
continue
pages.append(doc_to_fetched_page(result, include_html=req.include_html))
total_ms = round((time.perf_counter() - t0) * 1000, 1)
return FetchResponse(
request_id=request_id,
pages=pages,
total_fetched=len(pages),
total_failed=len(failed),
execution_time_ms=total_ms,
failed_urls=failed,
brain_meta=BrainMeta(
cache_status="HIT" if pages else "MISS",
api_version="v1",
implementation="didibrain",
evidence_sources=len(pages),
total_claim_atoms_matched=0,
),
)

View file

@ -0,0 +1,665 @@
"""POST /v1/gather — the main claim-to-evidence pipeline.
Flow:
1. semantic search with Type/Claim filter, pull top-K candidates
2. rerank with BGE cross-encoder against the input claim (precision)
3. group hits by parent document URL
4. fetch full parent doc + top claim atom bodies in parallel
5. emit one EvidenceItem per parent doc, sorted by best rerank score
6. shape the full GatherResponse with stages, stats, context
"""
from __future__ import annotations
import asyncio
import math
import time
import uuid
from datetime import datetime, timezone
from brain_api.schemas import (
BrainMeta,
EvidenceStats,
GatherRequest,
GatherResponse,
SearchContext,
SearchResultItem,
StageRecord,
)
from brain_api.services.mapping import (
detect_language_simple,
doc_to_search_result,
evidence_from_parent,
group_hits_by_parent,
parent_url_of,
parse_claim_atom_body,
)
from brain_api.services.nli import classify_batch
from brain_api.services import verification_cache as vcache
from brain_api.services import fact_status as fact_svc
from brain_api.services.classifier import EntityBinding
from shared.atomic_api import AtomicClient, SearchHit
from shared.embedding_client import EmbeddingClient
from shared.llm_client import LlmClient
from shared.logging import get_logger
from shared.taxonomy import TagResolver
log = get_logger(__name__)
FIRST_STAGE_LIMIT = 50 # how many claim hits to pull before reranking
SIMILARITY_FLOOR = 0.20 # below this we don't even consider a hit
RERANK_TOP_N = 15
# Phase B4 — recency boost configuration. Half-life is in days.
# Output of _combined_score = (1 - w_recency) * rerank + w_recency * recency,
# where recency = exp(-age_days / half_life). Stable claims skip the recency
# pass entirely; volatile claims weight recency heavily.
RECENCY_PROFILES: dict[str, tuple[float, float]] = {
# volatility → (w_recency, half_life_days)
"volatile": (0.50, 3.0), # heavy recency, fast decay
"evolving": (0.30, 30.0), # moderate recency, monthly half-life
"stable": (0.00, 9999.0), # ignore recency
}
# Default profile when no hint provided — light recency boost so older
# articles can't fully dominate even on stable topics, without harming them.
RECENCY_DEFAULT: tuple[float, float] = (0.15, 30.0)
# Hard recency filter (days) — only applied when volatility_hint=='volatile'.
DEFAULT_VOLATILE_RECENCY_DAYS = 7
async def gather(
req: GatherRequest,
*,
atomic: AtomicClient,
embed: EmbeddingClient,
llm: LlmClient,
resolver: TagResolver,
) -> GatherResponse:
t0 = time.perf_counter()
request_id = str(uuid.uuid4())
stages: list[StageRecord] = []
# ---- stage 1: context (very lightweight — no LLM for now) --------------
s1_t0 = time.perf_counter()
context = SearchContext(
primary_country="Global",
secondary_countries=[],
detected_language=req.language_hint or detect_language_simple(req.claim),
search_queries=[req.claim],
)
stages.append(
StageRecord(
stage="context",
success=True,
items_processed=1,
items_failed=0,
duration_ms=round((time.perf_counter() - s1_t0) * 1000, 1),
)
)
# ---- stage 2: first-stage retrieval against Type/Claim ----------------
s2_t0 = time.perf_counter()
type_claim_id = resolver.get("Type/Claim")
try:
# Atomic's /api/search currently does not accept a tag_id filter in the
# body — we filter client-side after the call. Pull enough results.
raw_hits = await atomic.search(
req.claim,
mode="semantic",
limit=FIRST_STAGE_LIMIT * 2,
threshold=SIMILARITY_FLOOR,
)
claim_hits: list[SearchHit] = [
h
for h in raw_hits
if any(t.get("name") == "Claim" for t in h.tags)
][:FIRST_STAGE_LIMIT]
stages.append(
StageRecord(
stage="retrieval",
success=True,
items_processed=len(claim_hits),
items_failed=0,
duration_ms=round((time.perf_counter() - s2_t0) * 1000, 1),
)
)
except Exception as e: # noqa: BLE001
log.error("gather_retrieval_failed", error=str(e))
stages.append(
StageRecord(
stage="retrieval",
success=False,
items_processed=0,
items_failed=1,
duration_ms=round((time.perf_counter() - s2_t0) * 1000, 1),
error=str(e),
)
)
return _empty_response(
request_id=request_id,
claim=req.claim,
context=context,
stages=stages,
started_at=t0,
)
if not claim_hits:
return _empty_response(
request_id=request_id,
claim=req.claim,
context=context,
stages=stages,
started_at=t0,
)
# ---- stage 3: rerank top-N with BGE cross-encoder --------------------
s3_t0 = time.perf_counter()
rerank_scores: dict[str, float] = {}
try:
# We need the CONTENT of each claim atom to rerank it; fetch in parallel.
full_claim_atoms = await _fetch_full_atoms(
atomic,
[h.atom_id for h in claim_hits[:RERANK_TOP_N]],
)
rerank_docs: list[str] = []
rerank_atom_ids: list[str] = []
for h in claim_hits[:RERANK_TOP_N]:
full = full_claim_atoms.get(h.atom_id) or {}
text = (full.get("content") or "")[:2000]
if text:
rerank_docs.append(text)
rerank_atom_ids.append(h.atom_id)
if rerank_docs:
reranked = await embed.rerank(req.claim, rerank_docs)
for r in reranked:
if 0 <= r.index < len(rerank_atom_ids):
rerank_scores[rerank_atom_ids[r.index]] = r.score
stages.append(
StageRecord(
stage="rerank",
success=True,
items_processed=len(rerank_docs),
items_failed=0,
duration_ms=round((time.perf_counter() - s3_t0) * 1000, 1),
)
)
except Exception as e: # noqa: BLE001
log.warning("gather_rerank_failed", error=str(e))
full_claim_atoms = {}
stages.append(
StageRecord(
stage="rerank",
success=False,
items_processed=0,
items_failed=len(claim_hits),
duration_ms=round((time.perf_counter() - s3_t0) * 1000, 1),
error=str(e),
)
)
# ---- stage 4: group by parent doc --------------------------------------
s4_t0 = time.perf_counter()
buckets = group_hits_by_parent(claim_hits, rerank_scores)
# Phase B4 — fetch a wider parent pool so the recency reranker has
# candidates to choose from. We oversample by 2x then trim to max_evidence
# AFTER recency reranking. For volatile claims we also need the broader
# pool so the hard recency filter doesn't leave us empty.
fetch_count = min(len(buckets), max(req.max_evidence * 2, req.max_evidence))
parent_urls = [url for url, _ in buckets][:fetch_count]
parent_atoms = await _fetch_parents(atomic, parent_urls)
# Apply recency reranking + filter, trim to max_evidence.
buckets = _apply_recency(
buckets=buckets,
parent_atoms=parent_atoms,
volatility_hint=req.volatility_hint,
recency_window_days=req.recency_window_days,
max_evidence=req.max_evidence,
)
# Fetch any missing full claim atom bodies (for non-reranked but still emitted)
need_more_claim_atoms = [
h.atom_id
for _, hs in buckets[: req.max_evidence]
for (h, _) in hs
if h.atom_id not in full_claim_atoms
]
if need_more_claim_atoms:
extra = await _fetch_full_atoms(atomic, need_more_claim_atoms)
full_claim_atoms.update(extra)
# ---- stage 5: NLI stance vs query (optional) --------------------------
# For each bucket that will become an evidence item, run NLI on the best
# matching claim atom (the one we surface as `summary`). Parallelized so
# ~10 calls complete in a couple of seconds rather than seconds per call.
nli_by_atom_id: dict[str, tuple[str, float, str | None]] = {}
if req.run_nli:
s5_t0 = time.perf_counter()
best_per_bucket: list[tuple[str, str]] = []
for parent_url, hits in buckets[: req.max_evidence]:
if parent_url not in parent_atoms:
continue
best_hit, _ = hits[0]
full = full_claim_atoms.get(best_hit.atom_id) or {}
claim_text, _stance, _parent_id = parse_claim_atom_body(
full.get("content") or ""
)
if claim_text:
best_per_bucket.append((best_hit.atom_id, claim_text))
if best_per_bucket:
nli_results = await classify_batch(
llm,
claim=req.claim,
evidence_texts=[text for _, text in best_per_bucket],
)
for (atom_id, _text), result in zip(
best_per_bucket, nli_results, strict=True
):
nli_by_atom_id[atom_id] = (
result.label,
result.confidence,
result.error,
)
failed = sum(1 for v in nli_by_atom_id.values() if v[2] is not None)
stages.append(
StageRecord(
stage="nli",
success=True,
items_processed=len(nli_by_atom_id),
items_failed=failed,
duration_ms=round((time.perf_counter() - s5_t0) * 1000, 1),
)
)
# ---- stage 6: build evidence list with NLI attached -------------------
evidence_items = []
for parent_url, hits in buckets[: req.max_evidence]:
parent_atom = parent_atoms.get(parent_url)
if not parent_atom:
continue
item = evidence_from_parent(
parent_atom=parent_atom,
claim_hits=hits,
parent_full_atoms=full_claim_atoms,
include_full_text=req.include_full_text,
nli_by_atom_id=nli_by_atom_id or None,
)
evidence_items.append(item)
stages.append(
StageRecord(
stage="evidence",
success=True,
items_processed=len(evidence_items),
items_failed=0,
duration_ms=round((time.perf_counter() - s4_t0) * 1000, 1),
)
)
# ---- shape response --------------------------------------------------
total_ms = round((time.perf_counter() - t0) * 1000, 1)
search_results: list[SearchResultItem] = []
for i, item in enumerate(evidence_items, 1):
search_results.append(
SearchResultItem(
query=req.claim,
url=item.url,
title=item.title,
snippet=item.snippet or "",
rank=i,
site=item.publisher,
published_at=item.published_at,
)
)
stats = EvidenceStats(
input_items=len(claim_hits),
after_dedup=len(buckets),
output_items=len(evidence_items),
duplicates_removed=max(0, len(claim_hits) - len(buckets)),
tokens_used=0,
)
# Cache status reflects BOTH presence and quality:
# - MISS if no evidence at all, or the best match is weak (< 0.3 rerank)
# - PARTIAL if we have evidence but best rerank is between 0.3 and 0.6
# - HIT when the brain actually has a strong, direct match (>= 0.6)
if not evidence_items:
cache_status = "MISS"
else:
top_relevance = max(
(e.relevance_score for e in evidence_items), default=0.0
)
if top_relevance < 0.30:
cache_status = "MISS"
elif top_relevance < 0.60:
cache_status = "PARTIAL"
else:
cache_status = "HIT"
brain_meta = BrainMeta(
cache_status=cache_status,
api_version="v1",
implementation="didibrain",
evidence_sources=len({e.url for e in evidence_items}),
total_claim_atoms_matched=len(claim_hits),
)
# ---- optional verification cache lookup ---------------------------------
# Lookup key is (claim_hash, tier) only. The evidence URLs the cache was
# written for may differ from `evidence_items` (different runs, different
# corpora). We surface the original URLs in metadata so backend can
# decide whether the cached verification applies to its current view.
if req.include_verification and req.tier:
try:
entry = await vcache.lookup(claim=req.claim, tier=req.tier)
except Exception as e: # noqa: BLE001
log.warning("verification_lookup_failed", error=f"{type(e).__name__}: {e}")
entry = None
staleness = vcache.decide_freshness(
entry,
current_prompt_hash=req.prompt_hash,
current_framework_version=req.framework_version,
)
# Pilon 11 — fact-status check. If the cache is otherwise fresh but
# one of its bound facts has flipped (e.g., "X is president of Y"
# was TRUE when cached, but brain_fact_status now says FALSE), demote
# to stale_evidence so the caller recomputes. We only run this when
# the cache would otherwise be served (fresh / stale_framework — the
# other states already force recompute).
if entry is not None and staleness in ("fresh", "stale_framework"):
try:
flipped = await _detect_flipped_facts(entry)
if flipped:
log.info(
"verification_facts_flipped",
flipped_count=len(flipped),
was=staleness,
)
staleness = "stale_evidence"
brain_meta.verification_facts_invalidated = flipped
except Exception as e: # noqa: BLE001
log.warning(
"verification_fact_check_failed",
error=f"{type(e).__name__}: {e}",
)
brain_meta.verification_staleness = staleness
if entry is not None:
brain_meta.verification_model = entry.model
brain_meta.verification_tier = entry.tier
brain_meta.verification_prompt_hash = entry.prompt_hash
brain_meta.verification_framework_version = entry.framework_version
brain_meta.verification_cached_at = entry.updated_at
brain_meta.verification_expires_at = entry.expires_at
brain_meta.verification_evidence_urls = entry.evidence_urls
brain_meta.verification_evidence_hash = entry.evidence_hash
if staleness == "fresh":
brain_meta.verification = entry.verification_processed
elif staleness == "stale_framework":
# Backend can recompute status from raw locally — no LLM call.
brain_meta.verification = entry.verification_raw
# stale_evidence / stale_prompt / miss → don't expose verification
# so caller is forced to recompute.
return GatherResponse(
request_id=request_id,
claim=req.claim,
evidence=evidence_items,
evidence_stats=stats,
search_context=context,
search_results=search_results,
stages=stages,
total_urls_found=len(claim_hits),
total_pages_fetched=len(evidence_items),
total_evidence_items=len(evidence_items),
execution_time_ms=total_ms,
brain_meta=brain_meta,
)
# --- helpers ---------------------------------------------------------------
def _published_at_of(parent_atom: dict | None) -> datetime | None:
"""Best-effort published_at extractor for a parent Document atom."""
if not parent_atom:
return None
raw = parent_atom.get("published_at") or parent_atom.get("created_at")
if not raw:
return None
if isinstance(raw, datetime):
return raw if raw.tzinfo else raw.replace(tzinfo=timezone.utc)
try:
s = str(raw).rstrip("Z")
# Tolerate trailing Z by using fromisoformat with tz-aware handling.
dt = datetime.fromisoformat(s)
return dt if dt.tzinfo else dt.replace(tzinfo=timezone.utc)
except (TypeError, ValueError):
return None
def _combined_score(
*,
rerank_score: float,
age_days: float | None,
volatility_hint: str | None,
) -> float:
"""Blend rerank with age-based recency. Stable / no-published_at = pure rerank.
For each volatility profile we use:
combined = (1 - w_recency) * rerank + w_recency * exp(-age_days/half_life)
rerank_score is on [0, 1] from the cross-encoder; recency is also [0, 1].
"""
if age_days is None or age_days < 0:
return rerank_score
profile = RECENCY_PROFILES.get(volatility_hint or "", RECENCY_DEFAULT)
w_recency, half_life = profile
if w_recency <= 0:
return rerank_score
recency = math.exp(-age_days / half_life)
return (1.0 - w_recency) * rerank_score + w_recency * recency
def _apply_recency(
*,
buckets: list[tuple[str, list[tuple[SearchHit, float]]]],
parent_atoms: dict[str, dict],
volatility_hint: str | None,
recency_window_days: int | None,
max_evidence: int,
) -> list[tuple[str, list[tuple[SearchHit, float]]]]:
"""Rerank buckets by combined (rerank + recency) score, trim to max_evidence.
For volatility_hint='volatile', also drops parents older than
recency_window_days (default 7) caller will see fewer evidence items
and can fall back to live web search.
Buckets without a fetched parent_atom (i.e., not in parent_atoms) drop
out entirely they were just placeholders for a wider fetch.
"""
now = datetime.now(tz=timezone.utc)
# Hard recency cut for volatile (or any volatility when caller pinned a
# window).
if volatility_hint == "volatile" or recency_window_days is not None:
window = recency_window_days or DEFAULT_VOLATILE_RECENCY_DAYS
def _within_window(parent_url: str) -> bool:
atom = parent_atoms.get(parent_url)
pub = _published_at_of(atom)
if pub is None:
# No publish date → keep only when no hard cut requested
# (volatile defaults to dropping unknowns to be safe).
return volatility_hint != "volatile"
return (now - pub).days <= window
buckets = [(u, hs) for (u, hs) in buckets if _within_window(u)]
# Score each surviving bucket using its top hit's rerank score and the
# parent's age, then re-sort. Buckets without parent_atoms are dropped.
scored: list[tuple[float, str, list[tuple[SearchHit, float]]]] = []
for parent_url, hits in buckets:
atom = parent_atoms.get(parent_url)
if not atom or not hits:
continue
top_rerank = float(hits[0][1])
pub = _published_at_of(atom)
age_days = (now - pub).days if pub else None
score = _combined_score(
rerank_score=top_rerank,
age_days=age_days,
volatility_hint=volatility_hint,
)
scored.append((score, parent_url, hits))
scored.sort(key=lambda t: t[0], reverse=True)
return [(url, hits) for (_score, url, hits) in scored[:max_evidence]]
def _extract_cached_truth(processed: dict) -> bool | None:
"""Map a cached verification_processed dict to a binary truth direction.
DIDI v1 schema uses "status": "TRUE" | "FALSE" | "UV" | "OP" | "MIXED".
Anything other than TRUE/FALSE returns None the cache didn't commit
to a direction so we can't compare it against fact_status.
"""
if not isinstance(processed, dict):
return None
status = processed.get("status")
if isinstance(status, str):
s = status.strip().upper()
if s in ("TRUE", "VERIFIED_TRUE", "VT"):
return True
if s in ("FALSE", "VERIFIED_FALSE", "VF"):
return False
return None
async def _detect_flipped_facts(entry: vcache.CacheEntry) -> list[dict]:
"""Return entity bindings whose current_truth contradicts the cached verdict.
Each returned dict mirrors the binding shape stored on the row, with
extra fields ``cached_assumes`` and ``current_truth`` so the caller
(admin dashboard, downstream backend) can show what changed.
Returns [] if the cache had no bindings, or if no current truth could be
extracted from verification_processed, or if no bound fact disagrees.
"""
if not entry.entity_bindings:
return []
cached_truth = _extract_cached_truth(entry.verification_processed)
if cached_truth is None:
return []
# Reconstruct EntityBinding instances from the JSONB row.
bindings: list[EntityBinding] = []
for raw in entry.entity_bindings:
if not isinstance(raw, dict):
continue
subj = str(raw.get("subject", "")).strip()
pred = str(raw.get("predicate", "")).strip()
obj_ = str(raw.get("object", "")).strip()
try:
conf = float(raw.get("confidence", 0.5))
except (TypeError, ValueError):
conf = 0.5
if subj and pred and obj_:
bindings.append(
EntityBinding(
subject=subj, predicate=pred, obj=obj_, confidence=conf
)
)
if not bindings:
return []
truth_map = await fact_svc.check_fact_validity(bindings)
flipped: list[dict] = []
for b in bindings:
canonical = fact_svc.canonicalize_triple(b.subject, b.predicate, b.obj)
ch = fact_svc.hash_canonical(canonical)
current = truth_map.get(ch)
# We only flag explicit disagreement; current=None means we have no
# opinion (yet) and falls through to the existing freshness checks.
if current is None:
continue
if current != cached_truth:
flipped.append({
"subject": b.subject,
"predicate": b.predicate,
"object": b.obj,
"canonical_form": canonical,
"cached_assumes": cached_truth,
"current_truth": current,
})
return flipped
async def _fetch_full_atoms(
atomic: AtomicClient, atom_ids: list[str]
) -> dict[str, dict]:
"""Parallel get_atom for a list of atom IDs. Missing atoms are dropped."""
if not atom_ids:
return {}
tasks = [atomic.get_atom(a) for a in atom_ids]
results = await asyncio.gather(*tasks, return_exceptions=True)
out: dict[str, dict] = {}
for a, r in zip(atom_ids, results, strict=True):
if isinstance(r, dict):
out[a] = r
return out
async def _fetch_parents(
atomic: AtomicClient, parent_urls: list[str]
) -> dict[str, dict]:
"""Parallel get_atom_by_source_url for parent document URLs."""
if not parent_urls:
return {}
tasks = [atomic.get_atom_by_source_url(u) for u in parent_urls]
results = await asyncio.gather(*tasks, return_exceptions=True)
out: dict[str, dict] = {}
for u, r in zip(parent_urls, results, strict=True):
if isinstance(r, dict):
out[u] = r
return out
def _empty_response(
*,
request_id: str,
claim: str,
context: SearchContext,
stages: list[StageRecord],
started_at: float,
) -> GatherResponse:
return GatherResponse(
request_id=request_id,
claim=claim,
evidence=[],
evidence_stats=EvidenceStats(),
search_context=context,
search_results=[],
stages=stages,
total_urls_found=0,
total_pages_fetched=0,
total_evidence_items=0,
execution_time_ms=round((time.perf_counter() - started_at) * 1000, 1),
brain_meta=BrainMeta(
cache_status="MISS",
api_version="v1",
implementation="didibrain",
evidence_sources=0,
total_claim_atoms_matched=0,
),
)

View file

@ -0,0 +1,268 @@
"""POST /v1/ingest — populate brain from Didi's web-module results.
When Didi's backend calls the live (expensive) web module and gets a fresh
GatherResponse, it can POST the same body here. We:
1. For each evidence item, dedup against existing atoms by canonical URL
2. Build proper Type/Document atoms with inferred tags:
- Credibility from credibility_score bucket
- Language from search_context.detected_language
- Country/Global by default (future: infer from publisher TLD)
- Any `default_tags` provided by the caller
3. Create atoms synchronously (returns quickly)
4. Optionally queue claim extraction in background via FastAPI BackgroundTasks
Response returns counts + the created atom IDs so the caller can correlate.
"""
from __future__ import annotations
import time
import uuid
from typing import Any
from fastapi import BackgroundTasks
from brain_api.schemas import (
EvidenceItem,
IngestRequest,
IngestResponse,
)
from brain_api.services.classifier import (
ClaimVolatility,
classify_claim_volatility,
)
from shared.atomic_api import AtomicApiError, AtomicClient
from shared.llm_client import LlmClient
from shared.logging import get_logger
from shared.taxonomy import TagResolver
log = get_logger(__name__)
# Credibility bucket boundaries mirror mapping.CREDIBILITY_SCORES inverse.
def credibility_score_to_tag_path(score: float) -> str:
if score >= 0.85:
return "Credibility/Tier1"
if score >= 0.60:
return "Credibility/Tier2"
if score >= 0.40:
return "Credibility/Tier3"
if score >= 0.25:
return "Credibility/StateAffiliated"
if score >= 0.01:
return "Credibility/KnownDisinfo"
return "Credibility/Unknown"
def detected_language_to_tag_path(lang: str | None) -> str:
if not lang:
return "Language/EN"
code = lang.strip().upper()[:2]
mapping = {
"RO": "Language/RO",
"EN": "Language/EN",
"RU": "Language/RU",
"UA": "Language/UA",
"FR": "Language/FR",
"DE": "Language/DE",
"ES": "Language/ES",
"IT": "Language/IT",
"PL": "Language/PL",
}
return mapping.get(code, "Language/EN")
def build_tag_ids_for_evidence(
*,
ev: EvidenceItem,
detected_language: str,
default_tags: list[str],
resolver: TagResolver,
) -> list[str]:
paths: list[str] = [
"Type/Document",
"SourceType/MainstreamMedia", # default — callers can override via default_tags
credibility_score_to_tag_path(ev.credibility_score),
detected_language_to_tag_path(detected_language),
"Country/Global",
]
# Append any caller-provided canonical tags, deduped
for p in default_tags:
if p and p not in paths:
paths.append(p)
return resolver.ids_for(paths, ignore_missing=True)
def evidence_to_markdown(ev: EvidenceItem) -> str:
"""Build the markdown body for a Type/Document atom from an EvidenceItem."""
title = ev.title or ev.url
body = ev.full_text or ev.summary or ev.snippet or ""
header_lines = [f"# {title}", ""]
if ev.published_at:
header_lines.append(f"**Published:** {ev.published_at.isoformat()}")
if ev.publisher:
header_lines.append(f"**Publisher:** {ev.publisher}")
if ev.published_at or ev.publisher:
header_lines.append("")
return "\n".join(header_lines) + body.strip() + "\n"
async def _classify_and_register_facts_async(
*, claim: str, llm: LlmClient, source_label: str
) -> None:
"""Classify the claim and register entity bindings in brain_fact_status.
Best-effort: any failure is logged and swallowed. Runs as a background
task triggered by FastAPI's BackgroundTasks queue, after the ingest
response has been returned to the caller.
"""
try:
classification = await classify_claim_volatility(llm, claim=claim)
except Exception as e: # noqa: BLE001
log.warning(
"ingest_classifier_failed",
error=f"{type(e).__name__}:{e}",
)
return
if not classification.entity_bindings:
return
try:
from brain_api.services.fact_status import register_facts_from_bindings
await register_facts_from_bindings(
classification.entity_bindings,
volatility=classification.volatility,
topic_codes=classification.topic_codes,
source_atom_id=source_label,
)
log.info(
"ingest_facts_registered",
source=source_label,
volatility=classification.volatility,
bindings=len(classification.entity_bindings),
)
except Exception as e: # noqa: BLE001
log.warning(
"ingest_fact_registration_failed",
error=f"{type(e).__name__}:{e}",
)
async def ingest(
req: IngestRequest,
*,
atomic: AtomicClient,
resolver: TagResolver,
background: BackgroundTasks | None = None,
llm: LlmClient | None = None,
) -> IngestResponse:
t0 = time.perf_counter()
request_id = str(uuid.uuid4())
accepted = 0
skipped = 0
errors = 0
warnings: list[str] = []
created_ids: list[str] = []
detected_language = "en"
# The caller may pass language hints inside default_tags or a nested context;
# we support both. Fall back to English.
for p in req.default_tags:
if p.startswith("Language/"):
detected_language = p.split("/", 1)[-1]
break
for ev in req.evidence:
if not ev.url:
warnings.append("evidence item missing url — skipped")
continue
# Dedup on canonical URL
try:
existing = await atomic.get_atom_by_source_url(ev.url)
except AtomicApiError as e:
errors += 1
warnings.append(f"dedup check failed for {ev.url}: {e.status}")
continue
if existing:
skipped += 1
continue
content = evidence_to_markdown(ev)
tag_ids = build_tag_ids_for_evidence(
ev=ev,
detected_language=detected_language,
default_tags=req.default_tags,
resolver=resolver,
)
published_iso = ev.published_at.isoformat() if ev.published_at else None
try:
atom = await atomic.create_atom(
content=content,
source_url=ev.url,
tag_ids=tag_ids,
published_at=published_iso,
)
atom_id = atom.get("id")
if atom_id:
created_ids.append(atom_id)
accepted += 1
else:
errors += 1
warnings.append(f"create_atom returned no id for {ev.url}")
except AtomicApiError as e:
errors += 1
warnings.append(f"create_atom failed for {ev.url}: {e.status} {e.body[:120]}")
extraction_queued = False
if req.run_extraction and created_ids and background is not None:
background.add_task(_run_extraction_background, created_ids)
extraction_queued = True
# Pilon 11: classify req.claim once and register entity bindings into
# brain_fact_status. Lazy (current_truth=NULL) until verified.
if (
req.claim
and req.claim.strip()
and llm is not None
and background is not None
):
background.add_task(
_classify_and_register_facts_async,
claim=req.claim,
llm=llm,
source_label=f"ingest:{request_id[:12]}",
)
return IngestResponse(
request_id=request_id,
accepted=accepted,
skipped_duplicate=skipped,
errors=errors,
created_atom_ids=created_ids,
extraction_queued=extraction_queued,
execution_time_ms=round((time.perf_counter() - t0) * 1000, 1),
warnings=warnings[:20],
)
async def _run_extraction_background(atom_ids: list[str]) -> None:
"""Fire-and-forget extraction for newly ingested atoms."""
from extractor.batch import run_batch
log.info("brain_ingest_extraction_start", count=len(atom_ids))
try:
stats = await run_batch(only_atom_ids=set(atom_ids))
log.info(
"brain_ingest_extraction_done",
processed=stats.docs_processed,
claims=stats.claims_created,
failed=stats.docs_failed,
)
except Exception as e: # noqa: BLE001
log.error("brain_ingest_extraction_failed", error=str(e))

View file

@ -0,0 +1,282 @@
"""Cache invalidation service — Pilon 8 of the cache freshness defense.
Mass-invalidates rows in brain_analysis_atom and brain_verification_cache
based on filters (topic, entity, since, content pattern). Used by:
- Admin dashboard "Flush topic" button (manual ops)
- didibrain-breaking-watcher (real-time, when a breaking story affects
a topic or entity)
- Daily auditor (when gold demotion cascades to dependent rows)
Invalidation = set expires_at = now() (soft delete, preserves row for audit).
Gold atoms in brain_analysis_atom are NOT invalidated by topic/entity filters
unless ``invalidate_gold=True`` is passed gold rows reflect human moderator
decisions and shouldn't be flushed by automated breaking news. Operators who
need to flush them must opt in explicitly.
Every invalidation logs a row in brain_audit_log with the filter spec and
counts so the admin dashboard can show recent flush operations.
"""
from __future__ import annotations
import json
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
from brain_api.db import db
from shared.logging import get_logger
log = get_logger(__name__)
@dataclass(slots=True)
class InvalidateFilter:
"""Selection criteria for a mass invalidation operation.
At least one of topic_codes / entity_canonicals / claim_pattern / since
must be non-empty calling with all-empty filters is rejected to avoid
accidental "flush everything".
Attributes:
topic_codes: Match rows whose ``topic_codes && this`` is true (any
overlap). E.g., ``["war", "elections"]``.
entity_canonicals: Match rows whose ``entity_bindings`` JSONB
includes a triple that lower-cases to one of these
``"<subject> <predicate> <object>"`` canonical strings.
Caller is responsible for normalizing (lowercasing, predicate
snake_case'd, etc.) — see ``fact_status.canonicalize_triple``.
claim_pattern: ILIKE pattern on content_preview / claim text.
since: Match rows created or updated AFTER this timestamp.
invalidate_gold: If true, also expire gold rows. Default false.
dry_run: If true, count matches without modifying anything.
"""
topic_codes: list[str] | None = None
entity_canonicals: list[str] | None = None
claim_pattern: str | None = None
since: datetime | None = None
invalidate_gold: bool = False
dry_run: bool = False
def is_empty(self) -> bool:
"""True if no filter criteria are set — caller must reject."""
return not (
self.topic_codes
or self.entity_canonicals
or self.claim_pattern
or self.since
)
@dataclass(slots=True)
class InvalidateResult:
"""Counts and metadata returned by ``invalidate_caches``."""
invalidated_atoms: int
invalidated_vcache: int
dry_run: bool
filters_applied: dict[str, Any]
executed_at: datetime
def _entity_match_clause(idx: int) -> str:
"""JSONB existence clause matching any binding's canonical form.
Caller pre-normalizes inputs to lowercase
``"<subject> <predicate_snake> <object>"`` and passes them as a text[].
PG just concatenates the binding fields with the same shape and
compares no extensions needed.
"""
return (
"EXISTS ("
" SELECT 1 FROM jsonb_array_elements(entity_bindings) AS b "
" WHERE lower(coalesce(b->>'subject','')) || ' ' || "
" lower(replace(coalesce(b->>'predicate',''), ' ', '_')) || ' ' || "
" lower(coalesce(b->>'object','')) "
f" = ANY(${idx}::text[])"
")"
)
def _build_atom_where_clause(
f: InvalidateFilter, params: list[Any]
) -> str:
"""Compose WHERE clause + side-effects on ``params`` for analysis atoms.
Returns a SQL fragment starting with ``WHERE`` (always non-empty since
is_empty() is checked upstream).
"""
clauses: list[str] = ["(expires_at IS NULL OR expires_at > now())"]
if not f.invalidate_gold:
clauses.append("cache_tier <> 'gold'")
if f.topic_codes:
params.append(f.topic_codes)
clauses.append(f"topic_codes && ${len(params)}::text[]")
if f.entity_canonicals:
# Caller-side canonicalization (lowercased "subject predicate object")
# — see fact_status.canonicalize_triple. PG just does string match
# against entity_bindings JSONB without needing unaccent/digest.
params.append(f.entity_canonicals)
clauses.append(_entity_match_clause(len(params)))
if f.claim_pattern:
params.append(f"%{f.claim_pattern}%")
clauses.append(f"content_preview ILIKE ${len(params)}")
if f.since is not None:
params.append(f.since)
clauses.append(f"updated_at >= ${len(params)}")
return "WHERE " + " AND ".join(clauses)
def _build_vcache_where_clause(
f: InvalidateFilter, params: list[Any]
) -> str:
"""Compose WHERE clause for verification cache (no cache_tier here)."""
clauses: list[str] = ["expires_at > now()"]
if f.topic_codes:
params.append(f.topic_codes)
clauses.append(f"topic_codes && ${len(params)}::text[]")
if f.entity_canonicals:
params.append(f.entity_canonicals)
clauses.append(_entity_match_clause(len(params)))
if f.claim_pattern:
params.append(f"%{f.claim_pattern}%")
# vcache stores the claim text only as a hash — match against the
# processed verification payload as a fallback.
clauses.append(
f"verification_processed::text ILIKE ${len(params)}"
)
if f.since is not None:
params.append(f.since)
clauses.append(f"updated_at >= ${len(params)}")
return "WHERE " + " AND ".join(clauses)
async def invalidate_caches(
f: InvalidateFilter,
*,
actor: str = "admin",
reason: str | None = None,
) -> InvalidateResult:
"""Invalidate rows in both cache tables matching the filter.
Args:
f: The selection criteria. Must not be empty (raises ValueError).
actor: Free-form label for the audit log (e.g., 'admin:foo@bar',
'breaking_watcher', 'auditor').
reason: Optional human-readable note for audit trail.
Returns:
InvalidateResult with counts and the filters that were applied.
Raises:
ValueError: If the filter is empty (no criteria set).
RuntimeError: If brain_db is not connected.
"""
if f.is_empty():
raise ValueError(
"invalidate filter is empty — refusing to flush everything"
)
if not db.pool:
raise RuntimeError("brain_db not connected")
now = datetime.now(tz=timezone.utc)
# Count first (always — even non-dry-run runs the count for the audit log).
atom_params: list[Any] = []
atom_where = _build_atom_where_clause(f, atom_params)
vcache_params: list[Any] = []
vcache_where = _build_vcache_where_clause(f, vcache_params)
async with db.pool.acquire() as conn:
# COUNT pre-update so we know how many rows we'll touch.
atom_count_row = await conn.fetchrow(
f"SELECT COUNT(*) AS c FROM brain_analysis_atom {atom_where}",
*atom_params,
)
vcache_count_row = await conn.fetchrow(
f"SELECT COUNT(*) AS c FROM brain_verification_cache {vcache_where}",
*vcache_params,
)
atom_count = int(atom_count_row["c"]) if atom_count_row else 0
vcache_count = int(vcache_count_row["c"]) if vcache_count_row else 0
if not f.dry_run and (atom_count > 0 or vcache_count > 0):
async with conn.transaction():
if atom_count > 0:
await conn.execute(
f"UPDATE brain_analysis_atom SET expires_at = now(), "
f"updated_at = now() {atom_where}",
*atom_params,
)
if vcache_count > 0:
await conn.execute(
f"UPDATE brain_verification_cache SET expires_at = now(), "
f"updated_at = now() {vcache_where}",
*vcache_params,
)
payload = {
"filter": {
"topic_codes": f.topic_codes,
"entity_canonicals_count": (
len(f.entity_canonicals)
if f.entity_canonicals
else 0
),
"claim_pattern": f.claim_pattern,
"since": f.since.isoformat() if f.since else None,
"invalidate_gold": f.invalidate_gold,
},
"counts": {
"atoms": atom_count,
"vcache": vcache_count,
},
"reason": reason,
}
await conn.execute(
"""
INSERT INTO brain_audit_log (action, target_table, target_id, actor, payload)
VALUES ('invalidate', 'multi', 'mass', $1, $2::jsonb)
""",
actor,
json.dumps(payload),
)
log.info(
"cache_invalidated",
atom_count=atom_count,
vcache_count=vcache_count,
dry_run=f.dry_run,
actor=actor,
topic_codes=f.topic_codes,
invalidate_gold=f.invalidate_gold,
)
return InvalidateResult(
invalidated_atoms=atom_count,
invalidated_vcache=vcache_count,
dry_run=f.dry_run,
filters_applied={
"topic_codes": f.topic_codes,
"entity_canonicals_count": (
len(f.entity_canonicals) if f.entity_canonicals else 0
),
"claim_pattern": f.claim_pattern,
"since": f.since.isoformat() if f.since else None,
"invalidate_gold": f.invalidate_gold,
},
executed_at=now,
)

View file

@ -0,0 +1,321 @@
"""Translate DidiBrain atoms into Didi's EvidenceItem / FetchedPage shapes.
The trick here is that the brain stores TWO kinds of atoms:
- Type/Document the full source article (parent)
- Type/Claim an atomic factual claim extracted from a parent
Didi's response shape expects evidence AT THE DOCUMENT LEVEL (url, title,
full_text). So we:
1. Run semantic + rerank at claim level (precision)
2. Group hits by parent document URL
3. Emit one EvidenceItem per distinct parent, with the best-scoring claim
attached as `summary` and supporting data in `brain_meta`
Credibility tags in our taxonomy (Tier1/Tier2/Tier3/StateAffiliated/KnownDisinfo)
map to numeric scores that mirror Didi's web-module output range (0..1).
"""
from __future__ import annotations
import hashlib
import re
from collections import defaultdict
from datetime import datetime, timezone
from typing import Any
from urllib.parse import unquote, urlparse
from brain_api.schemas import (
BrainEvidenceMeta,
EvidenceItem,
FetchedPage,
Provenance,
SearchResultItem,
)
from shared.atomic_api import SearchHit
# --- credibility mapping ---------------------------------------------------
CREDIBILITY_SCORES: dict[str, float] = {
"Tier1": 0.90,
"Tier2": 0.70,
"Tier3": 0.50,
"StateAffiliated": 0.40,
"KnownDisinfo": 0.15,
"Unknown": 0.50,
}
def tag_to_credibility_score(tags: list[dict[str, Any]]) -> float:
"""Pick the highest-priority credibility tag and map to a score."""
for t in tags:
name = (t.get("name") or "").strip()
if name in CREDIBILITY_SCORES:
return CREDIBILITY_SCORES[name]
return CREDIBILITY_SCORES["Unknown"]
# --- parent URL / publisher ------------------------------------------------
def parent_url_of(source_url: str | None) -> str:
"""Strip `#claim=...` fragment from a claim atom's source_url."""
if not source_url:
return ""
return source_url.split("#", 1)[0]
def publisher_of(url: str) -> str:
try:
host = urlparse(url).hostname or ""
except ValueError:
return ""
if host.startswith("www."):
host = host[4:]
return host
def title_from_url(url: str) -> str:
if not url:
return ""
slug = url.rstrip("/").rsplit("/", 1)[-1]
return unquote(slug).replace("_", " ")
# --- claim atom body parsing -----------------------------------------------
_CLAIM_BODY_RE = re.compile(r"^# Claim\s*\n+(.+?)\n+##", re.DOTALL | re.MULTILINE)
_STANCE_RE = re.compile(r"Stance in source:\s*(\w+)", re.IGNORECASE)
_PARENT_ID_RE = re.compile(r"Parent atom:\s*`([^`]+)`")
def parse_claim_atom_body(content: str) -> tuple[str, str, str]:
"""Return (claim_text, stance, parent_atom_id) from a Type/Claim markdown body."""
claim_text = ""
stance = "NEUTRAL"
parent_id = ""
m = _CLAIM_BODY_RE.search(content)
if m:
claim_text = m.group(1).strip()
s = _STANCE_RE.search(content)
if s:
stance = s.group(1).strip().upper()
p = _PARENT_ID_RE.search(content)
if p:
parent_id = p.group(1).strip()
return claim_text, stance, parent_id
# --- document title from content header -----------------------------------
def title_from_content(content: str | None, fallback_url: str = "") -> str:
if not content:
return title_from_url(fallback_url)
for line in content.lstrip().splitlines():
stripped = line.strip()
if stripped.startswith("# "):
return stripped[2:].strip()
if stripped:
break
return title_from_url(fallback_url)
def sha256_hex(text: str) -> str:
return hashlib.sha256(text.encode("utf-8", errors="replace")).hexdigest()
# --- aggregated per-parent ------------------------------------------------
def group_hits_by_parent(
hits: list[SearchHit], rerank_scores: dict[str, float]
) -> list[tuple[str, list[tuple[SearchHit, float]]]]:
"""Return ordered list of (parent_url, [(hit, rerank_score), ...]).
Order is by the best rerank_score within each parent, descending.
If a hit has no rerank score (wasn't in top-N), it falls to the end.
"""
buckets: dict[str, list[tuple[SearchHit, float]]] = defaultdict(list)
for h in hits:
parent = parent_url_of(h.source_url)
if not parent:
continue
rr = rerank_scores.get(h.atom_id, 0.0)
buckets[parent].append((h, rr))
# sort each bucket: highest rerank first, then highest embedding sim
for parent in buckets:
buckets[parent].sort(key=lambda x: (-x[1], -x[0].similarity))
# sort buckets by their top entry's rerank score, descending
ordered = sorted(
buckets.items(),
key=lambda kv: (-kv[1][0][1], -kv[1][0][0].similarity),
)
return ordered
# --- build EvidenceItem from a parent doc + claim matches ------------------
def evidence_from_parent(
*,
parent_atom: dict[str, Any],
claim_hits: list[tuple[SearchHit, float]],
parent_full_atoms: dict[str, dict[str, Any]],
include_full_text: bool,
nli_by_atom_id: dict[str, tuple[str, float, str | None]] | None = None,
) -> EvidenceItem:
"""Build an EvidenceItem given one parent document and its best-matching claim atoms.
`parent_atom` is the parent Type/Document atom (with full content).
`claim_hits` are (SearchHit, rerank_score) for claims belonging to this parent,
pre-sorted descending.
`parent_full_atoms` is an already-fetched map of full atom bodies so we can
pull the claim text from each matching claim atom.
"""
parent_url = parent_atom.get("source_url") or ""
parent_id = parent_atom.get("id") or ""
parent_content = parent_atom.get("content") or ""
title = title_from_content(parent_content, parent_url)
# Best matching claim (for summary + brain_meta)
best_hit, best_rerank = claim_hits[0]
best_full = parent_full_atoms.get(best_hit.atom_id) or {}
best_claim_text, best_stance, _best_parent = parse_claim_atom_body(
best_full.get("content") or ""
)
# Snippet: first paragraph of the parent, trimmed
snippet = (
parent_content.strip().split("\n\n", 1)[0][:300]
if parent_content
else None
)
# Dates — prefer published_at from the parent; fall back to created_at; never null
published_at = _parse_dt(parent_atom.get("published_at"))
retrieved_at = _parse_dt(parent_atom.get("created_at")) or datetime.now(timezone.utc)
# Credibility from parent's tag set
credibility = tag_to_credibility_score(parent_atom.get("tags") or [])
# NLI stance vs query — only attached to the BEST claim (the one we
# already surface as `summary`), since that's the one Didi will show.
nli_label = "UNKNOWN"
nli_conf = 0.0
nli_err: str | None = None
if nli_by_atom_id is not None:
entry = nli_by_atom_id.get(best_hit.atom_id)
if entry is not None:
nli_label, nli_conf, nli_err = entry
# Brain meta: one per evidence item, holds every matching claim's info
brain_meta = BrainEvidenceMeta(
parent_atom_id=parent_id,
matching_claim_atom_ids=[h.atom_id for h, _ in claim_hits],
best_claim_text=best_claim_text,
best_claim_stance_in_source=best_stance,
best_claim_hash=(best_hit.source_url or "").split("#claim=", 1)[-1][:16],
claim_count=len(claim_hits),
reranker_score=best_rerank,
embedding_similarity=best_hit.similarity,
stance_vs_query=nli_label,
nli_confidence=nli_conf,
nli_error=nli_err,
)
full_text = parent_content if include_full_text else None
return EvidenceItem(
url=parent_url,
canonical_url=parent_url or None,
title=title,
publisher=publisher_of(parent_url),
published_at=published_at,
retrieved_at=retrieved_at,
snippet=snippet,
summary=best_claim_text or None,
full_text=full_text,
full_text_hash=sha256_hex(parent_content) if parent_content else "",
provenance=Provenance(
extraction_method="brain",
fallback_chain=[],
brain_meta=brain_meta,
),
relevance_score=round(best_rerank or best_hit.similarity, 4),
credibility_score=credibility,
)
def _parse_dt(value: Any) -> datetime | None:
if not value:
return None
if isinstance(value, datetime):
return value
try:
text = str(value).replace("Z", "+00:00")
return datetime.fromisoformat(text)
except (ValueError, TypeError):
return None
# --- doc atom → FetchedPage ------------------------------------------------
def doc_to_fetched_page(
full_atom: dict[str, Any], *, include_html: bool = False
) -> FetchedPage:
url = full_atom.get("source_url") or ""
content = full_atom.get("content") or ""
return FetchedPage(
url=url,
canonical_url=url or None,
title=title_from_content(content, url),
text=content,
text_hash=sha256_hex(content),
html=None if not include_html else content,
extraction_method="brain",
fallback_chain=[],
published_at=_parse_dt(full_atom.get("published_at")),
retrieved_at=_parse_dt(full_atom.get("created_at")) or datetime.now(timezone.utc),
extraction_time_ms=0.0,
warnings=[],
needs_fallback=False,
status_code=200,
content_type="text/markdown",
)
# --- doc atom → SearchResultItem -------------------------------------------
def doc_to_search_result(
atom: dict[str, Any], *, query: str, rank: int
) -> SearchResultItem:
url = atom.get("source_url") or ""
title = title_from_content(atom.get("content"), url)
snippet = atom.get("snippet") or ""
if not snippet and atom.get("content"):
snippet = (atom["content"] or "").strip().split("\n\n", 1)[0][:200]
return SearchResultItem(
query=query,
url=url,
title=title,
snippet=snippet,
rank=rank,
site=publisher_of(url),
published_at=_parse_dt(atom.get("published_at")),
)
# --- language detection (tiny heuristic) ----------------------------------
_RO_CHARS = set("ăâîșțĂÂÎȘȚşţŞŢ")
def detect_language_simple(text: str) -> str:
if not text:
return "en"
if any(c in _RO_CHARS for c in text):
return "ro"
return "en"

View file

@ -0,0 +1,145 @@
"""NLI stance classification: is this evidence supporting, contradicting,
or neutral relative to the user's claim?
This is separate from `stance_in_source` (what the original source asserts
about itself). For disinfo analysis, the question Didi's backend really
needs answered is: "does this evidence back the user's claim or refute it?"
Implementation:
- one LLM call per (claim, evidence) pair
- async + parallel across top-N evidence items
- returns a stance label and a confidence
- uses the versioned prompt at brain_api/prompts/nli_v1.md
"""
from __future__ import annotations
import asyncio
from dataclasses import dataclass
from pathlib import Path
from shared.config import LlmRole
from shared.llm_client import LlmClient, LlmError
from shared.logging import get_logger
log = get_logger(__name__)
PROMPT_VERSION = "v1"
_PROMPT_PATH = Path(__file__).resolve().parent.parent / "prompts" / f"nli_{PROMPT_VERSION}.md"
ALLOWED_LABELS = {"SUPPORTS", "CONTRADICTS", "NEUTRAL"}
MAX_EVIDENCE_CHARS = 1500 # truncate long evidence before sending to the NLI model
# Match the effective llama.cpp backend concurrency: we have two instances
# behind the router (10.11.10.18 and 10.11.10.19), each serves one request
# at a time. Flooding with more parallel calls just queues them on the
# backend and hits our per-call timeout.
MAX_PARALLEL = 2
PER_CALL_TIMEOUT_S = 30.0 # generous; queuing + generation
TOTAL_TIMEOUT_S = 60.0 # wall-clock for the whole batch
_PROMPT_TEMPLATE: str | None = None
def _load_prompt() -> str:
global _PROMPT_TEMPLATE
if _PROMPT_TEMPLATE is None:
_PROMPT_TEMPLATE = _PROMPT_PATH.read_text(encoding="utf-8")
return _PROMPT_TEMPLATE
@dataclass(slots=True, frozen=True)
class NliResult:
label: str # SUPPORTS / CONTRADICTS / NEUTRAL
confidence: float # 0.0 - 1.0
error: str | None = None
async def classify_one(
llm: LlmClient, *, claim: str, evidence: str
) -> NliResult:
"""Classify a single (claim, evidence) pair. Never raises — returns
NeutralResult with error populated on failure so the caller can still
produce a response for that evidence item."""
if not evidence.strip():
return NliResult(label="NEUTRAL", confidence=0.0, error="empty_evidence")
truncated = evidence[:MAX_EVIDENCE_CHARS]
prompt = _load_prompt().replace("{claim}", claim).replace("{evidence}", truncated)
try:
result, _usage = await asyncio.wait_for(
llm.chat_json(
role=LlmRole.REASONING,
system=(
"You are an NLI classifier. Respond with strictly valid "
"JSON only, no commentary."
),
user=prompt,
max_tokens=120,
temperature=0.0,
),
timeout=PER_CALL_TIMEOUT_S,
)
except asyncio.TimeoutError:
return NliResult(label="NEUTRAL", confidence=0.0, error="timeout")
except LlmError as e:
return NliResult(label="NEUTRAL", confidence=0.0, error=f"llm:{e}")
except Exception as e: # noqa: BLE001
return NliResult(label="NEUTRAL", confidence=0.0, error=f"{type(e).__name__}:{e}")
if not isinstance(result, dict):
return NliResult(label="NEUTRAL", confidence=0.0, error="non_dict_response")
raw_label = (result.get("label") or "").strip().upper()
try:
conf = float(result.get("confidence", 0))
except (TypeError, ValueError):
conf = 0.0
conf = max(0.0, min(1.0, conf))
if raw_label not in ALLOWED_LABELS:
return NliResult(
label="NEUTRAL",
confidence=0.0,
error=f"bad_label:{raw_label[:40]}",
)
return NliResult(label=raw_label, confidence=conf)
async def classify_batch(
llm: LlmClient,
*,
claim: str,
evidence_texts: list[str],
max_parallel: int = MAX_PARALLEL,
) -> list[NliResult]:
"""Classify many evidence items in parallel, preserving input order.
Concurrency is bounded by `max_parallel` so we don't hammer the LLM
router. Individual failures produce a NeutralResult with an error field
(never raises to the caller).
"""
if not evidence_texts:
return []
sem = asyncio.Semaphore(max_parallel)
async def _guarded(ev: str) -> NliResult:
async with sem:
return await classify_one(llm, claim=claim, evidence=ev)
try:
results = await asyncio.wait_for(
asyncio.gather(*[_guarded(ev) for ev in evidence_texts]),
timeout=TOTAL_TIMEOUT_S,
)
except asyncio.TimeoutError:
log.warning("nli_batch_total_timeout", count=len(evidence_texts))
return [
NliResult(label="NEUTRAL", confidence=0.0, error="batch_timeout")
for _ in evidence_texts
]
return list(results)

View file

@ -0,0 +1,93 @@
"""POST /v1/search — thin retrieval that returns a flat list of results.
Unlike /v1/gather, we do NOT rerank or group just semantic search the brain
and convert each document-level hit into a SearchResultItem. This is the
equivalent of a search engine result list; callers that want ranked evidence
should hit /v1/gather instead.
"""
from __future__ import annotations
import asyncio
import time
import uuid
from brain_api.schemas import (
BrainMeta,
SearchRequest,
SearchResponse,
SearchResultItem,
)
from brain_api.services.mapping import doc_to_search_result, parent_url_of
from shared.atomic_api import AtomicClient
from shared.logging import get_logger
log = get_logger(__name__)
async def search(
req: SearchRequest, *, atomic: AtomicClient
) -> SearchResponse:
t0 = time.perf_counter()
request_id = str(uuid.uuid4())
# Run queries in parallel; merge into a flat ranked list.
tasks = [
atomic.search(q, mode="semantic", limit=req.max_results, threshold=0.2)
for q in req.queries
]
per_query_hits = await asyncio.gather(*tasks, return_exceptions=True)
# We need full document atoms (parents, de-duped by URL) to render results
seen_urls: set[str] = set()
results: list[SearchResultItem] = []
parent_atom_cache: dict[str, dict] = {}
for query_str, query_hits in zip(req.queries, per_query_hits, strict=True):
if isinstance(query_hits, Exception):
log.warning("search_query_failed", query=query_str, error=str(query_hits))
continue
# Collect parent URLs from this query in order
ordered_parents: list[str] = []
for h in query_hits:
parent = parent_url_of(h.source_url)
if not parent or parent in seen_urls:
continue
seen_urls.add(parent)
ordered_parents.append(parent)
if len(ordered_parents) >= req.max_results:
break
# Fetch any parent docs we haven't seen yet, in parallel
need = [p for p in ordered_parents if p not in parent_atom_cache]
if need:
atoms = await asyncio.gather(
*[atomic.get_atom_by_source_url(u) for u in need],
return_exceptions=True,
)
for url, atom in zip(need, atoms, strict=True):
if isinstance(atom, dict):
parent_atom_cache[url] = atom
for rank, url in enumerate(ordered_parents, start=len(results) + 1):
atom = parent_atom_cache.get(url)
if not atom:
continue
results.append(doc_to_search_result(atom, query=query_str, rank=rank))
total_ms = round((time.perf_counter() - t0) * 1000, 1)
return SearchResponse(
request_id=request_id,
results=results,
total_results=len(results),
execution_time_ms=total_ms,
queries_processed=len(req.queries),
brain_meta=BrainMeta(
cache_status="HIT" if results else "MISS",
api_version="v1",
implementation="didibrain",
evidence_sources=len({r.url for r in results}),
total_claim_atoms_matched=0,
),
)

View file

@ -0,0 +1,186 @@
"""Topic volatility overrides — Phase D1.
Reads admin-configured volatility/TTL/recency per topic from didiFramework's
sensitive_topic table (proxied via HTTP to keep brain free of a Redis
dependency). The classifier consults this map AFTER its LLM call: if any of
the LLM-derived topic_codes matches an admin-configured topic, the admin's
values override the LLM estimates for that topic.
Resolution order for a claim's effective TTL:
1. classifier returns volatility + estimated_validity_hours (LLM judgment)
2. for each LLM-detected topic_code, fetch admin override from this module
3. if admin override exists, use the more conservative of (LLM, admin) i.e.
pick the SHORTER TTL; admins can tighten brain's own estimate but never
loosen it (a stable claim that touches an admin-tagged 'volatile' topic
gets the volatile TTL)
Cache: in-process, 60s TTL. Failures (didiFramework down, HTTP timeout) leave
the cache empty so callers fall through to LLM-only behavior never blocks.
"""
from __future__ import annotations
import asyncio
import os
import time
from dataclasses import dataclass
import httpx
from shared.logging import get_logger
log = get_logger(__name__)
DIDI_FRAMEWORK_URL = os.environ.get(
"DIDI_FRAMEWORK_URL", "http://didi-framework:3005"
).rstrip("/")
CACHE_TTL_S = 60.0
HTTP_TIMEOUT_S = 5.0
@dataclass(slots=True, frozen=True)
class TopicConfig:
"""Admin-configured policy for one topic.
Attributes:
topic_code: Canonical code (matches classifier's topic_codes output).
volatility: One of "volatile", "evolving", "stable".
cache_ttl_hours: Hard cap on cache TTL for verdicts touching this topic.
recency_window_days: For volatile/evolving topics, drop evidence older
than this in /v1/gather.
half_life_days: Recency-boost half-life used in combined ranking.
"""
topic_code: str
volatility: str
cache_ttl_hours: int
recency_window_days: int
half_life_days: float
_cache: dict[str, TopicConfig] | None = None
_cache_loaded_at: float = 0.0
_cache_lock = asyncio.Lock()
async def _fetch_from_framework() -> dict[str, TopicConfig]:
"""Pull active topics from didiFramework. Empty dict on any failure."""
url = f"{DIDI_FRAMEWORK_URL}/api/sensitive-topics?active=true"
try:
async with httpx.AsyncClient(timeout=HTTP_TIMEOUT_S) as client:
resp = await client.get(url)
if resp.status_code >= 400:
log.debug(
"topic_overrides_http_error",
status=resp.status_code,
)
return {}
payload = resp.json()
except Exception as e: # noqa: BLE001
log.debug(
"topic_overrides_fetch_failed",
error=f"{type(e).__name__}:{e}",
)
return {}
if not isinstance(payload, dict) or not payload.get("success"):
return {}
rows = payload.get("data") or []
if not isinstance(rows, list):
return {}
out: dict[str, TopicConfig] = {}
for row in rows:
if not isinstance(row, dict):
continue
code = row.get("topic_code")
vol = row.get("volatility")
if not isinstance(code, str) or vol not in (
"volatile",
"evolving",
"stable",
):
continue
try:
out[code] = TopicConfig(
topic_code=code,
volatility=vol,
cache_ttl_hours=int(row.get("cache_ttl_hours") or 720),
recency_window_days=int(row.get("recency_window_days") or 30),
half_life_days=float(row.get("half_life_days") or 30.0),
)
except (TypeError, ValueError):
continue
return out
async def get_topic_overrides() -> dict[str, TopicConfig]:
"""Return current admin-configured topic policies (cached 60s).
Always returns a dict empty if didiFramework is unreachable or
sensitive_topic doesn't have the volatility columns yet (migration 012
not run). Callers can iterate over it freely.
"""
global _cache, _cache_loaded_at
now = time.time()
if _cache is not None and (now - _cache_loaded_at) < CACHE_TTL_S:
return _cache
async with _cache_lock:
# Double-check inside the lock.
if _cache is not None and (time.time() - _cache_loaded_at) < CACHE_TTL_S:
return _cache
fresh = await _fetch_from_framework()
_cache = fresh
_cache_loaded_at = time.time()
return fresh
def invalidate_cache() -> None:
"""Force a refetch on next ``get_topic_overrides()`` call.
Called by admin endpoints after a topic is mutated in didiFramework so
the change propagates without waiting for the 60s cache window.
"""
global _cache
_cache = None
def reconcile_with_classifier(
*,
classifier_volatility: str,
classifier_validity_hours: int,
classifier_topics: list[str],
overrides: dict[str, TopicConfig],
) -> tuple[str, int]:
"""Combine LLM classifier output with admin overrides.
Picks the MORE conservative (shorter) TTL when admin override exists.
Volatility ranking: volatile < evolving < stable (volatile = shorter
"shelf life"). If admin says 'volatile' for any matching topic, the
final volatility is 'volatile' regardless of what the classifier said.
Returns:
(effective_volatility, effective_ttl_hours).
"""
if not overrides or not classifier_topics:
return classifier_volatility, classifier_validity_hours
rank = {"volatile": 0, "evolving": 1, "stable": 2}
eff_vol = classifier_volatility
eff_ttl = classifier_validity_hours
for topic in classifier_topics:
cfg = overrides.get(topic)
if cfg is None:
continue
# Pick the more conservative volatility (lower rank wins).
if rank.get(cfg.volatility, 1) < rank.get(eff_vol, 1):
eff_vol = cfg.volatility
# Pick the shorter TTL.
if cfg.cache_ttl_hours < eff_ttl:
eff_ttl = cfg.cache_ttl_hours
return eff_vol, eff_ttl

View file

@ -0,0 +1,566 @@
"""Verification cache — store the LLM verification result per (claim, tier).
Contract agreed with didi-backend: backend runs its own LLM verification call
(prompt + model are owned by backend side via Redis config), then POSTs the
result here fire-and-forget. On the next /v1/gather for the same claim + tier,
brain returns the cached payload verbatim in brain_meta.
Schema v2 change:
- v1 keyed on (claim_hash, evidence_hash, tier) unreachable because
evidence URLs at gather read-time rarely match those at write-time
(brain's live search ranks/filters differently than backend's original
source list).
- v2 keys on (claim_hash, tier). evidence_hash + evidence_urls are kept
as metadata on the stored row; backend uses them at read-time to
decide overlap with its current evidence set.
This module owns:
- normalization + hashing (claim + urls; urls hash is metadata now)
- Upsert writes with TTL
- Staleness detection via prompt_hash and framework_version
"""
from __future__ import annotations
import asyncio
import hashlib
import json
import unicodedata
from dataclasses import dataclass, field
from datetime import datetime, timedelta, timezone
from typing import Any, Literal
from brain_api.db import db
from brain_api.services.classifier import (
ClaimVolatility,
classify_claim_volatility,
)
from shared.config import settings
from shared.llm_client import LlmClient
from shared.logging import get_logger
log = get_logger(__name__)
# ----------------------------------------------------------------------- hashing
def normalize_claim(s: str) -> str:
"""Hash-input normalization agreed with backend.
Same variations collapse to one bucket:
- "România a câștigat 9 medalii." (RO, punctuated)
- "romania a castigat 9 medalii" (stripped diacritics)
- "Romania a câștigat 9 medalii!" (extra whitespace, exclam)
Different intents stay separate:
- negation ("nu e sigur" vs "e sigur")
- numbers ("9 medalii" vs "10 medalii")
- middle punctuation ("X, ironic")
"""
s = unicodedata.normalize("NFKD", s)
s = "".join(c for c in s if not unicodedata.combining(c))
s = s.lower()
s = " ".join(s.split())
s = s.rstrip(".?!")
return s
def hash_claim(claim: str) -> str:
return hashlib.sha256(normalize_claim(claim).encode("utf-8")).hexdigest()
def hash_evidence_urls(urls: list[str]) -> str:
"""Order-independent, case-insensitive canonical hash over a URL set.
Used purely as metadata now (not part of the cache key) so backend can
detect corpus drift without needing to recompute locally.
"""
canonical = sorted({u.strip().lower().rstrip("/") for u in urls if u})
return hashlib.sha256("\n".join(canonical).encode("utf-8")).hexdigest()
# --------------------------------------------------------------------- dataclass
@dataclass(slots=True)
class CacheEntry:
claim_hash: str
tier: str
evidence_hash: str
evidence_urls: list[str]
model: str | None
prompt_hash: str
framework_version: str | None
schema_name: str
verification_raw: dict[str, Any] | None
verification_processed: dict[str, Any]
created_at: datetime
updated_at: datetime
expires_at: datetime
# Phase B1+B2 metadata (may be missing for legacy rows written before
# the volatility migration — defaults are conservative).
volatility: str | None = None
topic_codes: list[str] = field(default_factory=list)
entity_bindings: list[dict[str, Any]] = field(default_factory=list)
consecutive_audit_passes: int = 0
last_audited_at: datetime | None = None
# -------------------------------------------------------------------------- IO
async def _register_facts_async(
classification: ClaimVolatility | None, *, source_label: str
) -> None:
"""Best-effort fact registration after a successful verification cache write.
Lazy import avoids a circular dependency (fact_status imports classifier).
"""
if classification is None or not classification.entity_bindings:
return
try:
from brain_api.services.fact_status import register_facts_from_bindings
await register_facts_from_bindings(
classification.entity_bindings,
volatility=classification.volatility,
topic_codes=classification.topic_codes,
source_atom_id=source_label,
)
except Exception as e: # noqa: BLE001
log.warning(
"vcache_fact_registration_failed",
error=f"{type(e).__name__}:{e}",
)
def _resolve_ttl_days(
ttl_days: int | None, classification: ClaimVolatility | None
) -> int:
"""Pick the effective TTL in days, prefering classifier estimate.
When classification is fresh (not degraded), its hour estimate is
converted to whole days (rounded up) so a volatile 6h claim yields ttl=1
day, never 30 days. Otherwise we fall back to either the explicit
``ttl_days`` argument or the global verification_cache_ttl_days setting.
"""
if classification is not None and not classification.degraded:
# Round up to whole days, but never below 1 day.
days_from_classifier = max(
1, (classification.estimated_validity_hours + 23) // 24
)
return days_from_classifier
if ttl_days is not None:
return ttl_days
return settings.verification_cache_ttl_days
async def upsert(
*,
claim: str,
evidence_urls: list[str],
tier: Literal["free", "premium"],
prompt_hash: str,
verification_processed: dict[str, Any],
verification_raw: dict[str, Any] | None = None,
model: str | None = None,
framework_version: str | None = None,
schema_name: str = "didi-v1",
ttl_days: int | None = None,
classification: ClaimVolatility | None = None,
llm: LlmClient | None = None,
) -> CacheEntry:
"""Last-wins upsert on (claim_hash, tier), with volatility classification.
Pipeline:
1. If no ``classification`` provided and an ``llm`` client is, run the
volatility classifier on the claim text. This drives TTL and adds
topic_codes / entity_bindings metadata for invalidation by topic
and for fact_status registration.
2. UPSERT row, including the new metadata columns.
3. Fire-and-forget fact registration after successful write.
Multiple verification runs for the same claim+tier (different evidence
sets, rerun on model fallback) all write into the same row; the latest
successful verification wins.
Args:
claim: Original user claim text. Hashed for the cache key and also
passed to the classifier.
evidence_urls: URLs the LLM verification ran over (metadata only).
tier: free | premium.
prompt_hash: sha256[:12] of the verification prompt template.
verification_processed: Canonical mapped verdict (status, confidence,
etc.) what callers serve from cache.
verification_raw: Raw LLM response (used by stale_framework recompute).
model: LLM model identifier.
framework_version: sha256[:12] of relevant framework configs.
schema_name: Versioned schema label, default "didi-v1".
ttl_days: Caller-provided TTL override; ignored if a classification
with non-degraded estimate is available.
classification: Pre-computed ClaimVolatility from caller.
llm: LLM client for classifier. None disables classification.
"""
classification = await _maybe_classify(
classification=classification, llm=llm, claim=claim
)
ttl = _resolve_ttl_days(ttl_days, classification)
expires_at = datetime.now(tz=timezone.utc) + timedelta(days=ttl)
ch = hash_claim(claim)
eh = hash_evidence_urls(evidence_urls)
ev_urls_json = json.dumps(list(evidence_urls))
raw_json = (
json.dumps(verification_raw) if verification_raw is not None else None
)
processed_json = json.dumps(verification_processed)
volatility = classification.volatility if classification else None
topic_codes = classification.topic_codes if classification else []
entity_bindings_json = (
json.dumps(classification.entity_bindings_jsonb())
if classification
else "[]"
)
ttl_hours_used = ttl * 24
sql = """
INSERT INTO brain_verification_cache (
claim_hash, tier,
evidence_hash, evidence_urls,
model, prompt_hash, framework_version, schema_name,
verification_raw, verification_processed,
expires_at,
volatility, topic_codes, entity_bindings, ttl_hours_used
)
VALUES (
$1, $2, $3, $4::jsonb, $5, $6, $7, $8, $9::jsonb, $10::jsonb, $11,
$12, $13, $14::jsonb, $15
)
ON CONFLICT (claim_hash, tier) DO UPDATE SET
evidence_hash = EXCLUDED.evidence_hash,
evidence_urls = EXCLUDED.evidence_urls,
model = EXCLUDED.model,
prompt_hash = EXCLUDED.prompt_hash,
framework_version = EXCLUDED.framework_version,
schema_name = EXCLUDED.schema_name,
verification_raw = EXCLUDED.verification_raw,
verification_processed = EXCLUDED.verification_processed,
updated_at = now(),
expires_at = EXCLUDED.expires_at,
volatility = COALESCE(EXCLUDED.volatility, brain_verification_cache.volatility),
topic_codes = CASE
WHEN array_length(EXCLUDED.topic_codes, 1) > 0
THEN EXCLUDED.topic_codes
ELSE brain_verification_cache.topic_codes
END,
entity_bindings = CASE
WHEN jsonb_array_length(EXCLUDED.entity_bindings) > 0
THEN EXCLUDED.entity_bindings
ELSE brain_verification_cache.entity_bindings
END,
ttl_hours_used = EXCLUDED.ttl_hours_used
RETURNING
claim_hash, tier,
evidence_hash, evidence_urls,
model, prompt_hash, framework_version, schema_name,
verification_raw, verification_processed,
created_at, updated_at, expires_at
"""
async with db.pool.acquire() as conn:
row = await conn.fetchrow(
sql,
ch,
tier,
eh,
ev_urls_json,
model,
prompt_hash,
framework_version,
schema_name,
raw_json,
processed_json,
expires_at,
volatility,
topic_codes,
entity_bindings_json,
ttl_hours_used,
)
assert row is not None # UPSERT with RETURNING always yields a row
entry = _row_to_entry(row)
# Best-effort fact registration after the write succeeds.
if classification and classification.entity_bindings:
asyncio.create_task(
_register_facts_async(
classification, source_label=f"vcache:{ch[:12]}"
)
)
log.info(
"vcache_upsert_ok",
claim_hash=ch[:12],
tier=tier,
volatility=volatility,
ttl_days=ttl,
topic_codes=topic_codes,
binding_count=len(classification.entity_bindings) if classification else 0,
)
return entry
async def _maybe_classify(
*,
classification: ClaimVolatility | None,
llm: LlmClient | None,
claim: str,
) -> ClaimVolatility | None:
"""Use caller's classification or compute one. Never raises."""
if classification is not None:
return classification
if llm is None or not claim.strip():
return None
try:
return await classify_claim_volatility(llm, claim=claim)
except Exception as e: # noqa: BLE001
log.warning(
"vcache_classifier_failed",
error=f"{type(e).__name__}:{e}",
)
return None
# ============================================================================
# Phase B2: confidence decay + judge integration for verification_cache
# ============================================================================
# Same decay model as analysis_atom — keeps the two caches behaviorally
# consistent so callers don't have to special-case.
DECAY_HALF_LIVES_HOURS: dict[str, float] = {
"volatile": 24.0,
"evolving": 168.0,
"stable": float("inf"),
}
AUDIT_HISTORY_MAX = 50
AUDIT_PASS_MIN_INTERVAL_HOURS = 6.0
def compute_effective_confidence(
*,
base_confidence: float | None,
volatility: str | None,
age_hours: float,
consecutive_audit_passes: int = 0,
) -> float | None:
"""Decay base confidence by age, modulated by volatility and audit history.
Mirrors analysis_atom.compute_effective_confidence so the two cache
paths share the same model and admin tooling can reuse formulas.
"""
if base_confidence is None:
return None
half = DECAY_HALF_LIVES_HOURS.get(volatility or "evolving", 168.0)
if half == float("inf") or age_hours <= 0:
decay = 1.0
else:
decay = 0.5 ** (age_hours / half)
audit_boost = min(0.3, 0.03 * max(0, consecutive_audit_passes))
return float(base_confidence) * decay * (1.0 + audit_boost)
async def apply_judge_verdict(
*,
claim_hash: str,
tier: Literal["free", "premium"],
verdict: object, # JudgeVerdict — typed loosely to avoid circular import
) -> None:
"""Persist a JudgeVerdict to the verification_cache row.
Updates audit_history (capped), last_audited_at, consecutive_audit_passes
(rate-limited), and expires_at on INVALIDATE. Also writes a brain_audit_log
entry for telemetry.
"""
if not db.pool:
raise RuntimeError("brain_db not connected")
from brain_api.services.cache_judge import JudgeVerdict
if not isinstance(verdict, JudgeVerdict):
raise TypeError(
f"apply_judge_verdict: expected JudgeVerdict, got {type(verdict).__name__}"
)
audit_entry = verdict.to_audit_entry()
audit_json = json.dumps(audit_entry)
audit_array_json = json.dumps([audit_entry])
sql = """
UPDATE brain_verification_cache
SET
audit_history = (
SELECT jsonb_agg(elem)
FROM (
SELECT elem
FROM jsonb_array_elements(
COALESCE(audit_history, '[]'::jsonb) || $3::jsonb
) WITH ORDINALITY AS t(elem, ord)
ORDER BY ord DESC
LIMIT $4
) recent
),
last_audited_at = now(),
consecutive_audit_passes = CASE
WHEN $5 = 'KEEP_CACHE' AND (
last_audited_at IS NULL
OR last_audited_at < now() - ($6 || ' hours')::interval
)
THEN consecutive_audit_passes + 1
WHEN $5 = 'INVALIDATE' THEN 0
ELSE consecutive_audit_passes
END,
expires_at = CASE
WHEN $5 = 'INVALIDATE' THEN now()
ELSE expires_at
END,
updated_at = now()
WHERE claim_hash = $1 AND tier = $2
"""
async with db.pool.acquire() as conn:
await conn.execute(
sql,
claim_hash,
tier,
audit_array_json,
AUDIT_HISTORY_MAX,
verdict.decision,
str(int(AUDIT_PASS_MIN_INTERVAL_HOURS)),
)
await conn.execute(
"""
INSERT INTO brain_audit_log (action, target_table, target_id, actor, payload)
VALUES ($1, 'brain_verification_cache', $2, 'cache_judge', $3::jsonb)
""",
f"judge_{verdict.decision.lower()}",
f"{claim_hash[:12]}/{tier}",
audit_json,
)
async def lookup(
*,
claim: str,
tier: Literal["free", "premium"],
) -> CacheEntry | None:
"""Fetch the cached entry for (claim, tier) — evidence is metadata only."""
ch = hash_claim(claim)
sql = """
SELECT
claim_hash, tier,
evidence_hash, evidence_urls,
model, prompt_hash, framework_version, schema_name,
verification_raw, verification_processed,
created_at, updated_at, expires_at,
volatility, topic_codes, entity_bindings,
consecutive_audit_passes, last_audited_at
FROM brain_verification_cache
WHERE claim_hash = $1 AND tier = $2
AND expires_at > now()
"""
async with db.pool.acquire() as conn:
row = await conn.fetchrow(sql, ch, tier)
if not row:
return None
return _row_to_entry(row)
# ------------------------------------------------------------- staleness decider
StalenessStatus = Literal[
"fresh",
"stale_framework",
"stale_prompt",
"stale_evidence", # bound facts have flipped — caller must recompute
"miss",
]
def decide_freshness(
entry: CacheEntry | None,
current_prompt_hash: str | None,
current_framework_version: str | None,
) -> StalenessStatus:
"""Given a cached entry + the caller's current prompt/framework, decide.
- miss: no entry at all (or expired in DB)
- stale_prompt: prompt changed since cache was written (verification_raw
stances may differ semantically) caller should NOT use cache
- stale_framework: prompt unchanged but threshold config changed caller
CAN use verification_raw and recompute status locally
- fresh: everything matches; return verification_processed directly
"""
if entry is None:
return "miss"
if current_prompt_hash and entry.prompt_hash != current_prompt_hash:
return "stale_prompt"
if (
current_framework_version
and entry.framework_version
and entry.framework_version != current_framework_version
):
return "stale_framework"
return "fresh"
# ------------------------------------------------------------------- internal
def _row_to_entry(row) -> CacheEntry:
raw = row["verification_raw"]
processed = row["verification_processed"]
ev_urls = row["evidence_urls"]
# asyncpg decodes jsonb as str; json.loads needed
if isinstance(raw, str):
raw = json.loads(raw)
if isinstance(processed, str):
processed = json.loads(processed)
if isinstance(ev_urls, str):
ev_urls = json.loads(ev_urls)
# Optional B1 metadata — these may not be present on legacy rows or in
# callers that select an older column set.
def _opt(key: str, default: Any = None) -> Any:
try:
return row[key]
except (KeyError, IndexError):
return default
bindings = _opt("entity_bindings", [])
if isinstance(bindings, str):
bindings = json.loads(bindings)
topic_codes = _opt("topic_codes", []) or []
return CacheEntry(
claim_hash=row["claim_hash"],
tier=row["tier"],
evidence_hash=row["evidence_hash"],
evidence_urls=list(ev_urls) if ev_urls else [],
model=row["model"],
prompt_hash=row["prompt_hash"],
framework_version=row["framework_version"],
schema_name=row["schema_name"],
verification_raw=raw,
verification_processed=processed,
created_at=row["created_at"],
updated_at=row["updated_at"],
expires_at=row["expires_at"],
volatility=_opt("volatility"),
topic_codes=list(topic_codes) if topic_codes else [],
entity_bindings=list(bindings) if bindings else [],
consecutive_audit_passes=int(_opt("consecutive_audit_passes", 0) or 0),
last_audited_at=_opt("last_audited_at"),
)