matrix-mumu 0.1.1

Matrix operations for the mumu/lava language
Documentation
// matrix/src/lib.rs

use core_mumu::{
    parser::interpreter::Interpreter,
    parser::types::{Value, FunctionValue},
};
use std::sync::{Arc, Mutex};

fn to_float2d(v: &Value) -> Option<Vec<Vec<f64>>> {
    match v {
        Value::Float2DArray(rows) => Some(rows.clone()),
        Value::Int2DArray(rows) => Some(rows.iter().map(|r| r.iter().map(|&x| x as f64).collect()).collect()),
        Value::MixedArray(rows) => {
            let mut result = Vec::with_capacity(rows.len());
            for row in rows {
                match row {
                    Value::FloatArray(xs) => result.push(xs.clone()),
                    Value::IntArray(xs) => result.push(xs.iter().map(|&x| x as f64).collect()),
                    Value::MixedArray(inner) => {
                        // Accept [Int|Float|Mixed]Array for row
                        let mut vrow = Vec::with_capacity(inner.len());
                        for v in inner {
                            match v {
                                Value::Int(i) => vrow.push(*i as f64),
                                Value::Float(f) => vrow.push(*f),
                                _ => return None,
                            }
                        }
                        result.push(vrow);
                    }
                    _ => return None,
                }
            }
            Some(result)
        }
        _ => None,
    }
}

fn matrix_subtract(_interp: &mut Interpreter, args: Vec<Value>) -> Result<Value, String> {
    if args.len() != 2 {
        return Err("matrix:subtract expects 2 arguments".to_string());
    }
    // Try all combinations: prefer float, but fallback to int if possible
    match (&args[0], &args[1]) {
        // Int or float native types
        (Value::Int2DArray(xs), Value::Int2DArray(ys)) if xs.len() == ys.len() && xs[0].len() == ys[0].len() => {
            let result: Vec<Vec<i32>> = xs.iter().zip(ys.iter())
                .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a - b).collect())
                .collect();
            Ok(Value::Int2DArray(result))
        }
        (Value::Float2DArray(xs), Value::Float2DArray(ys)) if xs.len() == ys.len() && xs[0].len() == ys[0].len() => {
            let result: Vec<Vec<f64>> = xs.iter().zip(ys.iter())
                .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a - b).collect())
                .collect();
            Ok(Value::Float2DArray(result))
        }
        (Value::Int2DArray(xs), Value::Float2DArray(ys)) if xs.len() == ys.len() && xs[0].len() == ys[0].len() => {
            let result: Vec<Vec<f64>> = xs.iter().zip(ys.iter())
                .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| *a as f64 - b).collect())
                .collect();
            Ok(Value::Float2DArray(result))
        }
        (Value::Float2DArray(xs), Value::Int2DArray(ys)) if xs.len() == ys.len() && xs[0].len() == ys[0].len() => {
            let result: Vec<Vec<f64>> = xs.iter().zip(ys.iter())
                .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a - *b as f64).collect())
                .collect();
            Ok(Value::Float2DArray(result))
        }
        // Accept MixedArrays as long as numeric
        (a, b) => {
            let a_float = to_float2d(a).ok_or("Not numeric")?;
            let b_float = to_float2d(b).ok_or("Not numeric")?;
            if a_float.is_empty() || b_float.is_empty() || a_float.len() != b_float.len() || a_float[0].len() != b_float[0].len() {
                return Err("Matrix dimensions do not match".to_string());
            }
            let result: Vec<Vec<f64>> = a_float.iter().zip(b_float.iter())
                .map(|(row_a, row_b)| row_a.iter().zip(row_b.iter()).map(|(a, b)| a - b).collect())
                .collect();
            Ok(Value::Float2DArray(result))
        }
    }
}

