1use 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
110pub fn global_numeric_registry() -> &'static RwLock<NumericRegistry> {
112 GLOBAL_NUMERIC_REGISTRY.get_or_init(|| RwLock::new(NumericRegistry::default()))
113}
114
115pub 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
127pub 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
142pub 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}