use serde::{Deserialize, Serialize};
use wasm_bindgen::prelude::*;
use phop_core::{AnySolution, Config, DataSet, Discoverer, Solution};
use scirs2_core::ndarray::{Array1, Array2};
#[derive(Debug, Deserialize)]
struct DataJson {
x: Vec<Vec<f64>>,
y: Vec<f64>,
}
#[derive(Debug, Default, Deserialize)]
struct ConfigJson {
population: Option<usize>,
max_depth: Option<usize>,
max_epochs: Option<usize>,
learning_rate: Option<f64>,
seed: Option<u64>,
top_k: Option<usize>,
analyze: Option<bool>,
}
#[derive(Debug, Serialize)]
struct SolutionJson {
latex: String,
pretty: String,
mse: f64,
complexity: usize,
#[serde(skip_serializing_if = "Option::is_none")]
derivative: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
antiderivative: Option<String>,
}
#[derive(Debug, Serialize)]
struct ResultJson {
solutions: Vec<SolutionJson>,
}
#[derive(Debug, Serialize)]
struct ErrorJson {
error: String,
}
#[wasm_bindgen]
pub fn set_panic_hook() {
console_error_panic_hook::set_once();
}
#[wasm_bindgen]
pub fn discover_json(data_json: &str, config_json: &str) -> String {
match run(data_json, config_json) {
Ok(result) => serde_json::to_string(&result)
.unwrap_or_else(|e| format!("{{\"error\":\"serialization failed: {e}\"}}")),
Err(msg) => serde_json::to_string(&ErrorJson { error: msg })
.unwrap_or_else(|_| "{\"error\":\"unknown error\"}".to_string()),
}
}
fn run(data_json: &str, config_json: &str) -> Result<ResultJson, String> {
let data: DataJson = serde_json::from_str(data_json).map_err(|e| e.to_string())?;
let cfg_json: ConfigJson = if config_json.trim().is_empty() {
ConfigJson::default()
} else {
serde_json::from_str(config_json).map_err(|e| e.to_string())?
};
let rows = data.x.len();
if rows == 0 {
return Err("`x` must contain at least one row".to_string());
}
let cols = data.x[0].len();
if data.x.iter().any(|r| r.len() != cols) {
return Err("all rows of `x` must have the same length".to_string());
}
let flat: Vec<f64> = data.x.into_iter().flatten().collect();
let x_arr = Array2::from_shape_vec((rows, cols), flat).map_err(|e| e.to_string())?;
let y_arr = Array1::from(data.y);
let ds = DataSet::from_arrays(x_arr, y_arr).map_err(|e| e.to_string())?;
let mut config = Config::default();
if let Some(v) = cfg_json.population {
config = config.population(v);
}
if let Some(v) = cfg_json.max_depth {
config = config.max_depth(v);
}
if let Some(v) = cfg_json.max_epochs {
config = config.max_epochs(v);
}
if let Some(v) = cfg_json.learning_rate {
config = config.learning_rate(v);
}
if let Some(v) = cfg_json.seed {
config = config.seed(v);
}
let top_k = cfg_json.top_k.unwrap_or(config.top_k);
if let Some(v) = cfg_json.top_k {
config = config.top_k(v);
}
let front = Discoverer::new(config)
.fit(&ds)
.map_err(|e| e.to_string())?;
let analyze = cfg_json.analyze.unwrap_or(false);
let solutions = front
.pareto_top(top_k)
.into_iter()
.map(|s| {
let (derivative, antiderivative) = if analyze {
let a = s.analyze(0, 4);
(Some(a.derivative), a.antiderivative)
} else {
(None, None)
};
SolutionJson {
latex: s.latex(),
pretty: s.pretty(),
mse: s.mse,
complexity: s.complexity,
derivative,
antiderivative,
}
})
.collect();
Ok(ResultJson { solutions })
}
type DimVec = Vec<i32>;
#[derive(Debug, Default, Deserialize)]
struct VerifyConfigJson {
population: Option<usize>,
max_depth: Option<usize>,
max_epochs: Option<usize>,
learning_rate: Option<f64>,
seed: Option<u64>,
top_k: Option<usize>,
method: Option<String>,
analyze: Option<bool>,
canonical: Option<bool>,
certify: Option<bool>,
prove_no_root: Option<bool>,
units: Option<Vec<DimVec>>,
target_model: Option<String>,
}
#[derive(Debug, Serialize)]
struct VerifiedSolutionJson {
source: String,
latex: String,
pretty: String,
mse: f64,
r2: f64,
complexity: usize,
symbolic: bool,
#[serde(skip_serializing_if = "Option::is_none")]
derivative: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
antiderivative: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
canonical_latex: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
certified_range: Option<[f64; 2]>,
#[serde(skip_serializing_if = "Option::is_none")]
certified_root: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
proven_no_root: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
equivalent_to_target: Option<String>,
}
#[derive(Debug, Serialize)]
struct Capabilities {
analyze: bool,
certify: bool,
canonical: bool,
smt: bool,
}
#[derive(Debug, Serialize)]
struct VerifiedResultJson {
solutions: Vec<VerifiedSolutionJson>,
#[serde(skip_serializing_if = "Option::is_none")]
pi_groups: Option<Vec<DimVec>>,
capabilities: Capabilities,
}
#[wasm_bindgen]
pub fn capabilities() -> String {
let caps = Capabilities {
analyze: true,
certify: true,
canonical: cfg!(feature = "egraph"),
smt: cfg!(feature = "smt"),
};
serde_json::to_string(&caps).unwrap_or_else(|_| "{}".to_string())
}
#[wasm_bindgen]
pub fn discover_and_verify(data_json: &str, config_json: &str) -> String {
match verify_run(data_json, config_json) {
Ok(result) => serde_json::to_string(&result)
.unwrap_or_else(|e| format!("{{\"error\":\"serialization failed: {e}\"}}")),
Err(msg) => serde_json::to_string(&ErrorJson { error: msg })
.unwrap_or_else(|_| "{\"error\":\"unknown error\"}".to_string()),
}
}
fn r2_of(pred: &Array1<f64>, y: &Array1<f64>) -> f64 {
let n = y.len().max(1) as f64;
let mean = y.sum() / n;
let (mut sr, mut st) = (0.0, 0.0);
for (p, t) in pred.iter().zip(y.iter()) {
sr += (t - p) * (t - p);
st += (t - mean) * (t - mean);
}
if st == 0.0 {
f64::NAN
} else {
1.0 - sr / st
}
}
fn data_box(ds: &DataSet) -> Vec<(f64, f64)> {
(0..ds.n_vars())
.map(|j| {
let col = ds.x.column(j);
(
col.iter().copied().fold(f64::INFINITY, f64::min),
col.iter().copied().fold(f64::NEG_INFINITY, f64::max),
)
})
.collect()
}
fn canonical_of(_s: &Solution) -> Option<String> {
#[cfg(feature = "egraph")]
{
Some(_s.latex_egraph())
}
#[cfg(not(feature = "egraph"))]
{
None
}
}
fn prove_no_root_of(_s: &Solution, _bounds: &[(f64, f64)]) -> Option<String> {
#[cfg(feature = "smt")]
{
Some(format!("{:?}", _s.prove_no_root(_bounds)))
}
#[cfg(not(feature = "smt"))]
{
None
}
}
fn prove_equiv_of(_s: &Solution, _target: &Solution, _bounds: &[(f64, f64)]) -> Option<String> {
#[cfg(feature = "smt")]
{
Some(format!("{:?}", _s.prove_equivalent(_target, _bounds)))
}
#[cfg(not(feature = "smt"))]
{
None
}
}
fn verify_run(data_json: &str, config_json: &str) -> Result<VerifiedResultJson, String> {
let data: DataJson = serde_json::from_str(data_json).map_err(|e| e.to_string())?;
let cfg_json: VerifyConfigJson = if config_json.trim().is_empty() {
VerifyConfigJson::default()
} else {
serde_json::from_str(config_json).map_err(|e| e.to_string())?
};
let rows = data.x.len();
if rows == 0 {
return Err("`x` must contain at least one row".to_string());
}
let cols = data.x[0].len();
if data.x.iter().any(|r| r.len() != cols) {
return Err("all rows of `x` must have the same length".to_string());
}
let flat: Vec<f64> = data.x.into_iter().flatten().collect();
let x_arr = Array2::from_shape_vec((rows, cols), flat).map_err(|e| e.to_string())?;
let y_arr = Array1::from(data.y);
let mut ds = DataSet::from_arrays(x_arr, y_arr).map_err(|e| e.to_string())?;
let mut pi_groups: Option<Vec<DimVec>> = None;
if let Some(units) = &cfg_json.units {
let dims: Vec<phop_core::Dimension> = units
.iter()
.map(|v| {
<[i32; 7]>::try_from(v.clone())
.map_err(|_| "each `units` entry needs exactly 7 integer exponents".to_string())
})
.collect::<Result<_, _>>()?;
let (reduced, groups) = ds.to_dimensionless(&dims).map_err(|e| e.to_string())?;
ds = reduced;
pi_groups = Some(groups);
}
let mut config = Config::default();
if let Some(v) = cfg_json.population {
config = config.population(v);
}
if let Some(v) = cfg_json.max_depth {
config = config.max_depth(v);
}
if let Some(v) = cfg_json.max_epochs {
config = config.max_epochs(v);
}
if let Some(v) = cfg_json.learning_rate {
config = config.learning_rate(v);
}
if let Some(v) = cfg_json.seed {
config = config.seed(v);
}
let top_k = cfg_json.top_k.unwrap_or(config.top_k);
config = config.top_k(top_k);
let max_internal = cfg_json.max_depth.unwrap_or(3).clamp(1, 5);
const WASM_CAND_CAP: usize = 800;
let front: Vec<AnySolution> = match cfg_json.method.as_deref() {
Some("rich") => {
phop_core::discover_affine_pareto(&ds.x, &ds.y, max_internal, WASM_CAND_CAP)
.into_iter()
.map(AnySolution::Affine)
.collect()
}
Some("auto") => phop_core::discover_auto_all(&ds, &config, max_internal, WASM_CAND_CAP)
.map_err(|e| e.to_string())?,
_ => Discoverer::new(config)
.fit(&ds)
.map_err(|e| e.to_string())?
.solutions
.into_iter()
.map(AnySolution::Eml)
.collect(),
};
let mut scored: Vec<(AnySolution, f64)> = front
.into_iter()
.map(|s| {
let r2 = s
.predict(&ds.x)
.ok()
.map(|p| r2_of(&p, &ds.y))
.unwrap_or(f64::NAN);
(s, r2)
})
.collect();
const MIN_DISPLAY_R2: f64 = 0.5;
if scored
.iter()
.any(|(_, r2)| r2.is_finite() && *r2 >= MIN_DISPLAY_R2)
{
scored.retain(|(_, r2)| r2.is_finite() && *r2 >= MIN_DISPLAY_R2);
}
scored.sort_by(|a, b| {
let (ra, rb) = ((a.1 * 1e6).round(), (b.1 * 1e6).round());
rb.partial_cmp(&ra)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.0.complexity().cmp(&b.0.complexity()))
.then(
a.0.mse()
.partial_cmp(&b.0.mse())
.unwrap_or(std::cmp::Ordering::Equal),
)
});
scored.truncate(top_k);
let analyze = cfg_json.analyze.unwrap_or(false);
let canonical = cfg_json.canonical.unwrap_or(false);
let certify = cfg_json.certify.unwrap_or(false);
let prove_no_root = cfg_json.prove_no_root.unwrap_or(false);
let bounds = data_box(&ds);
let target: Option<Solution> = match &cfg_json.target_model {
Some(m) => Some(Solution::from_model_json(m).map_err(|e| e.to_string())?),
None => None,
};
let solutions = scored
.iter()
.map(|(s, r2)| {
let r2 = *r2;
let mut out = VerifiedSolutionJson {
source: s.source().to_string(),
latex: s.latex(),
pretty: s.expr(),
mse: s.mse(),
r2,
complexity: s.complexity(),
symbolic: s.is_symbolic(),
derivative: None,
antiderivative: None,
canonical_latex: None,
certified_range: None,
certified_root: None,
proven_no_root: None,
equivalent_to_target: None,
};
if let Some(e) = s.as_eml() {
if analyze {
let a = e.analyze(0, 4);
out.derivative = Some(a.derivative);
out.antiderivative = a.antiderivative;
}
if canonical {
out.canonical_latex = canonical_of(e);
}
if certify {
let (lo, hi) = e.certified_range(&bounds);
out.certified_range = Some([lo, hi]);
if let Some(&(x0lo, x0hi)) = bounds.first() {
let others: Vec<f64> =
bounds.iter().skip(1).map(|(a, b)| 0.5 * (a + b)).collect();
out.certified_root = Some(match e.certified_root(0, &others, x0lo, x0hi) {
Ok(cert) => format!("{cert:?}"),
Err(err) => format!("error: {err}"),
});
}
}
if prove_no_root {
out.proven_no_root = prove_no_root_of(e, &bounds);
}
if let Some(t) = &target {
out.equivalent_to_target = prove_equiv_of(e, t, &bounds);
}
}
out
})
.collect();
Ok(VerifiedResultJson {
solutions,
pi_groups,
capabilities: Capabilities {
analyze: true,
certify: true,
canonical: cfg!(feature = "egraph"),
smt: cfg!(feature = "smt"),
},
})
}
#[cfg(test)]
mod tests {
use super::*;
fn exp_dataset_json() -> String {
let xs: Vec<f64> = (0..21).map(|i| i as f64 * 0.1).collect();
let x: Vec<Vec<f64>> = xs.iter().map(|&v| vec![v]).collect();
let y: Vec<f64> = xs.iter().map(|&v| v.exp()).collect();
serde_json::to_string(&serde_json::json!({ "x": x, "y": y })).expect("dataset serializes")
}
#[test]
fn discover_json_recovers_exp() {
let data = exp_dataset_json();
let cfg = serde_json::to_string(&serde_json::json!({
"max_epochs": 100,
"max_depth": 1,
"seed": 0,
"top_k": 3,
}))
.expect("config serializes");
let out = discover_json(&data, &cfg);
assert!(
!out.contains("\"error\""),
"discovery should not return an error envelope, got: {out}"
);
let parsed: serde_json::Value =
serde_json::from_str(&out).expect("result must be valid JSON");
let obj = parsed.as_object().expect("result must be a JSON object");
let solutions = obj
.get("solutions")
.and_then(|v| v.as_array())
.expect("result must contain a `solutions` array");
assert!(
!solutions.is_empty(),
"the Pareto front must contain at least one solution"
);
let mse = solutions[0]
.get("mse")
.and_then(|v| v.as_f64())
.expect("first solution must report a numeric `mse`");
assert!(
mse < 0.5,
"best solution MSE should be small for y = exp(x), got {mse}"
);
}
#[test]
fn discover_json_includes_analysis_when_requested() {
let data = exp_dataset_json();
let cfg = serde_json::to_string(&serde_json::json!({
"max_epochs": 100, "max_depth": 1, "seed": 0, "top_k": 1, "analyze": true,
}))
.expect("config serializes");
let out = discover_json(&data, &cfg);
let parsed: serde_json::Value = serde_json::from_str(&out).expect("valid JSON");
let first = parsed["solutions"][0]
.as_object()
.expect("a first solution object");
assert!(
first.contains_key("derivative"),
"analyze=true must add a `derivative` field, got: {out}"
);
}
#[test]
fn verify_pipeline_runs_and_certifies() {
let data = exp_dataset_json();
let cfg = serde_json::json!({
"method": "enumerate", "max_epochs": 120, "max_depth": 1, "seed": 0, "top_k": 1,
"analyze": true, "certify": true,
})
.to_string();
let out = discover_and_verify(&data, &cfg);
assert!(!out.contains("\"error\""), "pipeline errored: {out}");
let v: serde_json::Value = serde_json::from_str(&out).expect("valid JSON");
let s0 = &v["solutions"][0];
assert!(
s0["certified_range"].is_array(),
"certify must add a certified_range, got: {out}"
);
assert!(
s0["derivative"].is_string(),
"analyze must add a derivative"
);
let r2 = s0["r2"].as_f64().expect("r2 present");
assert!(r2 > 0.9, "exp fit should be accurate, r2 = {r2}");
assert!(v["capabilities"]["certify"].as_bool().unwrap_or(false));
}
#[cfg(feature = "smt")]
#[test]
fn verify_pipeline_proves_no_root_with_smt() {
let data = exp_dataset_json();
let cfg = serde_json::json!({
"method": "enumerate", "max_epochs": 120, "max_depth": 1, "seed": 0, "top_k": 1,
"prove_no_root": true,
})
.to_string();
let out = discover_and_verify(&data, &cfg);
let v: serde_json::Value = serde_json::from_str(&out).expect("valid JSON");
assert!(
v["solutions"][0]["proven_no_root"].is_string(),
"smt build must add a proven_no_root verdict, got: {out}"
);
assert!(v["capabilities"]["smt"].as_bool().unwrap_or(false));
}
#[cfg(feature = "smt")]
#[test]
fn verify_proves_equivalence_after_integer_snap() {
let data = exp_dataset_json();
let cfg = serde_json::json!({
"method": "enumerate", "max_depth": 3, "max_epochs": 200, "seed": 0, "top_k": 1,
"target_model": "{\"root\":{\"Eml\":{\"left\":{\"Var\":0},\"right\":\"One\"}},\"num_vars\":1}",
})
.to_string();
let out = discover_and_verify(&data, &cfg);
let v: serde_json::Value = serde_json::from_str(&out).expect("valid JSON");
let s0 = &v["solutions"][0];
assert_eq!(
s0["equivalent_to_target"].as_str(),
Some("Proven"),
"≡ true law must be Proven after integer snapping, got: {out}"
);
assert_eq!(
s0["symbolic"].as_bool(),
Some(true),
"best law should be symbolic"
);
}
#[test]
fn discover_json_reports_malformed_input() {
let out = discover_json("{ this is not valid json", "{}");
let parsed: serde_json::Value =
serde_json::from_str(&out).expect("error envelope must itself be valid JSON");
let obj = parsed
.as_object()
.expect("error envelope must be an object");
assert!(
obj.contains_key("error"),
"malformed input must yield an `error` field, got: {out}"
);
}
}
#[cfg(target_arch = "wasm32")]
#[cfg(test)]
mod wasm_tests {
use super::*;
use wasm_bindgen_test::*;
wasm_bindgen_test_configure!(run_in_browser);
#[wasm_bindgen_test]
fn discover_json_runs_in_wasm() {
let data =
r#"{ "x": [[0.0],[0.1],[0.2],[0.3],[0.4]], "y": [1.0, 1.105, 1.221, 1.350, 1.492] }"#;
let cfg = r#"{ "max_epochs": 50, "max_depth": 1, "seed": 0, "top_k": 1 }"#;
let out = discover_json(data, cfg);
assert!(
!out.is_empty(),
"discover_json must return a non-empty string"
);
}
}