| from __future__ import annotations
|
|
|
| from collections import defaultdict
|
| from pathlib import Path
|
| from typing import Any
|
| from urllib.parse import unquote
|
|
|
| from utils import normalize_for_search
|
|
|
|
|
| MCKF_BASE = "https://mitranlil.ai/mckf/"
|
|
|
|
|
| class MCKFGraphRuntime:
|
| """Queryable normative graph with an RDF/SPARQL backend and deterministic fallback."""
|
|
|
| def __init__(
|
| self,
|
| evidence_spans: list[dict[str, Any]],
|
| cross_document_edges: list[dict[str, Any]],
|
| triples_path: Path | None = None,
|
| ) -> None:
|
| self.evidence_by_id = {
|
| str(item.get("evidence_id", "")): item
|
| for item in evidence_spans
|
| if item.get("evidence_id")
|
| }
|
| self.edges = cross_document_edges or []
|
| self.by_document_article: dict[str, list[str]] = defaultdict(list)
|
| self.by_concept: dict[str, list[str]] = defaultdict(list)
|
| self._build_direct_indexes(evidence_spans)
|
| self.graph = None
|
| self.backend = "semantic_graph_index"
|
| self.load_error = ""
|
| self._sparql_cache: dict[tuple[str, str], list[str]] = {}
|
| if triples_path and triples_path.exists():
|
| self._load_rdf_graph(triples_path)
|
|
|
| def status(self) -> dict[str, Any]:
|
| return {
|
| "backend": self.backend,
|
| "rdf_loaded": self.graph is not None,
|
| "triples": len(self.graph) if self.graph is not None else 0,
|
| "evidence_nodes": len(self.evidence_by_id),
|
| "load_error": self.load_error,
|
| }
|
|
|
| def retrieve(
|
| self,
|
| question_frame: dict[str, Any],
|
| requested_roles: set[str] | None = None,
|
| allowed_ids: set[str] | None = None,
|
| limit: int = 120,
|
| ) -> list[tuple[str, float, list[str]]]:
|
| requested_roles = requested_roles or set()
|
| scores: dict[str, float] = defaultdict(float)
|
| reasons: dict[str, list[str]] = defaultdict(list)
|
|
|
| values_by_predicate = {
|
| "hasActor": question_frame.get("actor", []) or [],
|
| "hasAction": question_frame.get("action", []) or [],
|
| "hasObject": question_frame.get("object", []) or [],
|
| "hasCompetentAuthority": question_frame.get("competent_authority", []) or [],
|
| "hasBeneficiary": question_frame.get("beneficiary", []) or [],
|
| }
|
| direct_weights = {
|
| "hasActor": 3.0,
|
| "hasAction": 3.2,
|
| "hasObject": 2.2,
|
| "hasCompetentAuthority": 3.4,
|
| "hasBeneficiary": 2.4,
|
| }
|
| for predicate, values in values_by_predicate.items():
|
| for value in values:
|
| key = (predicate, normalize_for_search(str(value)))
|
| for evidence_id in self.literal_index.get(key, []):
|
| if allowed_ids is not None and evidence_id not in allowed_ids:
|
| continue
|
| scores[evidence_id] += direct_weights[predicate]
|
| reasons[evidence_id].append(f"graph:{predicate}={value}")
|
|
|
| for role in requested_roles:
|
| for evidence_id in self.role_index.get(normalize_for_search(role), []):
|
| if allowed_ids is not None and evidence_id not in allowed_ids:
|
| continue
|
| scores[evidence_id] += 2.5
|
| reasons[evidence_id].append(f"graph:role={role}")
|
|
|
| if self.graph is not None:
|
| for evidence_id, predicate, value in self._sparql_matches(values_by_predicate):
|
| if allowed_ids is not None and evidence_id not in allowed_ids:
|
| continue
|
| scores[evidence_id] += 1.0
|
| reasons[evidence_id].append(f"sparql:{predicate}={value}")
|
|
|
| ranked = sorted(scores, key=lambda evidence_id: scores[evidence_id], reverse=True)
|
| return [(evidence_id, scores[evidence_id], reasons[evidence_id]) for evidence_id in ranked[:limit]]
|
|
|
| def expand(self, seed_ids: list[str], limit: int = 40) -> list[tuple[str, float, list[str]]]:
|
| scores: dict[str, float] = defaultdict(float)
|
| reasons: dict[str, list[str]] = defaultdict(list)
|
| for seed_id in seed_ids:
|
| seed = self.evidence_by_id.get(seed_id, {})
|
| if not seed:
|
| continue
|
| source_key = f"{seed.get('document_id')}::{seed.get('article_id')}"
|
| for edge in self.edges:
|
| edge_source = f"{edge.get('source_document_id')}::{edge.get('source_article_id')}"
|
| edge_target = f"{edge.get('target_document_id')}::{edge.get('target_article_id')}"
|
| target_key = edge_target if source_key == edge_source else edge_source if source_key == edge_target else ""
|
| if not target_key:
|
| continue
|
| for evidence_id in self.by_document_article.get(target_key, []):
|
| scores[evidence_id] += 1.5
|
| reasons[evidence_id].append(f"graph_edge:{edge.get('relation_type', 'related')}")
|
|
|
| concept_id = str(seed.get("_concept_id", "") or "")
|
| for evidence_id in self.by_concept.get(concept_id, []):
|
| if evidence_id == seed_id:
|
| continue
|
| evidence = self.evidence_by_id[evidence_id]
|
| roles = set(evidence.get("semantic_roles", []) or [])
|
| if roles & {"istisna", "sart", "sure", "tanim"}:
|
| scores[evidence_id] += 0.8
|
| reasons[evidence_id].append("graph:same_concept_qualifier")
|
|
|
| ranked = sorted(scores, key=lambda evidence_id: scores[evidence_id], reverse=True)
|
| return [(evidence_id, scores[evidence_id], reasons[evidence_id]) for evidence_id in ranked[:limit]]
|
|
|
| def _build_direct_indexes(self, evidence_spans: list[dict[str, Any]]) -> None:
|
| self.literal_index: dict[tuple[str, str], list[str]] = defaultdict(list)
|
| self.role_index: dict[str, list[str]] = defaultdict(list)
|
| for evidence in evidence_spans:
|
| evidence_id = str(evidence.get("evidence_id", "") or "")
|
| if not evidence_id:
|
| continue
|
| self.by_document_article[
|
| f"{evidence.get('document_id')}::{evidence.get('article_id')}"
|
| ].append(evidence_id)
|
| concept_id = str(evidence.get("_concept_id", "") or "")
|
| if concept_id:
|
| self.by_concept[concept_id].append(evidence_id)
|
| frame = evidence.get("semantic_frame", {}) or {}
|
| for predicate, field in (
|
| ("hasActor", "actor"),
|
| ("hasAction", "action"),
|
| ("hasObject", "object"),
|
| ("hasCompetentAuthority", "competent_authority"),
|
| ("hasBeneficiary", "beneficiary"),
|
| ):
|
| for value in frame.get(field, []) or []:
|
| self.literal_index[(predicate, normalize_for_search(str(value)))].append(evidence_id)
|
| for role in evidence.get("semantic_roles", []) or []:
|
| self.role_index[normalize_for_search(str(role))].append(evidence_id)
|
|
|
| def _load_rdf_graph(self, triples_path: Path) -> None:
|
| try:
|
| from rdflib import Graph
|
|
|
| graph = Graph()
|
| graph.parse(triples_path, format="nt")
|
| self.graph = graph
|
| self.backend = "rdflib_sparql"
|
| except Exception as exc:
|
| self.load_error = str(exc)
|
|
|
| def _sparql_matches(self, values_by_predicate: dict[str, list[Any]]) -> list[tuple[str, str, str]]:
|
| try:
|
| from rdflib import Literal
|
| except ImportError:
|
| return []
|
| rows = []
|
| for predicate, values in values_by_predicate.items():
|
| for value in values:
|
| cache_key = (predicate, str(value))
|
| if cache_key in self._sparql_cache:
|
| rows.extend((evidence_id, predicate, str(value)) for evidence_id in self._sparql_cache[cache_key])
|
| continue
|
| literal = Literal(str(value)).n3()
|
| query = f"""
|
| PREFIX mckf: <{MCKF_BASE}>
|
| SELECT ?unit WHERE {{ ?unit mckf:{predicate} {literal} . }}
|
| """
|
| try:
|
| matched_ids = []
|
| for result in self.graph.query(query):
|
| unit_uri = str(result[0])
|
| evidence_id = unquote(unit_uri.rsplit("/", 1)[-1])
|
| if evidence_id in self.evidence_by_id:
|
| rows.append((evidence_id, predicate, str(value)))
|
| matched_ids.append(evidence_id)
|
| self._sparql_cache[cache_key] = matched_ids
|
| except Exception as exc:
|
| self.load_error = str(exc)
|
| return rows
|
|
|