Skip to main content

sim_lib_numbers_numeric/
registry.rs

1//! The global numeric registry holding registered differentiator, quadrature,
2//! and ODE-solver plugins keyed by name and kind.
3
4use std::{
5    collections::BTreeMap,
6    sync::{Arc, OnceLock, RwLock},
7};
8
9use sim_kernel::{Error, Result, Symbol};
10
11use super::traits::{Differentiator, NumericKind, OdeSolver, Quadrature};
12
13#[derive(Default)]
14pub struct NumericRegistry {
15    differentiators: BTreeMap<Symbol, Arc<dyn Differentiator>>,
16    quadrature_fixed: BTreeMap<Symbol, Arc<dyn Quadrature>>,
17    quadrature_adaptive: BTreeMap<Symbol, Arc<dyn Quadrature>>,
18    ode_fixed: BTreeMap<Symbol, Arc<dyn OdeSolver>>,
19    ode_adaptive: BTreeMap<Symbol, Arc<dyn OdeSolver>>,
20}
21
22impl NumericRegistry {
23    pub fn register_differentiator(&mut self, plugin: Arc<dyn Differentiator>) -> Result<()> {
24        let name = plugin.name();
25        insert_plugin(
26            &mut self.differentiators,
27            NumericKind::Differentiator,
28            name,
29            plugin,
30        )
31    }
32
33    pub fn register_quadrature(&mut self, plugin: Arc<dyn Quadrature>) -> Result<()> {
34        let name = plugin.name();
35        match plugin.kind() {
36            NumericKind::QuadratureFixed => insert_plugin(
37                &mut self.quadrature_fixed,
38                NumericKind::QuadratureFixed,
39                name,
40                plugin,
41            ),
42            NumericKind::QuadratureAdaptive => insert_plugin(
43                &mut self.quadrature_adaptive,
44                NumericKind::QuadratureAdaptive,
45                name,
46                plugin,
47            ),
48            kind => Err(Error::Eval(format!(
49                "numeric plugin :{name} has invalid quadrature kind {kind:?}"
50            ))),
51        }
52    }
53
54    pub fn register_ode_solver(&mut self, plugin: Arc<dyn OdeSolver>) -> Result<()> {
55        let name = plugin.name();
56        match plugin.kind() {
57            NumericKind::OdeFixed => {
58                insert_plugin(&mut self.ode_fixed, NumericKind::OdeFixed, name, plugin)
59            }
60            NumericKind::OdeAdaptive => insert_plugin(
61                &mut self.ode_adaptive,
62                NumericKind::OdeAdaptive,
63                name,
64                plugin,
65            ),
66            kind => Err(Error::Eval(format!(
67                "numeric plugin :{name} has invalid ODE kind {kind:?}"
68            ))),
69        }
70    }
71
72    pub fn differentiator(&self, method: &Symbol) -> Option<Arc<dyn Differentiator>> {
73        self.differentiators.get(method).cloned()
74    }
75
76    pub fn quadrature_fixed(&self, method: &Symbol) -> Option<Arc<dyn Quadrature>> {
77        self.quadrature_fixed.get(method).cloned()
78    }
79
80    pub fn quadrature_adaptive(&self, method: &Symbol) -> Option<Arc<dyn Quadrature>> {
81        self.quadrature_adaptive.get(method).cloned()
82    }
83
84    pub fn ode_fixed(&self, method: &Symbol) -> Option<Arc<dyn OdeSolver>> {
85        self.ode_fixed.get(method).cloned()
86    }
87
88    pub fn ode_adaptive(&self, method: &Symbol) -> Option<Arc<dyn OdeSolver>> {
89        self.ode_adaptive.get(method).cloned()
90    }
91}
92
93fn insert_plugin<T: ?Sized>(
94    plugins: &mut BTreeMap<Symbol, Arc<T>>,
95    kind: NumericKind,
96    name: Symbol,
97    plugin: Arc<T>,
98) -> Result<()> {
99    if plugins.contains_key(&name) {
100        return Err(Error::Eval(format!(
101            "numeric {kind:?} plugin :{name} is already registered"
102        )));
103    }
104    plugins.insert(name, plugin);
105    Ok(())
106}
107
108static GLOBAL_NUMERIC_REGISTRY: OnceLock<RwLock<NumericRegistry>> = OnceLock::new();
109
110/// Returns the process-global numeric registry, creating it on first access.
111pub fn global_numeric_registry() -> &'static RwLock<NumericRegistry> {
112    GLOBAL_NUMERIC_REGISTRY.get_or_init(|| RwLock::new(NumericRegistry::default()))
113}
114
115/// Registers a differentiator backend in the global numeric registry.
116///
117/// # Errors
118///
119/// Returns an error if the global registry lock is poisoned.
120pub fn register_differentiator(plugin: Arc<dyn Differentiator>) -> Result<()> {
121    global_numeric_registry()
122        .write()
123        .map_err(|_| Error::PoisonedLock("numeric registry"))?
124        .register_differentiator(plugin)
125}
126
127/// Registers a quadrature backend in the global numeric registry.
128///
129/// The plugin is routed to the fixed or adaptive slot according to its
130/// [`NumericKind`](crate::NumericKind).
131///
132/// # Errors
133///
134/// Returns an error if the global registry lock is poisoned.
135pub fn register_quadrature(plugin: Arc<dyn Quadrature>) -> Result<()> {
136    global_numeric_registry()
137        .write()
138        .map_err(|_| Error::PoisonedLock("numeric registry"))?
139        .register_quadrature(plugin)
140}
141
142/// Registers an ODE-solver backend in the global numeric registry.
143///
144/// The plugin is routed to the fixed or adaptive slot according to its
145/// [`NumericKind`](crate::NumericKind).
146///
147/// # Errors
148///
149/// Returns an error if the global registry lock is poisoned.
150pub fn register_ode_solver(plugin: Arc<dyn OdeSolver>) -> Result<()> {
151    global_numeric_registry()
152        .write()
153        .map_err(|_| Error::PoisonedLock("numeric registry"))?
154        .register_ode_solver(plugin)
155}
156
157#[cfg(test)]
158mod tests {
159    use super::*;
160    use crate::implementation::traits::{
161        DiffOpts, NumericCallable, NumericPlugin, OdeOpts, OdeProblem, QuadOpts,
162    };
163    use sim_kernel::{Cx, Value};
164    use sim_lib_numbers_func::Func;
165
166    struct TestPlugin {
167        name: Symbol,
168        kind: NumericKind,
169    }
170
171    impl TestPlugin {
172        fn new(name: &str, kind: NumericKind) -> Self {
173            Self {
174                name: Symbol::new(name),
175                kind,
176            }
177        }
178    }
179
180    impl NumericPlugin for TestPlugin {
181        fn name(&self) -> Symbol {
182            self.name.clone()
183        }
184
185        fn kind(&self) -> NumericKind {
186            self.kind
187        }
188    }
189
190    impl Differentiator for TestPlugin {
191        fn diff_at(
192            &self,
193            _cx: &mut Cx,
194            _f: &Func,
195            _var: &Symbol,
196            point: &Value,
197            _opt: DiffOpts,
198        ) -> Result<Value> {
199            Ok(point.clone())
200        }
201    }
202
203    impl Quadrature for TestPlugin {
204        fn integrate(
205            &self,
206            _cx: &mut Cx,
207            _f: &NumericCallable,
208            _var: &Symbol,
209            lo: &Value,
210            _hi: &Value,
211            _opt: QuadOpts,
212        ) -> Result<Value> {
213            Ok(lo.clone())
214        }
215    }
216
217    impl OdeSolver for TestPlugin {
218        fn solve(
219            &self,
220            _cx: &mut Cx,
221            _problem: OdeProblem<'_>,
222            _opt: OdeOpts,
223        ) -> Result<Vec<(Value, Value)>> {
224            Ok(Vec::new())
225        }
226    }
227
228    fn differentiator(name: &str) -> Arc<dyn Differentiator> {
229        Arc::new(TestPlugin::new(name, NumericKind::Differentiator))
230    }
231
232    fn quadrature(name: &str, kind: NumericKind) -> Arc<dyn Quadrature> {
233        Arc::new(TestPlugin::new(name, kind))
234    }
235
236    fn ode_solver(name: &str, kind: NumericKind) -> Arc<dyn OdeSolver> {
237        Arc::new(TestPlugin::new(name, kind))
238    }
239
240    #[test]
241    fn duplicate_differentiator_registration_is_rejected() {
242        let mut registry = NumericRegistry::default();
243        registry
244            .register_differentiator(differentiator("central-5"))
245            .unwrap();
246
247        let err = registry
248            .register_differentiator(differentiator("central-5"))
249            .unwrap_err();
250
251        assert!(err.to_string().contains("already registered"), "{err}");
252    }
253
254    #[test]
255    fn duplicate_quadrature_registration_is_rejected_per_kind() {
256        let mut registry = NumericRegistry::default();
257        registry
258            .register_quadrature(quadrature("romberg", NumericKind::QuadratureFixed))
259            .unwrap();
260
261        let err = registry
262            .register_quadrature(quadrature("romberg", NumericKind::QuadratureFixed))
263            .unwrap_err();
264
265        assert!(err.to_string().contains("already registered"), "{err}");
266    }
267
268    #[test]
269    fn fixed_and_adaptive_quadrature_can_share_method_name() {
270        let mut registry = NumericRegistry::default();
271        let name = Symbol::new("romberg");
272
273        registry
274            .register_quadrature(quadrature("romberg", NumericKind::QuadratureFixed))
275            .unwrap();
276        registry
277            .register_quadrature(quadrature("romberg", NumericKind::QuadratureAdaptive))
278            .unwrap();
279
280        assert!(registry.quadrature_fixed(&name).is_some());
281        assert!(registry.quadrature_adaptive(&name).is_some());
282    }
283
284    #[test]
285    fn duplicate_ode_registration_is_rejected_per_kind() {
286        let mut registry = NumericRegistry::default();
287        registry
288            .register_ode_solver(ode_solver("rk4", NumericKind::OdeFixed))
289            .unwrap();
290
291        let err = registry
292            .register_ode_solver(ode_solver("rk4", NumericKind::OdeFixed))
293            .unwrap_err();
294
295        assert!(err.to_string().contains("already registered"), "{err}");
296    }
297}