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 UsesTrainingWeights,
44 UsesEarlyStopping,
46 PerformsInternalTuning,
48 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}