shape-jit 0.3.0

Tiered JIT compiler (Cranelift) for the Shape virtual machine
Documentation
//! Math FFI Symbol Registration
//!
//! This module handles registration and declaration of math-related FFI symbols
//! for the JIT compiler, including trigonometric functions, generic arithmetic,
//! series comparisons, and intrinsic statistics.

use cranelift::prelude::*;
use cranelift_jit::JITBuilder;
use cranelift_jit::JITModule;
use cranelift_module::{FuncId, Linkage, Module};
use std::collections::HashMap;

use super::super::ffi::math::{
    jit_acos, jit_asin, jit_atan, jit_cos, jit_exp, jit_ln, jit_log, jit_pow, jit_sin, jit_tan,
};
// W11-fup-A (Phase 3d, 2026-05-18): typed-pow FFI helpers for the JIT
// MIR-lowering path's `compile_binop_f64::BinOp::Pow` and
// `compile_binop_int64::BinOp::Pow` arms. `jit_pow_f64` lifts the
// pre-existing-but-unregistered helper at `ffi/v2_math.rs:302`;
// `jit_pow_i64` is new (this sub-cluster, `ffi/v2_math.rs::jit_pow_i64`).
use super::super::ffi::v2_math::{jit_pow_f64, jit_pow_i64};
// W15.2-LANG-6 jit-math-nan-poisoning (Phase 4b Round 2, 2026-05-18):
// typed-f64 single-arg math FFI helpers for the JIT MIR-lowering path's
// `MirConstant::Function(<math-builtin>)` Call-terminator interception.
// Each takes/returns native f64 (no NaN-box). Bodies live at
// `crates/shape-jit/src/ffi/v2_math.rs::jit_{sqrt,abs,floor,ceil,round,
// sin,cos,tan,asin,acos,atan,exp,ln}_f64`. The bytecode VM emits
// `OpCode::BuiltinCall(Sqrt|Abs|...)` for these names, but MIR lowering
// at `crates/shape-vm/src/mir/lowering/expr.rs:2003` leaks them as
// `MirConstant::Function(name)` — the JIT must intercept by name at the
// Call terminator and route to these typed FFI bodies (ADR-006 §2.7.5
// stamp-at-compile-time: the arg's `NativeKind::Float64` is the
// producing-site discriminator; no kind-blind fallback per §2.7.7 #9).
use super::super::ffi::v2_math::{
    jit_abs_f64, jit_acos_f64, jit_asin_f64, jit_atan_f64, jit_ceil_f64, jit_cos_f64, jit_exp_f64,
    jit_floor_f64, jit_ln_f64, jit_round_f64, jit_sin_f64, jit_sqrt_f64, jit_tan_f64,
};
use super::intrinsics::{
    jit_intrinsic_correlation, jit_intrinsic_covariance, jit_intrinsic_max, jit_intrinsic_mean,
    jit_intrinsic_median, jit_intrinsic_min, jit_intrinsic_percentile, jit_intrinsic_std,
    jit_intrinsic_sum, jit_intrinsic_variance, jit_series_broadcast, jit_series_clip,
    jit_series_cumprod, jit_series_diff, jit_series_ema, jit_series_highest_index,
    jit_series_lowest_index, jit_series_pct_change, jit_series_rolling_max, jit_series_rolling_min,
};

