22
33from __future__ import annotations
44
5+ from collections import Counter
6+
57from .diffing import make_unified_diff
68from .schemas import CandidatePrompt
79from .schemas import EvalResult
@@ -23,18 +25,30 @@ def propose(
2325 baseline_train : EvalResult ,
2426 failure_summary : dict [str , object ],
2527 ) -> list [CandidatePrompt ]:
26- failed_cases = [case for case in baseline_train .cases if not case .passed ]
27- if not failed_cases :
28- return []
28+ if not isinstance (failure_summary , dict ):
29+ raise TypeError ("failure_summary must be a dict" )
2930
30- observed_categories = {case .failure_category for case in failed_cases if case .failure_category }
31- by_category = failure_summary .get ("by_category" )
32- if isinstance (by_category , dict ):
33- observed_categories .update (
34- str (category ) for category , count in by_category .items () if _is_positive_count (count )
31+ if baseline_train .split != "train" :
32+ raise ValueError (
33+ "baseline_train.split must be 'train'; "
34+ f"got { baseline_train .split !r} "
3535 )
36+ for case in baseline_train .cases :
37+ if case .split != "train" :
38+ raise ValueError (
39+ f"baseline_train case { case .case_id !r} split must be 'train'; "
40+ f"got { case .split !r} "
41+ )
3642
37- targeted = sorted (observed_categories & _TARGET_FAILURE_CATEGORIES )
43+ failed_cases = [case for case in baseline_train .cases if not case .passed ]
44+ observed_counts = Counter (
45+ case .failure_category
46+ for case in failed_cases
47+ if case .failure_category
48+ )
49+ _validate_failure_summary (failure_summary , observed_counts )
50+
51+ targeted = sorted (observed_counts .keys () & _TARGET_FAILURE_CATEGORIES )
3852 if not targeted :
3953 return []
4054
@@ -75,5 +89,34 @@ def propose(
7589 ]
7690
7791
78- def _is_positive_count (value : object ) -> bool :
79- return not isinstance (value , bool ) and isinstance (value , (int , float )) and value > 0
92+ def _validate_failure_summary (
93+ failure_summary : dict [str , object ],
94+ observed_counts : Counter [str ],
95+ ) -> None :
96+ if "by_category" not in failure_summary :
97+ return
98+
99+ by_category = failure_summary ["by_category" ]
100+ if not isinstance (by_category , dict ):
101+ raise ValueError ("failure_summary['by_category'] must be a dict" )
102+
103+ summary_counts : dict [object , int ] = {}
104+ for category , count in by_category .items ():
105+ summary_counts [category ] = _normalize_positive_count (category , count )
106+
107+ if summary_counts != dict (observed_counts ):
108+ raise ValueError (
109+ "failure_summary['by_category'] must match failed train cases exactly; "
110+ f"summary={ summary_counts !r} , observed={ dict (observed_counts )!r} "
111+ )
112+
113+
114+ def _normalize_positive_count (category : object , value : object ) -> int :
115+ normalized = value if type (value ) is int and value > 0 else None
116+
117+ if normalized is None :
118+ raise ValueError (
119+ "failure_summary['by_category'] count must be a positive integer; "
120+ f"category={ category !r} , count={ value !r} "
121+ )
122+ return normalized
0 commit comments