sim_lib_numbers_cas_diff/implementation/
registry.rs1use std::{
6 collections::BTreeMap,
7 sync::{OnceLock, RwLock},
8};
9
10use sim_kernel::{Error, Result, Symbol};
11use sim_lib_numbers_cas::CasExpr;
12
13pub type DiffRule = Box<dyn Fn(&[CasExpr], &Symbol) -> Option<CasExpr> + Send + Sync>;
16
17#[derive(Default)]
43pub struct CasDiffRegistry {
44 pub rules: BTreeMap<Symbol, DiffRule>,
46}
47
48impl CasDiffRegistry {
49 pub fn new() -> Self {
51 Self::default()
52 }
53
54 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 pub fn override_rule(&mut self, symbol: Symbol, rule: DiffRule) -> Option<DiffRule> {
70 self.rules.insert(symbol, rule)
71 }
72
73 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
81pub fn global_diff_registry() -> &'static RwLock<CasDiffRegistry> {
83 REGISTRY.get_or_init(|| RwLock::new(CasDiffRegistry::new()))
84}
85
86pub 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
99pub 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}