use ocas_domain::{AlgebraicElement, AlgebraicNumberField, Rational, RationalDomain};
use ocas_poly::DenseUnivariatePolynomial;
use pyo3::exceptions::{PyTypeError, PyValueError};
use pyo3::prelude::*;
#[pyclass(name = "AlgebraicExtension")]
pub struct PyAlgebraicExtension {
pub(crate) field: AlgebraicNumberField,
}
#[pyclass(name = "AlgebraicElement")]
pub struct PyAlgebraicElement {
pub(crate) elem: AlgebraicElement<Rational>,
}
#[pyclass(name = "AlgebraicPolynomial", skip_from_py_object)]
#[derive(Clone)]
pub struct PyAlgebraicPolynomial {
pub(crate) inner: DenseUnivariatePolynomial<AlgebraicNumberField>,
}
#[pyclass(name = "AlgebraicFactor", skip_from_py_object)]
pub struct PyAlgebraicFactor {
#[pyo3(get)]
pub factor: PyAlgebraicPolynomial,
#[pyo3(get)]
pub multiplicity: usize,
}
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 parse_min_poly(coeffs: &Bound<'_, PyAny>) -> PyResult<Vec<Rational>> {
let iter = coeffs.try_iter().map_err(|_| {
PyTypeError::new_err(
"minimal polynomial coefficients must be a list of ints or (num, denom) tuples",
)
})?;
iter.map(|c| py_to_rational(&c?)).collect()
}
fn py_to_anf_element(
field: &AlgebraicNumberField,
obj: &Bound<'_, PyAny>,
) -> PyResult<AlgebraicElement<Rational>> {
if let Ok(r) = py_to_rational(obj) {
return Ok(field.from_base(r));
}
if let Ok(elem) = obj.extract::<PyRef<'_, PyAlgebraicElement>>() {
return Ok(elem.elem.clone());
}
if obj.try_iter().is_ok() {
let cs: PyResult<Vec<Rational>> = obj
.try_iter()
.unwrap()
.map(|c| py_to_rational(&c?))
.collect();
return Ok(field.element(cs?));
}
Err(PyTypeError::new_err(
"coefficient must be int, (num, denom), list, or AlgebraicElement",
))
}
#[pymethods]
impl PyAlgebraicExtension {
#[new]
fn new(min_poly: &Bound<'_, PyAny>) -> PyResult<Self> {
let coeffs = parse_min_poly(min_poly)?;
if coeffs.len() < 2 {
return Err(PyValueError::new_err(
"minimal polynomial must have degree at least 1",
));
}
if coeffs.last() != Some(&Rational::new(1, 1)) {
return Err(PyValueError::new_err("minimal polynomial must be monic"));
}
Ok(Self {
field: AlgebraicNumberField::new(RationalDomain, coeffs),
})
}
fn extension_degree(&self) -> usize {
self.field.extension_degree()
}
fn alpha(&self) -> PyAlgebraicElement {
PyAlgebraicElement {
elem: self.field.alpha(),
}
}
#[allow(clippy::wrong_self_convention)]
fn from_base(&self, c: &Bound<'_, PyAny>) -> PyResult<PyAlgebraicElement> {
let r = py_to_rational(c)?;
Ok(PyAlgebraicElement {
elem: self.field.from_base(r),
})
}
fn element(&self, coeffs: &Bound<'_, PyAny>) -> PyResult<PyAlgebraicElement> {
let iter = coeffs.try_iter().map_err(|_| {
PyTypeError::new_err(
"element coefficients must be a list of ints or (num, denom) tuples",
)
})?;
let cs: PyResult<Vec<Rational>> = iter.map(|c| py_to_rational(&c?)).collect();
Ok(PyAlgebraicElement {
elem: self.field.element(cs?),
})
}
fn __repr__(&self) -> String {
format!("AlgebraicExtension(deg={})", self.field.extension_degree())
}
}
#[pymethods]
impl PyAlgebraicElement {
fn coeffs(&self) -> Vec<String> {
self.elem.coeffs().iter().map(|c| c.to_string()).collect()
}
fn __str__(&self) -> String {
format!("{}", self.elem)
}
fn __repr__(&self) -> String {
format!("AlgebraicElement({})", self.elem)
}
}
#[pymethods]
impl PyAlgebraicPolynomial {
#[new]
fn new(field: PyRef<'_, PyAlgebraicExtension>, coeffs: &Bound<'_, PyAny>) -> PyResult<Self> {
let f = &field.field;
let iter = coeffs
.try_iter()
.map_err(|_| PyTypeError::new_err("polynomial coefficients must be a list"))?;
let mut out = Vec::new();
for c in iter {
out.push(py_to_anf_element(f, &c?)?);
}
Ok(Self {
inner: DenseUnivariatePolynomial::from_coeffs(f.clone(), out),
})
}
fn degree(&self) -> Option<usize> {
self.inner.degree()
}
fn len(&self) -> usize {
self.inner.coeffs().len()
}
fn is_zero(&self) -> bool {
self.inner.is_zero()
}
fn coeffs(&self) -> Vec<Vec<String>> {
self.inner
.coeffs()
.iter()
.map(|c| c.coeffs().iter().map(|r| r.to_string()).collect())
.collect()
}
fn __str__(&self) -> String {
format!("{}", PolyDisplay(&self.inner))
}
fn factor(&self) -> Vec<PyAlgebraicFactor> {
self.inner
.factor()
.into_iter()
.map(|(f, m)| PyAlgebraicFactor {
factor: PyAlgebraicPolynomial { inner: f },
multiplicity: m,
})
.collect()
}
}
struct PolyDisplay<'a>(&'a DenseUnivariatePolynomial<AlgebraicNumberField>);
impl std::fmt::Display for PolyDisplay<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let coeffs = self.0.coeffs();
if coeffs.is_empty() {
return write!(f, "0");
}
let mut first = true;
for (i, c) in coeffs.iter().enumerate() {
if c.coeffs().is_empty() {
continue;
}
if !first {
write!(f, " + ")?;
}
first = false;
match i {
0 => write!(f, "({})", c)?,
1 => write!(f, "({})*x", c)?,
_ => write!(f, "({})*x^{}", c, i)?,
}
}
if first {
write!(f, "0")?;
}
Ok(())
}
}