Skip to main content

tsift_quality/
dci_benchmark.rs

1use anyhow::{Context, Result, bail};
2use serde::{Deserialize, Serialize};
3use std::collections::{BTreeMap, BTreeSet};
4
5const EXPECTED_STRATEGIES: [&str; 3] = ["exact_chained_rg", "lexical_bm25", "hybrid"];
6
7#[derive(Debug, Clone, Deserialize)]
8pub struct DciBenchmarkFixture {
9    #[serde(default)]
10    pub description: Option<String>,
11    #[serde(default)]
12    pub expected_strategies: Option<Vec<String>>,
13    pub tasks: Vec<DciBenchmarkTask>,
14}
15
16#[derive(Debug, Clone, Deserialize)]
17pub struct DciBenchmarkTask {
18    pub id: String,
19    #[serde(default)]
20    pub label: Option<String>,
21    #[serde(default)]
22    pub target: Option<String>,
23    pub runs: Vec<DciBenchmarkRun>,
24}
25
26#[derive(Debug, Clone, Deserialize)]
27pub struct DciBenchmarkRun {
28    pub strategy: String,
29    pub localized: bool,
30    pub tool_calls: f64,
31    pub latency_ms: f64,
32    pub estimated_tokens: f64,
33    #[serde(default)]
34    pub useful_hits: Option<f64>,
35    #[serde(default)]
36    pub output_tokens: Option<f64>,
37    #[serde(default)]
38    pub zero_output: bool,
39    #[serde(default)]
40    pub notes: Option<String>,
41}
42
43#[derive(Debug, Clone, PartialEq, Serialize)]
44pub struct DciBenchmarkReport {
45    #[serde(skip_serializing_if = "Option::is_none")]
46    pub description: Option<String>,
47    pub tasks_loaded: usize,
48    pub strategies_compared: usize,
49    pub expected_strategies: Vec<String>,
50    pub strategy_summaries: Vec<DciStrategySummary>,
51    pub task_rows: Vec<DciTaskRow>,
52    #[serde(skip_serializing_if = "Option::is_none")]
53    pub memory_retrieval_gate: Option<MemoryRetrievalGate>,
54    #[serde(skip_serializing_if = "Vec::is_empty", default)]
55    pub warnings: Vec<String>,
56}
57
58#[derive(Debug, Clone, PartialEq, Serialize)]
59pub struct DciStrategySummary {
60    pub strategy: String,
61    pub task_runs: usize,
62    pub localized: usize,
63    pub localization_rate: f64,
64    pub useful_hits: f64,
65    pub avg_useful_hits: f64,
66    pub zero_output_failures: usize,
67    pub zero_output_rate: f64,
68    pub avg_tool_calls: f64,
69    pub avg_latency_ms: f64,
70    pub avg_estimated_tokens: f64,
71    pub avg_output_tokens: f64,
72    pub rank: usize,
73}
74
75#[derive(Debug, Clone, PartialEq, Serialize)]
76pub struct DciTaskRow {
77    pub task_id: String,
78    #[serde(skip_serializing_if = "Option::is_none")]
79    pub label: Option<String>,
80    #[serde(skip_serializing_if = "Option::is_none")]
81    pub target: Option<String>,
82    pub best_localization: Vec<String>,
83    pub most_useful_hits: Vec<String>,
84    pub lowest_tool_calls: Option<String>,
85    pub lowest_latency: Option<String>,
86    pub lowest_token_budget: Option<String>,
87    pub lowest_output_tokens: Option<String>,
88    pub zero_output_failures: Vec<String>,
89}
90
91#[derive(Debug, Clone, PartialEq, Serialize)]
92pub struct MemoryRetrievalGate {
93    pub decision: String,
94    pub baseline_strategy: String,
95    pub candidate_strategies: Vec<String>,
96    pub min_avg_useful_hits: f64,
97    pub max_zero_output_failures: usize,
98    pub rows: Vec<MemoryRetrievalGateRow>,
99    #[serde(skip_serializing_if = "Vec::is_empty", default)]
100    pub diagnostics: Vec<String>,
101}
102
103#[derive(Debug, Clone, PartialEq, Serialize)]
104pub struct MemoryRetrievalGateRow {
105    pub strategy: String,
106    pub avg_useful_hits: f64,
107    pub zero_output_failures: usize,
108    pub useful_hits_pass: bool,
109    pub zero_output_pass: bool,
110    pub status: String,
111}
112
113#[derive(Default)]
114struct Accumulator {
115    task_runs: usize,
116    localized: usize,
117    useful_hits: f64,
118    output_tokens: f64,
119    zero_output_failures: usize,
120    tool_calls: f64,
121    latency_ms: f64,
122    estimated_tokens: f64,
123}
124
125pub fn compute(input: &str) -> Result<DciBenchmarkReport> {
126    let fixture: DciBenchmarkFixture =
127        serde_json::from_str(input).context("parsing dci-benchmark fixture as JSON")?;
128    if fixture.tasks.is_empty() {
129        bail!("dci-benchmark fixture did not contain any tasks");
130    }
131    let expected_strategies = fixture.expected_strategies.clone().unwrap_or_else(|| {
132        EXPECTED_STRATEGIES
133            .iter()
134            .map(|strategy| strategy.to_string())
135            .collect()
136    });
137
138    let mut warnings = Vec::new();
139    let mut accumulators = BTreeMap::<String, Accumulator>::new();
140    let mut seen_strategies = BTreeSet::<String>::new();
141    let mut task_rows = Vec::new();
142
143    for task in &fixture.tasks {
144        if task.runs.is_empty() {
145            warnings.push(format!(
146                "task {} did not include any strategy runs",
147                task.id
148            ));
149            continue;
150        }
151
152        let mut localized = Vec::new();
153        let mut most_useful_hits = Vec::new();
154        let mut best_useful_hits = f64::NEG_INFINITY;
155        let mut lowest_tool_calls: Option<&DciBenchmarkRun> = None;
156        let mut lowest_latency: Option<&DciBenchmarkRun> = None;
157        let mut lowest_tokens: Option<&DciBenchmarkRun> = None;
158        let mut lowest_output_tokens: Option<&DciBenchmarkRun> = None;
159        let mut zero_output_failures = Vec::new();
160
161        for run in &task.runs {
162            if !run.tool_calls.is_finite() || run.tool_calls < 0.0 {
163                bail!(
164                    "task {} strategy {} has invalid tool_calls",
165                    task.id,
166                    run.strategy
167                );
168            }
169            if !run.latency_ms.is_finite() || run.latency_ms < 0.0 {
170                bail!(
171                    "task {} strategy {} has invalid latency_ms",
172                    task.id,
173                    run.strategy
174                );
175            }
176            if !run.estimated_tokens.is_finite() || run.estimated_tokens < 0.0 {
177                bail!(
178                    "task {} strategy {} has invalid estimated_tokens",
179                    task.id,
180                    run.strategy
181                );
182            }
183            if let Some(useful_hits) = run.useful_hits
184                && (!useful_hits.is_finite() || useful_hits < 0.0)
185            {
186                bail!(
187                    "task {} strategy {} has invalid useful_hits",
188                    task.id,
189                    run.strategy
190                );
191            }
192            if let Some(output_tokens) = run.output_tokens
193                && (!output_tokens.is_finite() || output_tokens < 0.0)
194            {
195                bail!(
196                    "task {} strategy {} has invalid output_tokens",
197                    task.id,
198                    run.strategy
199                );
200            }
201
202            seen_strategies.insert(run.strategy.clone());
203            if run.localized {
204                localized.push(run.strategy.clone());
205            }
206            let useful_hits = run_useful_hits(run);
207            if useful_hits > best_useful_hits {
208                best_useful_hits = useful_hits;
209                most_useful_hits.clear();
210                most_useful_hits.push(run.strategy.clone());
211            } else if (useful_hits - best_useful_hits).abs() < f64::EPSILON {
212                most_useful_hits.push(run.strategy.clone());
213            }
214            if run.zero_output {
215                zero_output_failures.push(run.strategy.clone());
216            }
217
218            lowest_tool_calls = choose_lowest(lowest_tool_calls, run, |value| value.tool_calls);
219            lowest_latency = choose_lowest(lowest_latency, run, |value| value.latency_ms);
220            lowest_tokens = choose_lowest(lowest_tokens, run, |value| value.estimated_tokens);
221            lowest_output_tokens = choose_lowest(lowest_output_tokens, run, run_output_tokens);
222
223            let acc = accumulators.entry(run.strategy.clone()).or_default();
224            acc.task_runs += 1;
225            acc.localized += usize::from(run.localized);
226            acc.useful_hits += useful_hits;
227            acc.output_tokens += run_output_tokens(run);
228            acc.zero_output_failures += usize::from(run.zero_output);
229            acc.tool_calls += run.tool_calls;
230            acc.latency_ms += run.latency_ms;
231            acc.estimated_tokens += run.estimated_tokens;
232        }
233
234        task_rows.push(DciTaskRow {
235            task_id: task.id.clone(),
236            label: task.label.clone(),
237            target: task.target.clone(),
238            best_localization: localized,
239            most_useful_hits,
240            lowest_tool_calls: lowest_tool_calls.map(|run| run.strategy.clone()),
241            lowest_latency: lowest_latency.map(|run| run.strategy.clone()),
242            lowest_token_budget: lowest_tokens.map(|run| run.strategy.clone()),
243            lowest_output_tokens: lowest_output_tokens.map(|run| run.strategy.clone()),
244            zero_output_failures,
245        });
246    }
247
248    for expected in &expected_strategies {
249        if !seen_strategies.contains(expected) {
250            warnings.push(format!("expected strategy {expected} was not present"));
251        }
252    }
253
254    let mut summaries = accumulators
255        .into_iter()
256        .map(|(strategy, acc)| {
257            let task_runs = acc.task_runs.max(1);
258            DciStrategySummary {
259                strategy,
260                task_runs: acc.task_runs,
261                localized: acc.localized,
262                localization_rate: acc.localized as f64 / task_runs as f64,
263                useful_hits: acc.useful_hits,
264                avg_useful_hits: acc.useful_hits / task_runs as f64,
265                zero_output_failures: acc.zero_output_failures,
266                zero_output_rate: acc.zero_output_failures as f64 / task_runs as f64,
267                avg_tool_calls: acc.tool_calls / task_runs as f64,
268                avg_latency_ms: acc.latency_ms / task_runs as f64,
269                avg_estimated_tokens: acc.estimated_tokens / task_runs as f64,
270                avg_output_tokens: acc.output_tokens / task_runs as f64,
271                rank: 0,
272            }
273        })
274        .collect::<Vec<_>>();
275    summaries.sort_by(strategy_rank);
276    for (index, summary) in summaries.iter_mut().enumerate() {
277        summary.rank = index + 1;
278    }
279    let memory_retrieval_gate =
280        build_memory_retrieval_gate(&summaries, &expected_strategies, fixture.tasks.len());
281
282    Ok(DciBenchmarkReport {
283        description: fixture.description,
284        tasks_loaded: fixture.tasks.len(),
285        strategies_compared: summaries.len(),
286        expected_strategies,
287        strategy_summaries: summaries,
288        task_rows,
289        memory_retrieval_gate,
290        warnings,
291    })
292}
293
294fn build_memory_retrieval_gate(
295    summaries: &[DciStrategySummary],
296    expected_strategies: &[String],
297    tasks_loaded: usize,
298) -> Option<MemoryRetrievalGate> {
299    const BASELINE: &str = "claude_mem_api";
300    const CANDIDATES: [&str; 2] = ["tsift_session_review_context_pack", "graph_db_related"];
301
302    let expected = expected_strategies
303        .iter()
304        .map(String::as_str)
305        .collect::<BTreeSet<_>>();
306    if !expected.contains(BASELINE)
307        || !CANDIDATES
308            .iter()
309            .all(|candidate| expected.contains(candidate))
310    {
311        return None;
312    }
313
314    let baseline = summaries
315        .iter()
316        .find(|summary| summary.strategy == BASELINE);
317    let min_avg_useful_hits = baseline
318        .map(|summary| summary.avg_useful_hits)
319        .unwrap_or_default();
320    let baseline_zero_output_failures = baseline
321        .map(|summary| summary.zero_output_failures)
322        .unwrap_or(tasks_loaded);
323    let max_zero_output_failures = if baseline_zero_output_failures == 0 {
324        0
325    } else {
326        baseline_zero_output_failures - 1
327    };
328
329    let mut diagnostics = Vec::new();
330    if baseline.is_none() {
331        diagnostics.push(format!(
332            "baseline strategy {BASELINE} was not present; memory retrieval gate blocks"
333        ));
334    }
335
336    let mut rows = Vec::new();
337    for candidate in CANDIDATES {
338        match summaries
339            .iter()
340            .find(|summary| summary.strategy == candidate)
341        {
342            Some(summary) => {
343                let useful_hits_pass =
344                    summary.avg_useful_hits + f64::EPSILON >= min_avg_useful_hits;
345                let zero_output_pass = if baseline_zero_output_failures == 0 {
346                    summary.zero_output_failures == 0
347                } else {
348                    summary.zero_output_failures < baseline_zero_output_failures
349                };
350                let status = if useful_hits_pass && zero_output_pass {
351                    "pass"
352                } else {
353                    "block"
354                };
355                if !useful_hits_pass {
356                    diagnostics.push(format!(
357                        "{candidate} avg useful hits {} is below baseline {}",
358                        format_number(summary.avg_useful_hits),
359                        format_number(min_avg_useful_hits)
360                    ));
361                }
362                if !zero_output_pass {
363                    diagnostics.push(format!(
364                        "{candidate} zero-output failures {} did not improve on baseline {}",
365                        summary.zero_output_failures, baseline_zero_output_failures
366                    ));
367                }
368                rows.push(MemoryRetrievalGateRow {
369                    strategy: candidate.to_string(),
370                    avg_useful_hits: summary.avg_useful_hits,
371                    zero_output_failures: summary.zero_output_failures,
372                    useful_hits_pass,
373                    zero_output_pass,
374                    status: status.to_string(),
375                });
376            }
377            None => {
378                diagnostics.push(format!(
379                    "candidate strategy {candidate} was not present; memory retrieval gate blocks"
380                ));
381                rows.push(MemoryRetrievalGateRow {
382                    strategy: candidate.to_string(),
383                    avg_useful_hits: 0.0,
384                    zero_output_failures: tasks_loaded,
385                    useful_hits_pass: false,
386                    zero_output_pass: false,
387                    status: "block".to_string(),
388                });
389            }
390        }
391    }
392
393    let decision = if diagnostics.is_empty() && rows.iter().all(|row| row.status == "pass") {
394        "pass"
395    } else {
396        "block"
397    };
398
399    Some(MemoryRetrievalGate {
400        decision: decision.to_string(),
401        baseline_strategy: BASELINE.to_string(),
402        candidate_strategies: CANDIDATES
403            .iter()
404            .map(|candidate| candidate.to_string())
405            .collect(),
406        min_avg_useful_hits,
407        max_zero_output_failures,
408        rows,
409        diagnostics,
410    })
411}
412
413fn run_useful_hits(run: &DciBenchmarkRun) -> f64 {
414    run.useful_hits
415        .unwrap_or(if run.localized { 1.0 } else { 0.0 })
416}
417
418fn run_output_tokens(run: &DciBenchmarkRun) -> f64 {
419    run.output_tokens.unwrap_or(run.estimated_tokens)
420}
421
422fn choose_lowest<'a, F>(
423    current: Option<&'a DciBenchmarkRun>,
424    candidate: &'a DciBenchmarkRun,
425    metric: F,
426) -> Option<&'a DciBenchmarkRun>
427where
428    F: Fn(&DciBenchmarkRun) -> f64,
429{
430    match current {
431        Some(existing) if metric(existing) <= metric(candidate) => Some(existing),
432        _ => Some(candidate),
433    }
434}
435
436fn strategy_rank(left: &DciStrategySummary, right: &DciStrategySummary) -> std::cmp::Ordering {
437    right
438        .localization_rate
439        .partial_cmp(&left.localization_rate)
440        .unwrap_or(std::cmp::Ordering::Equal)
441        .then_with(|| {
442            right
443                .avg_useful_hits
444                .partial_cmp(&left.avg_useful_hits)
445                .unwrap_or(std::cmp::Ordering::Equal)
446        })
447        .then_with(|| {
448            left.zero_output_rate
449                .partial_cmp(&right.zero_output_rate)
450                .unwrap_or(std::cmp::Ordering::Equal)
451        })
452        .then_with(|| {
453            left.avg_estimated_tokens
454                .partial_cmp(&right.avg_estimated_tokens)
455                .unwrap_or(std::cmp::Ordering::Equal)
456        })
457        .then_with(|| {
458            left.avg_output_tokens
459                .partial_cmp(&right.avg_output_tokens)
460                .unwrap_or(std::cmp::Ordering::Equal)
461        })
462        .then_with(|| {
463            left.avg_tool_calls
464                .partial_cmp(&right.avg_tool_calls)
465                .unwrap_or(std::cmp::Ordering::Equal)
466        })
467        .then_with(|| {
468            left.avg_latency_ms
469                .partial_cmp(&right.avg_latency_ms)
470                .unwrap_or(std::cmp::Ordering::Equal)
471        })
472        .then_with(|| left.strategy.cmp(&right.strategy))
473}
474
475pub fn format_number(value: f64) -> String {
476    if (value - value.round()).abs() < 0.005 {
477        format!("{}", value.round() as i64)
478    } else {
479        format!("{value:.2}")
480    }
481}
482
483#[cfg(test)]
484mod tests {
485    use super::*;
486
487    #[test]
488    fn ranks_localization_then_cost_metrics() {
489        let report = compute(
490            r#"{
491  "tasks": [
492    {
493      "id": "a",
494      "runs": [
495        {"strategy": "exact_chained_rg", "localized": true, "tool_calls": 3, "latency_ms": 120, "estimated_tokens": 500},
496        {"strategy": "lexical_bm25", "localized": true, "tool_calls": 5, "latency_ms": 800, "estimated_tokens": 900},
497        {"strategy": "hybrid", "localized": true, "tool_calls": 4, "latency_ms": 1800, "estimated_tokens": 750}
498      ]
499    },
500    {
501      "id": "b",
502      "runs": [
503        {"strategy": "exact_chained_rg", "localized": true, "tool_calls": 4, "latency_ms": 140, "estimated_tokens": 620},
504        {"strategy": "lexical_bm25", "localized": false, "tool_calls": 6, "latency_ms": 900, "estimated_tokens": 1100},
505        {"strategy": "hybrid", "localized": true, "tool_calls": 4, "latency_ms": 2100, "estimated_tokens": 820}
506      ]
507    }
508  ]
509}"#,
510        )
511        .unwrap();
512
513        assert_eq!(report.tasks_loaded, 2);
514        assert_eq!(report.strategy_summaries[0].strategy, "exact_chained_rg");
515        assert_eq!(report.strategy_summaries[0].localized, 2);
516        assert_eq!(
517            report.task_rows[0].lowest_token_budget.as_deref(),
518            Some("exact_chained_rg")
519        );
520        assert_eq!(report.strategy_summaries[0].useful_hits, 2.0);
521        assert_eq!(report.strategy_summaries[0].zero_output_failures, 0);
522        assert!(report.warnings.is_empty());
523    }
524
525    #[test]
526    fn supports_memory_retrieval_metrics() {
527        let report = compute(
528            r#"{
529  "expected_strategies": ["claude_mem_api", "tsift_session_review_context_pack", "graph_db_related"],
530  "tasks": [
531    {
532      "id": "observer-overflow",
533      "runs": [
534        {"strategy": "claude_mem_api", "localized": false, "useful_hits": 0, "zero_output": true, "tool_calls": 1, "latency_ms": 180, "estimated_tokens": 0, "output_tokens": 0},
535        {"strategy": "tsift_session_review_context_pack", "localized": true, "useful_hits": 2, "zero_output": false, "tool_calls": 2, "latency_ms": 510, "estimated_tokens": 1450, "output_tokens": 950},
536        {"strategy": "graph_db_related", "localized": true, "useful_hits": 3, "zero_output": false, "tool_calls": 2, "latency_ms": 430, "estimated_tokens": 880, "output_tokens": 620}
537      ]
538    }
539  ]
540}"#,
541        )
542        .unwrap();
543
544        assert_eq!(report.expected_strategies.len(), 3);
545        assert_eq!(report.strategy_summaries[0].strategy, "graph_db_related");
546        let graph = report
547            .strategy_summaries
548            .iter()
549            .find(|summary| summary.strategy == "graph_db_related")
550            .unwrap();
551        assert_eq!(graph.useful_hits, 3.0);
552        assert_eq!(graph.zero_output_failures, 0);
553        assert_eq!(graph.avg_output_tokens, 620.0);
554        let claude_mem = report
555            .strategy_summaries
556            .iter()
557            .find(|summary| summary.strategy == "claude_mem_api")
558            .unwrap();
559        assert_eq!(claude_mem.zero_output_rate, 1.0);
560        assert_eq!(
561            report.task_rows[0].most_useful_hits,
562            vec!["graph_db_related".to_string()]
563        );
564        assert_eq!(
565            report.task_rows[0].zero_output_failures,
566            vec!["claude_mem_api".to_string()]
567        );
568        let gate = report.memory_retrieval_gate.as_ref().unwrap();
569        assert_eq!(gate.decision, "pass");
570        assert_eq!(gate.baseline_strategy, "claude_mem_api");
571        assert_eq!(gate.min_avg_useful_hits, 0.0);
572        assert_eq!(gate.max_zero_output_failures, 0);
573        assert!(gate.rows.iter().all(|row| row.status == "pass"));
574        assert!(report.warnings.is_empty());
575    }
576
577    #[test]
578    fn memory_retrieval_gate_blocks_when_candidate_regresses() {
579        let report = compute(
580            r#"{
581  "expected_strategies": ["claude_mem_api", "tsift_session_review_context_pack", "graph_db_related"],
582  "tasks": [
583    {
584      "id": "regressed-cutover",
585      "runs": [
586        {"strategy": "claude_mem_api", "localized": true, "useful_hits": 2, "zero_output": false, "tool_calls": 1, "latency_ms": 180, "estimated_tokens": 800, "output_tokens": 600},
587        {"strategy": "tsift_session_review_context_pack", "localized": true, "useful_hits": 1, "zero_output": false, "tool_calls": 2, "latency_ms": 510, "estimated_tokens": 1450, "output_tokens": 950},
588        {"strategy": "graph_db_related", "localized": true, "useful_hits": 3, "zero_output": true, "tool_calls": 2, "latency_ms": 430, "estimated_tokens": 880, "output_tokens": 620}
589      ]
590    }
591  ]
592}"#,
593        )
594        .unwrap();
595
596        let gate = report.memory_retrieval_gate.as_ref().unwrap();
597        assert_eq!(gate.decision, "block");
598        assert!(
599            gate.diagnostics
600                .iter()
601                .any(|diagnostic| diagnostic.contains("avg useful hits"))
602        );
603        assert!(
604            gate.diagnostics
605                .iter()
606                .any(|diagnostic| diagnostic.contains("zero-output failures"))
607        );
608    }
609}