Skip to main content

sim_lib_numbers_cas_diff/implementation/
registry.rs

1//! The extensible differentiation-rule registry: a process-global map from an
2//! operator symbol to a custom `diff` rule, letting other libraries teach the
3//! differentiator new functions.
4
5use std::{
6    collections::BTreeMap,
7    sync::{OnceLock, RwLock},
8};
9
10use sim_kernel::{Error, Result, Symbol};
11use sim_lib_numbers_cas::CasExpr;
12
13/// A custom differentiation rule: maps an operator's arguments and the variable
14/// of differentiation to a derivative tree, or `None` to decline.
15pub type DiffRule = Box<dyn Fn(&[CasExpr], &Symbol) -> Option<CasExpr> + Send + Sync>;
16
17/// A registry of per-operator differentiation rules.
18///
19/// # Examples
20///
21/// ```
22/// use sim_kernel::Symbol;
23/// use sim_lib_numbers_cas::CasExpr;
24/// use sim_lib_numbers_cas_diff::CasDiffRegistry;
25///
26/// let mut registry = CasDiffRegistry::new();
27/// registry.register_rule(
28///     Symbol::new("id"),
29///     Box::new(|args: &[CasExpr], _var: &Symbol| args.first().cloned()),
30/// )
31/// .unwrap();
32///
33/// let out = registry.apply(
34///     &Symbol::new("id"),
35///     &[CasExpr::Var(Symbol::new("x"))],
36///     &Symbol::new("x"),
37/// );
38/// assert!(matches!(out, Some(CasExpr::Var(_))));
39/// // An unregistered operator yields `None`.
40/// assert!(registry.apply(&Symbol::new("nope"), &[], &Symbol::new("x")).is_none());
41/// ```
42#[derive(Default)]
43pub struct CasDiffRegistry {
44    /// The registered rules, keyed by operator symbol.
45    pub rules: BTreeMap<Symbol, DiffRule>,
46}
47
48impl CasDiffRegistry {
49    /// Construct an empty registry.
50    pub fn new() -> Self {
51        Self::default()
52    }
53
54    /// Register `rule` for `symbol`.
55    ///
56    /// Returns an error when a rule is already registered for the same symbol.
57    /// Use [`Self::override_rule`] when replacement is intentional.
58    pub fn register_rule(&mut self, symbol: Symbol, rule: DiffRule) -> Result<()> {
59        if self.rules.contains_key(&symbol) {
60            return Err(Error::Eval(format!(
61                "CAS diff rule for operator {symbol} is already registered"
62            )));
63        }
64        self.rules.insert(symbol, rule);
65        Ok(())
66    }
67
68    /// Replace the rule for `symbol`, returning any rule that was already there.
69    pub fn override_rule(&mut self, symbol: Symbol, rule: DiffRule) -> Option<DiffRule> {
70        self.rules.insert(symbol, rule)
71    }
72
73    /// Apply the rule registered for `symbol`, if any, to `args` and `var`.
74    pub fn apply(&self, symbol: &Symbol, args: &[CasExpr], var: &Symbol) -> Option<CasExpr> {
75        self.rules.get(symbol).and_then(|rule| rule(args, var))
76    }
77}
78
79static REGISTRY: OnceLock<RwLock<CasDiffRegistry>> = OnceLock::new();
80
81/// Access the process-global differentiation-rule registry.
82pub fn global_diff_registry() -> &'static RwLock<CasDiffRegistry> {
83    REGISTRY.get_or_init(|| RwLock::new(CasDiffRegistry::new()))
84}
85
86/// Register `rule` for `symbol` in the global registry.
87///
88/// # Errors
89///
90/// Returns an error if the global registry lock is poisoned or a rule already
91/// exists for `symbol`.
92pub fn register_diff_rule(symbol: Symbol, rule: DiffRule) -> Result<()> {
93    global_diff_registry()
94        .write()
95        .map_err(|_| Error::PoisonedLock("CAS diff registry"))?
96        .register_rule(symbol, rule)
97}
98
99/// Replace the global rule for `symbol`, returning any rule that was already
100/// registered.
101pub fn override_diff_rule(symbol: Symbol, rule: DiffRule) -> Option<DiffRule> {
102    global_diff_registry()
103        .write()
104        .expect("CAS diff registry should not be poisoned")
105        .override_rule(symbol, rule)
106}
107
108pub(crate) fn apply_registered_rule(
109    symbol: &Symbol,
110    args: &[CasExpr],
111    var: &Symbol,
112) -> Option<CasExpr> {
113    global_diff_registry()
114        .read()
115        .expect("CAS diff registry should not be poisoned")
116        .apply(symbol, args, var)
117}
118
119#[cfg(test)]
120mod tests {
121    use super::*;
122
123    fn identity_rule() -> DiffRule {
124        Box::new(|args: &[CasExpr], _var: &Symbol| args.first().cloned())
125    }
126
127    fn constant_rule(symbol: &'static str) -> DiffRule {
128        Box::new(move |_args: &[CasExpr], _var: &Symbol| Some(CasExpr::Var(Symbol::new(symbol))))
129    }
130
131    #[test]
132    fn duplicate_rule_registration_fails_closed() {
133        let mut registry = CasDiffRegistry::new();
134        registry
135            .register_rule(Symbol::new("custom"), identity_rule())
136            .unwrap();
137
138        let err = registry
139            .register_rule(Symbol::new("custom"), identity_rule())
140            .unwrap_err();
141
142        assert!(err.to_string().contains("already registered"), "{err}");
143    }
144
145    #[test]
146    fn explicit_override_replaces_existing_rule() {
147        let mut registry = CasDiffRegistry::new();
148        let symbol = Symbol::new("custom");
149        registry
150            .register_rule(symbol.clone(), constant_rule("before"))
151            .unwrap();
152
153        let replaced = registry.override_rule(symbol.clone(), constant_rule("after"));
154        let out = registry.apply(&symbol, &[], &Symbol::new("x"));
155
156        assert!(replaced.is_some());
157        assert!(matches!(out, Some(CasExpr::Var(value)) if value == Symbol::new("after")));
158    }
159}