1use std::collections::BTreeSet;
4use std::sync::Arc;
5
6use serde::{Deserialize, Serialize};
7
8use crate::aggregation::PredictionUnitId;
9use crate::criteria::{
10 builtin_metric_catalog, fingerprint_without, validate_fingerprint, validate_token,
11 CriterionInput, ImplementationCapability, ImplementationDescriptor, ImplementationSemanticKind,
12 LearningTaskKind, MetricDecomposition, MetricReduction, MetricReference, PortabilityClass,
13 ReplayabilityClass,
14};
15use crate::error::{DagMlError, Result};
16use crate::ids::{FoldId, GroupId, NodeId, ObservationId, SampleId, TargetId, VariantId};
17use crate::implementation_registry::LocalImplementationRegistry;
18use crate::metrics::{compute_metric_per_target, RegressionMetricKind};
19use crate::oof::PredictionPartition;
20use crate::policy::PredictionLevel;
21use crate::training::PredictionKind;
22
23pub const METRIC_EVALUATION_TASK_SCHEMA_VERSION: u32 = 1;
24pub const METRIC_EVALUATION_RESULT_SCHEMA_VERSION: u32 = 1;
25pub const METRIC_EVALUATION_TASK_SCHEMA_ID: &str =
26 "https://github.com/GBeurier/dag-ml/schemas/metric_evaluation_task.v1.schema.json";
27pub const METRIC_EVALUATION_RESULT_SCHEMA_ID: &str =
28 "https://github.com/GBeurier/dag-ml/schemas/metric_evaluation_result.v1.schema.json";
29
30const BUILTIN_METRIC_IMPLEMENTATION_VERSION: &str = "metrics-v1";
31const BUILTIN_METRIC_IMPLEMENTATION_FINGERPRINT: &str =
32 "0aa68bb906c00c9fa433f411f44004d46b5a7f196f9932ca9c313fd184d81d19";
33
34#[derive(Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Serialize, Deserialize)]
35#[serde(
36 rename_all = "snake_case",
37 tag = "level",
38 content = "id",
39 deny_unknown_fields
40)]
41pub enum MetricUnitId {
42 Observation(ObservationId),
43 Sample(SampleId),
44 Target(TargetId),
45 Group(GroupId),
46}
47
48impl MetricUnitId {
49 pub fn level(&self) -> PredictionLevel {
50 match self {
51 Self::Observation(_) => PredictionLevel::Observation,
52 Self::Sample(_) => PredictionLevel::Sample,
53 Self::Target(_) => PredictionLevel::Target,
54 Self::Group(_) => PredictionLevel::Group,
55 }
56 }
57}
58
59impl From<&PredictionUnitId> for MetricUnitId {
60 fn from(value: &PredictionUnitId) -> Self {
61 match value {
62 PredictionUnitId::Sample(id) => Self::Sample(id.clone()),
63 PredictionUnitId::Target(id) => Self::Target(id.clone()),
64 PredictionUnitId::Group(id) => Self::Group(id.clone()),
65 }
66 }
67}
68
69#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
70#[serde(deny_unknown_fields)]
71pub struct MetricEvaluationScope {
72 pub producer_node: NodeId,
73 #[serde(default, skip_serializing_if = "Option::is_none")]
74 pub producer_port: Option<String>,
75 #[serde(default, skip_serializing_if = "Option::is_none")]
76 pub prediction_id: Option<String>,
77 #[serde(default, skip_serializing_if = "Option::is_none")]
78 pub variant_id: Option<VariantId>,
79 pub partition: PredictionPartition,
80 #[serde(default, skip_serializing_if = "Option::is_none")]
81 pub fold_id: Option<FoldId>,
82 pub level: PredictionLevel,
83}
84
85impl MetricEvaluationScope {
86 fn validate(&self) -> Result<()> {
87 for (label, value) in [
88 ("metric producer_port", self.producer_port.as_deref()),
89 ("metric prediction_id", self.prediction_id.as_deref()),
90 ] {
91 if let Some(value) = value {
92 validate_token(label, value)?;
93 }
94 }
95 Ok(())
96 }
97}
98
99#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
100#[serde(deny_unknown_fields)]
101pub struct MetricEvaluationTask {
102 pub schema_version: u32,
103 pub request_id: String,
104 pub metric: MetricReference,
105 pub task_kind: LearningTaskKind,
106 pub prediction_kind: PredictionKind,
107 pub scope: MetricEvaluationScope,
108 pub unit_ids: Vec<MetricUnitId>,
109 pub predictions: Vec<Vec<f64>>,
110 pub targets: Vec<Vec<f64>>,
111 pub output_ids: Vec<String>,
112 #[serde(default, skip_serializing_if = "Option::is_none")]
113 pub sample_weights: Option<Vec<f64>>,
114 #[serde(default, skip_serializing_if = "Option::is_none")]
115 pub missing_mask: Option<Vec<Vec<bool>>>,
116 #[serde(default, skip_serializing_if = "Option::is_none")]
117 pub group_ids: Option<Vec<String>>,
118 pub task_fingerprint: String,
119}
120
121impl MetricEvaluationTask {
122 #[allow(clippy::too_many_arguments)]
123 pub fn new(
124 request_id: impl Into<String>,
125 metric: MetricReference,
126 task_kind: LearningTaskKind,
127 prediction_kind: PredictionKind,
128 scope: MetricEvaluationScope,
129 unit_ids: Vec<MetricUnitId>,
130 predictions: Vec<Vec<f64>>,
131 targets: Vec<Vec<f64>>,
132 output_ids: Vec<String>,
133 sample_weights: Option<Vec<f64>>,
134 missing_mask: Option<Vec<Vec<bool>>>,
135 group_ids: Option<Vec<String>>,
136 ) -> Result<Self> {
137 let mut task = Self {
138 schema_version: METRIC_EVALUATION_TASK_SCHEMA_VERSION,
139 request_id: request_id.into(),
140 metric,
141 task_kind,
142 prediction_kind,
143 scope,
144 unit_ids,
145 predictions,
146 targets,
147 output_ids,
148 sample_weights,
149 missing_mask,
150 group_ids,
151 task_fingerprint: String::new(),
152 };
153 task.task_fingerprint = task.compute_fingerprint()?;
154 task.validate()?;
155 Ok(task)
156 }
157
158 pub fn from_json(json: &str) -> Result<Self> {
159 let task: Self = crate::canonical::deserialize_external_contract(
160 json,
161 "metric evaluation task",
162 DagMlError::CampaignValidation,
163 )?;
164 task.validate()?;
165 Ok(task)
166 }
167
168 pub fn compute_fingerprint(&self) -> Result<String> {
169 fingerprint_without(self, "task_fingerprint", "metric evaluation task")
170 }
171
172 pub fn validate(&self) -> Result<()> {
173 if self.schema_version != METRIC_EVALUATION_TASK_SCHEMA_VERSION {
174 return task_error(format!(
175 "metric evaluation task schema_version {} is unsupported",
176 self.schema_version
177 ));
178 }
179 validate_token("metric request_id", &self.request_id)?;
180 self.metric.validate()?;
181 self.metric.spec.validate_compatibility(
182 self.task_kind,
183 self.prediction_kind,
184 self.scope.level,
185 )?;
186 self.scope.validate()?;
187 let row_count = self.unit_ids.len();
188 if row_count == 0 {
189 return task_error("metric evaluation task has no units");
190 }
191 if self
192 .unit_ids
193 .iter()
194 .any(|unit| unit.level() != self.scope.level)
195 {
196 return task_error("metric evaluation unit level does not match scope");
197 }
198 if self.unit_ids.iter().collect::<BTreeSet<_>>().len() != row_count {
199 return task_error("metric evaluation task contains duplicate unit ids");
200 }
201 let prediction_width =
202 validate_finite_matrix("metric predictions", &self.predictions, row_count)?;
203 let target_width = validate_finite_matrix("metric targets", &self.targets, row_count)?;
204 if self.output_ids.len() != target_width || self.output_ids.is_empty() {
205 return task_error(format!(
206 "metric output_ids length {} does not match target width {target_width}",
207 self.output_ids.len()
208 ));
209 }
210 let mut outputs = BTreeSet::new();
211 for output_id in &self.output_ids {
212 validate_token("metric output_id", output_id)?;
213 if !outputs.insert(output_id) {
214 return task_error(format!("duplicate metric output_id `{output_id}`"));
215 }
216 }
217 if matches!(
218 self.prediction_kind,
219 PredictionKind::RegressionPoint | PredictionKind::ClassLabel
220 ) && prediction_width != target_width
221 {
222 return task_error(format!(
223 "metric prediction width {prediction_width} does not match target width {target_width}"
224 ));
225 }
226 validate_optional_inputs(self, row_count, target_width)?;
227 validate_fingerprint("metric evaluation task", &self.task_fingerprint)?;
228 let expected = self.compute_fingerprint()?;
229 if self.task_fingerprint != expected {
230 return task_error(format!(
231 "metric evaluation task fingerprint mismatch: declared {}, expected {expected}",
232 self.task_fingerprint
233 ));
234 }
235 Ok(())
236 }
237}
238
239#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
240#[serde(deny_unknown_fields)]
241pub struct MetricEvaluationValue {
242 #[serde(default, skip_serializing_if = "Option::is_none")]
243 pub unit_id: Option<MetricUnitId>,
244 #[serde(default, skip_serializing_if = "Option::is_none")]
245 pub output_id: Option<String>,
246 pub value: f64,
247}
248
249#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
250#[serde(deny_unknown_fields)]
251pub struct MetricEvaluationResult {
252 pub schema_version: u32,
253 pub request_id: String,
254 pub semantic_id: String,
255 pub semantic_fingerprint: String,
256 pub implementation_fingerprint: String,
257 pub descriptor_fingerprint: String,
258 pub scope: MetricEvaluationScope,
259 pub values: Vec<MetricEvaluationValue>,
260 pub result_fingerprint: String,
261}
262
263impl MetricEvaluationResult {
264 pub fn for_task(
265 task: &MetricEvaluationTask,
266 values: Vec<MetricEvaluationValue>,
267 ) -> Result<Self> {
268 let mut result = Self {
269 schema_version: METRIC_EVALUATION_RESULT_SCHEMA_VERSION,
270 request_id: task.request_id.clone(),
271 semantic_id: task.metric.spec.metric_id.clone(),
272 semantic_fingerprint: task.metric.spec.spec_fingerprint.clone(),
273 implementation_fingerprint: task
274 .metric
275 .implementation
276 .implementation_fingerprint
277 .clone(),
278 descriptor_fingerprint: task.metric.implementation.descriptor_fingerprint.clone(),
279 scope: task.scope.clone(),
280 values,
281 result_fingerprint: String::new(),
282 };
283 result.result_fingerprint = result.compute_fingerprint()?;
284 result.validate_against(task)?;
285 Ok(result)
286 }
287
288 pub fn from_json_for_task(json: &str, task: &MetricEvaluationTask) -> Result<Self> {
289 let result: Self = crate::canonical::deserialize_external_contract(
290 json,
291 "metric evaluation result",
292 DagMlError::RuntimeValidation,
293 )?;
294 result.validate_against(task)?;
295 Ok(result)
296 }
297
298 pub fn compute_fingerprint(&self) -> Result<String> {
299 fingerprint_without(self, "result_fingerprint", "metric evaluation result")
300 }
301
302 pub fn validate_against(&self, task: &MetricEvaluationTask) -> Result<()> {
303 task.validate()?;
304 if self.schema_version != METRIC_EVALUATION_RESULT_SCHEMA_VERSION {
305 return result_error(format!(
306 "metric evaluation result schema_version {} is unsupported",
307 self.schema_version
308 ));
309 }
310 if self.request_id != task.request_id
311 || self.semantic_id != task.metric.spec.metric_id
312 || self.semantic_fingerprint != task.metric.spec.spec_fingerprint
313 || self.implementation_fingerprint
314 != task.metric.implementation.implementation_fingerprint
315 || self.descriptor_fingerprint != task.metric.implementation.descriptor_fingerprint
316 {
317 return result_error("metric provider identity/fingerprint does not match task");
318 }
319 if self.scope != task.scope {
320 return result_error("metric provider result scope does not match task");
321 }
322 if self.values.is_empty() {
323 return result_error("metric provider returned no values");
324 }
325 if self.values.iter().any(|value| !value.value.is_finite()) {
326 return result_error("metric provider returned a non-finite value");
327 }
328 validate_result_coverage(self, task)?;
329 validate_fingerprint("metric evaluation result", &self.result_fingerprint)
330 .map_err(|error| DagMlError::RuntimeValidation(error.to_string()))?;
331 let expected = self.compute_fingerprint()?;
332 if self.result_fingerprint != expected {
333 return result_error(format!(
334 "metric evaluation result fingerprint mismatch: declared {}, expected {expected}",
335 self.result_fingerprint
336 ));
337 }
338 Ok(())
339 }
340
341 pub fn aggregate_for_task(&self, task: &MetricEvaluationTask) -> Result<f64> {
342 self.validate_against(task)?;
343 self.reduce(task)
344 }
345
346 fn reduce(&self, task: &MetricEvaluationTask) -> Result<f64> {
347 let value = match task.metric.spec.reduction {
348 MetricReduction::Global => self.values[0].value,
349 MetricReduction::Mean => {
350 self.values.iter().map(|value| value.value).sum::<f64>() / self.values.len() as f64
351 }
352 MetricReduction::Sum => self.values.iter().map(|value| value.value).sum(),
353 MetricReduction::WeightedMean => {
354 let weights = task.sample_weights.as_ref().ok_or_else(|| {
355 DagMlError::RuntimeValidation(
356 "weighted metric reduction has no sample weights".to_string(),
357 )
358 })?;
359 let weighted_sum = self
360 .values
361 .iter()
362 .zip(weights)
363 .map(|(value, weight)| value.value * weight)
364 .sum::<f64>();
365 weighted_sum / weights.iter().sum::<f64>()
366 }
367 };
368 if !value.is_finite() {
369 return result_error("metric reduction produced a non-finite value");
370 }
371 Ok(value)
372 }
373}
374
375#[derive(Clone, Debug, PartialEq)]
376pub struct ValidatedMetricEvaluation {
377 pub result: MetricEvaluationResult,
378 pub aggregate: f64,
379}
380
381pub trait MetricProvider: Send + Sync {
382 fn evaluate(&self, task: &MetricEvaluationTask) -> Result<MetricEvaluationResult>;
383}
384
385#[derive(Default)]
386pub struct MetricProviderRegistry {
387 providers: LocalImplementationRegistry<Arc<dyn MetricProvider>>,
388}
389
390impl MetricProviderRegistry {
391 pub fn register(
392 &mut self,
393 descriptor: ImplementationDescriptor,
394 provider: Arc<dyn MetricProvider>,
395 ) -> Result<()> {
396 descriptor.validate()?;
397 if descriptor.semantic_kind != ImplementationSemanticKind::Metric {
398 return task_error("metric provider registry rejects non-metric descriptor");
399 }
400 self.providers.register(descriptor, provider)
401 }
402
403 pub fn evaluate(&self, task: &MetricEvaluationTask) -> Result<ValidatedMetricEvaluation> {
404 task.validate()?;
405 let provider = self.providers.resolve_metric(&task.metric)?;
406 let result = provider.evaluate(task)?;
407 let aggregate = result.aggregate_for_task(task)?;
408 Ok(ValidatedMetricEvaluation { result, aggregate })
409 }
410}
411
412pub fn builtin_metric_reference(metric: RegressionMetricKind) -> Result<MetricReference> {
413 let metric_id = format!("dagml.metric.{}@1", metric.name());
414 let spec = builtin_metric_catalog()?
415 .remove(&metric_id)
416 .ok_or_else(|| {
417 DagMlError::CampaignValidation(format!("missing `{metric_id}` catalog entry"))
418 })?;
419 let implementation = ImplementationDescriptor::new(
420 ImplementationSemanticKind::Metric,
421 &spec.metric_id,
422 &spec.spec_fingerprint,
423 "provider:dag-ml-core",
424 "binding:rust",
425 BUILTIN_METRIC_IMPLEMENTATION_VERSION,
426 BUILTIN_METRIC_IMPLEMENTATION_FINGERPRINT,
427 BTreeSet::new(),
428 BTreeSet::new(),
429 BTreeSet::from([ImplementationCapability::Deterministic]),
430 PortabilityClass::PortableBuiltIn,
431 ReplayabilityClass::Detached,
432 None,
433 )?;
434 let reference = MetricReference {
435 spec,
436 implementation,
437 };
438 reference.validate()?;
439 Ok(reference)
440}
441
442pub fn builtin_metric_registry() -> Result<MetricProviderRegistry> {
443 let mut registry = MetricProviderRegistry::default();
444 for metric in [
445 RegressionMetricKind::Mse,
446 RegressionMetricKind::Rmse,
447 RegressionMetricKind::Mae,
448 RegressionMetricKind::R2,
449 RegressionMetricKind::Accuracy,
450 RegressionMetricKind::BalancedAccuracy,
451 ] {
452 let reference = builtin_metric_reference(metric)?;
453 registry.register(
454 reference.implementation,
455 Arc::new(BuiltinMetricProvider { metric }),
456 )?;
457 }
458 Ok(registry)
459}
460
461struct BuiltinMetricProvider {
462 metric: RegressionMetricKind,
463}
464
465impl MetricProvider for BuiltinMetricProvider {
466 fn evaluate(&self, task: &MetricEvaluationTask) -> Result<MetricEvaluationResult> {
467 let expected_id = format!("dagml.metric.{}@1", self.metric.name());
468 if task.metric.spec.metric_id != expected_id {
469 return result_error(format!(
470 "built-in provider `{}` cannot evaluate `{}`",
471 self.metric.name(),
472 task.metric.spec.metric_id
473 ));
474 }
475 let predictions = task
476 .predictions
477 .iter()
478 .map(Vec::as_slice)
479 .collect::<Vec<_>>();
480 let targets = task.targets.iter().map(Vec::as_slice).collect::<Vec<_>>();
481 let values =
482 compute_metric_per_target(self.metric, task.output_ids.len(), &predictions, &targets)
483 .into_iter()
484 .zip(&task.output_ids)
485 .map(|(value, output_id)| MetricEvaluationValue {
486 unit_id: None,
487 output_id: Some(output_id.clone()),
488 value,
489 })
490 .collect();
491 MetricEvaluationResult::for_task(task, values)
492 }
493}
494
495fn validate_finite_matrix(label: &str, values: &[Vec<f64>], expected_rows: usize) -> Result<usize> {
496 if values.len() != expected_rows {
497 return task_error(format!(
498 "{label} has {} rows for {expected_rows} units",
499 values.len()
500 ));
501 }
502 let width = values.first().map_or(0, Vec::len);
503 if width == 0 || values.iter().any(|row| row.len() != width) {
504 return task_error(format!("{label} is empty or ragged"));
505 }
506 if values.iter().flatten().any(|value| !value.is_finite()) {
507 return task_error(format!("{label} contains non-finite values"));
508 }
509 Ok(width)
510}
511
512fn validate_optional_inputs(
513 task: &MetricEvaluationTask,
514 row_count: usize,
515 target_width: usize,
516) -> Result<()> {
517 let required = &task.metric.spec.required_inputs;
518 match &task.sample_weights {
519 Some(weights) => {
520 if !task
521 .metric
522 .spec
523 .capabilities
524 .contains(&crate::criteria::MetricCapability::SupportsSampleWeights)
525 {
526 return task_error("metric task supplies unsupported sample weights");
527 }
528 if weights.len() != row_count
529 || weights
530 .iter()
531 .any(|weight| !weight.is_finite() || *weight < 0.0)
532 || weights.iter().sum::<f64>() <= 0.0
533 {
534 return task_error("metric sample weights are invalid");
535 }
536 }
537 None if required.contains(&CriterionInput::SampleWeight) => {
538 return task_error("metric task is missing required sample weights");
539 }
540 None => {}
541 }
542 match &task.missing_mask {
543 Some(mask) => {
544 if !task
545 .metric
546 .spec
547 .capabilities
548 .contains(&crate::criteria::MetricCapability::SupportsMissingMask)
549 {
550 return task_error("metric task supplies unsupported missing mask");
551 }
552 if mask.len() != row_count || mask.iter().any(|row| row.len() != target_width) {
553 return task_error("metric missing mask shape does not match targets");
554 }
555 }
556 None if required.contains(&CriterionInput::MissingMask) => {
557 return task_error("metric task is missing required missing mask");
558 }
559 None => {}
560 }
561 match &task.group_ids {
562 Some(group_ids) => {
563 if !required.contains(&CriterionInput::Group) {
564 return task_error("metric task supplies undeclared group ids");
565 }
566 if group_ids.len() != row_count {
567 return task_error("metric group_ids length does not match units");
568 }
569 for group_id in group_ids {
570 validate_token("metric group_id", group_id)?;
571 }
572 }
573 None if required.contains(&CriterionInput::Group) => {
574 return task_error("metric task is missing required group ids");
575 }
576 None => {}
577 }
578 Ok(())
579}
580
581fn validate_result_coverage(
582 result: &MetricEvaluationResult,
583 task: &MetricEvaluationTask,
584) -> Result<()> {
585 match task.metric.spec.decomposition {
586 MetricDecomposition::Global => {
587 if result.values.len() != 1
588 || result.values[0].unit_id.is_some()
589 || result.values[0].output_id.is_some()
590 {
591 return result_error("global metric provider result has wrong coverage");
592 }
593 }
594 MetricDecomposition::PerOutput => {
595 if result.values.len() != task.output_ids.len() {
596 return result_error("per-output metric provider result has wrong coverage");
597 }
598 for (value, output_id) in result.values.iter().zip(&task.output_ids) {
599 if value.unit_id.is_some() || value.output_id.as_ref() != Some(output_id) {
600 return result_error("per-output metric provider result has wrong scope/order");
601 }
602 }
603 }
604 MetricDecomposition::PerUnit => {
605 if result.values.len() != task.unit_ids.len() {
606 return result_error("per-unit metric provider result has wrong coverage");
607 }
608 for (value, unit_id) in result.values.iter().zip(&task.unit_ids) {
609 if value.output_id.is_some() || value.unit_id.as_ref() != Some(unit_id) {
610 return result_error("per-unit metric provider result has wrong scope/order");
611 }
612 }
613 }
614 }
615 Ok(())
616}
617
618fn task_error<T>(message: impl Into<String>) -> Result<T> {
619 Err(DagMlError::CampaignValidation(message.into()))
620}
621
622fn result_error<T>(message: impl Into<String>) -> Result<T> {
623 Err(DagMlError::RuntimeValidation(message.into()))
624}
625
626#[cfg(test)]
627mod tests {
628 use serde_json::json;
629
630 use super::*;
631 use crate::criteria::{
632 ImplementationSemanticKind, MetricCapability, MetricSpec, SemanticSpecKind,
633 };
634 use crate::selection::MetricObjective;
635
636 fn sample_scope() -> MetricEvaluationScope {
637 MetricEvaluationScope {
638 producer_node: NodeId::new("model:custom").unwrap(),
639 producer_port: Some("prediction".to_string()),
640 prediction_id: Some("prediction:validation".to_string()),
641 variant_id: None,
642 partition: PredictionPartition::Validation,
643 fold_id: Some(FoldId::new("fold:0").unwrap()),
644 level: PredictionLevel::Sample,
645 }
646 }
647
648 fn custom_bias_reference() -> MetricReference {
649 let spec = MetricSpec::new(
650 "example.metric.bias@1",
651 SemanticSpecKind::Custom,
652 BTreeSet::from([LearningTaskKind::Regression]),
653 BTreeSet::from([PredictionKind::RegressionPoint]),
654 MetricObjective::Minimize,
655 BTreeSet::from([PredictionLevel::Sample]),
656 MetricDecomposition::PerUnit,
657 MetricReduction::Mean,
658 BTreeSet::from([CriterionInput::Target, CriterionInput::Prediction]),
659 BTreeSet::from([MetricCapability::Decomposable]),
660 json!({}),
661 )
662 .unwrap();
663 let implementation = ImplementationDescriptor::new(
664 ImplementationSemanticKind::Metric,
665 &spec.metric_id,
666 &spec.spec_fingerprint,
667 "provider:rust-local",
668 "binding:rust",
669 "1.0.0",
670 "4991854599d650fd613dfd02b10d90a649ad7fec85f20a027d5e7b2a553f628b",
671 BTreeSet::new(),
672 BTreeSet::new(),
673 BTreeSet::from([ImplementationCapability::Deterministic]),
674 PortabilityClass::HostLocal,
675 ReplayabilityClass::RegistryRequired,
676 Some("metric:run-123:bias".to_string()),
677 )
678 .unwrap();
679 MetricReference {
680 spec,
681 implementation,
682 }
683 }
684
685 fn custom_task() -> MetricEvaluationTask {
686 MetricEvaluationTask::new(
687 "metric-request:bias",
688 custom_bias_reference(),
689 LearningTaskKind::Regression,
690 PredictionKind::RegressionPoint,
691 sample_scope(),
692 vec![
693 MetricUnitId::Sample(SampleId::new("sample:0").unwrap()),
694 MetricUnitId::Sample(SampleId::new("sample:1").unwrap()),
695 ],
696 vec![vec![2.0], vec![5.0]],
697 vec![vec![1.0], vec![3.0]],
698 vec!["target".to_string()],
699 None,
700 None,
701 None,
702 )
703 .unwrap()
704 }
705
706 struct BiasProvider;
707
708 impl MetricProvider for BiasProvider {
709 fn evaluate(&self, task: &MetricEvaluationTask) -> Result<MetricEvaluationResult> {
710 let values = task
711 .unit_ids
712 .iter()
713 .zip(task.predictions.iter().zip(&task.targets))
714 .map(|(unit_id, (prediction, target))| MetricEvaluationValue {
715 unit_id: Some(unit_id.clone()),
716 output_id: None,
717 value: prediction[0] - target[0],
718 })
719 .collect();
720 MetricEvaluationResult::for_task(task, values)
721 }
722 }
723
724 #[test]
725 fn custom_metric_registry_executes_and_reduces_provider_values() {
726 let task = custom_task();
727 let mut registry = MetricProviderRegistry::default();
728 registry
729 .register(task.metric.implementation.clone(), Arc::new(BiasProvider))
730 .unwrap();
731 let evaluation = registry.evaluate(&task).unwrap();
732 assert_eq!(evaluation.aggregate, 1.5);
733 assert_eq!(evaluation.result.values.len(), 2);
734 }
735
736 #[test]
737 fn task_rejects_custom_metric_without_objective() {
738 let task = custom_task();
739 let mut value = serde_json::to_value(task).unwrap();
740 value["metric"]["spec"]
741 .as_object_mut()
742 .unwrap()
743 .remove("objective");
744 let error = MetricEvaluationTask::from_json(&value.to_string())
745 .unwrap_err()
746 .to_string();
747 assert!(error.contains("objective"));
748 }
749
750 #[test]
751 fn provider_result_rejects_nonfinite_wrong_scope_coverage_and_fingerprint() {
752 let task = custom_task();
753 let valid = BiasProvider.evaluate(&task).unwrap();
754
755 let mut nonfinite = valid.clone();
756 nonfinite.values[0].value = f64::NAN;
757 assert!(nonfinite
758 .validate_against(&task)
759 .unwrap_err()
760 .to_string()
761 .contains("non-finite"));
762
763 let mut wrong_scope = valid.clone();
764 wrong_scope.scope.partition = PredictionPartition::Test;
765 wrong_scope.result_fingerprint = wrong_scope.compute_fingerprint().unwrap();
766 assert!(wrong_scope
767 .validate_against(&task)
768 .unwrap_err()
769 .to_string()
770 .contains("scope"));
771
772 let mut wrong_coverage = valid.clone();
773 wrong_coverage.values.pop();
774 wrong_coverage.result_fingerprint = wrong_coverage.compute_fingerprint().unwrap();
775 assert!(wrong_coverage
776 .validate_against(&task)
777 .unwrap_err()
778 .to_string()
779 .contains("coverage"));
780
781 let mut wrong_fingerprint = valid;
782 wrong_fingerprint.implementation_fingerprint = "0".repeat(64);
783 wrong_fingerprint.result_fingerprint = wrong_fingerprint.compute_fingerprint().unwrap();
784 assert!(wrong_fingerprint
785 .validate_against(&task)
786 .unwrap_err()
787 .to_string()
788 .contains("identity/fingerprint"));
789 }
790
791 #[test]
792 fn built_in_registry_uses_existing_metric_kernel_and_per_output_reduction() {
793 let reference = builtin_metric_reference(RegressionMetricKind::Rmse).unwrap();
794 let task = MetricEvaluationTask::new(
795 "metric-request:rmse",
796 reference,
797 LearningTaskKind::Regression,
798 PredictionKind::RegressionPoint,
799 sample_scope(),
800 vec![
801 MetricUnitId::Sample(SampleId::new("sample:0").unwrap()),
802 MetricUnitId::Sample(SampleId::new("sample:1").unwrap()),
803 ],
804 vec![vec![2.0, 4.0], vec![4.0, 8.0]],
805 vec![vec![1.0, 2.0], vec![3.0, 6.0]],
806 vec!["a".to_string(), "b".to_string()],
807 None,
808 None,
809 None,
810 )
811 .unwrap();
812 let evaluation = builtin_metric_registry().unwrap().evaluate(&task).unwrap();
813 assert_eq!(evaluation.result.values[0].value, 1.0);
814 assert_eq!(evaluation.result.values[1].value, 2.0);
815 assert_eq!(evaluation.aggregate, 1.5);
816 }
817
818 #[test]
819 fn registry_rejects_descriptor_substitution_even_with_same_registry_key() {
820 let task = custom_task();
821 let mut substituted = task.metric.implementation.clone();
822 substituted.implementation_version = "2.0.0".to_string();
823 substituted.descriptor_fingerprint = substituted.compute_fingerprint().unwrap();
824 let mut registry = MetricProviderRegistry::default();
825 registry
826 .register(substituted, Arc::new(BiasProvider))
827 .unwrap();
828 assert!(registry
829 .evaluate(&task)
830 .unwrap_err()
831 .to_string()
832 .contains("descriptor"));
833 }
834
835 #[test]
836 fn published_provider_fixture_matches_rust_task_and_result_contracts() {
837 let fixture: serde_json::Value = serde_json::from_str(include_str!(
838 "../../../examples/fixtures/criteria/metric_provider_contracts.v1.json"
839 ))
840 .unwrap();
841 let task = MetricEvaluationTask::from_json(&fixture["valid"]["task"].to_string()).unwrap();
842 let result = MetricEvaluationResult::from_json_for_task(
843 &fixture["valid"]["result"].to_string(),
844 &task,
845 )
846 .unwrap();
847 assert_eq!(
848 result.reduce(&task).unwrap(),
849 fixture["valid"]["aggregate"].as_f64().unwrap()
850 );
851
852 for case in fixture["invalid"].as_array().unwrap() {
853 let document = case["document"].to_string();
854 let rejected = match case["contract"].as_str().unwrap() {
855 "metric_evaluation_task" => MetricEvaluationTask::from_json(&document).is_err(),
856 "metric_evaluation_result" => {
857 MetricEvaluationResult::from_json_for_task(&document, &task).is_err()
858 }
859 contract => panic!("unknown metric-provider fixture contract `{contract}`"),
860 };
861 assert!(rejected, "negative case `{}` was accepted", case["id"]);
862 }
863 }
864}