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 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}