Skip to main content

dag_ml_core/
controller.rs

1use std::collections::{BTreeMap, BTreeSet};
2
3use serde::{Deserialize, Serialize};
4
5use crate::data::ModelInputSpec;
6use crate::error::{DagMlError, Result};
7use crate::graph::{NodeKind, NodeSpec, PortKind, PortSpec};
8use crate::ids::ControllerId;
9use crate::phase::Phase;
10use crate::policy::FitInfluencePolicy;
11
12pub const CONTROLLER_MANIFEST_SCHEMA_VERSION: u32 = 1;
13pub const CONTROLLER_MANIFEST_SCHEMA_ID: &str =
14    "https://github.com/GBeurier/dag-ml/schemas/controller_manifest.v1.schema.json";
15
16#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Serialize, Deserialize)]
17#[serde(rename_all = "snake_case")]
18pub enum ControllerCapability {
19    Deterministic,
20    ThreadSafe,
21    ProcessSafe,
22    NeedsPythonGil,
23    EmitsPredictions,
24    ConsumesOofPredictions,
25    EmitsArtifacts,
26    Stateful,
27    EmitsRelation,
28    UsesCoreRng,
29    ShapeChanging,
30    GeneratesData,
31    GeneratesModel,
32    ExpandsVariants,
33    AggregatesPredictions,
34    SupportsSampleWeights,
35    SupportsRowResampling,
36    SupportsBackendLossWeights,
37    SupportsMissingMasks,
38    SupportsConfigurableLoss,
39    SupportsCustomLoss,
40    SupportsDifferentiableLoss,
41    /// Controller actively consumes non-uniform training influence, distinct
42    /// from merely supporting a weighting API.
43    UsesTrainingWeights,
44    /// Controller actively reads validation samples during fitting.
45    UsesEarlyStopping,
46    /// Controller performs a nested/internal candidate selection.
47    PerformsInternalTuning,
48    /// Prediction aggregation itself has fitted state.
49    TrainsAggregation,
50}
51
52#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Serialize, Deserialize)]
53#[serde(rename_all = "snake_case")]
54pub enum ControllerFitScope {
55    Stateless,
56    FoldTrain,
57    FullTrain,
58    InferenceOnly,
59}
60
61#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Serialize, Deserialize)]
62#[serde(rename_all = "snake_case")]
63pub enum RngPolicy {
64    UsesCoreSeed,
65    IgnoresSeed,
66    ExternallyDeterministic,
67    Nondeterministic,
68}
69
70#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Serialize, Deserialize)]
71#[serde(rename_all = "snake_case")]
72pub enum ArtifactPolicy {
73    Serializable,
74    HostOnly,
75    ContentAddressed,
76    ReplayRequired,
77}
78
79#[derive(Clone, Debug, Default, Eq, PartialEq, Serialize, Deserialize)]
80#[serde(deny_unknown_fields)]
81pub struct OperatorSelector {
82    #[serde(default, skip_serializing_if = "BTreeSet::is_empty")]
83    pub aliases: BTreeSet<String>,
84    #[serde(default, skip_serializing_if = "BTreeSet::is_empty")]
85    pub classes: BTreeSet<String>,
86    #[serde(default, skip_serializing_if = "BTreeSet::is_empty")]
87    pub class_prefixes: BTreeSet<String>,
88    #[serde(default, skip_serializing_if = "BTreeSet::is_empty")]
89    pub functions: BTreeSet<String>,
90    #[serde(default, skip_serializing_if = "BTreeSet::is_empty")]
91    pub refs: BTreeSet<String>,
92    #[serde(default, skip_serializing_if = "BTreeSet::is_empty")]
93    pub types: BTreeSet<String>,
94}
95
96impl OperatorSelector {
97    fn validate(&self, controller_id: &ControllerId) -> Result<()> {
98        if self.aliases.is_empty()
99            && self.classes.is_empty()
100            && self.class_prefixes.is_empty()
101            && self.functions.is_empty()
102            && self.refs.is_empty()
103            && self.types.is_empty()
104        {
105            return Err(DagMlError::ControllerValidation(format!(
106                "controller `{controller_id}` has an empty operator selector"
107            )));
108        }
109        for (field, values) in [
110            ("aliases", &self.aliases),
111            ("classes", &self.classes),
112            ("class_prefixes", &self.class_prefixes),
113            ("functions", &self.functions),
114            ("refs", &self.refs),
115            ("types", &self.types),
116        ] {
117            if values.iter().any(|value| value.trim().is_empty()) {
118                return Err(DagMlError::ControllerValidation(format!(
119                    "controller `{controller_id}` operator selector `{field}` contains an empty value"
120                )));
121            }
122        }
123        Ok(())
124    }
125}
126
127#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
128#[serde(deny_unknown_fields)]
129pub struct ControllerManifest {
130    pub controller_id: ControllerId,
131    pub controller_version: String,
132    pub operator_kind: NodeKind,
133    #[serde(default)]
134    pub priority: u32,
135    #[serde(default)]
136    pub supported_phases: BTreeSet<Phase>,
137    #[serde(default)]
138    pub input_ports: Vec<PortSpec>,
139    #[serde(default)]
140    pub output_ports: Vec<PortSpec>,
141    #[serde(default)]
142    pub data_requirements: Option<serde_json::Value>,
143    #[serde(default)]
144    pub capabilities: BTreeSet<ControllerCapability>,
145    #[serde(default, skip_serializing_if = "Vec::is_empty")]
146    pub operator_selectors: Vec<OperatorSelector>,
147    pub fit_scope: ControllerFitScope,
148    pub rng_policy: RngPolicy,
149    pub artifact_policy: ArtifactPolicy,
150}
151
152impl ControllerManifest {
153    pub fn validate(&self) -> Result<()> {
154        if self.controller_version.trim().is_empty() {
155            return Err(DagMlError::ControllerValidation(format!(
156                "controller `{}` has an empty version",
157                self.controller_id
158            )));
159        }
160        if self.supported_phases.is_empty() {
161            return Err(DagMlError::ControllerValidation(format!(
162                "controller `{}` supports no phases",
163                self.controller_id
164            )));
165        }
166        if let Some(model_input) = self.model_input_spec()? {
167            model_input.validate().map_err(|error| {
168                DagMlError::ControllerValidation(format!(
169                    "controller `{}` data_requirements are not a valid ModelInputSpec: {error}",
170                    self.controller_id
171                ))
172            })?;
173        }
174        validate_ports(&self.controller_id, "input", &self.input_ports)?;
175        validate_ports(&self.controller_id, "output", &self.output_ports)?;
176        for selector in &self.operator_selectors {
177            selector.validate(&self.controller_id)?;
178        }
179        if self.rng_policy == RngPolicy::Nondeterministic
180            && self
181                .capabilities
182                .contains(&ControllerCapability::Deterministic)
183        {
184            return Err(DagMlError::ControllerValidation(format!(
185                "controller `{}` cannot be deterministic with nondeterministic RNG",
186                self.controller_id
187            )));
188        }
189        if self.fit_scope == ControllerFitScope::InferenceOnly
190            && (self.supported_phases.contains(&Phase::FitCv)
191                || self.supported_phases.contains(&Phase::Refit))
192        {
193            return Err(DagMlError::ControllerValidation(format!(
194                "controller `{}` is inference_only but supports training phases",
195                self.controller_id
196            )));
197        }
198        if self.supported_phases.contains(&Phase::FitCv)
199            && matches!(
200                self.fit_scope,
201                ControllerFitScope::FullTrain | ControllerFitScope::InferenceOnly
202            )
203        {
204            return Err(DagMlError::ControllerValidation(format!(
205                "controller `{}` supports FIT_CV but has fit_scope {:?}",
206                self.controller_id, self.fit_scope
207            )));
208        }
209        if self
210            .output_ports
211            .iter()
212            .any(|port| port.kind == PortKind::Prediction)
213            && !self
214                .capabilities
215                .contains(&ControllerCapability::EmitsPredictions)
216        {
217            return Err(DagMlError::ControllerValidation(format!(
218                "controller `{}` has prediction output ports but lacks emits_predictions",
219                self.controller_id
220            )));
221        }
222        if self
223            .output_ports
224            .iter()
225            .any(|port| port.kind == PortKind::Artifact)
226            && !self
227                .capabilities
228                .contains(&ControllerCapability::EmitsArtifacts)
229        {
230            return Err(DagMlError::ControllerValidation(format!(
231                "controller `{}` has artifact output ports but lacks emits_artifacts",
232                self.controller_id
233            )));
234        }
235        let active_influence = [
236            ControllerCapability::UsesTrainingWeights,
237            ControllerCapability::UsesEarlyStopping,
238            ControllerCapability::PerformsInternalTuning,
239            ControllerCapability::TrainsAggregation,
240        ];
241        if matches!(
242            self.fit_scope,
243            ControllerFitScope::Stateless | ControllerFitScope::InferenceOnly
244        ) && active_influence
245            .iter()
246            .any(|capability| self.capabilities.contains(capability))
247        {
248            return Err(DagMlError::ControllerValidation(format!(
249                "controller `{}` has active training-influence capabilities with fit_scope {:?}",
250                self.controller_id, self.fit_scope
251            )));
252        }
253        if self
254            .capabilities
255            .contains(&ControllerCapability::UsesTrainingWeights)
256            && !self.capabilities.iter().any(|capability| {
257                matches!(
258                    capability,
259                    ControllerCapability::SupportsSampleWeights
260                        | ControllerCapability::SupportsRowResampling
261                        | ControllerCapability::SupportsBackendLossWeights
262                )
263            })
264        {
265            return Err(DagMlError::ControllerValidation(format!(
266                "controller `{}` uses training weights without a supported weighting mechanism",
267                self.controller_id
268            )));
269        }
270        if self
271            .capabilities
272            .contains(&ControllerCapability::TrainsAggregation)
273            && !self
274                .capabilities
275                .contains(&ControllerCapability::AggregatesPredictions)
276        {
277            return Err(DagMlError::ControllerValidation(format!(
278                "controller `{}` trains aggregation without aggregates_predictions",
279                self.controller_id
280            )));
281        }
282        if self
283            .capabilities
284            .contains(&ControllerCapability::SupportsCustomLoss)
285            && !self
286                .capabilities
287                .contains(&ControllerCapability::SupportsConfigurableLoss)
288        {
289            return Err(DagMlError::ControllerValidation(format!(
290                "controller `{}` supports custom loss without configurable loss",
291                self.controller_id
292            )));
293        }
294        if self
295            .capabilities
296            .contains(&ControllerCapability::SupportsDifferentiableLoss)
297            && !self
298                .capabilities
299                .contains(&ControllerCapability::SupportsConfigurableLoss)
300        {
301            return Err(DagMlError::ControllerValidation(format!(
302                "controller `{}` supports differentiable loss without configurable loss",
303                self.controller_id
304            )));
305        }
306        Ok(())
307    }
308
309    pub fn supports_phase(&self, phase: Phase) -> bool {
310        self.supported_phases.contains(&phase)
311    }
312
313    pub fn supports_parallel_invocation(&self) -> bool {
314        self.capabilities
315            .contains(&ControllerCapability::ThreadSafe)
316            || self
317                .capabilities
318                .contains(&ControllerCapability::ProcessSafe)
319    }
320
321    pub fn supports_fit_influence_policy(&self, policy: FitInfluencePolicy) -> bool {
322        capabilities_support_fit_influence(&self.capabilities, policy)
323    }
324
325    pub fn model_input_spec(&self) -> Result<Option<ModelInputSpec>> {
326        self.data_requirements
327            .as_ref()
328            .map(|value| {
329                serde_json::from_value::<ModelInputSpec>(value.clone()).map_err(|error| {
330                    DagMlError::ControllerValidation(format!(
331                        "controller `{}` data_requirements must be ModelInputSpec JSON: {error}",
332                        self.controller_id
333                    ))
334                })
335            })
336            .transpose()
337    }
338}
339
340pub fn capabilities_support_fit_influence(
341    capabilities: &BTreeSet<ControllerCapability>,
342    policy: FitInfluencePolicy,
343) -> bool {
344    match policy {
345        FitInfluencePolicy::Auto
346        | FitInfluencePolicy::UniformRows
347        | FitInfluencePolicy::ScorerOnly => true,
348        FitInfluencePolicy::EqualSampleInfluence => {
349            capabilities.contains(&ControllerCapability::SupportsSampleWeights)
350        }
351        FitInfluencePolicy::ResampleEqualized => {
352            capabilities.contains(&ControllerCapability::SupportsRowResampling)
353        }
354        FitInfluencePolicy::BackendLossWeight => {
355            capabilities.contains(&ControllerCapability::SupportsBackendLossWeights)
356        }
357        FitInfluencePolicy::StrictWeightSupport => {
358            capabilities.contains(&ControllerCapability::SupportsSampleWeights)
359                || capabilities.contains(&ControllerCapability::SupportsRowResampling)
360                || capabilities.contains(&ControllerCapability::SupportsBackendLossWeights)
361        }
362    }
363}
364
365#[derive(Clone, Debug, Default, Eq, PartialEq, Serialize, Deserialize)]
366pub struct ControllerRegistry {
367    manifests: BTreeMap<ControllerId, ControllerManifest>,
368}
369
370impl ControllerRegistry {
371    pub fn new() -> Self {
372        Self::default()
373    }
374
375    pub fn register(&mut self, manifest: ControllerManifest) -> Result<()> {
376        manifest.validate()?;
377        if self.manifests.contains_key(&manifest.controller_id) {
378            return Err(DagMlError::ControllerValidation(format!(
379                "duplicate controller id `{}`",
380                manifest.controller_id
381            )));
382        }
383        self.manifests
384            .insert(manifest.controller_id.clone(), manifest);
385        Ok(())
386    }
387
388    pub fn get(&self, controller_id: &ControllerId) -> Option<&ControllerManifest> {
389        self.manifests.get(controller_id)
390    }
391
392    pub fn manifests(&self) -> impl Iterator<Item = &ControllerManifest> {
393        self.manifests.values()
394    }
395
396    pub fn resolve_for_node(&self, node: &NodeSpec) -> Result<ControllerManifest> {
397        if let Some(requested) = requested_controller(node)? {
398            let manifest = self.get(&requested).ok_or_else(|| {
399                DagMlError::Planning(format!(
400                    "node `{}` requested unknown controller `{requested}`",
401                    node.id
402                ))
403            })?;
404            if manifest.operator_kind != node.kind {
405                return Err(DagMlError::Planning(format!(
406                    "node `{}` kind {:?} is incompatible with controller `{}` kind {:?}",
407                    node.id, node.kind, manifest.controller_id, manifest.operator_kind
408                )));
409            }
410            return Ok(manifest.clone());
411        }
412
413        let mut candidates = self
414            .manifests
415            .values()
416            .filter_map(|manifest| controller_candidate(manifest, node))
417            .collect::<Vec<_>>();
418        candidates.sort_by(|left, right| {
419            left.rank
420                .cmp(&right.rank)
421                .then_with(|| left.manifest.priority.cmp(&right.manifest.priority))
422                .then_with(|| {
423                    left.manifest
424                        .controller_id
425                        .cmp(&right.manifest.controller_id)
426                })
427        });
428        let Some(first) = candidates.first() else {
429            return Err(DagMlError::Planning(format!(
430                "no controller registered for node `{}` kind {:?}",
431                node.id, node.kind
432            )));
433        };
434        if candidates.get(1).is_some_and(|second| {
435            second.rank == first.rank && second.manifest.priority == first.manifest.priority
436        }) {
437            return Err(DagMlError::Planning(format!(
438                "node `{}` has ambiguous controllers for kind {:?}; set metadata.controller_id",
439                node.id, node.kind
440            )));
441        }
442        Ok(first.manifest.clone())
443    }
444
445    pub fn infer_operator_kind(&self, operator: &serde_json::Value) -> Result<Option<NodeKind>> {
446        let matches = self
447            .manifests
448            .values()
449            .filter(|manifest| {
450                !manifest.operator_selectors.is_empty()
451                    && manifest
452                        .operator_selectors
453                        .iter()
454                        .any(|selector| selector_matches_operator(selector, operator))
455            })
456            .collect::<Vec<_>>();
457        let Some(first) = matches.first() else {
458            return Ok(None);
459        };
460        let kind = first.operator_kind.clone();
461        let conflicting = matches
462            .iter()
463            .find(|manifest| manifest.operator_kind != kind);
464        if let Some(conflicting) = conflicting {
465            return Err(DagMlError::Planning(format!(
466                "minimal operator alias `{}` matches controllers with different node kinds ({:?} and {:?}); use explicit DSL syntax",
467                operator_label(operator),
468                kind,
469                conflicting.operator_kind
470            )));
471        }
472        Ok(Some(kind))
473    }
474}
475
476#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd)]
477enum ControllerMatchRank {
478    OperatorSelector,
479    GenericKind,
480}
481
482struct ControllerCandidate<'a> {
483    manifest: &'a ControllerManifest,
484    rank: ControllerMatchRank,
485}
486
487fn controller_candidate<'a>(
488    manifest: &'a ControllerManifest,
489    node: &NodeSpec,
490) -> Option<ControllerCandidate<'a>> {
491    if manifest.operator_kind != node.kind {
492        return None;
493    }
494    if manifest.operator_selectors.is_empty() {
495        return Some(ControllerCandidate {
496            manifest,
497            rank: ControllerMatchRank::GenericKind,
498        });
499    }
500    let operator = node.operator.as_ref()?;
501    manifest
502        .operator_selectors
503        .iter()
504        .any(|selector| selector_matches_operator(selector, operator))
505        .then_some(ControllerCandidate {
506            manifest,
507            rank: ControllerMatchRank::OperatorSelector,
508        })
509}
510
511fn selector_matches_operator(selector: &OperatorSelector, operator: &serde_json::Value) -> bool {
512    let descriptor = OperatorDescriptor::from_value(operator);
513    selector_matches_any(
514        &selector.aliases,
515        descriptor.alias_candidates.iter().copied(),
516    ) || descriptor
517        .class
518        .is_some_and(|class| selector_matches_exact(&selector.classes, class))
519        || descriptor.class.is_some_and(|class| {
520            selector
521                .class_prefixes
522                .iter()
523                .any(|prefix| normalized_starts_with(class, prefix))
524        })
525        || descriptor
526            .function
527            .is_some_and(|function| selector_matches_exact(&selector.functions, function))
528        || descriptor
529            .reference
530            .is_some_and(|reference| selector_matches_exact(&selector.refs, reference))
531        || descriptor
532            .operator_type
533            .is_some_and(|operator_type| selector_matches_exact(&selector.types, operator_type))
534}
535
536fn operator_label(operator: &serde_json::Value) -> String {
537    match operator {
538        serde_json::Value::String(value) => value.clone(),
539        serde_json::Value::Object(object) => ["type", "ref", "class", "function"]
540            .into_iter()
541            .find_map(|key| object.get(key).and_then(serde_json::Value::as_str))
542            .map(str::to_string)
543            .unwrap_or_else(|| operator.to_string()),
544        _ => operator.to_string(),
545    }
546}
547
548fn selector_matches_any<'a>(
549    values: &BTreeSet<String>,
550    mut candidates: impl Iterator<Item = &'a str>,
551) -> bool {
552    candidates.any(|candidate| selector_matches_exact(values, candidate))
553}
554
555fn selector_matches_exact(values: &BTreeSet<String>, candidate: &str) -> bool {
556    values
557        .iter()
558        .any(|value| normalized_eq(value.as_str(), candidate))
559}
560
561fn normalized_eq(left: &str, right: &str) -> bool {
562    left.trim().eq_ignore_ascii_case(right.trim())
563}
564
565fn normalized_starts_with(value: &str, prefix: &str) -> bool {
566    value
567        .trim()
568        .to_ascii_lowercase()
569        .starts_with(&prefix.trim().to_ascii_lowercase())
570}
571
572struct OperatorDescriptor<'a> {
573    class: Option<&'a str>,
574    function: Option<&'a str>,
575    reference: Option<&'a str>,
576    operator_type: Option<&'a str>,
577    alias_candidates: Vec<&'a str>,
578}
579
580impl<'a> OperatorDescriptor<'a> {
581    fn from_value(value: &'a serde_json::Value) -> Self {
582        let mut descriptor = Self {
583            class: None,
584            function: None,
585            reference: None,
586            operator_type: None,
587            alias_candidates: Vec::new(),
588        };
589        match value {
590            serde_json::Value::String(reference) => {
591                descriptor.reference = Some(reference);
592                descriptor.push_alias_candidates(reference);
593            }
594            serde_json::Value::Object(object) => {
595                descriptor.class = object.get("class").and_then(serde_json::Value::as_str);
596                descriptor.function = object.get("function").and_then(serde_json::Value::as_str);
597                descriptor.reference = object.get("ref").and_then(serde_json::Value::as_str);
598                descriptor.operator_type = object.get("type").and_then(serde_json::Value::as_str);
599                for value in [
600                    descriptor.operator_type,
601                    descriptor.reference,
602                    descriptor.class,
603                    descriptor.function,
604                ]
605                .into_iter()
606                .flatten()
607                {
608                    descriptor.push_alias_candidates(value);
609                }
610            }
611            _ => {}
612        }
613        descriptor
614    }
615
616    fn push_alias_candidates(&mut self, value: &'a str) {
617        self.alias_candidates.push(value);
618        if let Some(short) = value
619            .rsplit(['.', ':'])
620            .next()
621            .filter(|short| *short != value)
622        {
623            self.alias_candidates.push(short);
624        }
625    }
626}
627
628fn validate_ports(controller_id: &ControllerId, direction: &str, ports: &[PortSpec]) -> Result<()> {
629    let mut seen = BTreeSet::new();
630    for port in ports {
631        if port.name.trim().is_empty() {
632            return Err(DagMlError::ControllerValidation(format!(
633                "{direction} port on controller `{controller_id}` has an empty name"
634            )));
635        }
636        if !seen.insert(port.name.as_str()) {
637            return Err(DagMlError::ControllerValidation(format!(
638                "duplicate {direction} port `{}` on controller `{controller_id}`",
639                port.name
640            )));
641        }
642    }
643    Ok(())
644}
645
646fn requested_controller(node: &NodeSpec) -> Result<Option<ControllerId>> {
647    node.metadata
648        .get("controller_id")
649        .map(|value| {
650            value.as_str().ok_or_else(|| {
651                DagMlError::Planning(format!(
652                    "node `{}` metadata.controller_id must be a string",
653                    node.id
654                ))
655            })
656        })
657        .transpose()?
658        .map(ControllerId::new)
659        .transpose()
660}
661
662#[cfg(test)]
663mod tests {
664    use std::collections::{BTreeMap, BTreeSet};
665
666    use serde_json::json;
667
668    use super::*;
669    use crate::graph::{NodeSpec, PortCardinality, PortSchema};
670    use crate::ids::NodeId;
671
672    fn manifest(id: &str, kind: NodeKind, priority: u32) -> ControllerManifest {
673        ControllerManifest {
674            controller_id: ControllerId::new(id).unwrap(),
675            controller_version: "0.1.0".to_string(),
676            operator_kind: kind,
677            priority,
678            supported_phases: BTreeSet::from([Phase::FitCv]),
679            input_ports: Vec::new(),
680            output_ports: Vec::new(),
681            data_requirements: None,
682            capabilities: BTreeSet::from([ControllerCapability::Deterministic]),
683            operator_selectors: Vec::new(),
684            fit_scope: ControllerFitScope::FoldTrain,
685            rng_policy: RngPolicy::UsesCoreSeed,
686            artifact_policy: ArtifactPolicy::Serializable,
687        }
688    }
689
690    fn node(kind: NodeKind) -> NodeSpec {
691        NodeSpec {
692            id: NodeId::new("node:model").unwrap(),
693            kind,
694            operator: None,
695            params: BTreeMap::new(),
696            ports: PortSchema::default(),
697            metadata: BTreeMap::new(),
698            seed_label: None,
699        }
700    }
701
702    fn node_with_operator(kind: NodeKind, operator: serde_json::Value) -> NodeSpec {
703        NodeSpec {
704            operator: Some(operator),
705            ..node(kind)
706        }
707    }
708
709    fn alias_selector(alias: &str) -> OperatorSelector {
710        OperatorSelector {
711            aliases: BTreeSet::from([alias.to_string()]),
712            ..OperatorSelector::default()
713        }
714    }
715
716    #[test]
717    fn registry_resolves_lowest_priority_manifest() {
718        let mut registry = ControllerRegistry::new();
719        registry
720            .register(manifest("controller:slow", NodeKind::Model, 10))
721            .unwrap();
722        registry
723            .register(manifest("controller:fast", NodeKind::Model, 1))
724            .unwrap();
725
726        let resolved = registry.resolve_for_node(&node(NodeKind::Model)).unwrap();
727
728        assert_eq!(resolved.controller_id.as_str(), "controller:fast");
729    }
730
731    #[test]
732    fn explicit_controller_id_disambiguates() {
733        let mut registry = ControllerRegistry::new();
734        registry
735            .register(manifest("controller:a", NodeKind::Model, 1))
736            .unwrap();
737        registry
738            .register(manifest("controller:b", NodeKind::Model, 1))
739            .unwrap();
740        let mut node = node(NodeKind::Model);
741        node.metadata
742            .insert("controller_id".to_string(), json!("controller:b"));
743
744        let resolved = registry.resolve_for_node(&node).unwrap();
745
746        assert_eq!(resolved.controller_id.as_str(), "controller:b");
747    }
748
749    #[test]
750    fn equal_priority_requires_explicit_controller() {
751        let mut registry = ControllerRegistry::new();
752        registry
753            .register(manifest("controller:a", NodeKind::Model, 1))
754            .unwrap();
755        registry
756            .register(manifest("controller:b", NodeKind::Model, 1))
757            .unwrap();
758
759        assert!(registry.resolve_for_node(&node(NodeKind::Model)).is_err());
760    }
761
762    #[test]
763    fn operator_selector_prefers_specific_controller_over_generic() {
764        let mut registry = ControllerRegistry::new();
765        registry
766            .register(manifest(
767                "controller:transform.generic",
768                NodeKind::Transform,
769                0,
770            ))
771            .unwrap();
772        let mut specific = manifest("controller:transform.snv", NodeKind::Transform, 0);
773        specific.operator_selectors.push(alias_selector("SNV"));
774        registry.register(specific).unwrap();
775        let node = node_with_operator(NodeKind::Transform, json!("SNV"));
776
777        let resolved = registry.resolve_for_node(&node).unwrap();
778
779        assert_eq!(resolved.controller_id.as_str(), "controller:transform.snv");
780    }
781
782    #[test]
783    fn operator_selector_matches_plain_class_basename_alias() {
784        let mut registry = ControllerRegistry::new();
785        registry
786            .register(manifest(
787                "controller:transform.generic",
788                NodeKind::Transform,
789                0,
790            ))
791            .unwrap();
792        let mut specific = manifest("controller:transform.mixin", NodeKind::Transform, 0);
793        specific
794            .operator_selectors
795            .push(alias_selector("StandardScaler"));
796        registry.register(specific).unwrap();
797        let node = node_with_operator(
798            NodeKind::Transform,
799            json!({"class": "sklearn.preprocessing.StandardScaler"}),
800        );
801
802        let resolved = registry.resolve_for_node(&node).unwrap();
803
804        assert_eq!(
805            resolved.controller_id.as_str(),
806            "controller:transform.mixin"
807        );
808    }
809
810    #[test]
811    fn registry_infers_operator_kind_from_alias_selector() {
812        let mut registry = ControllerRegistry::new();
813        let mut model = manifest("controller:model.custom", NodeKind::Model, 0);
814        model
815            .operator_selectors
816            .push(alias_selector("ElasticSpectra"));
817        registry.register(model).unwrap();
818
819        let kind = registry
820            .infer_operator_kind(&json!("ElasticSpectra"))
821            .unwrap()
822            .unwrap();
823
824        assert_eq!(kind, NodeKind::Model);
825    }
826
827    #[test]
828    fn registry_refuses_cross_kind_alias_inference() {
829        let mut registry = ControllerRegistry::new();
830        let mut transform = manifest("controller:transform.custom", NodeKind::Transform, 0);
831        transform
832            .operator_selectors
833            .push(alias_selector("AmbiguousAlias"));
834        let mut model = manifest("controller:model.custom", NodeKind::Model, 0);
835        model
836            .operator_selectors
837            .push(alias_selector("AmbiguousAlias"));
838        registry.register(transform).unwrap();
839        registry.register(model).unwrap();
840
841        let error = registry
842            .infer_operator_kind(&json!("AmbiguousAlias"))
843            .unwrap_err()
844            .to_string();
845
846        assert!(error.contains("different node kinds"));
847    }
848
849    #[test]
850    fn operator_selector_matches_class_prefix() {
851        let mut registry = ControllerRegistry::new();
852        let mut sklearn = manifest("controller:sklearn.transform", NodeKind::Transform, 0);
853        sklearn.operator_selectors.push(OperatorSelector {
854            class_prefixes: BTreeSet::from(["sklearn.preprocessing.".to_string()]),
855            ..OperatorSelector::default()
856        });
857        registry.register(sklearn).unwrap();
858        let node = node_with_operator(
859            NodeKind::Transform,
860            json!({"class": "sklearn.preprocessing.MinMaxScaler"}),
861        );
862
863        let resolved = registry.resolve_for_node(&node).unwrap();
864
865        assert_eq!(
866            resolved.controller_id.as_str(),
867            "controller:sklearn.transform"
868        );
869    }
870
871    #[test]
872    fn equal_priority_operator_selector_matches_are_ambiguous() {
873        let mut registry = ControllerRegistry::new();
874        let mut first = manifest("controller:snv.a", NodeKind::Transform, 0);
875        first.operator_selectors.push(alias_selector("SNV"));
876        let mut second = manifest("controller:snv.b", NodeKind::Transform, 0);
877        second.operator_selectors.push(alias_selector("SNV"));
878        registry.register(first).unwrap();
879        registry.register(second).unwrap();
880        let node = node_with_operator(NodeKind::Transform, json!({"type": "SNV"}));
881
882        let error = registry.resolve_for_node(&node).unwrap_err().to_string();
883
884        assert!(error.contains("ambiguous controllers"));
885    }
886
887    #[test]
888    fn selector_only_controller_does_not_catch_unmatched_operator() {
889        let mut registry = ControllerRegistry::new();
890        let mut snv = manifest("controller:transform.snv", NodeKind::Transform, 0);
891        snv.operator_selectors.push(alias_selector("SNV"));
892        registry.register(snv).unwrap();
893        let node = node_with_operator(NodeKind::Transform, json!("MSC"));
894
895        let error = registry.resolve_for_node(&node).unwrap_err().to_string();
896
897        assert!(error.contains("no controller registered"));
898    }
899
900    #[test]
901    fn manifest_rejects_prediction_output_without_capability() {
902        let mut manifest = manifest("controller:predictor", NodeKind::Model, 0);
903        manifest.output_ports.push(PortSpec {
904            name: "pred".to_string(),
905            kind: PortKind::Prediction,
906            representation: None,
907            cardinality: PortCardinality::One,
908            unit_level: None,
909            alignment_key: None,
910            target_level: None,
911            description: String::new(),
912        });
913
914        let error = manifest.validate().unwrap_err().to_string();
915
916        assert!(error.contains("lacks emits_predictions"));
917    }
918
919    #[test]
920    fn manifest_rejects_training_phases_for_inference_only_controller() {
921        let mut manifest = manifest("controller:predict-only", NodeKind::Model, 0);
922        manifest.fit_scope = ControllerFitScope::InferenceOnly;
923
924        let error = manifest.validate().unwrap_err().to_string();
925
926        assert!(error.contains("inference_only"));
927    }
928
929    #[test]
930    fn manifest_requires_configurable_loss_for_specialized_loss_capabilities() {
931        for capability in [
932            ControllerCapability::SupportsCustomLoss,
933            ControllerCapability::SupportsDifferentiableLoss,
934        ] {
935            let mut manifest = manifest("controller:loss", NodeKind::Model, 0);
936            manifest.capabilities.insert(capability);
937            assert!(manifest
938                .validate()
939                .unwrap_err()
940                .to_string()
941                .contains("without configurable loss"));
942
943            manifest
944                .capabilities
945                .insert(ControllerCapability::SupportsConfigurableLoss);
946            manifest.validate().unwrap();
947        }
948    }
949
950    #[test]
951    fn manifest_validates_model_input_spec_data_requirements() {
952        let mut manifest = manifest("controller:data-aware", NodeKind::Model, 0);
953        manifest.data_requirements = Some(json!({
954            "schema_version": 1,
955            "ports": [{
956                "name": "x",
957                "accepted_representations": ["tabular_numeric"],
958                "accepted_types": ["f64"],
959                "rank": 2
960            }]
961        }));
962
963        let input_spec = manifest.model_input_spec().unwrap().unwrap();
964        assert_eq!(input_spec.ports[0].name, "x");
965        manifest.validate().unwrap();
966    }
967
968    #[test]
969    fn manifest_rejects_invalid_model_input_spec_data_requirements() {
970        let mut manifest = manifest("controller:data-aware", NodeKind::Model, 0);
971        manifest.data_requirements = Some(json!({
972            "schema_version": 1,
973            "ports": [{
974                "name": "x",
975                "accepted_representations": [],
976                "accepted_types": ["f64"]
977            }]
978        }));
979
980        let error = manifest.validate().unwrap_err().to_string();
981
982        assert!(error.contains("data_requirements"));
983        assert!(error.contains("accepted_representations"));
984    }
985
986    #[test]
987    fn manifest_rejects_empty_operator_selector() {
988        let mut manifest = manifest("controller:empty-selector", NodeKind::Transform, 0);
989        manifest
990            .operator_selectors
991            .push(OperatorSelector::default());
992
993        let error = manifest.validate().unwrap_err().to_string();
994
995        assert!(error.contains("empty operator selector"));
996    }
997
998    #[test]
999    fn manifest_reports_parallel_invocation_support() {
1000        let mut manifest = manifest("controller:parallel", NodeKind::Model, 0);
1001        assert!(!manifest.supports_parallel_invocation());
1002        manifest
1003            .capabilities
1004            .insert(ControllerCapability::ProcessSafe);
1005        assert!(manifest.supports_parallel_invocation());
1006    }
1007
1008    #[test]
1009    fn published_controller_manifest_schema_declares_current_contract() {
1010        let schema: serde_json::Value = serde_json::from_str(include_str!(
1011            "../../../docs/contracts/controller_manifest.schema.json"
1012        ))
1013        .unwrap();
1014
1015        assert_eq!(schema["$id"], CONTROLLER_MANIFEST_SCHEMA_ID);
1016        assert!(schema["required"]
1017            .as_array()
1018            .unwrap()
1019            .iter()
1020            .any(|field| field.as_str() == Some("controller_id")));
1021        assert!(schema["$defs"]["controller_capability"]["enum"]
1022            .as_array()
1023            .unwrap()
1024            .iter()
1025            .any(|capability| capability.as_str() == Some("emits_predictions")));
1026        assert!(schema["$defs"]["controller_capability"]["enum"]
1027            .as_array()
1028            .unwrap()
1029            .iter()
1030            .any(|capability| capability.as_str() == Some("aggregates_predictions")));
1031        assert!(schema["$defs"]["controller_capability"]["enum"]
1032            .as_array()
1033            .unwrap()
1034            .iter()
1035            .any(|capability| capability.as_str() == Some("supports_custom_loss")));
1036        assert!(schema["properties"]
1037            .as_object()
1038            .unwrap()
1039            .contains_key("operator_selectors"));
1040        assert_eq!(
1041            schema["$defs"]["model_input_spec"]["properties"]["schema_version"]["const"].as_u64(),
1042            Some(crate::data::MODEL_INPUT_SPEC_SCHEMA_VERSION as u64)
1043        );
1044    }
1045}