use std::{
collections::BTreeMap,
sync::{OnceLock, RwLock},
};
use sim_kernel::{Error, Result, Symbol};
use sim_lib_numbers_cas::CasExpr;
pub type DiffRule = Box<dyn Fn(&[CasExpr], &Symbol) -> Option<CasExpr> + Send + Sync>;
#[derive(Default)]
pub struct CasDiffRegistry {
pub rules: BTreeMap<Symbol, DiffRule>,
}
impl CasDiffRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register_rule(&mut self, symbol: Symbol, rule: DiffRule) -> Result<()> {
if self.rules.contains_key(&symbol) {
return Err(Error::Eval(format!(
"CAS diff rule for operator {symbol} is already registered"
)));
}
self.rules.insert(symbol, rule);
Ok(())
}
pub fn override_rule(&mut self, symbol: Symbol, rule: DiffRule) -> Option<DiffRule> {
self.rules.insert(symbol, rule)
}
pub fn apply(&self, symbol: &Symbol, args: &[CasExpr], var: &Symbol) -> Option<CasExpr> {
self.rules.get(symbol).and_then(|rule| rule(args, var))
}
}
static REGISTRY: OnceLock<RwLock<CasDiffRegistry>> = OnceLock::new();
pub fn global_diff_registry() -> &'static RwLock<CasDiffRegistry> {
REGISTRY.get_or_init(|| RwLock::new(CasDiffRegistry::new()))
}
pub fn register_diff_rule(symbol: Symbol, rule: DiffRule) -> Result<()> {
global_diff_registry()
.write()
.map_err(|_| Error::PoisonedLock("CAS diff registry"))?
.register_rule(symbol, rule)
}
pub fn override_diff_rule(symbol: Symbol, rule: DiffRule) -> Option<DiffRule> {
global_diff_registry()
.write()
.expect("CAS diff registry should not be poisoned")
.override_rule(symbol, rule)
}
pub(crate) fn apply_registered_rule(
symbol: &Symbol,
args: &[CasExpr],
var: &Symbol,
) -> Option<CasExpr> {
global_diff_registry()
.read()
.expect("CAS diff registry should not be poisoned")
.apply(symbol, args, var)
}
#[cfg(test)]
mod tests {
use super::*;
fn identity_rule() -> DiffRule {
Box::new(|args: &[CasExpr], _var: &Symbol| args.first().cloned())
}
fn constant_rule(symbol: &'static str) -> DiffRule {
Box::new(move |_args: &[CasExpr], _var: &Symbol| Some(CasExpr::Var(Symbol::new(symbol))))
}
#[test]
fn duplicate_rule_registration_fails_closed() {
let mut registry = CasDiffRegistry::new();
registry
.register_rule(Symbol::new("custom"), identity_rule())
.unwrap();
let err = registry
.register_rule(Symbol::new("custom"), identity_rule())
.unwrap_err();
assert!(err.to_string().contains("already registered"), "{err}");
}
#[test]
fn explicit_override_replaces_existing_rule() {
let mut registry = CasDiffRegistry::new();
let symbol = Symbol::new("custom");
registry
.register_rule(symbol.clone(), constant_rule("before"))
.unwrap();
let replaced = registry.override_rule(symbol.clone(), constant_rule("after"));
let out = registry.apply(&symbol, &[], &Symbol::new("x"));
assert!(replaced.is_some());
assert!(matches!(out, Some(CasExpr::Var(value)) if value == Symbol::new("after")));
}
}