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    pub tasks: Vec<DciBenchmarkTask>,
12}
13
14#[derive(Debug, Clone, Deserialize)]
15pub struct DciBenchmarkTask {
16    pub id: String,
17    #[serde(default)]
18    pub label: Option<String>,
19    #[serde(default)]
20    pub target: Option<String>,
21    pub runs: Vec<DciBenchmarkRun>,
22}
23
24#[derive(Debug, Clone, Deserialize)]
25pub struct DciBenchmarkRun {
26    pub strategy: String,
27    pub localized: bool,
28    pub tool_calls: f64,
29    pub latency_ms: f64,
30    pub estimated_tokens: f64,
31    #[serde(default)]
32    pub notes: Option<String>,
33}
34
35#[derive(Debug, Clone, PartialEq, Serialize)]
36pub struct DciBenchmarkReport {
37    #[serde(skip_serializing_if = "Option::is_none")]
38    pub description: Option<String>,
39    pub tasks_loaded: usize,
40    pub strategies_compared: usize,
41    pub expected_strategies: Vec<String>,
42    pub strategy_summaries: Vec<DciStrategySummary>,
43    pub task_rows: Vec<DciTaskRow>,
44    #[serde(skip_serializing_if = "Vec::is_empty", default)]
45    pub warnings: Vec<String>,
46}
47
48#[derive(Debug, Clone, PartialEq, Serialize)]
49pub struct DciStrategySummary {
50    pub strategy: String,
51    pub task_runs: usize,
52    pub localized: usize,
53    pub localization_rate: f64,
54    pub avg_tool_calls: f64,
55    pub avg_latency_ms: f64,
56    pub avg_estimated_tokens: f64,
57    pub rank: usize,
58}
59
60#[derive(Debug, Clone, PartialEq, Serialize)]
61pub struct DciTaskRow {
62    pub task_id: String,
63    #[serde(skip_serializing_if = "Option::is_none")]
64    pub label: Option<String>,
65    #[serde(skip_serializing_if = "Option::is_none")]
66    pub target: Option<String>,
67    pub best_localization: Vec<String>,
68    pub lowest_tool_calls: Option<String>,
69    pub lowest_latency: Option<String>,
70    pub lowest_token_budget: Option<String>,
71}
72
73#[derive(Default)]
74struct Accumulator {
75    task_runs: usize,
76    localized: usize,
77    tool_calls: f64,
78    latency_ms: f64,
79    estimated_tokens: f64,
80}
81
82pub fn compute(input: &str) -> Result<DciBenchmarkReport> {
83    let fixture: DciBenchmarkFixture =
84        serde_json::from_str(input).context("parsing dci-benchmark fixture as JSON")?;
85    if fixture.tasks.is_empty() {
86        bail!("dci-benchmark fixture did not contain any tasks");
87    }
88
89    let mut warnings = Vec::new();
90    let mut accumulators = BTreeMap::<String, Accumulator>::new();
91    let mut seen_strategies = BTreeSet::<String>::new();
92    let mut task_rows = Vec::new();
93
94    for task in &fixture.tasks {
95        if task.runs.is_empty() {
96            warnings.push(format!(
97                "task {} did not include any strategy runs",
98                task.id
99            ));
100            continue;
101        }
102
103        let mut localized = Vec::new();
104        let mut lowest_tool_calls: Option<&DciBenchmarkRun> = None;
105        let mut lowest_latency: Option<&DciBenchmarkRun> = None;
106        let mut lowest_tokens: Option<&DciBenchmarkRun> = None;
107
108        for run in &task.runs {
109            if !run.tool_calls.is_finite() || run.tool_calls < 0.0 {
110                bail!(
111                    "task {} strategy {} has invalid tool_calls",
112                    task.id,
113                    run.strategy
114                );
115            }
116            if !run.latency_ms.is_finite() || run.latency_ms < 0.0 {
117                bail!(
118                    "task {} strategy {} has invalid latency_ms",
119                    task.id,
120                    run.strategy
121                );
122            }
123            if !run.estimated_tokens.is_finite() || run.estimated_tokens < 0.0 {
124                bail!(
125                    "task {} strategy {} has invalid estimated_tokens",
126                    task.id,
127                    run.strategy
128                );
129            }
130
131            seen_strategies.insert(run.strategy.clone());
132            if run.localized {
133                localized.push(run.strategy.clone());
134            }
135
136            lowest_tool_calls = choose_lowest(lowest_tool_calls, run, |value| value.tool_calls);
137            lowest_latency = choose_lowest(lowest_latency, run, |value| value.latency_ms);
138            lowest_tokens = choose_lowest(lowest_tokens, run, |value| value.estimated_tokens);
139
140            let acc = accumulators.entry(run.strategy.clone()).or_default();
141            acc.task_runs += 1;
142            acc.localized += usize::from(run.localized);
143            acc.tool_calls += run.tool_calls;
144            acc.latency_ms += run.latency_ms;
145            acc.estimated_tokens += run.estimated_tokens;
146        }
147
148        task_rows.push(DciTaskRow {
149            task_id: task.id.clone(),
150            label: task.label.clone(),
151            target: task.target.clone(),
152            best_localization: localized,
153            lowest_tool_calls: lowest_tool_calls.map(|run| run.strategy.clone()),
154            lowest_latency: lowest_latency.map(|run| run.strategy.clone()),
155            lowest_token_budget: lowest_tokens.map(|run| run.strategy.clone()),
156        });
157    }
158
159    for expected in EXPECTED_STRATEGIES {
160        if !seen_strategies.contains(expected) {
161            warnings.push(format!("expected strategy {expected} was not present"));
162        }
163    }
164
165    let mut summaries = accumulators
166        .into_iter()
167        .map(|(strategy, acc)| {
168            let task_runs = acc.task_runs.max(1);
169            DciStrategySummary {
170                strategy,
171                task_runs: acc.task_runs,
172                localized: acc.localized,
173                localization_rate: acc.localized as f64 / task_runs as f64,
174                avg_tool_calls: acc.tool_calls / task_runs as f64,
175                avg_latency_ms: acc.latency_ms / task_runs as f64,
176                avg_estimated_tokens: acc.estimated_tokens / task_runs as f64,
177                rank: 0,
178            }
179        })
180        .collect::<Vec<_>>();
181    summaries.sort_by(strategy_rank);
182    for (index, summary) in summaries.iter_mut().enumerate() {
183        summary.rank = index + 1;
184    }
185
186    Ok(DciBenchmarkReport {
187        description: fixture.description,
188        tasks_loaded: fixture.tasks.len(),
189        strategies_compared: summaries.len(),
190        expected_strategies: EXPECTED_STRATEGIES
191            .iter()
192            .map(|strategy| strategy.to_string())
193            .collect(),
194        strategy_summaries: summaries,
195        task_rows,
196        warnings,
197    })
198}
199
200fn choose_lowest<'a, F>(
201    current: Option<&'a DciBenchmarkRun>,
202    candidate: &'a DciBenchmarkRun,
203    metric: F,
204) -> Option<&'a DciBenchmarkRun>
205where
206    F: Fn(&DciBenchmarkRun) -> f64,
207{
208    match current {
209        Some(existing) if metric(existing) <= metric(candidate) => Some(existing),
210        _ => Some(candidate),
211    }
212}
213
214fn strategy_rank(left: &DciStrategySummary, right: &DciStrategySummary) -> std::cmp::Ordering {
215    right
216        .localization_rate
217        .partial_cmp(&left.localization_rate)
218        .unwrap_or(std::cmp::Ordering::Equal)
219        .then_with(|| {
220            left.avg_estimated_tokens
221                .partial_cmp(&right.avg_estimated_tokens)
222                .unwrap_or(std::cmp::Ordering::Equal)
223        })
224        .then_with(|| {
225            left.avg_tool_calls
226                .partial_cmp(&right.avg_tool_calls)
227                .unwrap_or(std::cmp::Ordering::Equal)
228        })
229        .then_with(|| {
230            left.avg_latency_ms
231                .partial_cmp(&right.avg_latency_ms)
232                .unwrap_or(std::cmp::Ordering::Equal)
233        })
234        .then_with(|| left.strategy.cmp(&right.strategy))
235}
236
237pub fn format_number(value: f64) -> String {
238    if (value - value.round()).abs() < 0.005 {
239        format!("{}", value.round() as i64)
240    } else {
241        format!("{value:.2}")
242    }
243}
244
245#[cfg(test)]
246mod tests {
247    use super::*;
248
249    #[test]
250    fn ranks_localization_then_cost_metrics() {
251        let report = compute(
252            r#"{
253  "tasks": [
254    {
255      "id": "a",
256      "runs": [
257        {"strategy": "exact_chained_rg", "localized": true, "tool_calls": 3, "latency_ms": 120, "estimated_tokens": 500},
258        {"strategy": "lexical_bm25", "localized": true, "tool_calls": 5, "latency_ms": 800, "estimated_tokens": 900},
259        {"strategy": "hybrid", "localized": true, "tool_calls": 4, "latency_ms": 1800, "estimated_tokens": 750}
260      ]
261    },
262    {
263      "id": "b",
264      "runs": [
265        {"strategy": "exact_chained_rg", "localized": true, "tool_calls": 4, "latency_ms": 140, "estimated_tokens": 620},
266        {"strategy": "lexical_bm25", "localized": false, "tool_calls": 6, "latency_ms": 900, "estimated_tokens": 1100},
267        {"strategy": "hybrid", "localized": true, "tool_calls": 4, "latency_ms": 2100, "estimated_tokens": 820}
268      ]
269    }
270  ]
271}"#,
272        )
273        .unwrap();
274
275        assert_eq!(report.tasks_loaded, 2);
276        assert_eq!(report.strategy_summaries[0].strategy, "exact_chained_rg");
277        assert_eq!(report.strategy_summaries[0].localized, 2);
278        assert_eq!(
279            report.task_rows[0].lowest_token_budget.as_deref(),
280            Some("exact_chained_rg")
281        );
282        assert!(report.warnings.is_empty());
283    }
284}