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