Livrare LOT 1 - Didi
This commit is contained in:
commit
5380c3fc63
990 changed files with 133308 additions and 0 deletions
532
ai_platform/modules/didi_brain/brain_api/services/admin.py
Normal file
532
ai_platform/modules/didi_brain/brain_api/services/admin.py
Normal 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),
|
||||
}
|
||||
|
|
@ -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 silver→gold.
|
||||
|
||||
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
|
||||
(volatile≤48h, evolving≤720h, stable≤26280h) 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,
|
||||
}
|
||||
315
ai_platform/modules/didi_brain/brain_api/services/cache_judge.py
Normal file
315
ai_platform/modules/didi_brain/brain_api/services/cache_judge.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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
|
||||
350
ai_platform/modules/didi_brain/brain_api/services/classifier.py
Normal file
350
ai_platform/modules/didi_brain/brain_api/services/classifier.py
Normal 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
|
||||
758
ai_platform/modules/didi_brain/brain_api/services/fact_status.py
Normal file
758
ai_platform/modules/didi_brain/brain_api/services/fact_status.py
Normal 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"],
|
||||
)
|
||||
62
ai_platform/modules/didi_brain/brain_api/services/fetch.py
Normal file
62
ai_platform/modules/didi_brain/brain_api/services/fetch.py
Normal 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,
|
||||
),
|
||||
)
|
||||
665
ai_platform/modules/didi_brain/brain_api/services/gather.py
Normal file
665
ai_platform/modules/didi_brain/brain_api/services/gather.py
Normal 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,
|
||||
),
|
||||
)
|
||||
268
ai_platform/modules/didi_brain/brain_api/services/ingest.py
Normal file
268
ai_platform/modules/didi_brain/brain_api/services/ingest.py
Normal 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))
|
||||
|
|
@ -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,
|
||||
)
|
||||
321
ai_platform/modules/didi_brain/brain_api/services/mapping.py
Normal file
321
ai_platform/modules/didi_brain/brain_api/services/mapping.py
Normal 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"
|
||||
145
ai_platform/modules/didi_brain/brain_api/services/nli.py
Normal file
145
ai_platform/modules/didi_brain/brain_api/services/nli.py
Normal 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)
|
||||
93
ai_platform/modules/didi_brain/brain_api/services/search.py
Normal file
93
ai_platform/modules/didi_brain/brain_api/services/search.py
Normal 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,
|
||||
),
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
@ -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"),
|
||||
)
|
||||
Loading…
Add table
Add a link
Reference in a new issue