use std::collections::{HashMap, HashSet};
use std::error::Error;
use std::sync::atomic::{AtomicUsize, Ordering};
use serde_json::json;
use velesdb_memory::{ColumnFilter, ColumnOp, FusionOptions};
use crate::dataset::{Category, Qa};
use crate::dump::{self, QuestionTrace};
use crate::ingest::Store;
use crate::judge;
use crate::ollama_gen::{Generator, TokenUsage};
static GRAPH_INJECTED: AtomicUsize = AtomicUsize::new(0);
static GRAPH_ACTIVE_CONTEXTS: AtomicUsize = AtomicUsize::new(0);
pub fn graph_activity() -> (usize, usize) {
(
GRAPH_INJECTED.load(Ordering::Relaxed),
GRAPH_ACTIVE_CONTEXTS.load(Ordering::Relaxed),
)
}
#[allow(clippy::struct_excessive_bools)]
#[derive(Clone, Copy)]
pub struct EvalCfg {
pub k: usize,
pub graph_boost: f64,
pub hops: usize,
pub multihop_only: bool,
pub idf_weight: bool,
pub seed_breadth: usize,
pub date_context: bool,
pub date_routed: bool,
pub temporal_scaffold: bool,
pub cot: bool,
pub bm25: bool,
pub claude_judge: bool,
pub claude_gen: bool,
pub use_shipped_api: bool,
}
#[derive(Clone)]
pub(crate) struct RetrievedFact {
pub(crate) id: u64,
pub(crate) text: String,
pub(crate) dia_ids: Vec<String>,
pub(crate) score: f64,
pub(crate) graph_weight: f64,
pub(crate) ts: i64,
}
pub struct ModeResult {
pub evidence_hit: bool,
pub correct: bool,
pub f1: Option<f64>,
}
pub fn evaluate(
store: &Store,
generator: &Generator,
qa: &Qa,
cfg: EvalCfg,
graph_on: bool,
trace: Option<QuestionTrace<'_>>,
) -> Result<ModeResult, Box<dyn Error>> {
let (facts, raw) = retrieve(
store,
&qa.question,
cfg,
graph_on,
qa.category,
trace.is_some(),
)?;
let evidence_hit = any_evidence_hit(&facts, &qa.evidence);
let dated: Vec<(i64, String)> = facts.iter().map(|f| (f.ts, f.text.clone())).collect();
let temporal = temporal_flags(&qa.question, cfg);
let now_ts = store.latest_ts();
let (candidate, usage) = judge::answer(
generator,
&qa.question,
&dated,
temporal.date_on,
now_ts,
temporal.scaffold_on,
cfg.cot,
cfg.claude_gen,
)?;
let (correct, f1) = score(generator, qa, &candidate, cfg.claude_judge)?;
let result = ModeResult {
evidence_hit,
correct,
f1,
};
maybe_dump(
trace,
&record_inputs(
qa, cfg, graph_on, &raw, &facts, &candidate, usage, &result, &temporal, now_ts,
),
)?;
Ok(result)
}
#[allow(clippy::too_many_arguments)]
fn record_inputs<'a>(
qa: &'a Qa,
cfg: EvalCfg,
graph_on: bool,
raw: &'a [RetrievedFact],
reranked: &'a [RetrievedFact],
candidate: &'a str,
usage: Option<TokenUsage>,
verdict: &ModeResult,
temporal: &TemporalFlags,
latest_ts: i64,
) -> dump::RecordInputs<'a> {
dump::RecordInputs {
qa,
cfg,
graph_on,
raw,
reranked,
candidate,
usage,
correct: verdict.correct,
f1: verdict.f1,
evidence_hit: verdict.evidence_hit,
date_on: temporal.date_on,
scaffold_on: temporal.scaffold_on,
is_temporal_trigger: temporal.is_temporal_trigger,
latest_ts,
}
}
fn any_evidence_hit(facts: &[RetrievedFact], evidence: &[String]) -> bool {
facts
.iter()
.any(|f| f.dia_ids.iter().any(|id| evidence.contains(id)))
}
struct TemporalFlags {
is_temporal_trigger: bool,
date_on: bool,
scaffold_on: bool,
}
fn temporal_flags(question: &str, cfg: EvalCfg) -> TemporalFlags {
let is_temporal_trigger = is_temporal_question(question);
let date_on = cfg.date_context && (!cfg.date_routed || is_temporal_trigger);
let scaffold_on = cfg.temporal_scaffold && date_on;
TemporalFlags {
is_temporal_trigger,
date_on,
scaffold_on,
}
}
fn maybe_dump(
trace: Option<QuestionTrace<'_>>,
inputs: &dump::RecordInputs<'_>,
) -> Result<(), Box<dyn Error>> {
let Some(trace) = trace else {
return Ok(());
};
dump::write_record(trace, inputs)
}
pub fn retrieved_dia_ids(
store: &Store,
question: &str,
cfg: EvalCfg,
graph_on: bool,
category: Category,
) -> Result<Vec<Vec<String>>, Box<dyn Error>> {
Ok(retrieve(store, question, cfg, graph_on, category, false)?
.0
.into_iter()
.map(|f| f.dia_ids)
.collect())
}
fn is_temporal_question(question: &str) -> bool {
const CUES: &[&str] = &[
"when ",
"what year",
"which year",
"what month",
"which month",
"what date",
"what time",
"how long",
"how many days",
"how many weeks",
"how many months",
"how many years",
" ago",
"how often",
"how frequently",
];
let q = question.to_lowercase();
CUES.iter().any(|cue| q.contains(cue))
}
fn score(
generator: &Generator,
qa: &Qa,
candidate: &str,
claude_judge: bool,
) -> Result<(bool, Option<f64>), Box<dyn Error>> {
if qa.category.is_adversarial() {
return Ok((judge::abstained(candidate), None));
}
let Some(gold) = judge::gold_answer(qa) else {
return Ok((false, Some(0.0)));
};
let correct = judge::judge_correct(generator, &qa.question, gold, candidate, claude_judge)?;
Ok((correct, Some(judge::f1(candidate, gold))))
}
const POOL_FACTOR: usize = 8;
const POOL_MIN: usize = 64;
fn retrieve(
store: &Store,
question: &str,
cfg: EvalCfg,
graph_on: bool,
category: Category,
want_raw: bool,
) -> Result<(Vec<RetrievedFact>, Vec<RetrievedFact>), Box<dyn Error>> {
let use_graph = graph_on && (!cfg.multihop_only || matches!(category, Category::MultiHop));
if !use_graph {
let pool = vector_pool(store, question, pool_size(cfg.k), &[])?;
let raw = raw_if_wanted(&pool, &[], want_raw);
if cfg.bm25 {
return Ok((rrf_fuse(pool, store, question, cfg.k), raw));
}
return Ok((pool.into_iter().take(cfg.k).collect(), raw));
}
if cfg.use_shipped_api {
let opts = FusionOptions {
hops: cfg.hops,
graph_boost: cfg.graph_boost,
pool: None,
};
let fused = shipped_fused(store, question, cfg.k, opts)?;
let raw = raw_if_wanted(&fused, &[], want_raw);
return Ok((fused, raw));
}
let filters = temporal_filters(question);
let pool = vector_pool(store, question, pool_size(cfg.k), &filters)?;
let reached = graph_reached(store, question, cfg)?;
let raw = raw_if_wanted(&pool, &reached, want_raw);
Ok((fuse(pool, &reached, cfg), raw))
}
fn shipped_fused(
store: &Store,
question: &str,
k: usize,
opts: FusionOptions,
) -> Result<Vec<RetrievedFact>, Box<dyn Error>> {
let hits = store.svc.recall_fused(question, k, None, opts)?;
Ok(hits
.into_iter()
.filter(|hit| store.is_fact(hit.id))
.map(|hit| RetrievedFact {
id: hit.id,
text: hit.content,
dia_ids: store.dia_ids(hit.id).to_vec(),
score: f64::from(hit.score),
graph_weight: 0.0,
ts: store.fact_ts(hit.id),
})
.collect())
}
fn raw_if_wanted(
pool: &[RetrievedFact],
reached: &[RetrievedFact],
want_raw: bool,
) -> Vec<RetrievedFact> {
if !want_raw {
return Vec::new();
}
let mut all = pool.to_vec();
let present: HashSet<u64> = all.iter().map(|f| f.id).collect();
all.extend(reached.iter().filter(|f| !present.contains(&f.id)).cloned());
all
}
#[allow(clippy::similar_names)]
fn rrf_fuse(
pool: Vec<RetrievedFact>,
store: &Store,
question: &str,
k: usize,
) -> Vec<RetrievedFact> {
const RRF_K: f64 = 60.0;
const BM25_DEPTH: usize = 64;
let dense_rank: HashMap<u64, usize> = pool.iter().enumerate().map(|(i, f)| (f.id, i)).collect();
let bm25_rank: HashMap<u64, usize> = store
.bm25_search(question)
.into_iter()
.take(BM25_DEPTH)
.enumerate()
.map(|(i, id)| (id, i))
.collect();
let mut ids: HashSet<u64> = dense_rank.keys().copied().collect();
ids.extend(bm25_rank.keys().copied());
let mut scored: Vec<(u64, f64)> = ids
.into_iter()
.map(|id| {
let mut score = 0.0;
if let Some(&r) = dense_rank.get(&id) {
score += 1.0 / (RRF_K + rank_f(r));
}
if let Some(&r) = bm25_rank.get(&id) {
score += 1.0 / (RRF_K + rank_f(r));
}
(id, score)
})
.collect();
scored.sort_by(|a, b| b.1.total_cmp(&a.1));
scored.truncate(k);
let mut by_id: HashMap<u64, RetrievedFact> = pool.into_iter().map(|f| (f.id, f)).collect();
scored
.into_iter()
.map(|(id, _)| {
by_id.remove(&id).unwrap_or_else(|| RetrievedFact {
id,
text: store.fact_text(id).to_string(),
dia_ids: store.dia_ids(id).to_vec(),
score: 0.0,
graph_weight: 0.0,
ts: store.fact_ts(id),
})
})
.collect()
}
fn rank_f(rank: usize) -> f64 {
f64::from(u32::try_from(rank).unwrap_or(u32::MAX))
}
fn temporal_filters(question: &str) -> Vec<ColumnFilter> {
let Some(year) = find_year(question) else {
return Vec::new();
};
let low = year * 10_000;
vec![
ColumnFilter {
field: "ts".to_string(),
op: ColumnOp::Ge,
value: json!(low),
},
ColumnFilter {
field: "ts".to_string(),
op: ColumnOp::Le,
value: json!(low + 9_999),
},
]
}
fn find_year(question: &str) -> Option<i64> {
question
.split(|c: char| !c.is_ascii_digit())
.filter_map(|token| token.parse::<i64>().ok())
.find(|year| (1990..=2099).contains(year))
}
fn pool_size(k: usize) -> usize {
k.saturating_mul(POOL_FACTOR).max(POOL_MIN)
}
fn vector_pool(
store: &Store,
question: &str,
want: usize,
filters: &[ColumnFilter],
) -> Result<Vec<RetrievedFact>, Box<dyn Error>> {
let oversample = want.saturating_mul(2).saturating_add(16);
let hits = if filters.is_empty() {
store.svc.recall(question, oversample, None)?
} else {
store.svc.recall_where(question, oversample, filters)?
};
let mut out = Vec::with_capacity(want);
for hit in hits {
if store.is_fact(hit.id) {
out.push(RetrievedFact {
id: hit.id,
text: hit.content,
dia_ids: store.dia_ids(hit.id).to_vec(),
score: f64::from(hit.score),
graph_weight: 0.0,
ts: store.fact_ts(hit.id),
});
if out.len() == want {
break;
}
}
}
Ok(out)
}
fn graph_reached(
store: &Store,
question: &str,
cfg: EvalCfg,
) -> Result<Vec<RetrievedFact>, Box<dyn Error>> {
if cfg.seed_breadth > 1 {
return graph_reached_multiseed(store, question, cfg);
}
let explanation = store.svc.why(question, cfg.hops, None)?;
let seed_entities = if cfg.idf_weight {
seed_entity_set(store, question, cfg.k)?
} else {
HashSet::new()
};
Ok(explanation
.nodes
.into_iter()
.filter(|node| node.hop >= 1 && store.is_fact(node.id))
.map(|node| {
let weight = if cfg.idf_weight {
connection_weight(store, node.id, &seed_entities)
} else {
1.0
};
RetrievedFact {
id: node.id,
dia_ids: store.dia_ids(node.id).to_vec(),
text: node.content,
score: 0.0,
graph_weight: weight,
ts: store.fact_ts(node.id),
}
})
.filter(|fact| fact.graph_weight > 0.0)
.collect())
}
fn graph_reached_multiseed(
store: &Store,
question: &str,
cfg: EvalCfg,
) -> Result<Vec<RetrievedFact>, Box<dyn Error>> {
let seeds = top_vector_fact_ids(store, question, cfg.seed_breadth)?;
let seed_set: HashSet<u64> = seeds.iter().copied().collect();
let mut weights: HashMap<u64, f64> = HashMap::new();
for &sid in &seeds {
for &eid in store.fact_entity_ids(sid) {
let w = if cfg.idf_weight {
store.entity_idf(eid)
} else {
1.0
};
if w <= 0.0 {
continue;
}
for &fid in store.entity_fact_ids(eid) {
if seed_set.contains(&fid) {
continue;
}
let slot = weights.entry(fid).or_insert(0.0);
*slot = slot.max(w);
}
}
}
Ok(weights
.into_iter()
.map(|(id, weight)| RetrievedFact {
id,
dia_ids: store.dia_ids(id).to_vec(),
text: store.fact_text(id).to_string(),
score: 0.0,
graph_weight: weight,
ts: store.fact_ts(id),
})
.collect())
}
fn top_vector_fact_ids(
store: &Store,
question: &str,
n: usize,
) -> Result<Vec<u64>, Box<dyn Error>> {
let hits = store
.svc
.recall(question, n.saturating_mul(3).saturating_add(8), None)?;
let mut ids = Vec::with_capacity(n);
for hit in hits {
if !store.is_fact(hit.id) {
continue;
}
ids.push(hit.id);
if ids.len() == n {
break;
}
}
Ok(ids)
}
fn seed_entity_set(
store: &Store,
question: &str,
k: usize,
) -> Result<HashSet<u64>, Box<dyn Error>> {
let mut entities = HashSet::new();
for fid in top_vector_fact_ids(store, question, k)? {
entities.extend(store.fact_entity_ids(fid).iter().copied());
}
Ok(entities)
}
fn connection_weight(store: &Store, fact_id: u64, seed_entities: &HashSet<u64>) -> f64 {
store
.fact_entity_ids(fact_id)
.iter()
.filter(|eid| seed_entities.contains(eid))
.map(|eid| store.entity_idf(*eid))
.fold(0.0, f64::max)
}
fn fuse(pool: Vec<RetrievedFact>, reached: &[RetrievedFact], cfg: EvalCfg) -> Vec<RetrievedFact> {
let weights: HashMap<u64, f64> = reached.iter().map(|f| (f.id, f.graph_weight)).collect();
let vector_top: HashSet<u64> = pool.iter().take(cfg.k).map(|f| f.id).collect();
let max_score = pool
.iter()
.map(|f| f.score)
.fold(f64::MIN, f64::max)
.max(f64::EPSILON);
let mut candidates: Vec<RetrievedFact> = pool;
let present: HashSet<u64> = candidates.iter().map(|f| f.id).collect();
candidates.extend(reached.iter().filter(|f| !present.contains(&f.id)).cloned());
for candidate in &mut candidates {
if let Some(&weight) = weights.get(&candidate.id) {
candidate.graph_weight = weight;
}
}
candidates.sort_by(|a, b| {
fused_score(b, &weights, max_score, cfg)
.total_cmp(&fused_score(a, &weights, max_score, cfg))
});
candidates.truncate(cfg.k);
record_injection(
candidates
.iter()
.filter(|f| !vector_top.contains(&f.id))
.count(),
);
candidates
}
fn fused_score(
fact: &RetrievedFact,
weights: &HashMap<u64, f64>,
max_score: f64,
cfg: EvalCfg,
) -> f64 {
let weight = weights.get(&fact.id).copied().unwrap_or(0.0);
fact.score / max_score + cfg.graph_boost * weight
}
fn record_injection(injected: usize) {
if injected > 0 {
GRAPH_INJECTED.fetch_add(injected, Ordering::Relaxed);
GRAPH_ACTIVE_CONTEXTS.fetch_add(1, Ordering::Relaxed);
}
}
#[cfg(test)]
#[path = "eval/eval_tests.rs"]
mod tests;