Skip to main content

onnx_runtime_ep_api/
registry.rs

1//! Op → kernel-factory registry and the EP registry (§4.3, §4.6).
2
3use std::collections::HashMap;
4use std::path::Path;
5
6use onnx_runtime_ir::{DataType, Node, Shape, TensorLayout};
7
8use crate::error::Result;
9use crate::kernel::{Kernel, KernelMatch};
10use crate::provider::{EpConfig, EpId, ExecutionProvider};
11
12/// Registry key: an operator identity plus the opset version it was introduced.
13#[derive(Clone, PartialEq, Eq, Hash, Debug)]
14pub struct OpKey {
15    pub op_type: String,
16    pub domain: String,
17    pub since_version: u64,
18}
19
20impl OpKey {
21    pub fn new(op_type: impl Into<String>, domain: impl Into<String>, since_version: u64) -> Self {
22        Self {
23            op_type: op_type.into(),
24            domain: domain.into(),
25            since_version,
26        }
27    }
28}
29
30/// Normalise the default ONNX domain: the empty string and `"ai.onnx"` name the
31/// same (standard) domain. Contrib domains (e.g. `"com.microsoft"`) are left
32/// untouched. Keeps dispatch keyed on `(op_type, domain)` model-agnostically.
33fn norm_domain(domain: &str) -> &str {
34    if domain == "ai.onnx" { "" } else { domain }
35}
36
37/// Creates kernels for a specific op.
38pub trait KernelFactory: Send + Sync {
39    fn create(&self, node: &Node, input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>>;
40}
41
42/// Maps `(op_type, domain, opset)` → kernel factory (§4.3).
43#[derive(Default)]
44pub struct OpRegistry {
45    entries: HashMap<OpKey, Box<dyn KernelFactory>>,
46    /// Normalized domain → op type → sorted registered `since_version`s.
47    by_op: HashMap<String, HashMap<String, Vec<u64>>>,
48}
49
50impl OpRegistry {
51    pub fn new() -> Self {
52        Self::default()
53    }
54
55    /// Register a factory under `key`.
56    pub fn register(&mut self, mut key: OpKey, factory: Box<dyn KernelFactory>) {
57        key.domain = norm_domain(&key.domain).to_owned();
58        let versions = self
59            .by_op
60            .entry(key.domain.clone())
61            .or_default()
62            .entry(key.op_type.clone())
63            .or_default();
64        if let Err(index) = versions.binary_search(&key.since_version) {
65            versions.insert(index, key.since_version);
66        }
67        self.entries.insert(key, factory);
68    }
69
70    /// Look up the best matching factory: the highest `since_version` that is
71    /// `<= opset` for the given `(op_type, domain)`.
72    pub fn lookup(&self, op_type: &str, domain: &str, opset: u64) -> Option<&dyn KernelFactory> {
73        let domain = norm_domain(domain);
74        let versions = self.by_op.get(domain)?.get(op_type)?;
75        let index = versions.partition_point(|&version| version <= opset);
76        let since_version = *versions.get(index.checked_sub(1)?)?;
77        self.entries
78            .get(&OpKey::new(op_type, domain, since_version))
79            .map(Box::as_ref)
80    }
81
82    /// Whether a factory is registered for `(op_type, domain)` at or before
83    /// `opset`.
84    pub fn supports(&self, op_type: &str, domain: &str, opset: u64) -> bool {
85        let domain = norm_domain(domain);
86        self.by_op
87            .get(domain)
88            .and_then(|ops| ops.get(op_type))
89            .and_then(|versions| versions.first())
90            .is_some_and(|&since_version| since_version <= opset)
91    }
92
93    /// Earliest registered opset for `(op_type, domain)`, if the EP knows the
94    /// operator at any version. Used only to make decline diagnostics actionable.
95    pub fn earliest_since_version(&self, op_type: &str, domain: &str) -> Option<u64> {
96        let domain = norm_domain(domain);
97        self.by_op.get(domain)?.get(op_type)?.first().copied()
98    }
99
100    /// Number of registered entries.
101    pub fn len(&self) -> usize {
102        self.entries.len()
103    }
104
105    /// Whether the registry is empty.
106    pub fn is_empty(&self) -> bool {
107        self.entries.is_empty()
108    }
109}
110
111#[cfg(test)]
112mod tests {
113    use super::*;
114
115    struct DummyFactory(u64);
116
117    impl KernelFactory for DummyFactory {
118        fn create(&self, _node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
119            let _ = self.0;
120            unreachable!("registry tests do not create kernels")
121        }
122    }
123
124    #[test]
125    fn indexed_queries_match_linear_reference() {
126        let mut registry = OpRegistry::new();
127        let mut state = 0x9e37_79b9_u64;
128        let ops = ["Add", "Mul", "Gemm", "Attention"];
129        let domains = ["", "ai.onnx", "com.microsoft", "pkg.nxrt"];
130
131        for factory_id in 0..256 {
132            state = state
133                .wrapping_mul(6_364_136_223_846_793_005)
134                .wrapping_add(1);
135            let op_type = ops[(state as usize) % ops.len()];
136            state = state
137                .wrapping_mul(6_364_136_223_846_793_005)
138                .wrapping_add(1);
139            let domain = domains[(state as usize) % domains.len()];
140            state = state
141                .wrapping_mul(6_364_136_223_846_793_005)
142                .wrapping_add(1);
143            let since_version = state % 25;
144            registry.register(
145                OpKey::new(op_type, domain, since_version),
146                Box::new(DummyFactory(factory_id)),
147            );
148        }
149
150        for _ in 0..512 {
151            state = state
152                .wrapping_mul(6_364_136_223_846_793_005)
153                .wrapping_add(1);
154            let op_type = ops[(state as usize) % ops.len()];
155            state = state
156                .wrapping_mul(6_364_136_223_846_793_005)
157                .wrapping_add(1);
158            let domain = domains[(state as usize) % domains.len()];
159            state = state
160                .wrapping_mul(6_364_136_223_846_793_005)
161                .wrapping_add(1);
162            let opset = state % 30;
163            let domain = norm_domain(domain);
164
165            let linear_lookup = registry
166                .entries
167                .iter()
168                .filter(|(key, _)| {
169                    key.op_type == op_type && key.domain == domain && key.since_version <= opset
170                })
171                .max_by_key(|(key, _)| key.since_version)
172                .map(|(_, factory)| factory.as_ref());
173            match (registry.lookup(op_type, domain, opset), linear_lookup) {
174                (Some(indexed), Some(linear)) => assert!(std::ptr::eq(indexed, linear)),
175                (None, None) => {}
176                _ => panic!("indexed lookup differed from linear reference"),
177            }
178
179            let linear_supports = registry.entries.keys().any(|key| {
180                key.op_type == op_type && key.domain == domain && key.since_version <= opset
181            });
182            assert_eq!(registry.supports(op_type, domain, opset), linear_supports);
183
184            let linear_earliest = registry
185                .entries
186                .keys()
187                .filter(|key| key.op_type == op_type && key.domain == domain)
188                .map(|key| key.since_version)
189                .min();
190            assert_eq!(
191                registry.earliest_since_version(op_type, domain),
192                linear_earliest
193            );
194        }
195    }
196}
197
198/// Ordered set of execution providers with a priority list (§4.6).
199#[derive(Default)]
200pub struct EpRegistry {
201    eps: Vec<Box<dyn ExecutionProvider>>,
202    /// Priority order as indices into `eps` (front = highest priority).
203    priority: Vec<EpId>,
204}
205
206impl EpRegistry {
207    pub fn new() -> Self {
208        Self::default()
209    }
210
211    /// Register an EP, returning its [`EpId`]. Appended to the priority list.
212    pub fn register(&mut self, ep: Box<dyn ExecutionProvider>) -> EpId {
213        let id = EpId(self.eps.len() as u32);
214        self.eps.push(ep);
215        self.priority.push(id);
216        id
217    }
218
219    /// Load a legacy ORT plugin EP from a shared library (Phase 2).
220    pub fn load_legacy(&mut self, path: &Path, config: &EpConfig) -> Result<EpId> {
221        let _ = (path, config);
222        todo!("ort2-ep-api Phase 2: dlopen legacy ORT plugin EP and adapt its vtable")
223    }
224
225    /// Override the priority order.
226    pub fn set_priority(&mut self, order: Vec<EpId>) {
227        self.priority = order;
228    }
229
230    /// Borrow an EP by id.
231    pub fn get(&self, id: EpId) -> Option<&dyn ExecutionProvider> {
232        self.eps.get(id.0 as usize).map(|b| b.as_ref())
233    }
234
235    /// The priority order.
236    pub fn priority(&self) -> &[EpId] {
237        &self.priority
238    }
239
240    /// All EPs (in priority order) that can handle `op`, with their match info.
241    pub fn candidates_for_op(
242        &self,
243        op: &Node,
244        opset: u64,
245        shapes: &[Shape],
246        input_dtypes: &[DataType],
247        layouts: &[TensorLayout],
248    ) -> Vec<(EpId, KernelMatch)> {
249        let mut out = Vec::new();
250        for &id in &self.priority {
251            if let Some(ep) = self.get(id) {
252                let m = ep.supports_op(op, opset, shapes, input_dtypes, layouts);
253                if m.is_supported() {
254                    out.push((id, m));
255                }
256            }
257        }
258        out
259    }
260}