204 lines
6.4 KiB
Python
204 lines
6.4 KiB
Python
"""Single-document claim extraction logic.
|
|
|
|
Given one Type/Document atom, this module:
|
|
|
|
1. Builds the extraction prompt from the versioned template
|
|
2. Calls Qwen 397B (REASONING role) for structured JSON output
|
|
3. Parses + validates each claim:
|
|
- Required fields present and right types
|
|
- Stance is in the allowed enum
|
|
- Confidence above threshold
|
|
- Quote is a verbatim substring of the source (programmatic check —
|
|
this is the cheap defense against LLM hallucination)
|
|
4. Returns a list of ExtractedClaim dataclasses ready for the pusher.
|
|
|
|
The pusher (push.py) takes ExtractedClaim and creates a Type/Claim atom in
|
|
Atomic, with the proper tag inheritance and a stable hash-based source_url.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import re
|
|
from dataclasses import dataclass
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
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 / "prompts" / f"claim_extraction_{PROMPT_VERSION}.md"
|
|
ALLOWED_STANCES = {"ASSERTS", "REPORTS", "REFUTES", "QUESTIONS", "NEUTRAL"}
|
|
MIN_CONFIDENCE = 0.7
|
|
MIN_CLAIM_CHARS = 20
|
|
MIN_QUOTE_CHARS = 10
|
|
MAX_QUOTE_CHARS = 600
|
|
MAX_INPUT_CHARS = 120_000 # ~30K tokens, well under Qwen's 262K context
|
|
|
|
_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
|
|
|
|
|
|
# ============================================================ data structures
|
|
|
|
|
|
@dataclass(slots=True, frozen=True)
|
|
class ExtractedClaim:
|
|
claim: str
|
|
quote: str
|
|
stance: str # uppercase: ASSERTS / REPORTS / REFUTES / QUESTIONS / NEUTRAL
|
|
confidence: float
|
|
|
|
def stable_hash(self) -> str:
|
|
"""8-char SHA-1 of canonicalized claim text — used in source_url fragment."""
|
|
canonical = " ".join(self.claim.lower().strip().split())
|
|
return hashlib.sha1(canonical.encode("utf-8")).hexdigest()[:8]
|
|
|
|
|
|
@dataclass(slots=True)
|
|
class ExtractionResult:
|
|
"""What `extract_claims_from_atom` returns."""
|
|
|
|
raw_count: int # how many items the LLM returned
|
|
valid: list[ExtractedClaim]
|
|
rejected: dict[str, int] # reason → count
|
|
backend: str # which backend served the request
|
|
|
|
|
|
# ============================================================ extraction core
|
|
|
|
|
|
async def extract_claims_from_atom(
|
|
llm: LlmClient,
|
|
*,
|
|
title: str,
|
|
language: str,
|
|
content: str,
|
|
max_tokens_out: int = 6000,
|
|
) -> ExtractionResult:
|
|
"""Run the extraction LLM call and validate the results.
|
|
|
|
Raises LlmError on transport / JSON parse failure (caller decides whether
|
|
to retry or skip the document).
|
|
"""
|
|
prompt = _load_prompt().replace("{title}", title)
|
|
prompt = prompt.replace("{language}", language)
|
|
truncated = content[:MAX_INPUT_CHARS]
|
|
if len(content) > MAX_INPUT_CHARS:
|
|
truncated += "\n\n[document truncated for extraction]"
|
|
prompt = prompt.replace("{content}", truncated)
|
|
|
|
result, usage = await llm.chat_json(
|
|
role=LlmRole.REASONING,
|
|
system=(
|
|
"You are a precise information extractor. Output strictly valid "
|
|
"JSON, no commentary, no markdown fences."
|
|
),
|
|
user=prompt,
|
|
max_tokens=max_tokens_out,
|
|
temperature=0.0,
|
|
)
|
|
|
|
if not isinstance(result, dict) or "claims" not in result:
|
|
raise LlmError(
|
|
f"unexpected response shape: keys={list(result.keys()) if isinstance(result, dict) else type(result).__name__}"
|
|
)
|
|
|
|
raw_claims = result.get("claims") or []
|
|
if not isinstance(raw_claims, list):
|
|
raise LlmError(f"`claims` is not a list: {type(raw_claims).__name__}")
|
|
|
|
valid: list[ExtractedClaim] = []
|
|
rejected: dict[str, int] = {}
|
|
|
|
norm_source = _normalize_for_match(content)
|
|
|
|
for raw in raw_claims:
|
|
outcome = _parse_one(raw, norm_source)
|
|
if isinstance(outcome, ExtractedClaim):
|
|
valid.append(outcome)
|
|
else:
|
|
rejected[outcome] = rejected.get(outcome, 0) + 1
|
|
|
|
# Dedup within this batch by stable hash
|
|
seen: set[str] = set()
|
|
deduped: list[ExtractedClaim] = []
|
|
for c in valid:
|
|
h = c.stable_hash()
|
|
if h in seen:
|
|
rejected["intra_batch_duplicate"] = rejected.get("intra_batch_duplicate", 0) + 1
|
|
continue
|
|
seen.add(h)
|
|
deduped.append(c)
|
|
|
|
return ExtractionResult(
|
|
raw_count=len(raw_claims),
|
|
valid=deduped,
|
|
rejected=rejected,
|
|
backend=str(usage.get("backend", "")),
|
|
)
|
|
|
|
|
|
# ============================================================ validation
|
|
|
|
|
|
def _parse_one(raw: Any, norm_source: str) -> ExtractedClaim | str:
|
|
"""Validate one raw item from the LLM. Returns ExtractedClaim or error reason str."""
|
|
if not isinstance(raw, dict):
|
|
return "not_a_dict"
|
|
|
|
claim = (raw.get("claim") or "").strip()
|
|
quote = (raw.get("quote") or "").strip()
|
|
stance = (raw.get("stance") or "").strip().upper()
|
|
try:
|
|
confidence = float(raw.get("confidence", 0))
|
|
except (TypeError, ValueError):
|
|
return "bad_confidence_type"
|
|
|
|
if len(claim) < MIN_CLAIM_CHARS:
|
|
return "claim_too_short"
|
|
if len(quote) < MIN_QUOTE_CHARS:
|
|
return "quote_too_short"
|
|
if len(quote) > MAX_QUOTE_CHARS:
|
|
return "quote_too_long"
|
|
if stance not in ALLOWED_STANCES:
|
|
return "bad_stance"
|
|
if confidence < MIN_CONFIDENCE:
|
|
return "low_confidence"
|
|
|
|
# The critical anti-hallucination check: the quote must actually appear
|
|
# in the source document (after whitespace normalization).
|
|
if not _quote_in_source(quote, norm_source):
|
|
return "quote_not_in_source"
|
|
|
|
return ExtractedClaim(
|
|
claim=claim,
|
|
quote=quote,
|
|
stance=stance,
|
|
confidence=confidence,
|
|
)
|
|
|
|
|
|
_WHITESPACE_RE = re.compile(r"\s+")
|
|
|
|
|
|
def _normalize_for_match(s: str) -> str:
|
|
"""Collapse runs of whitespace and normalize quote chars for substring match."""
|
|
s = s.replace("\u2018", "'").replace("\u2019", "'")
|
|
s = s.replace("\u201c", '"').replace("\u201d", '"')
|
|
s = s.replace("\u2013", "-").replace("\u2014", "-")
|
|
return _WHITESPACE_RE.sub(" ", s).strip()
|
|
|
|
|
|
def _quote_in_source(quote: str, norm_source: str) -> bool:
|
|
return _normalize_for_match(quote) in norm_source
|