Skip to main content

dag_ml_core/
metric_provider.rs

1//! Typed metric-provider dispatch shared by native and binding-local implementations.
2
3use 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}