55from typing import Any
66
77from .schemas import CaseDelta
8+ from .schemas import CostSummary
89from .schemas import EvalResult
910from .schemas import GateDecision
1011
@@ -35,17 +36,32 @@ def decide(
3536 candidate_train : EvalResult ,
3637 candidate_validation : EvalResult ,
3738 deltas : list [CaseDelta ],
39+ cost_summary : CostSummary ,
3840 cumulative_cost : float = 0.0 ,
3941 ) -> GateDecision :
4042 train_delta = round (candidate_train .score - baseline_train .score , 6 )
4143 val_delta = round (candidate_validation .score - baseline_validation .score , 6 )
4244 candidate_cost = round (candidate_train .cost + candidate_validation .cost , 6 )
4345 reasons : list [str ] = []
4446
47+ train_comparable = _append_comparability_reasons (
48+ reasons ,
49+ split = "train" ,
50+ baseline = baseline_train ,
51+ candidate = candidate_train ,
52+ )
53+ validation_comparable = _append_comparability_reasons (
54+ reasons ,
55+ split = "validation" ,
56+ baseline = baseline_validation ,
57+ candidate = candidate_validation ,
58+ )
59+
4560 overfit_detected = train_delta > 0 and val_delta <= 0
4661 if overfit_detected :
4762 reasons .append (
48- "reject: overfit detected because train score improved but validation score regressed or did not improve "
63+ "reject: overfit detected because train score improved but "
64+ "validation score regressed or did not improve "
4965 f"({ train_delta :+.3f} train, { val_delta :+.3f} validation)"
5066 )
5167
@@ -56,20 +72,32 @@ def decide(
5672 f"{ val_delta :+.3f} is below required { min_val_improvement :+.3f} "
5773 )
5874
59- baseline_validation_by_id = baseline_validation .by_case_id ()
60- candidate_validation_by_id = candidate_validation .by_case_id ()
61- validation_new_failures = [
62- case_id
63- for case_id , candidate_case in sorted (candidate_validation_by_id .items ())
64- if not candidate_case .passed and baseline_validation_by_id .get (case_id )
65- and baseline_validation_by_id [case_id ].passed
66- ]
67- new_hard_failures = [
68- case_id
69- for case_id , candidate_case in sorted (candidate_validation_by_id .items ())
70- if candidate_case .hard_failed and baseline_validation_by_id .get (case_id )
71- and baseline_validation_by_id [case_id ].passed
72- ]
75+ validation_new_failures : list [str ] = []
76+ if validation_comparable :
77+ baseline_validation_by_id = baseline_validation .by_case_id ()
78+ candidate_validation_by_id = candidate_validation .by_case_id ()
79+ validation_new_failures = [
80+ case_id
81+ for case_id , candidate_case in sorted (candidate_validation_by_id .items ())
82+ if not candidate_case .passed and baseline_validation_by_id [case_id ].passed
83+ ]
84+ if validation_new_failures :
85+ reasons .append (f"reject: new validation failures appeared: { validation_new_failures } " )
86+
87+ new_hard_failures : list [str ] = []
88+ for baseline_result , candidate_result , comparable in (
89+ (baseline_train , candidate_train , train_comparable ),
90+ (baseline_validation , candidate_validation , validation_comparable ),
91+ ):
92+ if not comparable :
93+ continue
94+ baseline_by_id = baseline_result .by_case_id ()
95+ new_hard_failures .extend (
96+ case_id
97+ for case_id , candidate_case in sorted (candidate_result .by_case_id ().items ())
98+ if candidate_case .hard_failed and not baseline_by_id [case_id ].hard_failed
99+ )
100+ new_hard_failures = sorted (set (new_hard_failures ))
73101 if new_hard_failures and not bool (self .config ["allow_new_hard_fail" ]):
74102 reasons .append (f"reject: new hard failures appeared: { new_hard_failures } " )
75103
@@ -91,10 +119,16 @@ def decide(
91119 if excessive_drops :
92120 reasons .append (f"reject: per-case validation score drops exceed { max_drop :.3f} : { excessive_drops } " )
93121
94- max_total_cost = float (self .config ["max_total_cost" ])
95122 total_run_cost = round (cumulative_cost + candidate_cost , 6 )
96- if total_run_cost > max_total_cost :
97- reasons .append (f"reject: total run cost { total_run_cost :.3f} exceeds budget { max_total_cost :.3f} " )
123+ configured_max_total_cost = self .config ["max_total_cost" ]
124+ if configured_max_total_cost is not None :
125+ max_total_cost = float (configured_max_total_cost )
126+ if not cost_summary .complete :
127+ reasons .append ("reject: cost_unavailable for configured max_total_cost" )
128+ elif total_run_cost > max_total_cost :
129+ reasons .append (
130+ f"reject: total run cost { total_run_cost :.3f} exceeds budget { max_total_cost :.3f} "
131+ )
98132
99133 accepted = not any (reason .startswith ("reject:" ) for reason in reasons )
100134 if accepted :
@@ -118,4 +152,61 @@ def decide(
118152 cumulative_cost = round (cumulative_cost , 6 ),
119153 total_run_cost = total_run_cost ,
120154 cost = candidate_cost ,
155+ gate_status = "applied" ,
156+ gate_not_applied_reason = None ,
157+ not_applied_checks = [],
158+ )
159+
160+
161+ def _append_comparability_reasons (
162+ reasons : list [str ],
163+ * ,
164+ split : str ,
165+ baseline : EvalResult ,
166+ candidate : EvalResult ,
167+ ) -> bool :
168+ """Reject malformed result pairs before any dict conversion can hide evidence."""
169+
170+ comparable = True
171+ baseline_ids = [case .case_id for case in baseline .cases ]
172+ candidate_ids = [case .case_id for case in candidate .cases ]
173+ if not baseline_ids :
174+ reasons .append (f"reject: baseline { split } has empty case results" )
175+ comparable = False
176+ if not candidate_ids :
177+ reasons .append (f"reject: candidate { split } has empty case results" )
178+ comparable = False
179+
180+ baseline_duplicates = _duplicates (baseline_ids )
181+ candidate_duplicates = _duplicates (candidate_ids )
182+ if baseline_duplicates :
183+ reasons .append (
184+ f"reject: baseline { split } has duplicate case IDs: { baseline_duplicates } "
185+ )
186+ comparable = False
187+ if candidate_duplicates :
188+ reasons .append (
189+ f"reject: candidate { split } has duplicate case IDs: { candidate_duplicates } "
190+ )
191+ comparable = False
192+
193+ baseline_set = set (baseline_ids )
194+ candidate_set = set (candidate_ids )
195+ if baseline_set != candidate_set :
196+ reasons .append (
197+ f"reject: { split } case ID set mismatch; "
198+ f"missing={ sorted (baseline_set - candidate_set )} , "
199+ f"extra={ sorted (candidate_set - baseline_set )} "
121200 )
201+ comparable = False
202+ return comparable
203+
204+
205+ def _duplicates (case_ids : list [str ]) -> list [str ]:
206+ seen : set [str ] = set ()
207+ duplicates : set [str ] = set ()
208+ for case_id in case_ids :
209+ if case_id in seen :
210+ duplicates .add (case_id )
211+ seen .add (case_id )
212+ return sorted (duplicates )
0 commit comments