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