#![allow(dead_code)]
use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet};
use std::path::Path;
use fathomdb_engine::{rerank_fused, Engine};
use serde_json::{json, Value};
pub const K_LADDER: [usize; 4] = [5, 10, 20, 50];
pub const HEADLINE_K: usize = 10;
pub const DEFAULT_FANOUT: usize = 50;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Necessity {
Required,
Supporting,
}
impl Necessity {
fn parse(s: &str) -> Option<Self> {
match s {
"required" => Some(Self::Required),
"supporting" => Some(Self::Supporting),
_ => None,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub enum QueryClass {
Commitment,
Action,
ExactFact,
Preference,
Exploratory,
Negative,
}
impl QueryClass {
pub fn label(&self) -> &'static str {
match self {
Self::Commitment => "commitment",
Self::Action => "action",
Self::ExactFact => "exact_fact",
Self::Preference => "preference",
Self::Exploratory => "exploratory",
Self::Negative => "negative",
}
}
fn parse(s: &str) -> Option<Self> {
Some(match s {
"commitment" => Self::Commitment,
"action" => Self::Action,
"exact_fact" => Self::ExactFact,
"preference" => Self::Preference,
"exploratory" => Self::Exploratory,
"negative" => Self::Negative,
_ => return None,
})
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum QueryOrigin {
HumanDataset,
LlmGenerated,
Templated,
}
impl QueryOrigin {
pub fn label(&self) -> &'static str {
match self {
Self::HumanDataset => "human_dataset",
Self::LlmGenerated => "llm_generated",
Self::Templated => "templated",
}
}
fn parse(s: &str) -> Option<Self> {
Some(match s {
"human_dataset" => Self::HumanDataset,
"llm_generated" => Self::LlmGenerated,
"templated" => Self::Templated,
_ => return None,
})
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Span {
pub doc_id: String,
pub start: usize,
pub end: usize,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Locator {
pub kind: String,
pub spans: Option<Vec<Span>>,
}
#[derive(Clone, Debug)]
pub struct EvidenceUnit {
pub evidence_id: String,
pub doc_id: String,
pub necessity: Necessity,
pub locator: Option<Locator>,
}
#[derive(Clone, Debug)]
pub struct GoldQuery {
pub query: String,
pub query_id: Option<String>,
pub query_class: QueryClass,
pub required_evidence: Vec<EvidenceUnit>,
pub expected_top_k_doc_ids: Vec<String>,
pub relation_type: Option<String>,
pub chain_shape: Option<String>,
pub source: Option<String>,
pub answer_type: Option<String>,
pub query_origin: QueryOrigin,
}
#[derive(Clone, Debug)]
pub struct GoldSet {
pub corpus_hash: String,
pub qrels_version: String,
pub note: Option<String>,
pub queries: Vec<GoldQuery>,
}
pub const UNPINNED_PLACEHOLDER: &str = "TODO(COR-2-freeze)";
pub fn load_gold_set(path: &Path) -> Result<GoldSet, String> {
let text =
std::fs::read_to_string(path).map_err(|e| format!("read {}: {e}", path.display()))?;
parse_gold_set(&text)
}
pub fn parse_gold_set(text: &str) -> Result<GoldSet, String> {
let v: Value = serde_json::from_str(text).map_err(|e| format!("gold set JSON: {e}"))?;
let corpus_hash = v.get("corpus_hash").and_then(Value::as_str).unwrap_or_default().to_string();
let qrels_version =
v.get("qrels_version").and_then(Value::as_str).unwrap_or_default().to_string();
let note = v.get("note").and_then(Value::as_str).map(str::to_string);
let arr = v
.get("queries")
.and_then(Value::as_array)
.ok_or_else(|| "gold set: missing `queries` array".to_string())?;
let mut queries = Vec::with_capacity(arr.len());
for (i, q) in arr.iter().enumerate() {
queries.push(parse_query(q).map_err(|e| format!("queries[{i}]: {e}"))?);
}
Ok(GoldSet { corpus_hash, qrels_version, note, queries })
}
fn parse_query(q: &Value) -> Result<GoldQuery, String> {
let query = q.get("query").and_then(Value::as_str).ok_or("missing `query`")?.to_string();
let query_id = q.get("query_id").and_then(Value::as_str).map(str::to_string);
let class_str = q.get("query_class").and_then(Value::as_str).ok_or("missing `query_class`")?;
let query_class =
QueryClass::parse(class_str).ok_or_else(|| format!("unknown query_class `{class_str}`"))?;
let mut required_evidence = Vec::new();
if let Some(units) = q.get("required_evidence").and_then(Value::as_array) {
for (j, u) in units.iter().enumerate() {
required_evidence
.push(parse_evidence(u).map_err(|e| format!("required_evidence[{j}]: {e}"))?);
}
}
let expected_top_k_doc_ids = q
.get("expected_top_k_doc_ids")
.and_then(Value::as_array)
.map(|a| a.iter().filter_map(|x| x.as_str().map(str::to_string)).collect())
.unwrap_or_default();
let relation_type = q.get("relation_type").and_then(Value::as_str).map(str::to_string);
let chain_shape = q.get("chain_shape").and_then(Value::as_str).map(str::to_string);
let source =
q.get("source").or_else(|| q.get("_source")).and_then(Value::as_str).map(str::to_string);
let answer_type = q
.get("answer_type")
.or_else(|| q.get("_answer_type"))
.and_then(Value::as_str)
.map(str::to_string);
let query_origin = match q.get("query_origin").and_then(Value::as_str) {
Some(s) => QueryOrigin::parse(s).ok_or_else(|| format!("unknown query_origin `{s}`"))?,
None => QueryOrigin::HumanDataset,
};
Ok(GoldQuery {
query,
query_id,
query_class,
required_evidence,
expected_top_k_doc_ids,
relation_type,
chain_shape,
source,
answer_type,
query_origin,
})
}
fn parse_evidence(u: &Value) -> Result<EvidenceUnit, String> {
let evidence_id =
u.get("evidence_id").and_then(Value::as_str).ok_or("missing `evidence_id`")?.to_string();
let doc_id = u.get("doc_id").and_then(Value::as_str).ok_or("missing `doc_id`")?.to_string();
let nec_str = u.get("necessity").and_then(Value::as_str).ok_or("missing `necessity`")?;
let necessity =
Necessity::parse(nec_str).ok_or_else(|| format!("unknown necessity `{nec_str}`"))?;
let locator = u.get("locator").and_then(parse_locator);
Ok(EvidenceUnit { evidence_id, doc_id, necessity, locator })
}
fn parse_locator(v: &Value) -> Option<Locator> {
let m = v.as_object()?;
let kind = m.get("kind").and_then(Value::as_str)?.to_string();
let spans = m
.get("spans")
.and_then(Value::as_array)
.map(|arr| arr.iter().filter_map(parse_span).collect::<Vec<Span>>());
Some(Locator { kind, spans })
}
fn parse_span(v: &Value) -> Option<Span> {
let m = v.as_object()?;
let doc_id = m.get("doc_id").and_then(Value::as_str)?.to_string();
let start = m.get("start").and_then(Value::as_u64)? as usize;
let end = m.get("end").and_then(Value::as_u64)? as usize;
Some(Span { doc_id, start, end })
}
pub fn required_doc_ids(q: &GoldQuery) -> BTreeSet<String> {
if q.required_evidence.is_empty() {
return q.expected_top_k_doc_ids.iter().cloned().collect();
}
q.required_evidence
.iter()
.filter(|e| e.necessity == Necessity::Required)
.map(|e| e.doc_id.clone())
.collect()
}
pub fn supporting_doc_ids(q: &GoldQuery) -> BTreeSet<String> {
q.required_evidence
.iter()
.filter(|e| e.necessity == Necessity::Supporting)
.map(|e| e.doc_id.clone())
.collect()
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct PerQueryRecall {
pub strict: f64,
pub graded: f64,
pub supporting_coverage: f64,
pub required_n: usize,
pub required_hits: usize,
}
pub fn evidence_recall_at_k(
q: &GoldQuery,
retrieved_doc_ids: &[String],
k: usize,
) -> PerQueryRecall {
let topk: HashSet<&String> = retrieved_doc_ids.iter().take(k).collect();
let required = required_doc_ids(q);
let supporting = supporting_doc_ids(q);
let required_n = required.len();
let required_hits = required.iter().filter(|d| topk.contains(d)).count();
let strict = if required_n == 0 || required_hits == required_n { 1.0 } else { 0.0 };
let graded = if required_n == 0 { 1.0 } else { required_hits as f64 / required_n as f64 };
let sup_hits = supporting.iter().filter(|d| topk.contains(d)).count();
let supporting_coverage =
if supporting.is_empty() { 0.0 } else { sup_hits as f64 / supporting.len() as f64 };
PerQueryRecall { strict, graded, supporting_coverage, required_n, required_hits }
}
pub fn negative_abstained(retrieved_doc_ids: &[String], k: usize) -> bool {
retrieved_doc_ids.iter().take(k).next().is_none()
}
#[derive(Clone, Debug, Default)]
pub struct ClassAgg {
pub n: usize,
pub strict_sum: f64,
pub graded_sum: f64,
pub supporting_sum: f64,
}
impl ClassAgg {
fn add(&mut self, m: &PerQueryRecall) {
self.n += 1;
self.strict_sum += m.strict;
self.graded_sum += m.graded;
self.supporting_sum += m.supporting_coverage;
}
pub fn strict(&self) -> f64 {
if self.n == 0 {
0.0
} else {
self.strict_sum / self.n as f64
}
}
pub fn graded(&self) -> f64 {
if self.n == 0 {
0.0
} else {
self.graded_sum / self.n as f64
}
}
pub fn supporting(&self) -> f64 {
if self.n == 0 {
0.0
} else {
self.supporting_sum / self.n as f64
}
}
}
#[derive(Clone, Debug, Default)]
pub struct NegativeAgg {
pub n: usize,
pub abstained: usize,
}
impl NegativeAgg {
pub fn false_positive_rate(&self) -> f64 {
if self.n == 0 {
0.0
} else {
(self.n - self.abstained) as f64 / self.n as f64
}
}
}
#[derive(Clone, Debug)]
pub struct KResult {
pub k: usize,
pub overall: ClassAgg,
pub per_class: BTreeMap<QueryClass, ClassAgg>,
pub negative: NegativeAgg,
}
impl KResult {
fn new(k: usize) -> Self {
Self {
k,
overall: ClassAgg::default(),
per_class: BTreeMap::new(),
negative: NegativeAgg::default(),
}
}
}
pub fn evaluate_gold_set<F>(
gold: &GoldSet,
ladder: &[usize],
mut retrieve: F,
) -> Result<BTreeMap<usize, KResult>, String>
where
F: FnMut(&GoldQuery) -> Result<Vec<String>, String>,
{
let mut out: BTreeMap<usize, KResult> = ladder.iter().map(|&k| (k, KResult::new(k))).collect();
for q in &gold.queries {
let retrieved = retrieve(q)?;
for &k in ladder {
let r = out.get_mut(&k).expect("ladder key");
if q.query_class == QueryClass::Negative {
r.negative.n += 1;
if negative_abstained(&retrieved, k) {
r.negative.abstained += 1;
}
} else {
let m = evidence_recall_at_k(q, &retrieved, k);
r.overall.add(&m);
r.per_class.entry(q.query_class).or_default().add(&m);
}
}
}
Ok(out)
}
pub fn validate_gold_set(gold: &GoldSet) -> Vec<String> {
let mut issues = Vec::new();
if gold.corpus_hash.trim().is_empty() {
issues.push("corpus_hash missing (pinning principle §(f))".to_string());
} else if gold.corpus_hash == UNPINNED_PLACEHOLDER {
issues.push(format!(
"corpus_hash is the `{UNPINNED_PLACEHOLDER}` placeholder — fixture-only, NOT pinned to a frozen snapshot"
));
}
if gold.qrels_version.trim().is_empty() {
issues.push("qrels_version missing (pinning principle §(f))".to_string());
}
let mut seen_qids: HashSet<&str> = HashSet::new();
for (i, q) in gold.queries.iter().enumerate() {
let qid = q.query_id.as_deref().unwrap_or("<no query_id>");
let where_ = format!("query[{i}] ({qid})");
if q.query.trim().is_empty() {
issues.push(format!("{where_}: empty query text"));
}
if let Some(id) = q.query_id.as_deref() {
if !seen_qids.insert(id) {
issues.push(format!("{where_}: duplicate query_id `{id}`"));
}
}
let mut seen_ev: HashSet<&str> = HashSet::new();
for e in &q.required_evidence {
if e.doc_id.trim().is_empty() {
issues.push(format!("{where_}: evidence `{}` has empty doc_id", e.evidence_id));
}
if !seen_ev.insert(&e.evidence_id) {
issues.push(format!("{where_}: duplicate evidence_id `{}`", e.evidence_id));
}
for s in e.locator.iter().flat_map(|l| l.spans.iter().flatten()) {
if s.end < s.start {
issues.push(format!(
"{where_}: evidence `{}` span has end<start ({}..{})",
e.evidence_id, s.start, s.end
));
}
if s.doc_id != e.doc_id {
issues.push(format!(
"{where_}: evidence `{}` span doc_id `{}` != evidence doc_id `{}`",
e.evidence_id, s.doc_id, e.doc_id
));
}
}
}
let req = required_doc_ids(q);
match q.query_class {
QueryClass::Negative => {
if !req.is_empty() {
issues.push(format!(
"{where_}: negative class must have an EMPTY required denominator (abstention)"
));
}
}
_ => {
if req.is_empty() {
issues.push(format!(
"{where_}: non-negative class `{}` has an EMPTY required denominator",
q.query_class.label()
));
}
}
}
}
issues
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub enum RetrievalMode {
RrfHybrid,
VectorOnly,
RerankStub,
FtsWriteCursor,
Bm25Fts,
}
impl RetrievalMode {
pub fn label(&self) -> &'static str {
match self {
Self::RrfHybrid => "rrf_hybrid",
Self::VectorOnly => "vector_only",
Self::RerankStub => "rerank_stub",
Self::FtsWriteCursor => "fts_write_cursor",
Self::Bm25Fts => "bm25_fts",
}
}
pub fn is_runnable_now(&self) -> bool {
matches!(self, Self::RrfHybrid | Self::VectorOnly | Self::RerankStub)
}
}
pub const RUNNABLE_NOW_MODES: [RetrievalMode; 3] =
[RetrievalMode::RrfHybrid, RetrievalMode::VectorOnly, RetrievalMode::RerankStub];
pub fn run_mode_bodies(
engine: &Engine,
query: &str,
mode: RetrievalMode,
) -> Result<Vec<String>, String> {
match mode {
RetrievalMode::RrfHybrid => search_bodies(engine, query),
RetrievalMode::VectorOnly => {
engine.set_vector_stage_only_for_test(true);
let r = search_bodies(engine, query);
engine.set_vector_stage_only_for_test(false);
r
}
RetrievalMode::RerankStub => {
let res = engine.search(query).map_err(|e| format!("search: {e:?}"))?;
Ok(rerank_fused(query, res.results, 0, 0.3, 0).into_iter().map(|h| h.body).collect())
}
RetrievalMode::FtsWriteCursor | RetrievalMode::Bm25Fts => Err(format!(
"mode `{}` deferred: TODO(COR-2-freeze) — needs harness FTS5 SQL + frozen corpus",
mode.label()
)),
}
}
fn search_bodies(engine: &Engine, query: &str) -> Result<Vec<String>, String> {
Ok(engine
.search(query)
.map_err(|e| format!("search: {e:?}"))?
.results
.into_iter()
.map(|h| h.body)
.collect())
}
pub struct ExperimentResult {
pub fanout: usize,
pub per_mode: BTreeMap<RetrievalMode, BTreeMap<usize, KResult>>,
pub deferred_modes: Vec<RetrievalMode>,
}
pub fn run_experiment(
engine: &Engine,
gold: &GoldSet,
body_to_doc_id: &HashMap<String, String>,
modes: &[RetrievalMode],
ladder: &[usize],
) -> Result<ExperimentResult, String> {
let deepest = ladder.iter().copied().max().unwrap_or(HEADLINE_K);
let fanout = deepest.max(DEFAULT_FANOUT);
engine.set_search_limit_for_test(fanout);
let mut per_mode = BTreeMap::new();
let mut deferred_modes = Vec::new();
for &mode in modes {
if !mode.is_runnable_now() {
deferred_modes.push(mode);
continue;
}
let result = evaluate_gold_set(gold, ladder, |q| {
let bodies = run_mode_bodies(engine, &q.query, mode)?;
Ok(map_bodies_to_doc_ids(&bodies, body_to_doc_id))
})?;
per_mode.insert(mode, result);
}
Ok(ExperimentResult { fanout, per_mode, deferred_modes })
}
pub fn map_bodies_to_doc_ids(bodies: &[String], map: &HashMap<String, String>) -> Vec<String> {
bodies.iter().filter_map(|b| map.get(b).cloned()).collect()
}
fn round4(x: f64) -> f64 {
(x * 10_000.0).round() / 10_000.0
}
pub fn experiment_to_json(gold: &GoldSet, result: &ExperimentResult) -> Value {
let per_mode: serde_json::Map<String, Value> = result
.per_mode
.iter()
.map(|(mode, by_k)| {
let k_obj: serde_json::Map<String, Value> = by_k
.iter()
.map(|(k, r)| {
let per_class: serde_json::Map<String, Value> = r
.per_class
.iter()
.map(|(cls, agg)| {
(
cls.label().to_string(),
json!({
"n": agg.n,
"strict_evidence_recall": round4(agg.strict()),
"graded_evidence_recall": round4(agg.graded()),
"supporting_coverage": round4(agg.supporting()),
}),
)
})
.collect();
(
k.to_string(),
json!({
"overall": {
"n": r.overall.n,
"strict_evidence_recall": round4(r.overall.strict()),
"graded_evidence_recall": round4(r.overall.graded()),
"supporting_coverage": round4(r.overall.supporting()),
},
"per_class": per_class,
"negative_class": {
"n": r.negative.n,
"abstained": r.negative.abstained,
"false_positive_rate": round4(r.negative.false_positive_rate()),
},
}),
)
})
.collect();
(mode.label().to_string(), Value::Object(k_obj))
})
.collect();
json!({
"_comment": "IR-B (IR-1 Phase 2) Evidence Recall@K — STRUCTURE only. \
No thresholds, no verdict (Phase 4 / IR-2 / HITL). Real-corpus \
numbers are DEFERRED to the COR-2 freeze (IR-C).",
"measure": "evidence_recall_at_k",
"headline_k": HEADLINE_K,
"k_ladder": K_LADDER,
"fanout": result.fanout,
"corpus_hash": gold.corpus_hash,
"qrels_version": gold.qrels_version,
"query_count": gold.queries.len(),
"deferred_modes": result.deferred_modes.iter().map(|m| m.label()).collect::<Vec<_>>(),
"per_mode": per_mode,
})
}