1use agentforge_core::{DimensionScores, FailureCluster, Trace, TraceStatus, TraceStep};
2
3pub fn classify_failure_cluster(
5 trace: &Trace,
6 scores: &DimensionScores,
7 failure_reasons: &[String],
8) -> FailureCluster {
9 if trace.status == TraceStatus::Pass {
10 return FailureCluster::NoFailure;
11 }
12
13 if trace.status == TraceStatus::Error {
14 return FailureCluster::ApiError;
16 }
17
18 if scores.schema_compliance < 0.3 {
22 return FailureCluster::SchemaViolation;
23 }
24
25 if scores.argument_correctness < 0.3 {
27 return FailureCluster::HallucinatedArgument;
28 }
29
30 if detect_loop(trace) {
32 return FailureCluster::Looping;
33 }
34
35 if scores.path_efficiency < 0.1 {
37 return FailureCluster::PrematureStop;
38 }
39
40 let failure_text = failure_reasons.join(" ").to_lowercase();
42 if failure_text.contains("wrong_tool") || failure_text.contains("missing required tools") {
43 return FailureCluster::WrongTool;
44 }
45 if failure_text.contains("argument") || failure_text.contains("hallucinated") {
46 return FailureCluster::HallucinatedArgument;
47 }
48 if failure_text.contains("schema") {
49 return FailureCluster::SchemaViolation;
50 }
51 if failure_text.contains("constraint") || failure_text.contains("instruction adherence") {
52 return FailureCluster::ConstraintBreach;
53 }
54
55 let candidates = [
59 (scores.task_completion, FailureCluster::PrematureStop),
60 (scores.tool_selection, FailureCluster::WrongTool),
61 (
62 scores.instruction_adherence,
63 FailureCluster::ConstraintBreach,
64 ),
65 ];
66
67 candidates
68 .iter()
69 .min_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(std::cmp::Ordering::Equal))
70 .map(|(_, cluster)| cluster.clone())
71 .unwrap_or(FailureCluster::Unknown)
72}
73
74fn detect_loop(trace: &Trace) -> bool {
76 let llm_count = trace
77 .steps
78 .iter()
79 .filter(|s| matches!(s, TraceStep::LlmCall(_)))
80 .count();
81 let tool_count = trace
82 .steps
83 .iter()
84 .filter(|s| matches!(s, TraceStep::ToolCall(_)))
85 .count();
86
87 llm_count > 5 && tool_count <= 1
89}
90
91#[cfg(test)]
92mod tests {
93 use super::*;
94 use agentforge_core::{FailureCluster, TraceStatus};
95
96 fn make_scores(
97 tool: f64,
98 args: f64,
99 schema: f64,
100 adherence: f64,
101 efficiency: f64,
102 ) -> DimensionScores {
103 DimensionScores {
104 task_completion: 0.5,
105 tool_selection: tool,
106 argument_correctness: args,
107 schema_compliance: schema,
108 instruction_adherence: adherence,
109 path_efficiency: efficiency,
110 }
111 }
112
113 fn make_empty_trace(status: TraceStatus) -> Trace {
114 Trace {
115 id: uuid::Uuid::new_v4(),
116 run_id: uuid::Uuid::new_v4(),
117 scenario_id: uuid::Uuid::new_v4(),
118 status,
119 steps: vec![],
120 final_output: None,
121 scores: None,
122 aggregate_score: None,
123 failure_cluster: FailureCluster::Unknown,
124 failure_reason: None,
125 review_needed: false,
126 llm_calls: 0,
127 tool_invocations: 0,
128 input_tokens: 0,
129 output_tokens: 0,
130 latency_ms: 0,
131 retry_count: 0,
132 seed: 0,
133 created_at: chrono::Utc::now(),
134 }
135 }
136
137 #[test]
138 fn pass_returns_no_failure() {
139 let trace = make_empty_trace(TraceStatus::Pass);
140 let scores = make_scores(1.0, 1.0, 1.0, 1.0, 1.0);
141 assert_eq!(
142 classify_failure_cluster(&trace, &scores, &[]),
143 FailureCluster::NoFailure
144 );
145 }
146
147 #[test]
148 fn error_returns_api_error() {
149 let trace = make_empty_trace(TraceStatus::Error);
150 let scores = make_scores(0.0, 0.0, 0.0, 0.0, 0.0);
151 assert_eq!(
152 classify_failure_cluster(&trace, &scores, &[]),
153 FailureCluster::ApiError
154 );
155 }
156
157 #[test]
158 fn low_schema_compliance_is_schema_violation() {
159 let trace = make_empty_trace(TraceStatus::Fail);
160 let scores = make_scores(1.0, 1.0, 0.1, 1.0, 1.0);
161 assert_eq!(
162 classify_failure_cluster(&trace, &scores, &[]),
163 FailureCluster::SchemaViolation
164 );
165 }
166
167 #[test]
168 fn low_tool_selection_is_wrong_tool() {
169 let trace = make_empty_trace(TraceStatus::Fail);
170 let scores = make_scores(0.1, 1.0, 1.0, 1.0, 1.0);
171 assert_eq!(
172 classify_failure_cluster(&trace, &scores, &[]),
173 FailureCluster::WrongTool
174 );
175 }
176
177 #[test]
178 fn low_args_is_hallucinated_argument() {
179 let trace = make_empty_trace(TraceStatus::Fail);
180 let scores = make_scores(1.0, 0.1, 1.0, 1.0, 1.0);
181 assert_eq!(
182 classify_failure_cluster(&trace, &scores, &[]),
183 FailureCluster::HallucinatedArgument
184 );
185 }
186
187 #[test]
188 fn low_constraint_is_breach() {
189 let trace = make_empty_trace(TraceStatus::Fail);
190 let scores = make_scores(1.0, 1.0, 1.0, 0.1, 1.0);
191 assert_eq!(
192 classify_failure_cluster(&trace, &scores, &[]),
193 FailureCluster::ConstraintBreach
194 );
195 }
196
197 #[test]
200 fn fail_trace_with_moderate_scores_never_gets_no_failure() {
201 let trace = make_empty_trace(TraceStatus::Fail);
202 let scores = DimensionScores {
205 task_completion: 0.55,
206 tool_selection: 0.65,
207 argument_correctness: 0.70,
208 schema_compliance: 0.60,
209 instruction_adherence: 0.70,
210 path_efficiency: 0.75,
211 };
212 let cluster = classify_failure_cluster(&trace, &scores, &[]);
213 assert_ne!(
214 cluster,
215 FailureCluster::NoFailure,
216 "Fail trace must never get NoFailure cluster"
217 );
218 }
219
220 #[test]
221 fn review_needed_trace_never_gets_no_failure() {
222 let trace = make_empty_trace(TraceStatus::ReviewNeeded);
223 let scores = DimensionScores {
224 task_completion: 0.5,
225 tool_selection: 0.9,
226 argument_correctness: 0.9,
227 schema_compliance: 0.9,
228 instruction_adherence: 0.9,
229 path_efficiency: 0.9,
230 };
231 let cluster = classify_failure_cluster(&trace, &scores, &[]);
232 assert_ne!(
233 cluster,
234 FailureCluster::NoFailure,
235 "ReviewNeeded trace must never get NoFailure cluster"
236 );
237 }
238
239 #[test]
242 fn weakest_task_completion_yields_premature_stop() {
243 let trace = make_empty_trace(TraceStatus::Fail);
244 let scores = DimensionScores {
246 task_completion: 0.3,
247 tool_selection: 0.6,
248 argument_correctness: 0.7,
249 schema_compliance: 0.7,
250 instruction_adherence: 0.6,
251 path_efficiency: 0.5,
252 };
253 let cluster = classify_failure_cluster(&trace, &scores, &[]);
254 assert_eq!(
255 cluster,
256 FailureCluster::PrematureStop,
257 "Weakest task_completion should yield PrematureStop"
258 );
259 }
260
261 #[test]
262 fn weakest_tool_selection_yields_wrong_tool() {
263 let trace = make_empty_trace(TraceStatus::Fail);
264 let scores = DimensionScores {
265 task_completion: 0.6,
266 tool_selection: 0.3,
267 argument_correctness: 0.7,
268 schema_compliance: 0.7,
269 instruction_adherence: 0.6,
270 path_efficiency: 0.5,
271 };
272 let cluster = classify_failure_cluster(&trace, &scores, &[]);
273 assert_eq!(
274 cluster,
275 FailureCluster::WrongTool,
276 "Weakest tool_selection should yield WrongTool"
277 );
278 }
279
280 #[test]
281 fn weakest_instruction_adherence_yields_constraint_breach() {
282 let trace = make_empty_trace(TraceStatus::Fail);
283 let scores = DimensionScores {
284 task_completion: 0.6,
285 tool_selection: 0.6,
286 argument_correctness: 0.7,
287 schema_compliance: 0.7,
288 instruction_adherence: 0.2,
289 path_efficiency: 0.6,
290 };
291 let cluster = classify_failure_cluster(&trace, &scores, &[]);
292 assert_eq!(
293 cluster,
294 FailureCluster::ConstraintBreach,
295 "Weakest instruction_adherence should yield ConstraintBreach"
296 );
297 }
298
299 #[test]
302 fn schema_violation_beats_moderate_task_completion() {
303 let trace = make_empty_trace(TraceStatus::Fail);
304 let scores = DimensionScores {
306 task_completion: 0.2,
307 tool_selection: 0.9,
308 argument_correctness: 0.9,
309 schema_compliance: 0.2, instruction_adherence: 0.9,
311 path_efficiency: 0.9,
312 };
313 assert_eq!(
314 classify_failure_cluster(&trace, &scores, &[]),
315 FailureCluster::SchemaViolation,
316 "Schema violation should take priority over weak task_completion"
317 );
318 }
319
320 #[test]
321 fn hallucinated_arg_beats_weak_dimensions() {
322 let trace = make_empty_trace(TraceStatus::Fail);
323 let scores = DimensionScores {
324 task_completion: 0.3,
325 tool_selection: 0.3,
326 argument_correctness: 0.1, schema_compliance: 0.9,
328 instruction_adherence: 0.3,
329 path_efficiency: 0.3,
330 };
331 assert_eq!(
332 classify_failure_cluster(&trace, &scores, &[]),
333 FailureCluster::HallucinatedArgument
334 );
335 }
336
337 #[test]
340 fn wrong_tool_keyword_triggers_cluster() {
341 let trace = make_empty_trace(TraceStatus::Fail);
342 let scores = make_scores(0.7, 0.9, 0.9, 0.9, 0.7);
343 assert_eq!(
344 classify_failure_cluster(&trace, &scores, &["wrong_tool".to_string()]),
345 FailureCluster::WrongTool
346 );
347 }
348
349 #[test]
350 fn missing_required_tools_keyword_triggers_wrong_tool() {
351 let trace = make_empty_trace(TraceStatus::Fail);
352 let scores = make_scores(0.7, 0.9, 0.9, 0.9, 0.7);
353 assert_eq!(
354 classify_failure_cluster(&trace, &scores, &["missing required tools".to_string()]),
355 FailureCluster::WrongTool
356 );
357 }
358
359 #[test]
360 fn argument_keyword_triggers_hallucinated_argument() {
361 let trace = make_empty_trace(TraceStatus::Fail);
362 let scores = make_scores(0.7, 0.9, 0.9, 0.9, 0.7);
363 assert_eq!(
364 classify_failure_cluster(&trace, &scores, &["argument mismatch".to_string()]),
365 FailureCluster::HallucinatedArgument
366 );
367 }
368
369 #[test]
370 fn constraint_keyword_triggers_constraint_breach() {
371 let trace = make_empty_trace(TraceStatus::Fail);
372 let scores = make_scores(0.7, 0.9, 0.9, 0.9, 0.7);
373 assert_eq!(
374 classify_failure_cluster(
375 &trace,
376 &scores,
377 &["instruction adherence failed".to_string()]
378 ),
379 FailureCluster::ConstraintBreach
380 );
381 }
382
383 #[test]
386 fn many_llm_calls_few_tools_triggers_looping() {
387 let mut trace = make_empty_trace(TraceStatus::Fail);
388 use agentforge_core::{LlmCallStep, TraceStep};
389 use chrono::Utc;
390 for i in 0..6 {
391 trace.steps.push(TraceStep::LlmCall(LlmCallStep {
392 index: i,
393 model: "gpt-4o".to_string(),
394 messages: vec![],
395 response: serde_json::json!({}),
396 input_tokens: 50,
397 output_tokens: 20,
398 latency_ms: 500,
399 timestamp: Utc::now(),
400 }));
401 }
402 use agentforge_core::{ToolCallStep, TraceStep as TS};
404 trace.steps.push(TS::ToolCall(ToolCallStep {
405 index: 6,
406 tool_name: "search".to_string(),
407 call_id: "c1".to_string(),
408 arguments: serde_json::json!({}),
409 timestamp: Utc::now(),
410 }));
411 let scores = make_scores(0.5, 0.5, 0.9, 0.9, 0.2);
412 assert_eq!(
413 classify_failure_cluster(&trace, &scores, &[]),
414 FailureCluster::Looping
415 );
416 }
417
418 #[test]
419 fn few_llm_calls_does_not_trigger_looping() {
420 let mut trace = make_empty_trace(TraceStatus::Fail);
421 use agentforge_core::{LlmCallStep, TraceStep};
422 use chrono::Utc;
423 for i in 0..3 {
424 trace.steps.push(TraceStep::LlmCall(LlmCallStep {
425 index: i,
426 model: "gpt-4o".to_string(),
427 messages: vec![],
428 response: serde_json::json!({}),
429 input_tokens: 50,
430 output_tokens: 20,
431 latency_ms: 500,
432 timestamp: Utc::now(),
433 }));
434 }
435 let scores = make_scores(0.5, 0.5, 0.9, 0.9, 0.2);
436 assert_ne!(
437 classify_failure_cluster(&trace, &scores, &[]),
438 FailureCluster::Looping
439 );
440 }
441
442 #[test]
445 fn very_low_path_efficiency_is_premature_stop() {
446 let trace = make_empty_trace(TraceStatus::Fail);
447 let scores = DimensionScores {
449 task_completion: 0.7,
450 tool_selection: 0.7,
451 argument_correctness: 0.9,
452 schema_compliance: 0.9,
453 instruction_adherence: 0.7,
454 path_efficiency: 0.05,
455 };
456 assert_eq!(
457 classify_failure_cluster(&trace, &scores, &[]),
458 FailureCluster::PrematureStop
459 );
460 }
461
462 #[test]
465 fn error_trace_ignores_zero_scores() {
466 let trace = make_empty_trace(TraceStatus::Error);
468 let scores = make_scores(0.0, 0.0, 0.0, 0.0, 0.0);
469 assert_eq!(
470 classify_failure_cluster(&trace, &scores, &[]),
471 FailureCluster::ApiError
472 );
473 }
474
475 #[test]
476 fn error_trace_ignores_high_scores() {
477 let trace = make_empty_trace(TraceStatus::Error);
479 let scores = make_scores(1.0, 1.0, 1.0, 1.0, 1.0);
480 assert_eq!(
481 classify_failure_cluster(&trace, &scores, &[]),
482 FailureCluster::ApiError
483 );
484 }
485
486 #[test]
487 fn pass_trace_ignores_zero_scores() {
488 let trace = make_empty_trace(TraceStatus::Pass);
490 let scores = make_scores(0.0, 0.0, 0.0, 0.0, 0.0);
491 assert_eq!(
492 classify_failure_cluster(&trace, &scores, &[]),
493 FailureCluster::NoFailure
494 );
495 }
496
497 #[test]
498 fn schema_at_exactly_boundary_is_not_violation() {
499 let trace = make_empty_trace(TraceStatus::Fail);
501 let scores = DimensionScores {
502 task_completion: 0.5,
503 tool_selection: 0.5,
504 argument_correctness: 0.5,
505 schema_compliance: 0.3,
506 instruction_adherence: 0.5,
507 path_efficiency: 0.5,
508 };
509 assert_ne!(
510 classify_failure_cluster(&trace, &scores, &[]),
511 FailureCluster::SchemaViolation,
512 "schema_compliance == 0.3 should NOT trigger SchemaViolation (threshold is < 0.3)"
513 );
514 }
515
516 #[test]
517 fn args_at_exactly_boundary_is_not_hallucinated() {
518 let trace = make_empty_trace(TraceStatus::Fail);
520 let scores = DimensionScores {
521 task_completion: 0.5,
522 tool_selection: 0.5,
523 argument_correctness: 0.3,
524 schema_compliance: 0.5,
525 instruction_adherence: 0.5,
526 path_efficiency: 0.5,
527 };
528 assert_ne!(
529 classify_failure_cluster(&trace, &scores, &[]),
530 FailureCluster::HallucinatedArgument
531 );
532 }
533
534 #[test]
535 fn path_efficiency_at_exactly_boundary_does_not_use_hard_threshold() {
536 let trace = make_empty_trace(TraceStatus::Fail);
540 let scores = DimensionScores {
541 task_completion: 0.9,
542 tool_selection: 0.3, argument_correctness: 0.5,
544 schema_compliance: 0.5,
545 instruction_adherence: 0.9,
546 path_efficiency: 0.1, };
548 assert_eq!(
550 classify_failure_cluster(&trace, &scores, &[]),
551 FailureCluster::WrongTool
552 );
553 }
554
555 #[test]
556 fn schema_keyword_triggers_schema_violation() {
557 let trace = make_empty_trace(TraceStatus::Fail);
558 let scores = make_scores(0.7, 0.9, 0.9, 0.9, 0.7);
559 assert_eq!(
560 classify_failure_cluster(
561 &trace,
562 &scores,
563 &["output schema validation failed".to_string()]
564 ),
565 FailureCluster::SchemaViolation
566 );
567 }
568
569 #[test]
570 fn hallucinated_keyword_triggers_hallucinated_argument() {
571 let trace = make_empty_trace(TraceStatus::Fail);
572 let scores = make_scores(0.7, 0.9, 0.9, 0.9, 0.7);
573 assert_eq!(
574 classify_failure_cluster(&trace, &scores, &["hallucinated field value".to_string()]),
575 FailureCluster::HallucinatedArgument
576 );
577 }
578
579 #[test]
580 fn constraint_keyword_triggers_constraint_breach_variant() {
581 let trace = make_empty_trace(TraceStatus::Fail);
582 let scores = make_scores(0.7, 0.9, 0.9, 0.9, 0.7);
583 assert_eq!(
584 classify_failure_cluster(&trace, &scores, &["constraint violated".to_string()]),
585 FailureCluster::ConstraintBreach
586 );
587 }
588
589 #[test]
590 fn review_needed_with_low_tool_selection_yields_wrong_tool() {
591 let trace = make_empty_trace(TraceStatus::ReviewNeeded);
593 let scores = DimensionScores {
594 task_completion: 0.6,
595 tool_selection: 0.2,
596 argument_correctness: 0.8,
597 schema_compliance: 0.8,
598 instruction_adherence: 0.8,
599 path_efficiency: 0.8,
600 };
601 let cluster = classify_failure_cluster(&trace, &scores, &[]);
602 assert_ne!(cluster, FailureCluster::NoFailure);
603 }
604
605 #[test]
606 fn loop_detection_exactly_five_llm_calls_does_not_trigger() {
607 let mut trace = make_empty_trace(TraceStatus::Fail);
609 use agentforge_core::{LlmCallStep, TraceStep};
610 use chrono::Utc;
611 for i in 0..5u32 {
612 trace.steps.push(TraceStep::LlmCall(LlmCallStep {
613 index: i,
614 model: "gpt-4o".to_string(),
615 messages: vec![],
616 response: serde_json::json!({}),
617 input_tokens: 50,
618 output_tokens: 20,
619 latency_ms: 500,
620 timestamp: Utc::now(),
621 }));
622 }
623 let scores = make_scores(0.4, 0.4, 0.9, 0.9, 0.2);
624 assert_ne!(
625 classify_failure_cluster(&trace, &scores, &[]),
626 FailureCluster::Looping,
627 "Exactly 5 LLM calls should NOT trigger looping (threshold is >5)"
628 );
629 }
630
631 #[test]
632 fn loop_detection_six_llm_calls_two_tools_does_not_trigger() {
633 let mut trace = make_empty_trace(TraceStatus::Fail);
635 use agentforge_core::{LlmCallStep, ToolCallStep, TraceStep};
636 use chrono::Utc;
637 for i in 0..6u32 {
638 trace.steps.push(TraceStep::LlmCall(LlmCallStep {
639 index: i,
640 model: "gpt-4o".to_string(),
641 messages: vec![],
642 response: serde_json::json!({}),
643 input_tokens: 50,
644 output_tokens: 20,
645 latency_ms: 500,
646 timestamp: Utc::now(),
647 }));
648 }
649 for i in 6..8u32 {
651 trace.steps.push(TraceStep::ToolCall(ToolCallStep {
652 index: i,
653 tool_name: "search".to_string(),
654 call_id: format!("c{i}"),
655 arguments: serde_json::json!({}),
656 timestamp: Utc::now(),
657 }));
658 }
659 let scores = make_scores(0.5, 0.5, 0.9, 0.9, 0.3);
660 assert_ne!(
661 classify_failure_cluster(&trace, &scores, &[]),
662 FailureCluster::Looping
663 );
664 }
665
666 #[test]
667 fn all_perfect_scores_fail_trace_gets_meaningful_cluster() {
668 let trace = make_empty_trace(TraceStatus::Fail);
671 let scores = make_scores(1.0, 1.0, 1.0, 1.0, 1.0);
672 let cluster = classify_failure_cluster(&trace, &scores, &[]);
673 assert_ne!(cluster, FailureCluster::NoFailure);
674 }
675
676 #[test]
677 fn schema_violation_takes_priority_over_low_args() {
678 let trace = make_empty_trace(TraceStatus::Fail);
680 let scores = DimensionScores {
681 task_completion: 0.8,
682 tool_selection: 0.8,
683 argument_correctness: 0.2, schema_compliance: 0.1, instruction_adherence: 0.8,
686 path_efficiency: 0.5,
687 };
688 assert_eq!(
689 classify_failure_cluster(&trace, &scores, &[]),
690 FailureCluster::SchemaViolation,
691 "Schema violation should take priority over HallucinatedArgument"
692 );
693 }
694
695 #[test]
696 fn empty_failure_reasons_falls_back_to_dimension_analysis() {
697 let trace = make_empty_trace(TraceStatus::Fail);
699 let scores = DimensionScores {
700 task_completion: 0.4,
701 tool_selection: 0.8,
702 argument_correctness: 0.8,
703 schema_compliance: 0.8,
704 instruction_adherence: 0.8,
705 path_efficiency: 0.8,
706 };
707 assert_eq!(
709 classify_failure_cluster(&trace, &scores, &[]),
710 FailureCluster::PrematureStop
711 );
712 }
713}