use std::collections::HashMap;
use crate::opset::{resolve_opset, LATEST_ONNX_OPSET};
#[derive(Clone, Debug)]
pub struct DomainInfo {
pub name: String,
pub default_opset: u64,
}
#[derive(Debug)]
pub struct DomainRegistry {
domains: HashMap<String, DomainInfo>,
}
impl DomainRegistry {
pub fn new() -> Self {
let mut reg = Self {
domains: HashMap::new(),
};
reg.register("", LATEST_ONNX_OPSET);
reg.register("ai.onnx.ml", 3);
reg.register("com.microsoft", 1);
reg
}
pub fn register(&mut self, domain: &str, default_opset: u64) {
self.domains.insert(
domain.to_string(),
DomainInfo {
name: domain.to_string(),
default_opset,
},
);
}
pub fn default_opset(&self, domain: &str) -> u64 {
self.domains
.get(domain)
.map(|d| d.default_opset)
.unwrap_or(LATEST_ONNX_OPSET)
}
pub fn resolve_opset(&self, domain: &str, explicit: Option<u64>) -> u64 {
resolve_opset(self.default_opset(domain), explicit)
}
pub fn contains(&self, domain: &str) -> bool {
self.domains.contains_key(domain)
}
pub fn domains(&self) -> Vec<(String, u64)> {
let mut out: Vec<(String, u64)> = self
.domains
.values()
.map(|d| (d.name.clone(), d.default_opset))
.collect();
out.sort_by(|a, b| a.0.cmp(&b.0));
out
}
}
impl Default for DomainRegistry {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn defaults_are_registered() {
let reg = DomainRegistry::new();
assert_eq!(reg.default_opset(""), LATEST_ONNX_OPSET);
assert_eq!(reg.default_opset("ai.onnx.ml"), 3);
assert_eq!(reg.default_opset("com.microsoft"), 1);
}
#[test]
fn unregistered_domain_defaults_to_latest() {
let reg = DomainRegistry::new();
assert!(!reg.contains("com.acme"));
assert_eq!(reg.default_opset("com.acme"), LATEST_ONNX_OPSET);
}
#[test]
fn explicit_opset_overrides_domain_default() {
let reg = DomainRegistry::new();
assert_eq!(reg.resolve_opset("com.microsoft", Some(2)), 2);
assert_eq!(reg.resolve_opset("com.microsoft", None), 1);
}
#[test]
fn custom_domain_registration() {
let mut reg = DomainRegistry::new();
reg.register("com.acme", 5);
assert_eq!(reg.default_opset("com.acme"), 5);
assert!(reg.domains().contains(&("com.acme".to_string(), 5)));
}
}