use std::collections::HashMap;
use symplex::prelude::*;
const TOLERANCE: f64 = 1e-6;
#[derive(serde::Deserialize)]
struct FixtureFile {
generated_by: String,
fixture_count: usize,
fixtures: Vec<Fixture>,
}
#[derive(serde::Deserialize)]
#[allow(dead_code)]
struct Fixture {
id: usize,
category: String,
#[serde(default)]
subcategory: Option<String>,
#[serde(default)]
input: Option<String>,
#[serde(default)]
operation: Option<String>,
#[serde(default)]
variable: Option<String>,
#[serde(default)]
params: Option<serde_json::Value>,
#[serde(default)]
sympy_result: Option<String>,
#[serde(default)]
eval_points: Option<Vec<EvalPoint>>,
#[serde(default)]
sympy_roots: Option<Vec<Root>>,
#[serde(default)]
value: Option<NumValue>,
#[serde(default)]
matrix: Option<Vec<Vec<serde_json::Value>>>,
#[serde(default)]
matrix_a: Option<Vec<Vec<serde_json::Value>>>,
#[serde(default)]
matrix_b: Option<Vec<Vec<serde_json::Value>>>,
#[serde(default)]
result_matrix: Option<Vec<Vec<serde_json::Value>>>,
#[serde(default)]
eigenvalues: Option<Vec<Eigenvalue>>,
#[serde(default)]
lower: Option<String>,
#[serde(default)]
upper: Option<String>,
#[serde(default)]
point: Option<String>,
#[serde(default)]
order: Option<u32>,
#[serde(default)]
diff_var: Option<String>,
#[serde(default)]
digits: Option<u32>,
#[serde(default)]
sympy_timeout: bool,
}
#[derive(serde::Deserialize)]
struct EvalPoint {
subs: HashMap<String, f64>,
value: Option<NumValue>,
}
#[derive(serde::Deserialize, Clone, Debug)]
struct NumValue {
re: f64,
#[allow(dead_code)]
im: f64,
}
#[derive(serde::Deserialize)]
#[allow(dead_code)]
struct Root {
symbolic: String,
re: f64,
im: f64,
}
#[derive(serde::Deserialize)]
#[allow(dead_code)]
struct Eigenvalue {
symbolic: String,
re: f64,
im: f64,
multiplicity: Option<u32>,
}
enum Status {
Pass,
Fail(String),
NotImplemented(String),
UnsupportedApi(String),
SkippedOracle(String),
}
#[derive(Default, Clone)]
struct CategoryStats {
passed: usize,
failed: usize,
not_impl: usize,
no_api: usize,
skipped_oracle: usize,
}
impl CategoryStats {
fn total(&self) -> usize {
self.passed + self.failed + self.not_impl + self.no_api + self.skipped_oracle
}
}
fn approx_eq(a: f64, b: f64) -> bool {
if a.is_nan() && b.is_nan() {
return true;
}
if a.is_infinite() && b.is_infinite() {
return a.signum() == b.signum();
}
if a.is_nan() || b.is_nan() || a.is_infinite() || b.is_infinite() {
return false;
}
let diff = (a - b).abs();
let denom = a.abs().max(b.abs()).max(1e-15);
diff < TOLERANCE || diff / denom < TOLERANCE
}
fn eval_at_point(expr: &Ex, ctx: &Context, subs: &HashMap<String, f64>) -> Option<Complex64> {
let mut result = expr.clone();
for (var_name, val) in subs {
let var = ctx.symbol(var_name);
let point_str = if *val == val.floor() && val.abs() < 1e15 {
format!("{}", *val as i64)
} else {
format!("{}", val)
};
let point = symplex::parse::parse(ctx, &point_str).ok()?;
result = result.subs(&var, &point);
}
result.eval_complex64().ok()
}
fn values_match(actual: Complex64, expected: &NumValue) -> bool {
approx_eq(actual.re, expected.re) && approx_eq(actual.im, expected.im)
}
fn parse_point(ctx: &Context, s: &str) -> Option<Ex> {
match s.trim() {
"oo" | "inf" | "Infinity" => Some(ctx.infinity()),
"-oo" | "-inf" | "-Infinity" => Some(ctx.neg_infinity()),
other => symplex::parse::parse(ctx, other).ok(),
}
}
fn fixture_desc(fixture: &Fixture) -> String {
let input = fixture.input.as_deref().unwrap_or("(none)");
let cat = &fixture.category;
let subcat = fixture.subcategory.as_deref().unwrap_or("");
if subcat.is_empty() {
format!("[{}] {}", cat, input)
} else {
format!("[{}:{}] {}", cat, subcat, input)
}
}
fn cat_key(fixture: &Fixture) -> String {
let cat = fixture.category.as_str();
let subcat = fixture.subcategory.as_deref().unwrap_or("");
if subcat.is_empty() {
cat.to_string()
} else {
format!("{}:{}", cat, subcat)
}
}
fn check_eval_points_strict(
result: &Ex,
ctx: &Context,
eval_points: &[EvalPoint],
fixture: &Fixture,
) -> Status {
let mut any_evaluated = false;
let mut mismatches: Vec<String> = Vec::new();
for pt in eval_points {
let expected = match &pt.value {
Some(v) => v,
None => continue,
};
if expected.re.is_nan() || expected.re.is_infinite() {
continue;
}
match eval_at_point(result, ctx, &pt.subs) {
Some(val) => {
any_evaluated = true;
if !values_match(val, expected) {
mismatches.push(format!(
"at {:?}: symplex=({}, {}i), sympy=({}, {}i)",
pt.subs, val.re, val.im, expected.re, expected.im
));
}
}
None => {
any_evaluated = true;
mismatches.push(format!(
"at {:?}: symplex=<eval failed>, sympy=({}, {}i)",
pt.subs, expected.re, expected.im
));
}
}
}
if !any_evaluated {
return Status::NotImplemented(format!("no evaluable points for id={}", fixture.id));
}
if mismatches.is_empty() {
Status::Pass
} else {
Status::Fail(mismatches.join("; "))
}
}
fn check_eval_points_fixture_strict(result: &Ex, ctx: &Context, fixture: &Fixture) -> Status {
match &fixture.eval_points {
Some(points) if !points.is_empty() => {
check_eval_points_strict(result, ctx, points, fixture)
}
_ => {
if let Some(expected) = &fixture.value {
let evaled = result.clone().eval();
match evaled.eval_complex64() {
Ok(val) => {
if values_match(val, expected) {
Status::Pass
} else {
Status::Fail(format!(
"value mismatch: symplex=({}, {}i), sympy=({}, {}i)",
val.re, val.im, expected.re, expected.im
))
}
}
Err(e) => Status::NotImplemented(format!("evalf failed: {}", e)),
}
} else {
Status::NotImplemented("no eval points or expected value in fixture".into())
}
}
}
}
fn check_eval_points_integration(
result: &Ex,
ctx: &Context,
eval_points: &[EvalPoint],
fixture: &Fixture,
) -> Status {
let result_str = format!("{}", result);
if result_str.contains("Integral") || result_str.contains("integral") {
return Status::NotImplemented(format!(
"integration returned unevaluated form: {}",
if result_str.len() > 80 {
format!("{}...", &result_str[..80])
} else {
result_str
}
));
}
let mut pairs: Vec<(f64, f64, HashMap<String, f64>)> = Vec::new();
for pt in eval_points {
let expected = match &pt.value {
Some(v) if !v.re.is_nan() && !v.re.is_infinite() => v.re,
_ => continue,
};
match eval_at_point(result, ctx, &pt.subs) {
Some(val) => pairs.push((val.re, expected, pt.subs.clone())),
None => {
}
}
}
if pairs.is_empty() {
return Status::NotImplemented(format!(
"could not evaluate integration result at any point for id={}",
fixture.id
));
}
if pairs.len() == 1 {
let (got, exp, ref subs) = pairs[0];
if approx_eq(got, exp) {
return Status::Pass;
}
return Status::NotImplemented(format!(
"only 1 eval point; cannot determine constant offset (symplex={}, sympy={}, at {:?})",
got, exp, subs
));
}
let (ref_symplex, ref_sympy, _) = pairs[0];
let mut mismatches: Vec<String> = Vec::new();
for (i, &(sx, sp, ref subs)) in pairs.iter().enumerate().skip(1) {
let symplex_diff = sx - ref_symplex;
let sympy_diff = sp - ref_sympy;
if !approx_eq(symplex_diff, sympy_diff) {
mismatches.push(format!(
"point {}: symplex_diff={}, sympy_diff={} (at {:?})",
i, symplex_diff, sympy_diff, subs
));
}
}
if mismatches.is_empty() {
Status::Pass
} else {
Status::Fail(format!(
"integration difference method mismatch: {}",
mismatches.join("; ")
))
}
}
#[test]
fn cross_validate_against_sympy() {
let json_str = include_str!("../fixtures/sympy_cross_validation.json");
let file: FixtureFile =
serde_json::from_str(json_str).expect("Failed to parse sympy fixture JSON");
println!("\n=== SymPy Cross-Validation ({}) ===\n", file.generated_by);
assert_eq!(file.fixture_count, file.fixtures.len());
let ctx = Context::new();
let mut stats: HashMap<String, CategoryStats> = HashMap::new();
let mut failures: Vec<String> = Vec::new();
let mut not_impls: Vec<String> = Vec::new();
for fixture in &file.fixtures {
let key = cat_key(fixture);
let entry = stats.entry(key.clone()).or_default();
let desc = fixture_desc(fixture);
let result = if fixture.sympy_timeout {
Status::SkippedOracle("SymPy timed out while generating this fixture".into())
} else {
process_fixture(&ctx, fixture)
};
match result {
Status::Pass => {
entry.passed += 1;
}
Status::SkippedOracle(reason) => {
entry.skipped_oracle += 1;
let msg = format!("SKIPPED_ORACLE id={} {} — {}", fixture.id, desc, reason);
not_impls.push(msg);
}
Status::Fail(reason) => {
entry.failed += 1;
let msg = format!("FAIL id={} {} — {}", fixture.id, desc, reason);
failures.push(msg);
}
Status::NotImplemented(reason) => {
entry.not_impl += 1;
let msg = format!("NOT_IMPL id={} {} — {}", fixture.id, desc, reason);
not_impls.push(msg);
}
Status::UnsupportedApi(reason) => {
entry.no_api += 1;
let msg = format!("NO_API id={} {} — {}", fixture.id, desc, reason);
not_impls.push(msg);
}
}
}
println!("\n{}", "=".repeat(80));
println!(
"{:<30} {:>6} {:>6} {:>8} {:>7} {:>9}",
"Category", "Pass", "Fail", "NotImpl", "NoAPI", "SkipOrcl"
);
println!("{}", "-".repeat(80));
let mut sorted_keys: Vec<String> = stats.keys().cloned().collect();
sorted_keys.sort();
let mut total = CategoryStats::default();
for key in &sorted_keys {
let s = &stats[key];
println!(
"{:<30} {:>6} {:>6} {:>8} {:>7} {:>9}",
key, s.passed, s.failed, s.not_impl, s.no_api, s.skipped_oracle
);
total.passed += s.passed;
total.failed += s.failed;
total.not_impl += s.not_impl;
total.no_api += s.no_api;
total.skipped_oracle += s.skipped_oracle;
}
println!("{}", "-".repeat(80));
println!(
"{:<30} {:>6} {:>6} {:>8} {:>7} {:>9}",
"TOTAL", total.passed, total.failed, total.not_impl, total.no_api, total.skipped_oracle
);
println!("{}", "=".repeat(80));
println!(
"Total fixtures processed: {} / {}",
total.total(),
file.fixture_count
);
println!();
if !failures.is_empty() {
println!("=== FAILURES (must fix) ===");
for msg in &failures {
println!("{}", msg);
}
println!();
}
if !not_impls.is_empty() {
println!("=== NOT IMPLEMENTED / NO API (future work) ===");
for msg in ¬_impls {
println!("{}", msg);
}
println!();
}
assert_eq!(
total.failed, 0,
"\n{} cross-validation fixture(s) FAILED.\n\
NotImplemented={} and NoAPI={} are informational (not failures).\n\
See FAILURES list above for details.",
total.failed, total.not_impl, total.no_api
);
}
fn process_fixture(ctx: &Context, fixture: &Fixture) -> Status {
let cat = fixture.category.as_str();
let subcat = fixture.subcategory.as_deref().unwrap_or("");
match cat {
"diff" => process_diff(ctx, fixture, subcat),
"integrate" => process_integrate(ctx, fixture),
"definite_integral" => process_definite_integral(ctx, fixture),
"simplify" => process_simplify(ctx, fixture),
"expand" => process_expand(ctx, fixture, subcat),
"solve" => process_solve(ctx, fixture),
"eval" => process_eval(ctx, fixture),
"series" => process_series(ctx, fixture),
"limit" => process_limit(ctx, fixture),
"matrix" => process_matrix(ctx, fixture, subcat),
"algebra" => process_algebra(ctx, fixture, subcat),
"evalf" => process_evalf(ctx, fixture),
"special_func" => process_special_func(ctx, fixture, subcat),
other => Status::UnsupportedApi(format!("unknown category '{}'", other)),
}
}
fn process_diff(ctx: &Context, fixture: &Fixture, subcat: &str) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let expr = match symplex::parse::parse(ctx, input_str) {
Ok(e) => e,
Err(e) => return Status::NotImplemented(format!("parse error: {}", e)),
};
let result = match subcat {
"higher_order" => {
let var_name = fixture.variable.as_deref().unwrap_or("x");
let var = ctx.symbol(var_name);
let order = fixture.order.unwrap_or(1) as usize;
expr.diff_n(&var, order)
}
"partial" => {
let var_name = fixture
.diff_var
.as_deref()
.or(fixture.variable.as_deref())
.unwrap_or("x");
let var = ctx.symbol(var_name);
expr.diff(&var)
}
_ => {
let var_name = fixture.variable.as_deref().unwrap_or("x");
let var = ctx.symbol(var_name);
expr.diff(&var)
}
};
check_eval_points_fixture_strict(&result, ctx, fixture)
}
fn process_integrate(ctx: &Context, fixture: &Fixture) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let expr = match symplex::parse::parse(ctx, input_str) {
Ok(e) => e,
Err(e) => return Status::NotImplemented(format!("parse error: {}", e)),
};
let var_name = fixture.variable.as_deref().unwrap_or("x");
let var = ctx.symbol(var_name);
let result = expr.integrate(&var);
let result_str = format!("{}", result);
if result_str.contains("Integral") || result_str.contains("integral") {
return Status::NotImplemented(format!(
"integration returned unevaluated form: {}",
if result_str.len() > 100 {
format!("{}...", &result_str[..100])
} else {
result_str
}
));
}
match &fixture.eval_points {
Some(points) if !points.is_empty() => {
check_eval_points_integration(&result, ctx, points, fixture)
}
_ => {
if let Some(expected) = &fixture.value {
let evaled = result.eval();
match evaled.eval_f64() {
Ok(val) if !val.is_nan() => {
if approx_eq(val, expected.re) {
Status::Pass
} else {
Status::Fail(format!(
"integrate value: symplex={}, sympy={}",
val, expected.re
))
}
}
Ok(_) => Status::NotImplemented("integrate evalf returned NaN".into()),
Err(e) => Status::NotImplemented(format!("integrate evalf failed: {}", e)),
}
} else {
Status::NotImplemented("no eval points or expected value".into())
}
}
}
}
fn process_definite_integral(ctx: &Context, fixture: &Fixture) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let expr = match symplex::parse::parse(ctx, input_str) {
Ok(e) => e,
Err(e) => return Status::NotImplemented(format!("parse error: {}", e)),
};
let var_name = fixture.variable.as_deref().unwrap_or("x");
let var = ctx.symbol(var_name);
let lower_str = match &fixture.lower {
Some(s) => s,
None => return Status::NotImplemented("no lower bound in fixture".into()),
};
let upper_str = match &fixture.upper {
Some(s) => s,
None => return Status::NotImplemented("no upper bound in fixture".into()),
};
let lower = match parse_point(ctx, lower_str) {
Some(e) => e,
None => return Status::NotImplemented(format!("can't parse lower bound: {}", lower_str)),
};
let upper = match parse_point(ctx, upper_str) {
Some(e) => e,
None => return Status::NotImplemented(format!("can't parse upper bound: {}", upper_str)),
};
let result = expr.integrate_definite(&var, &lower, &upper);
let result_str = format!("{}", result);
if result_str.contains("Integral") || result_str.contains("integral") {
return Status::NotImplemented(format!(
"definite integral returned unevaluated form: {}",
if result_str.len() > 100 {
format!("{}...", &result_str[..100])
} else {
result_str
}
));
}
let result = result.eval().simplify();
if let Some(expected) = &fixture.value {
match result.eval_f64() {
Ok(val) if !val.is_nan() => {
if approx_eq(val, expected.re) {
Status::Pass
} else {
Status::Fail(format!(
"definite_integral: symplex={}, sympy={}",
val, expected.re
))
}
}
Ok(_) => Status::NotImplemented("definite integral evalf returned NaN".into()),
Err(e) => Status::NotImplemented(format!("definite integral evalf failed: {}", e)),
}
} else {
Status::NotImplemented("no expected value in fixture".into())
}
}
fn process_simplify(ctx: &Context, fixture: &Fixture) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let expr = match symplex::parse::parse(ctx, input_str) {
Ok(e) => e,
Err(e) => return Status::NotImplemented(format!("parse error: {}", e)),
};
let result = expr.simplify();
check_eval_points_fixture_strict(&result, ctx, fixture)
}
fn process_expand(ctx: &Context, fixture: &Fixture, subcat: &str) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let expr = match symplex::parse::parse(ctx, input_str) {
Ok(e) => e,
Err(e) => return Status::NotImplemented(format!("parse error: {}", e)),
};
let result = match subcat {
"trig" => expr.expand_trig(),
_ => expr.expand(),
};
check_eval_points_fixture_strict(&result, ctx, fixture)
}
fn process_solve(ctx: &Context, fixture: &Fixture) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let expr = match symplex::parse::parse(ctx, input_str) {
Ok(e) => e,
Err(e) => return Status::NotImplemented(format!("parse error: {}", e)),
};
let var_name = fixture.variable.as_deref().unwrap_or("x");
let var = ctx.symbol(var_name);
let sympy_roots = match &fixture.sympy_roots {
Some(r) => r,
None => return Status::NotImplemented("no sympy_roots in fixture".into()),
};
let roots = expr.solve_or_empty(&var);
let mut bad_roots: Vec<String> = Vec::new();
for root in &roots {
let residual = expr.subs(&var, root).eval();
if let Ok(r) = residual.eval_f64()
&& !r.is_nan()
&& r.abs() > TOLERANCE
{
bad_roots.push(format!("root {} has residual {} (should be ~0)", root, r));
}
}
let real_sympy_roots: Vec<_> = sympy_roots.iter().filter(|r| r.im.abs() < 1e-10).collect();
if roots.len() < real_sympy_roots.len() {
println!(
" INFO [solve] id={} '{}': symplex found {}/{} real roots",
fixture.id,
input_str,
roots.len(),
real_sympy_roots.len()
);
}
if bad_roots.is_empty() {
Status::Pass
} else {
Status::Fail(format!("bad roots: {}", bad_roots.join("; ")))
}
}
fn process_eval(ctx: &Context, fixture: &Fixture) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let expr = match symplex::parse::parse(ctx, input_str) {
Ok(e) => e,
Err(e) => return Status::NotImplemented(format!("parse error: {}", e)),
};
let result = expr.eval();
if let Some(expected) = &fixture.value {
match result.eval_f64() {
Ok(val) if !val.is_nan() => {
if approx_eq(val, expected.re) {
Status::Pass
} else {
Status::Fail(format!("eval: symplex={}, sympy={}", val, expected.re))
}
}
Ok(_) => Status::NotImplemented("eval evalf returned NaN".into()),
Err(e) => Status::NotImplemented(format!("eval evalf failed: {}", e)),
}
} else {
Status::NotImplemented("no expected value in fixture".into())
}
}
fn process_series(ctx: &Context, fixture: &Fixture) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let expr = match symplex::parse::parse(ctx, input_str) {
Ok(e) => e,
Err(e) => return Status::NotImplemented(format!("parse error: {}", e)),
};
let var_name = fixture.variable.as_deref().unwrap_or("x");
let var = ctx.symbol(var_name);
let order = fixture.order.unwrap_or(6);
let point_str = fixture.point.as_deref().unwrap_or("0");
let result = if point_str == "0" {
let s = expr.maclaurin(&var, order);
if s.has_unevaluated() {
return Status::NotImplemented("maclaurin returned unevaluated form".into());
}
s
} else {
let point = match parse_point(ctx, point_str) {
Some(p) => p,
None => {
return Status::NotImplemented(format!("can't parse series point: {}", point_str));
}
};
let s = expr.series(&var, &point, order);
if s.has_unevaluated() {
return Status::NotImplemented("series returned unevaluated form".into());
}
s
};
check_eval_points_fixture_strict(&result, ctx, fixture)
}
fn process_limit(ctx: &Context, fixture: &Fixture) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let expr = match symplex::parse::parse(ctx, input_str) {
Ok(e) => e,
Err(e) => return Status::NotImplemented(format!("parse error: {}", e)),
};
let var_name = fixture.variable.as_deref().unwrap_or("x");
let var = ctx.symbol(var_name);
let point_str = match &fixture.point {
Some(s) => s.as_str(),
None => return Status::NotImplemented("no limit point in fixture".into()),
};
let point = match parse_point(ctx, point_str) {
Some(p) => p,
None => return Status::NotImplemented(format!("can't parse limit point: {}", point_str)),
};
let expected = match &fixture.value {
Some(v) => v,
None => return Status::NotImplemented("no expected value for limit".into()),
};
let limit_result = expr.limit(&var, &point);
if limit_result.has_unevaluated() {
return Status::NotImplemented(format!(
"limit returned unevaluated form (input: {}, point: {})",
input_str, point_str
));
}
let limit_result = limit_result.eval();
match limit_result.eval_f64() {
Ok(val) => {
if val.is_nan() {
Status::NotImplemented(format!(
"limit result evaluated to NaN (input: {}, point: {})",
input_str, point_str
))
} else if approx_eq(val, expected.re) {
Status::Pass
} else {
Status::Fail(format!(
"lim({}, {}->{}): symplex={}, sympy={}",
input_str, var_name, point_str, val, expected.re
))
}
}
Err(e) => Status::NotImplemented(format!(
"limit result couldn't be evaluated: {} (input: {}, point: {})",
e, input_str, point_str
)),
}
}
fn parse_matrix_from_json(
ctx: &Context,
rows: &[Vec<serde_json::Value>],
) -> Option<symplex::matrix::Matrix> {
let nrows = rows.len();
if nrows == 0 {
return None;
}
let ncols = rows[0].len();
if ncols == 0 {
return None;
}
let mut mat_rows: Vec<Vec<Ex>> = Vec::with_capacity(nrows);
for row in rows {
if row.len() != ncols {
return None;
}
let mut mat_row: Vec<Ex> = Vec::with_capacity(ncols);
for val in row {
let expr = match val {
serde_json::Value::Number(n) => {
if let Some(i) = n.as_i64() {
ctx.int(i)
} else {
let f = n.as_f64()?;
let s = format!("{}", f);
symplex::parse::parse(ctx, &s).ok()?
}
}
serde_json::Value::String(s) => symplex::parse::parse(ctx, s).ok()?,
_ => return None,
};
mat_row.push(expr);
}
mat_rows.push(mat_row);
}
Some(symplex::matrix::Matrix::new(mat_rows).unwrap())
}
fn process_matrix(ctx: &Context, fixture: &Fixture, subcat: &str) -> Status {
match subcat {
"det" => process_matrix_det(ctx, fixture),
"trace" => process_matrix_trace(ctx, fixture),
"inverse" => process_matrix_inverse(ctx, fixture),
"eigenvalue" => process_matrix_eigenvalue(ctx, fixture),
"multiply" => process_matrix_multiply(ctx, fixture),
other => Status::UnsupportedApi(format!("matrix subcategory '{}' not supported", other)),
}
}
fn process_matrix_det(ctx: &Context, fixture: &Fixture) -> Status {
let rows = match &fixture.matrix {
Some(m) => m,
None => return Status::NotImplemented("no matrix in fixture".into()),
};
let mat = match parse_matrix_from_json(ctx, rows) {
Some(m) => m,
None => return Status::NotImplemented("can't parse matrix".into()),
};
let det = mat.det().unwrap().eval();
if let Some(expected) = &fixture.value {
match det.eval_f64() {
Ok(val) if !val.is_nan() => {
if approx_eq(val, expected.re) {
Status::Pass
} else {
Status::Fail(format!(
"matrix det: symplex={}, sympy={}",
val, expected.re
))
}
}
Ok(_) => Status::NotImplemented("det evalf returned NaN".into()),
Err(e) => Status::NotImplemented(format!("det evalf failed: {}", e)),
}
} else {
Status::NotImplemented("no expected value for det".into())
}
}
fn process_matrix_trace(ctx: &Context, fixture: &Fixture) -> Status {
let rows = match &fixture.matrix {
Some(m) => m,
None => return Status::NotImplemented("no matrix in fixture".into()),
};
let mat = match parse_matrix_from_json(ctx, rows) {
Some(m) => m,
None => return Status::NotImplemented("can't parse matrix".into()),
};
let tr = mat.trace().unwrap().eval();
if let Some(expected) = &fixture.value {
match tr.eval_f64() {
Ok(val) if !val.is_nan() => {
if approx_eq(val, expected.re) {
Status::Pass
} else {
Status::Fail(format!(
"matrix trace: symplex={}, sympy={}",
val, expected.re
))
}
}
Ok(_) => Status::NotImplemented("trace evalf returned NaN".into()),
Err(e) => Status::NotImplemented(format!("trace evalf failed: {}", e)),
}
} else {
Status::NotImplemented("no expected value for trace".into())
}
}
fn process_matrix_multiply(ctx: &Context, fixture: &Fixture) -> Status {
let rows_a = match &fixture.matrix_a {
Some(m) => m,
None => return Status::NotImplemented("no matrix_a in fixture".into()),
};
let rows_b = match &fixture.matrix_b {
Some(m) => m,
None => return Status::NotImplemented("no matrix_b in fixture".into()),
};
let mat_a = match parse_matrix_from_json(ctx, rows_a) {
Some(m) => m,
None => return Status::NotImplemented("can't parse matrix_a".into()),
};
let mat_b = match parse_matrix_from_json(ctx, rows_b) {
Some(m) => m,
None => return Status::NotImplemented("can't parse matrix_b".into()),
};
let product = mat_a.matmul(&mat_b).unwrap();
if let Some(expected_rows) = &fixture.result_matrix {
let expected_mat = match parse_matrix_from_json(ctx, expected_rows) {
Some(m) => m,
None => return Status::NotImplemented("can't parse expected result_matrix".into()),
};
let (nr, nc) = product.shape();
let (enr, enc) = expected_mat.shape();
if nr != enr || nc != enc {
return Status::Fail(format!(
"multiply shape mismatch: ({},{}) vs ({},{})",
nr, nc, enr, enc
));
}
for i in 0..nr {
for j in 0..nc {
let got = product.get(i, j).eval();
let exp = expected_mat.get(i, j).eval();
match (got.eval_f64(), exp.eval_f64()) {
(Ok(g), Ok(e)) if !g.is_nan() && !e.is_nan() => {
if !approx_eq(g, e) {
return Status::Fail(format!(
"multiply[{},{}]: symplex={}, sympy={}",
i, j, g, e
));
}
}
_ => {
return Status::NotImplemented(format!(
"can't evalf multiply[{},{}]",
i, j
));
}
}
}
}
Status::Pass
} else {
Status::NotImplemented("no expected result_matrix in fixture".into())
}
}
fn process_matrix_inverse(ctx: &Context, fixture: &Fixture) -> Status {
let rows = match &fixture.matrix {
Some(m) => m,
None => return Status::NotImplemented("no matrix in fixture".into()),
};
let mat = match parse_matrix_from_json(ctx, rows) {
Some(m) => m,
None => return Status::NotImplemented("can't parse matrix".into()),
};
let inv = match mat.inv() {
Ok(m) => m,
Err(_) => return Status::NotImplemented("matrix is singular (inv returned Err)".into()),
};
if let Some(expected_rows) = &fixture.result_matrix {
let expected_mat = match parse_matrix_from_json(ctx, expected_rows) {
Some(m) => m,
None => return Status::NotImplemented("can't parse expected result_matrix".into()),
};
let (nr, nc) = inv.shape();
let (enr, enc) = expected_mat.shape();
if nr != enr || nc != enc {
return Status::Fail(format!(
"inverse shape mismatch: ({},{}) vs ({},{})",
nr, nc, enr, enc
));
}
for i in 0..nr {
for j in 0..nc {
let got = inv.get(i, j).eval();
let exp = expected_mat.get(i, j).eval();
match (got.eval_f64(), exp.eval_f64()) {
(Ok(g), Ok(e)) if !g.is_nan() && !e.is_nan() => {
if !approx_eq(g, e) {
return Status::Fail(format!(
"inverse[{},{}]: symplex={}, sympy={}",
i, j, g, e
));
}
}
_ => {
return Status::NotImplemented(format!("can't evalf inverse[{},{}]", i, j));
}
}
}
}
Status::Pass
} else {
Status::NotImplemented("no expected result_matrix for inverse".into())
}
}
fn process_matrix_eigenvalue(ctx: &Context, fixture: &Fixture) -> Status {
let rows = match &fixture.matrix {
Some(m) => m,
None => return Status::NotImplemented("no matrix in fixture".into()),
};
let mat = match parse_matrix_from_json(ctx, rows) {
Some(m) => m,
None => return Status::NotImplemented("can't parse matrix".into()),
};
let expected_eigs = match &fixture.eigenvalues {
Some(e) => e,
None => return Status::NotImplemented("no eigenvalues in fixture".into()),
};
let computed = mat.eigenvals().unwrap();
if computed.is_empty() {
return Status::NotImplemented("eigenvalue solver returned empty".into());
}
let mut computed_vals: Vec<f64> = Vec::new();
for ev in &computed {
match ev.eval_f64() {
Ok(v) if !v.is_nan() => computed_vals.push(v),
_ => {
return Status::NotImplemented(format!("can't evalf eigenvalue: {}", ev));
}
}
}
computed_vals.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let mut expected_vals: Vec<f64> = Vec::new();
for ev in expected_eigs {
let mult = ev.multiplicity.unwrap_or(1);
for _ in 0..mult {
expected_vals.push(ev.re);
}
}
expected_vals.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
if computed_vals.len() != expected_vals.len() {
return Status::Fail(format!(
"eigenvalue count mismatch: symplex={}, sympy={}",
computed_vals.len(),
expected_vals.len()
));
}
for (i, (got, exp)) in computed_vals.iter().zip(expected_vals.iter()).enumerate() {
if !approx_eq(*got, *exp) {
return Status::Fail(format!("eigenvalue[{}]: symplex={}, sympy={}", i, got, exp));
}
}
Status::Pass
}
fn process_algebra(ctx: &Context, fixture: &Fixture, subcat: &str) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let expr = match symplex::parse::parse(ctx, input_str) {
Ok(e) => e,
Err(e) => return Status::NotImplemented(format!("parse error: {}", e)),
};
let var_name = fixture.variable.as_deref().unwrap_or("x");
let var = ctx.symbol(var_name);
let result = match subcat {
"factor" => expr.factor(&var),
"collect" => expr.collect(&var),
"together" => expr.together(),
"cancel" => expr.cancel(&var),
"apart" => expr.partial_fractions(&var),
other => {
return Status::UnsupportedApi(format!(
"algebra subcategory '{}' not supported",
other
));
}
};
check_eval_points_fixture_strict(&result, ctx, fixture)
}
fn process_evalf(ctx: &Context, fixture: &Fixture) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let expr = match symplex::parse::parse(ctx, input_str) {
Ok(e) => e,
Err(e) => return Status::NotImplemented(format!("parse error: {}", e)),
};
let digits = fixture.digits.unwrap_or(15);
let sympy_result = match &fixture.sympy_result {
Some(s) => s,
None => return Status::NotImplemented("no sympy_result for evalf".into()),
};
match expr.eval_decimal(digits) {
Ok(result_str) => {
let symplex_val: f64 = match result_str.parse() {
Ok(v) => v,
Err(_) => {
return Status::NotImplemented(format!(
"can't parse symplex evalf result as f64: {}",
result_str
));
}
};
let sympy_val: f64 = match sympy_result.parse() {
Ok(v) => v,
Err(_) => {
return Status::NotImplemented(format!(
"can't parse sympy result as f64: {}",
sympy_result
));
}
};
if symplex_val.is_nan() || sympy_val.is_nan() {
return Status::NotImplemented("NaN in evalf comparison".into());
}
let evalf_tol = 10.0_f64.powi(-(digits.min(15) as i32) + 2);
let diff = (symplex_val - sympy_val).abs();
let rel = diff / sympy_val.abs().max(1e-30);
if diff < evalf_tol || rel < evalf_tol {
Status::Pass
} else {
Status::Fail(format!(
"evalf({} digits): symplex='{}', sympy='{}'",
digits, result_str, sympy_result
))
}
}
Err(e) => Status::NotImplemented(format!("evalf({}) failed: {}", digits, e)),
}
}
fn process_special_func(ctx: &Context, fixture: &Fixture, subcat: &str) -> Status {
match subcat {
"factorial" => process_factorial(ctx, fixture),
other => Status::UnsupportedApi(format!(
"special_func subcategory '{}' not supported",
other
)),
}
}
fn process_factorial(ctx: &Context, fixture: &Fixture) -> Status {
let input_str = match &fixture.input {
Some(s) => s,
None => return Status::NotImplemented("no input field in fixture".into()),
};
let n: i64 = match input_str.parse() {
Ok(n) => n,
Err(_) => {
return Status::NotImplemented(format!(
"can't parse factorial input as integer: {}",
input_str
));
}
};
let n_expr = ctx.int(n);
let result = n_expr.factorial().eval();
if let Some(expected) = &fixture.value {
match result.eval_f64() {
Ok(val) if !val.is_nan() => {
if approx_eq(val, expected.re) {
Status::Pass
} else {
Status::Fail(format!(
"factorial({}): symplex={}, sympy={}",
n, val, expected.re
))
}
}
Ok(_) => Status::NotImplemented("factorial evalf returned NaN".into()),
Err(e) => Status::NotImplemented(format!("factorial evalf failed: {}", e)),
}
} else {
Status::NotImplemented("no expected value for factorial".into())
}
}