/// Register math FFI symbols with the JIT builder
pub fn register_math_symbols(builder: &mut JITBuilder) {
    // Trigonometric functions
    builder.symbol("jit_sin", jit_sin as *const u8);
    builder.symbol("jit_cos", jit_cos as *const u8);
    builder.symbol("jit_tan", jit_tan as *const u8);
    builder.symbol("jit_asin", jit_asin as *const u8);
    builder.symbol("jit_acos", jit_acos as *const u8);
    builder.symbol("jit_atan", jit_atan as *const u8);

    // Exponential and logarithmic functions
    builder.symbol("jit_exp", jit_exp as *const u8);
    builder.symbol("jit_ln", jit_ln as *const u8);
    builder.symbol("jit_log", jit_log as *const u8);
    builder.symbol("jit_pow", jit_pow as *const u8);
    // W11-fup-A typed-pow helpers (native f64 / i64 ABI, distinct from
    // the NaN-boxed `jit_pow` above).
    builder.symbol("jit_pow_f64", jit_pow_f64 as *const u8);
    builder.symbol("jit_pow_i64", jit_pow_i64 as *const u8);

    // W15.2-LANG-6 jit-math-nan-poisoning typed-f64 single-arg math
    // helpers (Phase 4b Round 2, 2026-05-18). Distinct from the NaN-boxed
    // `jit_sin` / `jit_cos` / `jit_tan` / `jit_asin` / `jit_acos` /
    // `jit_atan` / `jit_exp` / `jit_ln` registrations above which take
    // and return u64 NaN-boxed bits; these take native f64 and return
    // native f64.
    builder.symbol("jit_sqrt_f64", jit_sqrt_f64 as *const u8);
    builder.symbol("jit_abs_f64", jit_abs_f64 as *const u8);
    builder.symbol("jit_floor_f64", jit_floor_f64 as *const u8);
    builder.symbol("jit_ceil_f64", jit_ceil_f64 as *const u8);
    builder.symbol("jit_round_f64", jit_round_f64 as *const u8);
    builder.symbol("jit_sin_f64", jit_sin_f64 as *const u8);
    builder.symbol("jit_cos_f64", jit_cos_f64 as *const u8);
    builder.symbol("jit_tan_f64", jit_tan_f64 as *const u8);
    builder.symbol("jit_asin_f64", jit_asin_f64 as *const u8);
    builder.symbol("jit_acos_f64", jit_acos_f64 as *const u8);
    builder.symbol("jit_atan_f64", jit_atan_f64 as *const u8);
    builder.symbol("jit_exp_f64", jit_exp_f64 as *const u8);
    builder.symbol("jit_ln_f64", jit_ln_f64 as *const u8);

    // R7.1: the 11 `jit_generic_*` dispatch-fallback trampolines were
    // removed together with their `FFIFuncRefs` fields and Cranelift
    // signatures after R5 retargeted every dynamic arithmetic /
    // comparison path to typed opcodes or `CallMethod`.

    // Series comparison functions
    builder.symbol(
        "jit_series_gt",
        super::super::ffi::math::jit_series_gt as *const u8,
    );
    builder.symbol(
        "jit_series_lt",
        super::super::ffi::math::jit_series_lt as *const u8,
    );
    builder.symbol(
        "jit_series_gte",
        super::super::ffi::math::jit_series_gte as *const u8,
    );
    builder.symbol(
        "jit_series_lte",
        super::super::ffi::math::jit_series_lte as *const u8,
    );

    // Intrinsic aggregation functions
    builder.symbol("jit_intrinsic_sum", jit_intrinsic_sum as *const u8);
    builder.symbol("jit_intrinsic_mean", jit_intrinsic_mean as *const u8);
    builder.symbol("jit_intrinsic_min", jit_intrinsic_min as *const u8);
    builder.symbol("jit_intrinsic_max", jit_intrinsic_max as *const u8);
    builder.symbol("jit_intrinsic_std", jit_intrinsic_std as *const u8);
    builder.symbol(
        "jit_intrinsic_variance",
        jit_intrinsic_variance as *const u8,
    );
    builder.symbol("jit_intrinsic_median", jit_intrinsic_median as *const u8);
    builder.symbol(
        "jit_intrinsic_percentile",
        jit_intrinsic_percentile as *const u8,
    );
    builder.symbol(
        "jit_intrinsic_correlation",
        jit_intrinsic_correlation as *const u8,
    );
    builder.symbol(
        "jit_intrinsic_covariance",
        jit_intrinsic_covariance as *const u8,
    );

    // Series operations
    builder.symbol(
        "jit_series_rolling_min",
        jit_series_rolling_min as *const u8,
    );
    builder.symbol(
        "jit_series_rolling_max",
        jit_series_rolling_max as *const u8,
    );
    builder.symbol("jit_series_ema", jit_series_ema as *const u8);
    builder.symbol("jit_series_diff", jit_series_diff as *const u8);
    builder.symbol("jit_series_pct_change", jit_series_pct_change as *const u8);
    builder.symbol("jit_series_cumprod", jit_series_cumprod as *const u8);
    builder.symbol("jit_series_clip", jit_series_clip as *const u8);
    builder.symbol("jit_series_broadcast", jit_series_broadcast as *const u8);
    builder.symbol(
        "jit_series_highest_index",
        jit_series_highest_index as *const u8,
    );
    builder.symbol(
        "jit_series_lowest_index",
        jit_series_lowest_index as *const u8,
    );
}

