use ocas_atom::{Atom, AtomArena, Symbol, normalize::normalize};
use ocas_calc::{diff, integrate, integrate_heuristic, substitute, taylor};
use ocas_core::arena::Arena;
use ocas_parse::parse;
use ocas_rewrite::rules::default_rules;
use ocas_rewrite::simplify::simplify;
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
struct ExprInner {
arena_ptr: *mut Arena,
ctx_ptr: *mut AtomArena<'static>,
atom: Atom<'static>,
}
unsafe impl Send for ExprInner {}
unsafe impl Sync for ExprInner {}
impl Drop for ExprInner {
fn drop(&mut self) {
unsafe {
let _ = Box::from_raw(self.ctx_ptr);
let _ = Box::from_raw(self.arena_ptr);
}
}
}
struct ArenaGuard {
arena_ptr: *mut Arena,
ctx_ptr: *mut AtomArena<'static>,
armed: bool,
}
impl ArenaGuard {
fn new(arena_ptr: *mut Arena, ctx_ptr: *mut AtomArena<'static>) -> Self {
ArenaGuard {
arena_ptr,
ctx_ptr,
armed: true,
}
}
fn disarm(&mut self) {
self.armed = false;
}
}
impl Drop for ArenaGuard {
fn drop(&mut self) {
if self.armed {
unsafe {
let _ = Box::from_raw(self.ctx_ptr);
let _ = Box::from_raw(self.arena_ptr);
}
}
}
}
impl ExprInner {
fn ctx(&self) -> &'static AtomArena<'static> {
unsafe { &*self.ctx_ptr }
}
fn new_pair() -> (*mut Arena, *mut AtomArena<'static>) {
let arena_box: Box<Arena> = Box::new(Arena::new());
let arena_ptr = Box::into_raw(arena_box);
let arena_ref: &'static Arena = unsafe { &*arena_ptr };
let ctx = AtomArena::new(arena_ref);
let ctx_ptr = Box::into_raw(Box::new(ctx));
(arena_ptr, ctx_ptr)
}
fn build<F>(f: F) -> PyResult<Box<Self>>
where
F: FnOnce(&'static AtomArena<'static>) -> Result<Atom<'static>, String>,
{
let (arena_ptr, ctx_ptr) = Self::new_pair();
let mut guard = ArenaGuard::new(arena_ptr, ctx_ptr);
let ctx = unsafe { &*ctx_ptr };
let atom = f(ctx).map_err(PyValueError::new_err)?;
let normalized = normalize(ctx, atom);
guard.disarm();
Ok(Box::new(ExprInner {
arena_ptr,
ctx_ptr,
atom: normalized,
}))
}
fn from_str(input: &str) -> PyResult<Box<Self>> {
Self::build(|ctx| match parse(ctx, input) {
Ok(a) => Ok(a),
Err(e) => Err(format!("parse error: {e}")),
})
}
fn from_string_src(src: String) -> PyResult<Box<Self>> {
Self::build(|ctx| match parse(ctx, &src) {
Ok(a) => Ok(a),
Err(e) => Err(format!("parse error: {e}")),
})
}
}
#[pyclass(name = "Expression")]
pub struct Expression {
inner: Box<ExprInner>,
}
impl Expression {
pub(crate) fn ctx_ref(&self) -> &'static AtomArena<'static> {
self.inner.ctx()
}
pub(crate) fn atom(&self) -> Atom<'static> {
self.inner.atom
}
}
#[pymethods]
impl Expression {
#[new]
fn new(input: &str) -> PyResult<Self> {
Ok(Expression {
inner: ExprInner::from_str(input)?,
})
}
fn __str__(&self) -> String {
self.inner.atom.to_string()
}
fn __repr__(&self) -> String {
format!("Expression({:?})", self.inner.atom.to_string())
}
fn __add__(&self, other: &Expression) -> PyResult<Expression> {
let left = self.inner.atom.to_string();
let right = other.inner.atom.to_string();
let combined = format!("({left}) + ({right})");
Ok(Expression {
inner: ExprInner::from_string_src(combined)?,
})
}
fn __sub__(&self, other: &Expression) -> PyResult<Expression> {
let left = self.inner.atom.to_string();
let right = other.inner.atom.to_string();
let combined = format!("({left}) + (-1)*({right})");
Ok(Expression {
inner: ExprInner::from_string_src(combined)?,
})
}
fn __mul__(&self, other: &Expression) -> PyResult<Expression> {
let left = self.inner.atom.to_string();
let right = other.inner.atom.to_string();
let combined = format!("({left})*({right})");
Ok(Expression {
inner: ExprInner::from_string_src(combined)?,
})
}
fn __pow__(&self, other: &Expression, _modulo: Option<&Expression>) -> PyResult<Expression> {
let left = self.inner.atom.to_string();
let right = other.inner.atom.to_string();
let combined = format!("({left})^({right})");
Ok(Expression {
inner: ExprInner::from_string_src(combined)?,
})
}
fn __neg__(&self) -> PyResult<Expression> {
let src = self.inner.atom.to_string();
Ok(Expression {
inner: ExprInner::from_string_src(format!("(-1)*({src})"))?,
})
}
fn __eq__(&self, other: &Expression) -> bool {
let a = normalize(self.inner.ctx(), self.inner.atom);
let b = normalize(other.inner.ctx(), other.inner.atom);
a.to_string() == b.to_string()
}
fn __hash__(&self) -> u64 {
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
let mut h = DefaultHasher::new();
self.inner.atom.to_string().hash(&mut h);
h.finish()
}
fn clone(&self) -> PyResult<Expression> {
let src = self.inner.atom.to_string();
Ok(Expression {
inner: ExprInner::from_string_src(src)?,
})
}
fn simplify(&self) -> PyResult<Expression> {
let src = self.inner.atom.to_string();
ExprInner::build(|ctx| {
let a = parse(ctx, &src).map_err(|e| e.to_string())?;
let rules = default_rules(ctx, &());
Ok(simplify(ctx, a, &rules, 20))
})
.map(|inner| Expression { inner })
}
fn diff(&self, var: &str) -> PyResult<Expression> {
let src = self.inner.atom.to_string();
let var_sym = Symbol::new(var);
ExprInner::build(|ctx| match parse(ctx, &src) {
Ok(a) => Ok(diff(ctx, a, var_sym)),
Err(e) => Err(e.to_string()),
})
.map(|inner| Expression { inner })
}
fn integrate(&self, var: &str) -> PyResult<Expression> {
let src = self.inner.atom.to_string();
let var_sym = Symbol::new(var);
ExprInner::build(|ctx| match parse(ctx, &src) {
Ok(a) => Ok(integrate(ctx, a, var_sym)),
Err(e) => Err(e.to_string()),
})
.map(|inner| Expression { inner })
}
fn integrate_heuristic(&self, var: &str) -> PyResult<Expression> {
let src = self.inner.atom.to_string();
let var_sym = Symbol::new(var);
ExprInner::build(|ctx| match parse(ctx, &src) {
Ok(a) => Ok(integrate_heuristic(ctx, a, var_sym)),
Err(e) => Err(e.to_string()),
})
.map(|inner| Expression { inner })
}
fn taylor(&self, var: &str, point: &Expression, order: usize) -> PyResult<Expression> {
let expr_src = self.inner.atom.to_string();
let point_src = point.inner.atom.to_string();
let var_sym = Symbol::new(var);
ExprInner::build(|ctx| {
let e = parse(ctx, &expr_src).map_err(|e| e.to_string())?;
let p = parse(ctx, &point_src).map_err(|e| e.to_string())?;
Ok(taylor(ctx, e, var_sym, p, order))
})
.map(|inner| Expression { inner })
}
fn substitute(&self, var: &str, replacement: &Expression) -> PyResult<Expression> {
let expr_src = self.inner.atom.to_string();
let repl_src = replacement.inner.atom.to_string();
let var_sym = Symbol::new(var);
ExprInner::build(|ctx| {
let e = parse(ctx, &expr_src).map_err(|e| e.to_string())?;
let r = parse(ctx, &repl_src).map_err(|e| e.to_string())?;
Ok(substitute(ctx, e, var_sym, r))
})
.map(|inner| Expression { inner })
}
}