1use std::collections::BTreeSet;
2
3use serde::Serialize;
4
5use crate::graph_index::GraphIndex;
6use crate::retrieval::ground_subgraph;
7use crate::schema::Graph;
8use crate::workflow::{
9 TaskContextItem, TaskOperator, TaskRun, TaskState, draft_task_dag, initial_run, ready_tasks,
10 validate_dag,
11};
12
13use super::{GoldenCase, kind_label};
14
15#[derive(Serialize, Clone, Debug)]
17pub struct WorkflowCaseResult {
18 pub query: String,
19 #[serde(skip_serializing_if = "Option::is_none")]
20 pub expected_profile: Option<String>,
21 #[serde(skip_serializing_if = "Option::is_none")]
22 pub profile: Option<String>,
23 #[serde(skip_serializing_if = "Option::is_none")]
24 pub profile_ok: Option<bool>,
25 pub valid: bool,
26 pub bounded: bool,
27 pub required_chain: bool,
28 pub initial_ready_context_only: bool,
29 pub inspection_recall: Option<f64>,
31 pub required_check_recall: Option<f64>,
33 pub risk_hint_recall: Option<f64>,
35 pub node_count: usize,
36 pub ready: Vec<String>,
37 pub missing_inspections: Vec<String>,
38 pub missing_required_checks: Vec<String>,
39 pub missing_risk_hints: Vec<String>,
40 pub error: Option<String>,
41}
42
43#[derive(Serialize, Clone, Debug)]
45pub struct WorkflowReport {
46 #[serde(rename = "brief_cases")]
47 pub workflow_cases: usize,
48 pub valid_rate: f64,
49 pub bounded_rate: f64,
50 pub required_chain_rate: f64,
51 pub initial_ready_rate: f64,
52 pub inspection_recall: Option<f64>,
54 #[serde(skip_serializing_if = "Option::is_none")]
56 pub profile_match_rate: Option<f64>,
57 #[serde(skip_serializing_if = "Option::is_none")]
59 pub required_check_recall: Option<f64>,
60 #[serde(skip_serializing_if = "Option::is_none")]
62 pub risk_hint_recall: Option<f64>,
63 pub mean_node_count: f64,
64 pub cases: Vec<WorkflowCaseResult>,
65}
66
67#[derive(Clone, Debug)]
68pub struct WorkflowInspectionJudgment {
69 pub recall: Option<f64>,
70 pub missing: Vec<String>,
71}
72
73pub fn evaluate_workflow(
79 graph: &Graph,
80 cases: &[GoldenCase],
81 limit: usize,
82 depth: usize,
83 width: usize,
84 max_nodes: usize,
85) -> WorkflowReport {
86 let mut results = Vec::new();
87 let index = GraphIndex::build(graph);
88
89 for c in cases.iter().filter(|case| !case.garbage) {
90 let sg = ground_subgraph(graph, &c.query, limit, depth, width);
91 let context = sg
92 .context_order
93 .iter()
94 .filter_map(|id| {
95 let node = index.node(id)?;
96 Some(TaskContextItem {
97 id: node.id.clone(),
98 title: node.title.clone(),
99 kind: kind_label(node.kind),
100 source: index.source_of(&node.id).unwrap_or_default(),
101 })
102 })
103 .collect::<Vec<_>>();
104 let dag = draft_task_dag(&c.query, &context);
105 let node_count = dag.nodes.len();
106 let valid = validate_dag(&dag).is_ok();
107 let bounded = node_count <= max_nodes;
108 let required_chain = has_required_operator_chain(&dag.nodes);
109 let (ready, initial_ready_context_only, error) = match initial_run(&dag) {
110 Ok(run) => match ready_tasks(&dag, &run) {
111 Ok(ready) => {
112 let ok = workflow_initial_ready_context_only(&run, &ready);
113 (ready, ok, None)
114 }
115 Err(err) => (Vec::new(), false, Some(format!("{err:?}"))),
116 },
117 Err(err) => (Vec::new(), false, Some(format!("{err:?}"))),
118 };
119
120 let inspected = dag
121 .nodes
122 .iter()
123 .flat_map(|node| node.inputs.iter().map(String::as_str))
124 .collect::<BTreeSet<_>>();
125 let inspection = judge_workflow_inspections(c, &inspected);
126 let required_checks = dag
127 .nodes
128 .iter()
129 .flat_map(|node| node.required_checks.iter().map(String::as_str))
130 .collect::<BTreeSet<_>>();
131 let risk_hints = dag
132 .nodes
133 .iter()
134 .filter(|node| node.risk != crate::workflow::TaskRisk::Low)
135 .map(|node| node.title.as_str())
136 .collect::<BTreeSet<_>>();
137 let required_check_judgment = judge_required_checks(c, &required_checks);
138 let risk_hint_judgment = judge_risk_hints(c, &risk_hints);
139
140 results.push(WorkflowCaseResult {
141 query: c.query.clone(),
142 expected_profile: None,
143 profile: None,
144 profile_ok: None,
145 valid,
146 bounded,
147 required_chain,
148 initial_ready_context_only,
149 inspection_recall: inspection.recall,
150 required_check_recall: required_check_judgment.recall,
151 risk_hint_recall: risk_hint_judgment.recall,
152 node_count,
153 ready,
154 missing_inspections: inspection.missing,
155 missing_required_checks: required_check_judgment.missing,
156 missing_risk_hints: risk_hint_judgment.missing,
157 error,
158 });
159 }
160
161 summarize_workflow(results)
162}
163
164pub fn summarize_workflow(results: Vec<WorkflowCaseResult>) -> WorkflowReport {
165 let workflow_cases = results.len();
166 let frac = |n: usize| {
167 if workflow_cases == 0 {
168 0.0
169 } else {
170 n as f64 / workflow_cases as f64
171 }
172 };
173 let valid_n = results.iter().filter(|case| case.valid).count();
174 let bounded_n = results.iter().filter(|case| case.bounded).count();
175 let required_chain_n = results.iter().filter(|case| case.required_chain).count();
176 let initial_ready_n = results
177 .iter()
178 .filter(|case| case.initial_ready_context_only)
179 .count();
180 let inspection_recall = mean_optional(results.iter().map(|case| case.inspection_recall));
181 let profile_match_rate = mean_bool(results.iter().filter_map(|case| case.profile_ok));
182 let required_check_recall =
183 mean_optional(results.iter().map(|case| case.required_check_recall));
184 let risk_hint_recall = mean_optional(results.iter().map(|case| case.risk_hint_recall));
185 let node_count_sum = results.iter().map(|case| case.node_count).sum::<usize>();
186
187 WorkflowReport {
188 workflow_cases,
189 valid_rate: frac(valid_n),
190 bounded_rate: frac(bounded_n),
191 required_chain_rate: frac(required_chain_n),
192 initial_ready_rate: frac(initial_ready_n),
193 inspection_recall,
194 profile_match_rate,
195 required_check_recall,
196 risk_hint_recall,
197 mean_node_count: if workflow_cases == 0 {
198 0.0
199 } else {
200 node_count_sum as f64 / workflow_cases as f64
201 },
202 cases: results,
203 }
204}
205
206pub fn judge_workflow_inspections(
207 case: &GoldenCase,
208 inspected: &BTreeSet<&str>,
209) -> WorkflowInspectionJudgment {
210 let missing = case
211 .context_must
212 .iter()
213 .filter(|id| !inspected.contains(id.as_str()))
214 .cloned()
215 .collect::<Vec<_>>();
216 let recall = if case.context_must.is_empty() {
217 None
218 } else {
219 Some((case.context_must.len() - missing.len()) as f64 / case.context_must.len() as f64)
220 };
221 WorkflowInspectionJudgment { recall, missing }
222}
223
224pub fn judge_required_checks(
225 case: &GoldenCase,
226 checks: &BTreeSet<&str>,
227) -> WorkflowInspectionJudgment {
228 let missing = case
229 .expected_required_checks
230 .iter()
231 .filter(|check| !checks.contains(check.as_str()))
232 .cloned()
233 .collect::<Vec<_>>();
234 let recall = if case.expected_required_checks.is_empty() {
235 None
236 } else {
237 Some(
238 (case.expected_required_checks.len() - missing.len()) as f64
239 / case.expected_required_checks.len() as f64,
240 )
241 };
242 WorkflowInspectionJudgment { recall, missing }
243}
244
245pub fn judge_risk_hints(case: &GoldenCase, hints: &BTreeSet<&str>) -> WorkflowInspectionJudgment {
246 let missing = case
247 .expected_risk_hints
248 .iter()
249 .filter(|hint| !hints.contains(hint.as_str()))
250 .cloned()
251 .collect::<Vec<_>>();
252 let recall = if case.expected_risk_hints.is_empty() {
253 None
254 } else {
255 Some(
256 (case.expected_risk_hints.len() - missing.len()) as f64
257 / case.expected_risk_hints.len() as f64,
258 )
259 };
260 WorkflowInspectionJudgment { recall, missing }
261}
262
263pub fn workflow_initial_ready_context_only(run: &TaskRun, ready: &[String]) -> bool {
264 ready == ["context"] && run.states.get("context") == Some(&TaskState::Ready)
265}
266
267pub fn has_required_operator_chain(nodes: &[crate::workflow::TaskNode]) -> bool {
268 let has = |operator| nodes.iter().any(|node| node.operator == operator);
269 has(TaskOperator::ReadContext)
270 && has(TaskOperator::Decompose)
271 && (has(TaskOperator::Edit) || has(TaskOperator::ProposeDoc))
272 && has(TaskOperator::Check)
273 && has(TaskOperator::Review)
274 && has(TaskOperator::HumanApproval)
275}
276
277fn mean_optional(values: impl Iterator<Item = Option<f64>>) -> Option<f64> {
278 let mut sum = 0.0;
279 let mut count = 0usize;
280 for value in values.flatten() {
281 sum += value;
282 count += 1;
283 }
284 (count > 0).then_some(sum / count as f64)
285}
286
287fn mean_bool(values: impl Iterator<Item = bool>) -> Option<f64> {
288 let mut ok = 0usize;
289 let mut count = 0usize;
290 for value in values {
291 ok += value as usize;
292 count += 1;
293 }
294 (count > 0).then_some(ok as f64 / count as f64)
295}
296
297#[cfg(test)]
298mod tests {
299 use super::*;
300
301 #[test]
302 fn summarize_workflow_reports_profile_match_rate_for_judged_cases() {
303 let report = summarize_workflow(vec![
304 workflow_case(Some(true), Some(1.0), 10),
305 workflow_case(Some(false), Some(0.5), 20),
306 workflow_case(None, None, 30),
307 ]);
308
309 assert_eq!(report.workflow_cases, 3);
310 assert_eq!(report.profile_match_rate, Some(0.5));
311 assert_eq!(report.inspection_recall, Some(0.75));
312 assert_eq!(report.mean_node_count, 20.0);
313 }
314
315 #[test]
316 fn summarize_workflow_keeps_unjudged_inspection_recall_unmeasured() {
317 let report = summarize_workflow(vec![workflow_case(None, None, 10)]);
318
319 assert_eq!(report.workflow_cases, 1);
320 assert_eq!(report.inspection_recall, None);
321 }
322
323 fn workflow_case(
324 profile_ok: Option<bool>,
325 inspection_recall: Option<f64>,
326 node_count: usize,
327 ) -> WorkflowCaseResult {
328 WorkflowCaseResult {
329 query: "task".to_string(),
330 expected_profile: profile_ok.map(|_| "safe-change".to_string()),
331 profile: profile_ok.map(|ok| {
332 if ok {
333 "safe-change".to_string()
334 } else {
335 "other".to_string()
336 }
337 }),
338 profile_ok,
339 valid: true,
340 bounded: true,
341 required_chain: true,
342 initial_ready_context_only: true,
343 inspection_recall,
344 required_check_recall: None,
345 risk_hint_recall: None,
346 node_count,
347 ready: vec!["context".to_string()],
348 missing_inspections: Vec::new(),
349 missing_required_checks: Vec::new(),
350 missing_risk_hints: Vec::new(),
351 error: None,
352 }
353 }
354}