fn matrix_add(_interp: &mut Interpreter, args: Vec<Value>) -> Result<Value, String> {
    if args.len() != 2 {
        return Err("matrix:add expects 2 arguments".to_string());
    }
    match (&args[0], &args[1]) {
        (Value::Int2DArray(xs), Value::Int2DArray(ys)) if xs.len() == ys.len() && xs[0].len() == ys[0].len() => {
            let result: Vec<Vec<i32>> = xs.iter().zip(ys.iter())
                .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a + b).collect())
                .collect();
            Ok(Value::Int2DArray(result))
        }
        (Value::Float2DArray(xs), Value::Float2DArray(ys)) if xs.len() == ys.len() && xs[0].len() == ys[0].len() => {
            let result: Vec<Vec<f64>> = xs.iter().zip(ys.iter())
                .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a + b).collect())
                .collect();
            Ok(Value::Float2DArray(result))
        }
        (Value::Int2DArray(xs), Value::Float2DArray(ys)) if xs.len() == ys.len() && xs[0].len() == ys[0].len() => {
            let result: Vec<Vec<f64>> = xs.iter().zip(ys.iter())
                .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| *a as f64 + b).collect())
                .collect();
            Ok(Value::Float2DArray(result))
        }
        (Value::Float2DArray(xs), Value::Int2DArray(ys)) if xs.len() == ys.len() && xs[0].len() == ys[0].len() => {
            let result: Vec<Vec<f64>> = xs.iter().zip(ys.iter())
                .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a + *b as f64).collect())
                .collect();
            Ok(Value::Float2DArray(result))
        }
        (a, b) => {
            let a_float = to_float2d(a).ok_or("Not numeric")?;
            let b_float = to_float2d(b).ok_or("Not numeric")?;
            if a_float.is_empty() || b_float.is_empty() || a_float.len() != b_float.len() || a_float[0].len() != b_float[0].len() {
                return Err("Matrix dimensions do not match".to_string());
            }
            let result: Vec<Vec<f64>> = a_float.iter().zip(b_float.iter())
                .map(|(row_a, row_b)| row_a.iter().zip(row_b.iter()).map(|(a, b)| a + b).collect())
                .collect();
            Ok(Value::Float2DArray(result))
        }
    }
}

fn matrix_multiply(_interp: &mut Interpreter, args: Vec<Value>) -> Result<Value, String> {
    if args.len() != 2 {
        return Err("matrix:multiply expects 2 arguments".to_string());
    }
    match (&args[0], &args[1]) {
        (Value::Int2DArray(a), Value::Int2DArray(b)) => matrix_mul_int(a, b),
        (Value::Float2DArray(a), Value::Float2DArray(b)) => matrix_mul_float(a, b),
        (Value::Int2DArray(a), Value::Float2DArray(b)) => matrix_mul_float2d_cast(a, b),
        (Value::Float2DArray(a), Value::Int2DArray(b)) => matrix_mul_float2d_cast_rev(a, b),
        (a, b) => {
            let a_float = to_float2d(a).ok_or("Not numeric")?;
            let b_float = to_float2d(b).ok_or("Not numeric")?;
            let n = a_float.len();
            let m = if n > 0 { a_float[0].len() } else { 0 };
            let p = if !b_float.is_empty() { b_float[0].len() } else { 0 };
            if b_float.len() != m {
                return Err(format!("dimension mismatch: left={}x{}, right={}x{}", n, m, b_float.len(), p));
            }
            let mut result = vec![vec![0.0f64; p]; n];
            for i in 0..n {
                for j in 0..p {
                    for k in 0..m {
                        result[i][j] += a_float[i][k] * b_float[k][j];
                    }
                }
            }
            Ok(Value::Float2DArray(result))
        }
    }
}

fn matrix_mul_int(a: &Vec<Vec<i32>>, b: &Vec<Vec<i32>>) -> Result<Value, String> {
    let n = a.len();
    let m = a[0].len();
    let p = b[0].len();
    if b.len() != m {
        return Err(format!("dimension mismatch: left={}x{}, right={}x{}", n, m, b.len(), p));
    }
    let mut result = vec![vec![0i32; p]; n];
    for i in 0..n {
        for j in 0..p {
            for k in 0..m {
                result[i][j] += a[i][k] * b[k][j];
            }
        }
    }
    Ok(Value::Int2DArray(result))
}