/// Declare math FFI function signatures in the module
pub fn declare_math_functions(module: &mut JITModule, ffi_funcs: &mut HashMap<String, FuncId>) {
    // Math functions (single param)
    for name in [
        "jit_sin", "jit_cos", "jit_tan", "jit_asin", "jit_acos", "jit_atan", "jit_exp", "jit_ln",
    ] {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64));
        sig.returns.push(AbiParam::new(types::I64));
        let func_id = module
            .declare_function(name, Linkage::Import, &sig)
            .unwrap_or_else(|_| panic!("Failed to declare {}", name));
        ffi_funcs.insert(name.to_string(), func_id);
    }

    // jit_log(value_bits, base_bits) -> u64
    {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64)); // value
        sig.params.push(AbiParam::new(types::I64)); // base
        sig.returns.push(AbiParam::new(types::I64));
        let func_id = module
            .declare_function("jit_log", Linkage::Import, &sig)
            .expect("Failed to declare jit_log");
        ffi_funcs.insert("jit_log".to_string(), func_id);
    }

    // jit_pow(base_bits, exp_bits) -> u64
    {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64)); // base
        sig.params.push(AbiParam::new(types::I64)); // exp
        sig.returns.push(AbiParam::new(types::I64));
        let func_id = module
            .declare_function("jit_pow", Linkage::Import, &sig)
            .expect("Failed to declare jit_pow");
        ffi_funcs.insert("jit_pow".to_string(), func_id);
    }

    // W11-fup-A (Phase 3d, 2026-05-18): typed-pow signatures for the
    // MIR `BinOp::Pow` JIT path. F64 ABI for `compile_binop_f64`,
    // I64 ABI for `compile_binop_int64`. Distinct from the NaN-boxed
    // `jit_pow` above which takes/returns I64-bit-pattern operands.
    // jit_pow_f64(a: f64, b: f64) -> f64
    {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::F64));
        sig.params.push(AbiParam::new(types::F64));
        sig.returns.push(AbiParam::new(types::F64));
        let func_id = module
            .declare_function("jit_pow_f64", Linkage::Import, &sig)
            .expect("Failed to declare jit_pow_f64");
        ffi_funcs.insert("jit_pow_f64".to_string(), func_id);
    }
    // jit_pow_i64(base: i64, exp: i64) -> i64
    {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64));
        sig.params.push(AbiParam::new(types::I64));
        sig.returns.push(AbiParam::new(types::I64));
        let func_id = module
            .declare_function("jit_pow_i64", Linkage::Import, &sig)
            .expect("Failed to declare jit_pow_i64");
        ffi_funcs.insert("jit_pow_i64".to_string(), func_id);
    }

    // W15.2-LANG-6 jit-math-nan-poisoning (Phase 4b Round 2, 2026-05-18):
    // typed-f64 single-arg math FFI signatures for the JIT Call-terminator
    // interception of `MirConstant::Function("sqrt"|"abs"|...)`. Native
    // f64 ABI: `extern "C" fn(f64) -> f64`. Bodies at
    // `ffi/v2_math.rs::jit_{sqrt,abs,...}_f64`. Distinct from the NaN-boxed
    // `jit_sin` / `jit_cos` / `jit_tan` / `jit_asin` / `jit_acos` /
    // `jit_atan` / `jit_exp` / `jit_ln` signatures registered above
    // (single I64 param / I64 return) which the MIR consumer does not call.
    for name in [
        "jit_sqrt_f64",
        "jit_abs_f64",
        "jit_floor_f64",
        "jit_ceil_f64",
        "jit_round_f64",
        "jit_sin_f64",
        "jit_cos_f64",
        "jit_tan_f64",
        "jit_asin_f64",
        "jit_acos_f64",
        "jit_atan_f64",
        "jit_exp_f64",
        "jit_ln_f64",
    ] {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::F64));
        sig.returns.push(AbiParam::new(types::F64));
        let func_id = module
            .declare_function(name, Linkage::Import, &sig)
            .unwrap_or_else(|_| panic!("Failed to declare {}", name));
        ffi_funcs.insert(name.to_string(), func_id);
    }

    // R7.1: Generic binary op declarations (11 `jit_generic_*` names) were
    // removed here together with their Rust bodies and `FFIFuncRefs`
    // fields. MIR no longer emits a fully dynamic arithmetic / comparison
    // binop after R5.

    // Series comparison functions (a_bits: u64, b_bits: u64) -> u64
    for name in [
        "jit_series_gt",
        "jit_series_lt",
        "jit_series_gte",
        "jit_series_lte",
    ] {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64)); // a_bits
        sig.params.push(AbiParam::new(types::I64)); // b_bits
        sig.returns.push(AbiParam::new(types::I64)); // result
        let func_id = module
            .declare_function(name, Linkage::Import, &sig)
            .unwrap_or_else(|_| panic!("Failed to declare {}", name));
        ffi_funcs.insert(name.to_string(), func_id);
    }

    // Intrinsic aggregation functions (series_bits: u64) -> u64
    for name in [
        "jit_intrinsic_sum",
        "jit_intrinsic_mean",
        "jit_intrinsic_min",
        "jit_intrinsic_max",
        "jit_intrinsic_std",
        "jit_intrinsic_variance",
        "jit_intrinsic_median",
        "jit_series_cumprod",
    ] {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64)); // series_bits
        sig.returns.push(AbiParam::new(types::I64)); // result
        let func_id = module
            .declare_function(name, Linkage::Import, &sig)
            .unwrap_or_else(|_| panic!("Failed to declare {}", name));
        ffi_funcs.insert(name.to_string(), func_id);
    }

    // Intrinsic two-arg functions (series_bits: u64, arg: u64) -> u64
    for name in [
        "jit_intrinsic_percentile",
        "jit_series_rolling_min",
        "jit_series_rolling_max",
        "jit_series_ema",
        "jit_series_diff",
        "jit_series_pct_change",
    ] {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64)); // series_bits
        sig.params.push(AbiParam::new(types::I64)); // second arg
        sig.returns.push(AbiParam::new(types::I64)); // result
        let func_id = module
            .declare_function(name, Linkage::Import, &sig)
            .unwrap_or_else(|_| panic!("Failed to declare {}", name));
        ffi_funcs.insert(name.to_string(), func_id);
    }

    // Two-series functions (a_bits: u64, b_bits: u64) -> u64
    for name in ["jit_intrinsic_correlation", "jit_intrinsic_covariance"] {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64)); // a_bits
        sig.params.push(AbiParam::new(types::I64)); // b_bits
        sig.returns.push(AbiParam::new(types::I64)); // result
        let func_id = module
            .declare_function(name, Linkage::Import, &sig)
            .unwrap_or_else(|_| panic!("Failed to declare {}", name));
        ffi_funcs.insert(name.to_string(), func_id);
    }

    // Clip function (series_bits: u64, min: u64, max: u64) -> u64
    {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64)); // series_bits
        sig.params.push(AbiParam::new(types::I64)); // min
        sig.params.push(AbiParam::new(types::I64)); // max
        sig.returns.push(AbiParam::new(types::I64)); // result
        let func_id = module
            .declare_function("jit_series_clip", Linkage::Import, &sig)
            .expect("Failed to declare jit_series_clip");
        ffi_funcs.insert("jit_series_clip".to_string(), func_id);
    }

    // jit_series_broadcast(value_bits: u64, len_bits: u64) -> u64
    {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64)); // value_bits
        sig.params.push(AbiParam::new(types::I64)); // len_bits
        sig.returns.push(AbiParam::new(types::I64)); // result
        let func_id = module
            .declare_function("jit_series_broadcast", Linkage::Import, &sig)
            .expect("Failed to declare jit_series_broadcast");
        ffi_funcs.insert("jit_series_broadcast".to_string(), func_id);
    }

    // jit_series_highest_index(series_bits: u64) -> u64
    {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64)); // series_bits
        sig.returns.push(AbiParam::new(types::I64)); // result
        let func_id = module
            .declare_function("jit_series_highest_index", Linkage::Import, &sig)
            .expect("Failed to declare jit_series_highest_index");
        ffi_funcs.insert("jit_series_highest_index".to_string(), func_id);
    }

    // jit_series_lowest_index(series_bits: u64) -> u64
    {
        let mut sig = module.make_signature();
        sig.params.push(AbiParam::new(types::I64)); // series_bits
        sig.returns.push(AbiParam::new(types::I64)); // result
        let func_id = module
            .declare_function("jit_series_lowest_index", Linkage::Import, &sig)
            .expect("Failed to declare jit_series_lowest_index");
        ffi_funcs.insert("jit_series_lowest_index".to_string(), func_id);
    }
}