Skip to main content

dag_ml_data_core/
adapter.rs

1use std::cmp::Ordering;
2use std::collections::{BTreeMap, BTreeSet, BinaryHeap};
3
4use serde::{Deserialize, Serialize};
5
6use crate::error::{DataError, Result};
7use crate::ids::{RepresentationId, TypeId};
8use crate::plan::{FitScope, PlanIssue};
9
10#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
11pub struct InputPortSpec {
12    pub name: String,
13    pub accepted_representations: Vec<RepresentationId>,
14    pub accepted_types: Vec<TypeId>,
15    pub rank: Option<usize>,
16    #[serde(default)]
17    pub multi_source: bool,
18    #[serde(default)]
19    pub optional: bool,
20}
21
22impl InputPortSpec {
23    pub fn validate(&self) -> Result<()> {
24        validate_name("input port", &self.name)?;
25        if self.accepted_representations.is_empty() {
26            return Err(DataError::Validation(format!(
27                "input port `{}` accepts no representations",
28                self.name
29            )));
30        }
31        if self.accepted_types.is_empty() {
32            return Err(DataError::Validation(format!(
33                "input port `{}` accepts no types",
34                self.name
35            )));
36        }
37        Ok(())
38    }
39}
40
41#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
42pub struct ModelInputSpec {
43    pub ports: Vec<InputPortSpec>,
44    #[serde(default)]
45    pub default_fusion: Option<serde_json::Value>,
46}
47
48impl ModelInputSpec {
49    pub fn validate(&self) -> Result<()> {
50        if self.ports.is_empty() {
51            return Err(DataError::Validation(
52                "model input spec contains no ports".to_string(),
53            ));
54        }
55        let mut names = BTreeSet::new();
56        for port in &self.ports {
57            port.validate()?;
58            if !names.insert(port.name.as_str()) {
59                return Err(DataError::Validation(format!(
60                    "duplicate model input port `{}`",
61                    port.name
62                )));
63            }
64        }
65        Ok(())
66    }
67}
68
69#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
70pub struct AdapterSpec {
71    pub id: String,
72    pub version: String,
73    pub input_type: TypeId,
74    pub input_representation: RepresentationId,
75    pub output_type: TypeId,
76    pub output_representation: RepresentationId,
77    pub cost: u64,
78    #[serde(default)]
79    pub lossy: bool,
80    #[serde(default)]
81    pub supervised: bool,
82    #[serde(default)]
83    pub stateful: bool,
84    #[serde(default = "default_true")]
85    pub deterministic: bool,
86    pub fit_scope: FitScope,
87    #[serde(default)]
88    pub params: BTreeMap<String, serde_json::Value>,
89}
90
91fn default_true() -> bool {
92    true
93}
94
95impl AdapterSpec {
96    pub fn validate(&self) -> Result<()> {
97        validate_name("adapter", &self.id)?;
98        validate_name("adapter version", &self.version)?;
99        if !self.deterministic {
100            return Err(DataError::Validation(format!(
101                "adapter `{}` is not deterministic",
102                self.id
103            )));
104        }
105        if self.stateful && self.fit_scope == FitScope::Stateless {
106            return Err(DataError::Validation(format!(
107                "stateful adapter `{}` cannot use stateless fit scope",
108                self.id
109            )));
110        }
111        Ok(())
112    }
113
114    fn source(&self) -> RepresentationNode {
115        RepresentationNode {
116            type_id: self.input_type.clone(),
117            representation_id: self.input_representation.clone(),
118        }
119    }
120
121    fn target(&self) -> RepresentationNode {
122        RepresentationNode {
123            type_id: self.output_type.clone(),
124            representation_id: self.output_representation.clone(),
125        }
126    }
127}
128
129#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
130pub struct PlanningPolicy {
131    #[serde(default)]
132    pub allow_lossy: bool,
133    #[serde(default)]
134    pub allow_stateful: bool,
135    #[serde(default)]
136    pub allow_supervised: bool,
137    #[serde(default)]
138    pub forbidden_adapters: BTreeSet<String>,
139    #[serde(default)]
140    pub preferred_adapters: BTreeSet<String>,
141    #[serde(default = "default_true")]
142    pub require_user_choice_on_ambiguity: bool,
143    pub max_hops: Option<usize>,
144}
145
146impl Default for PlanningPolicy {
147    fn default() -> Self {
148        Self {
149            allow_lossy: false,
150            allow_stateful: false,
151            allow_supervised: false,
152            forbidden_adapters: BTreeSet::new(),
153            preferred_adapters: BTreeSet::new(),
154            require_user_choice_on_ambiguity: true,
155            max_hops: None,
156        }
157    }
158}
159
160#[derive(Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash, Serialize, Deserialize)]
161pub struct RepresentationNode {
162    pub type_id: TypeId,
163    pub representation_id: RepresentationId,
164}
165
166#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
167pub struct AdapterPath {
168    pub adapters: Vec<AdapterSpec>,
169    pub total_cost: u64,
170    pub effective_score: u64,
171}
172
173#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
174pub struct PathResolution {
175    pub path: Option<AdapterPath>,
176    #[serde(default)]
177    pub requires_user_choice: bool,
178    #[serde(default)]
179    pub issues: Vec<PlanIssue>,
180}
181
182impl PathResolution {
183    pub fn resolved(path: AdapterPath) -> Self {
184        Self {
185            path: Some(path),
186            requires_user_choice: false,
187            issues: Vec::new(),
188        }
189    }
190
191    pub fn unresolved(code: &str, message: String, choices: Vec<String>) -> Self {
192        Self {
193            path: None,
194            requires_user_choice: !choices.is_empty(),
195            issues: vec![PlanIssue {
196                code: code.to_string(),
197                message,
198                choices,
199            }],
200        }
201    }
202}
203
204#[derive(Clone, Debug, Default)]
205pub struct AdapterRegistry {
206    adapters: BTreeMap<String, AdapterSpec>,
207}
208
209#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
210pub struct AdapterRegistrySpec {
211    #[serde(default)]
212    pub adapters: Vec<AdapterSpec>,
213}
214
215impl AdapterRegistry {
216    pub fn new() -> Self {
217        Self::default()
218    }
219
220    pub fn from_spec(spec: AdapterRegistrySpec) -> Result<Self> {
221        let mut registry = Self::new();
222        for adapter in spec.adapters {
223            registry.register_adapter(adapter)?;
224        }
225        Ok(registry)
226    }
227
228    pub fn register_adapter(&mut self, adapter: AdapterSpec) -> Result<()> {
229        adapter.validate()?;
230        if self.adapters.contains_key(&adapter.id) {
231            return Err(DataError::Validation(format!(
232                "duplicate adapter id `{}`",
233                adapter.id
234            )));
235        }
236        self.adapters.insert(adapter.id.clone(), adapter);
237        Ok(())
238    }
239
240    pub fn adapters(&self) -> impl Iterator<Item = &AdapterSpec> {
241        self.adapters.values()
242    }
243
244    pub fn find_path(
245        &self,
246        source_type: &TypeId,
247        source_representation: &RepresentationId,
248        target_type: &TypeId,
249        target_representation: &RepresentationId,
250        policy: &PlanningPolicy,
251    ) -> PathResolution {
252        let start = RepresentationNode {
253            type_id: source_type.clone(),
254            representation_id: source_representation.clone(),
255        };
256        let goal = RepresentationNode {
257            type_id: target_type.clone(),
258            representation_id: target_representation.clone(),
259        };
260        if start == goal {
261            return PathResolution::resolved(AdapterPath {
262                adapters: Vec::new(),
263                total_cost: 0,
264                effective_score: 0,
265            });
266        }
267
268        let mut edges: BTreeMap<RepresentationNode, Vec<&AdapterSpec>> = BTreeMap::new();
269        for adapter in self.adapters.values() {
270            if policy.forbidden_adapters.contains(&adapter.id) {
271                continue;
272            }
273            if adapter.lossy && !policy.allow_lossy {
274                continue;
275            }
276            if adapter.stateful && !policy.allow_stateful {
277                continue;
278            }
279            if adapter.supervised && !policy.allow_supervised {
280                continue;
281            }
282            edges.entry(adapter.source()).or_default().push(adapter);
283        }
284
285        let mut heap = BinaryHeap::new();
286        heap.push(SearchState {
287            score: 0,
288            raw_cost: 0,
289            hops: 0,
290            node: start.clone(),
291            adapter_ids: Vec::new(),
292        });
293
294        // Cost dominance is valid only at the same hop budget. A cheaper
295        // arrival with more hops may have no remaining budget to reach the goal.
296        let mut best_seen: BTreeMap<(RepresentationNode, usize), u64> = BTreeMap::new();
297        best_seen.insert((start.clone(), 0), 0);
298        let mut cost_overflow = false;
299        let mut best_goal: Option<(u64, usize, u64)> = None;
300        let mut goal_paths = Vec::new();
301
302        while let Some(state) = heap.pop() {
303            if let Some((best_score, best_hops, _)) = best_goal {
304                if (state.score, state.hops) > (best_score, best_hops) {
305                    break;
306                }
307            }
308            if state.node == goal {
309                best_goal.get_or_insert((state.score, state.hops, state.raw_cost));
310                goal_paths.push(state.adapter_ids);
311                continue;
312            }
313            if policy
314                .max_hops
315                .is_some_and(|max_hops| state.hops >= max_hops)
316            {
317                continue;
318            }
319            let Some(next_edges) = edges.get(&state.node) else {
320                continue;
321            };
322            for adapter in next_edges {
323                if state.adapter_ids.iter().any(|id| id == &adapter.id) {
324                    continue;
325                }
326                let next = adapter.target();
327                let hops = state.hops + 1;
328                let Some(score) =
329                    adapter_score(adapter, policy).and_then(|score| state.score.checked_add(score))
330                else {
331                    cost_overflow = true;
332                    continue;
333                };
334                let Some(raw_cost) = state.raw_cost.checked_add(adapter.cost) else {
335                    cost_overflow = true;
336                    continue;
337                };
338                let key = (next.clone(), hops);
339                if best_seen
340                    .get(&key)
341                    .is_some_and(|best_score| score > *best_score)
342                {
343                    continue;
344                }
345                best_seen.insert(key, score);
346                let mut adapter_ids = state.adapter_ids.clone();
347                adapter_ids.push(adapter.id.clone());
348                heap.push(SearchState {
349                    score,
350                    raw_cost,
351                    hops,
352                    node: next,
353                    adapter_ids,
354                });
355            }
356        }
357
358        if goal_paths.is_empty() {
359            return PathResolution::unresolved(
360                if cost_overflow {
361                    "cost_overflow"
362                } else {
363                    "no_path"
364                },
365                format!(
366                    "no adapter path from `{}/{}` to `{}/{}`",
367                    source_type, source_representation, target_type, target_representation
368                ),
369                Vec::new(),
370            );
371        }
372
373        goal_paths.sort();
374        goal_paths.dedup();
375        if goal_paths.len() > 1 && policy.require_user_choice_on_ambiguity {
376            let choices = goal_paths
377                .iter()
378                .map(|path| path.join(" -> "))
379                .collect::<Vec<_>>();
380            return PathResolution::unresolved(
381                "ambiguous_path",
382                "multiple equivalent adapter paths require user choice".to_string(),
383                choices,
384            );
385        }
386
387        let adapter_ids = goal_paths.remove(0);
388        let adapters = adapter_ids
389            .iter()
390            .map(|id| self.adapters.get(id).expect("path adapter exists").clone())
391            .collect::<Vec<_>>();
392        // Every queued state was checked; use the checked totals of the
393        // selected path as well, without repeating unchecked arithmetic.
394        let total_cost = adapters
395            .iter()
396            .try_fold(0u64, |sum, adapter| sum.checked_add(adapter.cost));
397        let effective_score = adapters.iter().try_fold(0u64, |sum, adapter| {
398            sum.checked_add(adapter_score(adapter, policy)?)
399        });
400        let (Some(total_cost), Some(effective_score)) = (total_cost, effective_score) else {
401            return PathResolution::unresolved(
402                "cost_overflow",
403                "adapter path cost exceeds u64".into(),
404                Vec::new(),
405            );
406        };
407        PathResolution::resolved(AdapterPath {
408            adapters,
409            total_cost,
410            effective_score,
411        })
412    }
413}
414
415#[derive(Clone, Debug, Eq, PartialEq)]
416struct SearchState {
417    score: u64,
418    raw_cost: u64,
419    hops: usize,
420    node: RepresentationNode,
421    adapter_ids: Vec<String>,
422}
423
424impl Ord for SearchState {
425    fn cmp(&self, other: &Self) -> Ordering {
426        other
427            .score
428            .cmp(&self.score)
429            .then_with(|| other.hops.cmp(&self.hops))
430            .then_with(|| other.raw_cost.cmp(&self.raw_cost))
431            .then_with(|| other.node.cmp(&self.node))
432            .then_with(|| other.adapter_ids.cmp(&self.adapter_ids))
433    }
434}
435
436impl PartialOrd for SearchState {
437    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
438        Some(self.cmp(other))
439    }
440}
441
442fn adapter_score(adapter: &AdapterSpec, policy: &PlanningPolicy) -> Option<u64> {
443    let mut score = u128::from(adapter.cost.max(1));
444    if adapter.lossy {
445        score += 1_000_000;
446    }
447    if adapter.stateful {
448        score += 100_000;
449    }
450    if adapter.supervised {
451        score += 100_000;
452    }
453    if policy.preferred_adapters.contains(&adapter.id) {
454        score = score.saturating_sub(1);
455    }
456    u64::try_from(score).ok()
457}
458
459fn validate_name(kind: &str, value: &str) -> Result<()> {
460    if value.trim().is_empty() {
461        return Err(DataError::Validation(format!("{kind} name is empty")));
462    }
463    if !value
464        .bytes()
465        .all(|b| b.is_ascii_alphanumeric() || matches!(b, b'_' | b'-' | b'.' | b':' | b'/'))
466    {
467        return Err(DataError::Validation(format!(
468            "{kind} `{value}` contains unsupported characters"
469        )));
470    }
471    Ok(())
472}
473
474#[cfg(test)]
475mod tests {
476    use super::*;
477
478    fn tid(value: &str) -> TypeId {
479        TypeId::new(value).unwrap()
480    }
481
482    fn rid(value: &str) -> RepresentationId {
483        RepresentationId::new(value).unwrap()
484    }
485
486    fn adapter(id: &str, input: &str, output: &str, cost: u64) -> AdapterSpec {
487        AdapterSpec {
488            id: id.to_string(),
489            version: "0.1.0".to_string(),
490            input_type: tid("dense_signal"),
491            input_representation: rid(input),
492            output_type: if output == "tabular_numeric" {
493                tid("table")
494            } else {
495                tid("dense_signal")
496            },
497            output_representation: rid(output),
498            cost,
499            lossy: false,
500            supervised: false,
501            stateful: false,
502            deterministic: true,
503            fit_scope: FitScope::Stateless,
504            params: BTreeMap::new(),
505        }
506    }
507
508    #[test]
509    fn validates_model_input_ports() {
510        let spec = ModelInputSpec {
511            ports: vec![InputPortSpec {
512                name: "X".to_string(),
513                accepted_representations: vec![rid("tabular_numeric")],
514                accepted_types: vec![tid("table")],
515                rank: Some(2),
516                multi_source: true,
517                optional: false,
518            }],
519            default_fusion: None,
520        };
521
522        assert!(spec.validate().is_ok());
523    }
524
525    #[test]
526    fn rejects_duplicate_adapter_ids() {
527        let mut registry = AdapterRegistry::new();
528        registry
529            .register_adapter(adapter(
530                "spectra.flatten",
531                "signal_1d",
532                "tabular_numeric",
533                1,
534            ))
535            .unwrap();
536
537        assert!(registry
538            .register_adapter(adapter(
539                "spectra.flatten",
540                "signal_1d",
541                "tabular_numeric",
542                1
543            ))
544            .is_err());
545    }
546
547    #[test]
548    fn path_selection_is_registration_order_independent() {
549        let mut left = AdapterRegistry::new();
550        left.register_adapter(adapter("a.to_mid", "signal_1d", "signal_mid", 1))
551            .unwrap();
552        left.register_adapter(adapter("b.to_tabular", "signal_mid", "tabular_numeric", 1))
553            .unwrap();
554        left.register_adapter(adapter("c.direct", "signal_1d", "tabular_numeric", 10))
555            .unwrap();
556
557        let mut right = AdapterRegistry::new();
558        right
559            .register_adapter(adapter("c.direct", "signal_1d", "tabular_numeric", 10))
560            .unwrap();
561        right
562            .register_adapter(adapter("b.to_tabular", "signal_mid", "tabular_numeric", 1))
563            .unwrap();
564        right
565            .register_adapter(adapter("a.to_mid", "signal_1d", "signal_mid", 1))
566            .unwrap();
567
568        let policy = PlanningPolicy::default();
569        let left_path = left
570            .find_path(
571                &tid("dense_signal"),
572                &rid("signal_1d"),
573                &tid("table"),
574                &rid("tabular_numeric"),
575                &policy,
576            )
577            .path
578            .unwrap();
579        let right_path = right
580            .find_path(
581                &tid("dense_signal"),
582                &rid("signal_1d"),
583                &tid("table"),
584                &rid("tabular_numeric"),
585                &policy,
586            )
587            .path
588            .unwrap();
589
590        assert_eq!(
591            left_path
592                .adapters
593                .iter()
594                .map(|adapter| adapter.id.as_str())
595                .collect::<Vec<_>>(),
596            vec!["a.to_mid", "b.to_tabular"]
597        );
598        assert_eq!(left_path, right_path);
599    }
600
601    #[test]
602    fn lossy_paths_are_refused_unless_allowed() {
603        let mut lossy = adapter("image.embedding", "signal_1d", "tabular_numeric", 1);
604        lossy.lossy = true;
605
606        let mut registry = AdapterRegistry::new();
607        registry.register_adapter(lossy).unwrap();
608
609        let refused = registry.find_path(
610            &tid("dense_signal"),
611            &rid("signal_1d"),
612            &tid("table"),
613            &rid("tabular_numeric"),
614            &PlanningPolicy::default(),
615        );
616        assert!(refused.path.is_none());
617
618        let allowed = registry.find_path(
619            &tid("dense_signal"),
620            &rid("signal_1d"),
621            &tid("table"),
622            &rid("tabular_numeric"),
623            &PlanningPolicy {
624                allow_lossy: true,
625                ..PlanningPolicy::default()
626            },
627        );
628        assert_eq!(allowed.path.unwrap().adapters[0].id, "image.embedding");
629    }
630
631    #[test]
632    fn equivalent_best_paths_require_user_choice() {
633        let mut registry = AdapterRegistry::new();
634        registry
635            .register_adapter(adapter("a.flatten", "signal_1d", "tabular_numeric", 1))
636            .unwrap();
637        registry
638            .register_adapter(adapter("b.flatten", "signal_1d", "tabular_numeric", 1))
639            .unwrap();
640
641        let resolution = registry.find_path(
642            &tid("dense_signal"),
643            &rid("signal_1d"),
644            &tid("table"),
645            &rid("tabular_numeric"),
646            &PlanningPolicy::default(),
647        );
648
649        assert!(resolution.path.is_none());
650        assert!(resolution.requires_user_choice);
651        assert_eq!(resolution.issues[0].code, "ambiguous_path");
652        assert_eq!(resolution.issues[0].choices.len(), 2);
653    }
654}