fn matrix_mul_float(a: &Vec<Vec<f64>>, b: &Vec<Vec<f64>>) -> Result<Value, String> {
    let n = a.len();
    let m = a[0].len();
    let p = b[0].len();
    if b.len() != m {
        return Err(format!("dimension mismatch: left={}x{}, right={}x{}", n, m, b.len(), p));
    }
    let mut result = vec![vec![0.0f64; p]; n];
    for i in 0..n {
        for j in 0..p {
            for k in 0..m {
                result[i][j] += a[i][k] * b[k][j];
            }
        }
    }
    Ok(Value::Float2DArray(result))
}

fn matrix_mul_float2d_cast(a: &Vec<Vec<i32>>, b: &Vec<Vec<f64>>) -> Result<Value, String> {
    let n = a.len();
    let m = a[0].len();
    let p = b[0].len();
    if b.len() != m {
        return Err(format!("dimension mismatch: left={}x{}, right={}x{}", n, m, b.len(), p));
    }
    let mut result = vec![vec![0.0f64; p]; n];
    for i in 0..n {
        for j in 0..p {
            for k in 0..m {
                result[i][j] += a[i][k] as f64 * b[k][j];
            }
        }
    }
    Ok(Value::Float2DArray(result))
}

fn matrix_mul_float2d_cast_rev(a: &Vec<Vec<f64>>, b: &Vec<Vec<i32>>) -> Result<Value, String> {
    let n = a.len();
    let m = a[0].len();
    let p = b[0].len();
    if b.len() != m {
        return Err(format!("dimension mismatch: left={}x{}, right={}x{}", n, m, b.len(), p));
    }
    let mut result = vec![vec![0.0f64; p]; n];
    for i in 0..n {
        for j in 0..p {
            for k in 0..m {
                result[i][j] += a[i][k] * b[k][j] as f64;
            }
        }
    }
    Ok(Value::Float2DArray(result))
}

fn transpose2d<T: Clone>(m: &Vec<Vec<T>>) -> Vec<Vec<T>> {
    let rows = m.len();
    let cols = if rows > 0 { m[0].len() } else { 0 };
    (0..cols)
        .map(|c| (0..rows).map(|r| m[r][c].clone()).collect())
        .collect()
}

fn matrix_transpose(_interp: &mut Interpreter, args: Vec<Value>) -> Result<Value, String> {
    if args.len() != 1 {
        return Err("matrix:transpose expects 1 argument".to_string());
    }
    match &args[0] {
        Value::Int2DArray(xs) => Ok(Value::Int2DArray(transpose2d(xs))),
        Value::Float2DArray(xs) => Ok(Value::Float2DArray(transpose2d(xs))),
        Value::MixedArray(rows) => {
            let as_float2d = to_float2d(&Value::MixedArray(rows.clone())).ok_or("Not numeric")?;
            Ok(Value::Float2DArray(transpose2d(&as_float2d)))
        }
        _ => Err("Not numeric".to_string()),
    }
}

#[no_mangle]
pub unsafe extern "C" fn Cargo_lock(
    interp_ptr: *mut std::ffi::c_void,
    _extra_str: *const std::ffi::c_void,
) -> i32 {
    let interp = &mut *(interp_ptr as *mut Interpreter);

    macro_rules! reg {
        ($name:expr, $f:expr) => {{
            let func = Arc::new(Mutex::new($f));
            interp.register_dynamic_function($name, func);
            interp.set_variable($name, Value::Function(Box::new(FunctionValue::Named($name.into()))));
        }};
    }
    reg!("matrix:add", matrix_add);
    reg!("matrix:subtract", matrix_subtract);
    reg!("matrix:multiply", matrix_multiply);
    reg!("matrix:transpose", matrix_transpose);

    0
}