use std::sync::Arc;
use ocas_domain::Rational;
use ocas_domain::dual::{DualShape, HyperDual, new_first_order};
use pyo3::exceptions::{PyTypeError, PyValueError};
use pyo3::prelude::*;
fn py_to_rational(obj: &Bound<'_, PyAny>) -> PyResult<Rational> {
if let Ok(n) = obj.extract::<i64>() {
Ok(Rational::new(n, 1))
} else if let Ok((num, den)) = obj.extract::<(i64, i64)>() {
if den == 0 {
Err(PyValueError::new_err("rational denominator cannot be zero"))
} else {
Ok(Rational::new(num, den))
}
} else {
Err(PyTypeError::new_err("expected int or (num, denom) tuple"))
}
}
fn rational_to_string(r: &Rational) -> String {
let n = r.numer().to_i64().unwrap_or(0);
let d = r.denom().to_i64().unwrap_or(0);
if d == 1 {
n.to_string()
} else {
format!("{n}/{d}")
}
}
#[pyclass(name = "DualShape")]
pub struct PyDualShape {
shape: Arc<DualShape>,
}
#[pymethods]
impl PyDualShape {
#[staticmethod]
fn first_order(n_vars: usize) -> PyResult<Self> {
if n_vars == 0 {
return Err(PyValueError::new_err("n_vars must be >= 1"));
}
Ok(PyDualShape {
shape: new_first_order::<Rational>(n_vars),
})
}
#[getter]
fn n_vars(&self) -> usize {
self.shape.n_vars()
}
#[getter]
fn n_components(&self) -> usize {
self.shape.n_components()
}
fn __repr__(&self) -> String {
format!(
"DualShape(n_vars={}, n_components={})",
self.n_vars(),
self.n_components()
)
}
}
#[pyclass(name = "HyperDual")]
pub struct PyHyperDual {
inner: HyperDual<Rational>,
shape: Arc<DualShape>,
}
impl PyHyperDual {
fn new(inner: HyperDual<Rational>) -> Self {
let shape = inner.shape().clone();
PyHyperDual { inner, shape }
}
}
#[pymethods]
impl PyHyperDual {
#[staticmethod]
fn variable(shape: &PyDualShape, i: usize, value: &Bound<'_, PyAny>) -> PyResult<Self> {
if i >= shape.shape.n_vars() {
return Err(PyValueError::new_err(format!(
"variable index {i} out of range (n_vars = {})",
shape.shape.n_vars()
)));
}
let v = py_to_rational(value)?;
Ok(PyHyperDual::new(HyperDual::variable(&shape.shape, i, v)))
}
#[staticmethod]
fn constant(shape: &PyDualShape, value: &Bound<'_, PyAny>) -> PyResult<Self> {
let v = py_to_rational(value)?;
Ok(PyHyperDual::new(HyperDual::constant(&shape.shape, v)))
}
fn value(&self) -> String {
rational_to_string(self.inner.value())
}
fn deriv(&self, i: usize) -> Option<String> {
self.inner.deriv(i).map(rational_to_string)
}
#[getter]
fn n_vars(&self) -> usize {
self.shape.n_vars()
}
fn __repr__(&self) -> String {
format!(
"HyperDual(value={}, n_vars={})",
self.value(),
self.n_vars()
)
}
fn __add__(&self, other: &PyHyperDual) -> PyResult<PyHyperDual> {
if !Arc::ptr_eq(&self.shape, &other.shape) {
return Err(PyValueError::new_err(
"cannot add HyperDuals with different shapes",
));
}
Ok(PyHyperDual::new(self.inner.clone() + other.inner.clone()))
}
fn __sub__(&self, other: &PyHyperDual) -> PyResult<PyHyperDual> {
if !Arc::ptr_eq(&self.shape, &other.shape) {
return Err(PyValueError::new_err(
"cannot subtract HyperDuals with different shapes",
));
}
Ok(PyHyperDual::new(self.inner.clone() - other.inner.clone()))
}
fn __mul__(&self, other: &PyHyperDual) -> PyResult<PyHyperDual> {
if !Arc::ptr_eq(&self.shape, &other.shape) {
return Err(PyValueError::new_err(
"cannot multiply HyperDuals with different shapes",
));
}
Ok(PyHyperDual::new(self.inner.clone() * other.inner.clone()))
}
fn __truediv__(&self, other: &PyHyperDual) -> PyResult<PyHyperDual> {
if !Arc::ptr_eq(&self.shape, &other.shape) {
return Err(PyValueError::new_err(
"cannot divide HyperDuals with different shapes",
));
}
if other.inner.value() == &Rational::new(0, 1) {
return Err(PyValueError::new_err(
"division by zero (value component is zero)",
));
}
Ok(PyHyperDual::new(self.inner.clone() / other.inner.clone()))
}
fn __neg__(&self) -> PyHyperDual {
PyHyperDual::new(-self.inner.clone())
}
}