onnx_runtime_eager/
domain.rs1use std::collections::HashMap;
14
15use crate::opset::{LATEST_ONNX_OPSET, resolve_opset};
16
17#[derive(Clone, Debug)]
19pub struct DomainInfo {
20 pub name: String,
22 pub default_opset: u64,
24}
25
26#[derive(Debug)]
29pub struct DomainRegistry {
30 domains: HashMap<String, DomainInfo>,
31}
32
33impl DomainRegistry {
34 pub fn new() -> Self {
38 let mut reg = Self {
39 domains: HashMap::new(),
40 };
41 reg.register("", LATEST_ONNX_OPSET);
42 reg.register("ai.onnx.ml", 3);
43 reg.register("com.microsoft", 1);
44 reg
45 }
46
47 pub fn register(&mut self, domain: &str, default_opset: u64) {
50 self.domains.insert(
51 domain.to_string(),
52 DomainInfo {
53 name: domain.to_string(),
54 default_opset,
55 },
56 );
57 }
58
59 pub fn default_opset(&self, domain: &str) -> u64 {
62 self.domains
63 .get(domain)
64 .map(|d| d.default_opset)
65 .unwrap_or(LATEST_ONNX_OPSET)
66 }
67
68 pub fn resolve_opset(&self, domain: &str, explicit: Option<u64>) -> u64 {
71 resolve_opset(self.default_opset(domain), explicit)
72 }
73
74 pub fn contains(&self, domain: &str) -> bool {
76 self.domains.contains_key(domain)
77 }
78
79 pub fn domains(&self) -> Vec<(String, u64)> {
82 let mut out: Vec<(String, u64)> = self
83 .domains
84 .values()
85 .map(|d| (d.name.clone(), d.default_opset))
86 .collect();
87 out.sort_by(|a, b| a.0.cmp(&b.0));
88 out
89 }
90}
91
92impl Default for DomainRegistry {
93 fn default() -> Self {
94 Self::new()
95 }
96}
97
98#[cfg(test)]
99mod tests {
100 use super::*;
101
102 #[test]
103 fn defaults_are_registered() {
104 let reg = DomainRegistry::new();
105 assert_eq!(reg.default_opset(""), LATEST_ONNX_OPSET);
106 assert_eq!(reg.default_opset("ai.onnx.ml"), 3);
107 assert_eq!(reg.default_opset("com.microsoft"), 1);
108 }
109
110 #[test]
111 fn unregistered_domain_defaults_to_latest() {
112 let reg = DomainRegistry::new();
113 assert!(!reg.contains("com.acme"));
114 assert_eq!(reg.default_opset("com.acme"), LATEST_ONNX_OPSET);
115 }
116
117 #[test]
118 fn explicit_opset_overrides_domain_default() {
119 let reg = DomainRegistry::new();
120 assert_eq!(reg.resolve_opset("com.microsoft", Some(2)), 2);
121 assert_eq!(reg.resolve_opset("com.microsoft", None), 1);
122 }
123
124 #[test]
125 fn custom_domain_registration() {
126 let mut reg = DomainRegistry::new();
127 reg.register("com.acme", 5);
128 assert_eq!(reg.default_opset("com.acme"), 5);
129 assert!(reg.domains().contains(&("com.acme".to_string(), 5)));
130 }
131}