Skip to main content

shape_jit/compiler/
strategy.rs

1//! Strategy compilation
2//!
3//! This module provides two compilation modes:
4//!
5//! 1. **Standard ABI** (`compile_strategy`): Uses `fn(*mut JITContext) -> i32`
6//!    - Full access to VM features (closures, FFI, etc.)
7//!    - Suitable for general-purpose JIT compilation
8//!
9//! 2. **Kernel ABI** (`compile_simulation_kernel`): Uses `fn(usize, *const *const f64, *mut u8) -> i32`
10//!    - Zero-allocation hot path for simulation
11//!    - Direct memory access to series data and state
12//!    - Enables >10M ticks/sec performance
13
14use cranelift::codegen::ir::FuncRef;
15use cranelift::prelude::*;
16use cranelift_module::{Linkage, Module};
17use std::collections::HashMap;
18
19use super::setup::JITCompiler;
20use crate::context::{
21    CorrelatedKernelFn, JittedStrategyFn, SimulationKernelConfig, SimulationKernelFn,
22};
23use shape_vm::bytecode::BytecodeProgram;
24
25impl JITCompiler {
26    #[inline(always)]
27    pub fn compile_strategy(
28        &mut self,
29        name: &str,
30        program: &BytecodeProgram,
31    ) -> Result<JittedStrategyFn, String> {
32        // MirToIR is the ONLY compilation path.
33        let mir_data = program.top_level_mir.as_ref().ok_or_else(|| {
34            "MirToIR: top-level code has no MIR data".to_string()
35        })?;
36        let preflight = crate::mir_compiler::preflight(mir_data);
37        if !preflight.can_compile {
38            return Err(format!(
39                "MirToIR: top-level preflight failed: {}",
40                preflight.blockers.join("; ")
41            ));
42        }
43
44        let mut sig = self.module.make_signature();
45        sig.params.push(AbiParam::new(types::I64));
46        sig.returns.push(AbiParam::new(types::I32));
47
48        let func_id = self
49            .module
50            .declare_function(name, Linkage::Export, &sig)
51            .map_err(|e| format!("Failed to declare function: {}", e))?;
52
53        let mut ctx = self.module.make_context();
54        ctx.func.signature = sig;
55
56        let mut func_builder_ctx = FunctionBuilderContext::new();
57        {
58            let mut builder = FunctionBuilder::new(&mut ctx.func, &mut func_builder_ctx);
59            let entry_block = builder.create_block();
60            builder.append_block_params_for_function_params(entry_block);
61            builder.switch_to_block(entry_block);
62            builder.seal_block(entry_block);
63
64            let ctx_ptr = builder.block_params(entry_block)[0];
65
66            let ffi = self.build_ffi_refs(&mut builder)?;
67
68            {
69                let slot_kinds: Vec<Option<shape_vm::type_tracking::NativeKind>> = program
70                    .top_level_frame
71                    .as_ref()
72                    .map(|fd| fd.slots.iter().copied().map(Some).collect())
73                    .unwrap_or_default();
74                // ADR-006 §2.7.5 conduit: thread the bytecode compiler's
75                // proven per-slot `ConcreteType` for top-level locals into
76                // MirToIR. The bytecode compiler stamps the side-table at
77                // `populate_program_storage_hints` time from
78                // `local_array_element_types`, `local_map_key_value_types`,
79                // and the type-tracker's schema registry (W12-top-level-
80                // concrete-types-conduit close, 2026-05-12). MirToIR's v2
81                // fast path uses `Array<scalar>` / `Struct(_)` /
82                // `HashMap(K, V)` slot kinds to bypass `Rvalue::Aggregate`
83                // surface-and-stop and the kind-blind ObjectStore path.
84                // Empty vec (no top-level code) → MirToIR falls through to
85                // the legacy NaN-boxed path naturally.
86                let concrete_types: Vec<shape_value::v2::ConcreteType> =
87                    program.top_level_local_concrete_types.clone();
88                let function_indices: std::collections::HashMap<String, u16> = program
89                    .functions
90                    .iter()
91                    .enumerate()
92                    .map(|(i, f)| (f.name.clone(), i as u16))
93                    .collect();
94                let closure_function_layouts: HashMap<
95                    u16,
96                    std::sync::Arc<shape_value::v2::closure_layout::ClosureLayout>,
97                > = program
98                    .closure_function_layouts
99                    .iter()
100                    .enumerate()
101                    .filter_map(|(i, opt)| opt.as_ref().map(|l| (i as u16, l.clone())))
102                    .collect();
103                let mut mir_compiler = crate::mir_compiler::MirToIR::new_with_closure_layouts(
104                    &mut builder,
105                    ctx_ptr,
106                    ffi,
107                    mir_data,
108                    slot_kinds,
109                    concrete_types,
110                    &program.strings,
111                    entry_block,
112                    &function_indices,
113                    HashMap::new(),
114                    HashMap::new(),
115                    closure_function_layouts,
116                );
117                // V3-S6c-jit-method-monomorph-routing: thread the V3-S6b
118                // side-table for top-level (`__main__`) code. Caller id is
119                // `None` per the bytecode compiler's convention
120                // (`self.current_function == None` when compiling
121                // top-level statements at
122                // `expressions/function_calls.rs:3278`).
123                mir_compiler.set_monomorph_routing_context(
124                    program.monomorphized_method_call_sites.clone(),
125                    None,
126                );
127                // W10 jit-call-method-user-trait-fix (2026-05-17): top-
128                // level mirror of the per-user-function threading at
129                // `compiler/program.rs::compile_function_with_user_funcs`.
130                mir_compiler.set_operator_trait_dispatch_sites(
131                    program.operator_trait_dispatch_sites.clone(),
132                );
133                // Bounds-check elision: scan the MIR for trusted index
134                // accesses and install the plan before compile_body. The
135                // analyzer is conservative; an empty plan preserves the
136                // bounds-checked path for every access.
137                let elision_plan =
138                    crate::mir_compiler::bounds_elision::analyze(&mir_data.mir);
139                mir_compiler.set_bounds_elision_plan(elision_plan);
140                // W14.2-E-followup SURFACE-A2 fix (2026-05-19): same as
141                // the per-user-function path at `program.rs` — pre-
142                // populate `field_byte_offsets` from the schema registry
143                // so top-level field reads resolve at JIT-compile time.
144                mir_compiler
145                    .populate_field_byte_offsets_from_schemas(&program.type_schema_registry);
146                mir_compiler.create_blocks();
147                mir_compiler.declare_locals();
148                mir_compiler.initialize_locals();
149                // Session 1 Commit 3: allocate Arc<SharedCell>s for
150                // every SharedCow local slot before the body runs.
151                mir_compiler.initialize_shared_local_slots();
152                mir_compiler.compile_body()?;
153            }
154            builder.finalize();
155        }
156
157        self.module
158            .define_function(func_id, &mut ctx)
159            .map_err(|e| format!("Failed to define function (strategy): {:?}", e))?;
160
161        self.module.clear_context(&mut ctx);
162        self.module
163            .finalize_definitions()
164            .map_err(|e| format!("Failed to finalize (strategy): {:?}", e))?;
165
166        let code_ptr = self.module.get_finalized_function(func_id);
167        self.compiled_functions.insert(name.to_string(), code_ptr);
168
169        Ok(unsafe { std::mem::transmute(code_ptr) })
170    }
171
172    #[inline(always)]
173    pub(super) fn compile_strategy_with_user_funcs(
174        &mut self,
175        name: &str,
176        program: &BytecodeProgram,
177        user_func_ids: &HashMap<u16, cranelift_module::FuncId>,
178        user_func_arities: &HashMap<u16, u16>,
179    ) -> Result<cranelift_module::FuncId, String> {
180        let mut sig = self.module.make_signature();
181        sig.params.push(AbiParam::new(types::I64));
182        sig.returns.push(AbiParam::new(types::I32));
183
184        let func_id = self
185            .module
186            .declare_function(name, Linkage::Export, &sig)
187            .map_err(|e| format!("Failed to declare function: {}", e))?;
188
189        let mut ctx = self.module.make_context();
190        ctx.func.signature = sig;
191
192        // MirToIR is the ONLY JIT compilation path (Phase 4: BytecodeToIR removed).
193        let mir_data = program.top_level_mir.as_ref().ok_or_else(|| {
194            "MirToIR: top-level code has no MIR data".to_string()
195        })?;
196        let preflight = crate::mir_compiler::preflight(mir_data);
197        if !preflight.can_compile {
198            return Err(format!(
199                "MirToIR: top-level preflight failed: {}",
200                preflight.blockers.join("; ")
201            ));
202        }
203
204        let mut func_builder_ctx = FunctionBuilderContext::new();
205        {
206            let mut builder = FunctionBuilder::new(&mut ctx.func, &mut func_builder_ctx);
207            let entry_block = builder.create_block();
208            builder.append_block_params_for_function_params(entry_block);
209            builder.switch_to_block(entry_block);
210            builder.seal_block(entry_block);
211
212            let ctx_ptr = builder.block_params(entry_block)[0];
213
214            let mut user_func_refs: HashMap<u16, FuncRef> = HashMap::new();
215            for (&fn_idx, &fn_id) in user_func_ids {
216                let func_ref = self.module.declare_func_in_func(fn_id, builder.func);
217                user_func_refs.insert(fn_idx, func_ref);
218            }
219
220            let ffi = self.build_ffi_refs(&mut builder)?;
221
222            {
223                let slot_kinds: Vec<Option<shape_vm::type_tracking::NativeKind>> = program
224                    .top_level_frame
225                    .as_ref()
226                    .map(|fd| fd.slots.iter().copied().map(Some).collect())
227                    .unwrap_or_default();
228                // ADR-006 §2.7.5 conduit: thread the bytecode compiler's
229                // proven per-slot `ConcreteType` for top-level locals into
230                // MirToIR (W12-top-level-concrete-types-conduit close,
231                // 2026-05-12). Same source as the no-user-funcs path
232                // above; see the populate_program_storage_hints
233                // commentary in `crates/shape-vm/src/compiler/helpers.rs`.
234                let concrete_types: Vec<shape_value::v2::ConcreteType> =
235                    program.top_level_local_concrete_types.clone();
236                let function_indices: std::collections::HashMap<String, u16> = program
237                    .functions
238                    .iter()
239                    .enumerate()
240                    .map(|(i, f)| (f.name.clone(), i as u16))
241                    .collect();
242                let closure_function_layouts: HashMap<
243                    u16,
244                    std::sync::Arc<shape_value::v2::closure_layout::ClosureLayout>,
245                > = program
246                    .closure_function_layouts
247                    .iter()
248                    .enumerate()
249                    .filter_map(|(i, opt)| opt.as_ref().map(|l| (i as u16, l.clone())))
250                    .collect();
251                let mut mir_compiler = crate::mir_compiler::MirToIR::new_with_closure_layouts(
252                    &mut builder,
253                    ctx_ptr,
254                    ffi,
255                    mir_data,
256                    slot_kinds,
257                    concrete_types,
258                    &program.strings,
259                    entry_block,
260                    &function_indices,
261                    user_func_refs.clone(),
262                    user_func_arities.clone(),
263                    closure_function_layouts,
264                );
265                // V3-S6c-jit-method-monomorph-routing: top-level path with
266                // user-funcs visible. Caller id = None per the same
267                // convention as the no-user-funcs path above.
268                mir_compiler.set_monomorph_routing_context(
269                    program.monomorphized_method_call_sites.clone(),
270                    None,
271                );
272                // W10 jit-call-method-user-trait-fix (2026-05-17): same
273                // top-level mirror as the no-user-funcs branch above.
274                mir_compiler.set_operator_trait_dispatch_sites(
275                    program.operator_trait_dispatch_sites.clone(),
276                );
277                let elision_plan =
278                    crate::mir_compiler::bounds_elision::analyze(&mir_data.mir);
279                mir_compiler.set_bounds_elision_plan(elision_plan);
280                // W14.2-E-followup SURFACE-A2 fix (2026-05-19): top-level
281                // with user-funcs path — same schema pre-population as
282                // the sibling top-level no-user-funcs branch above.
283                mir_compiler
284                    .populate_field_byte_offsets_from_schemas(&program.type_schema_registry);
285                mir_compiler.create_blocks();
286                mir_compiler.declare_locals();
287                mir_compiler.initialize_locals();
288                // Session 1 Commit 3: allocate Arc<SharedCell>s for
289                // every SharedCow local slot before the body runs.
290                mir_compiler.initialize_shared_local_slots();
291                mir_compiler.compile_body()?;
292                tracing::debug!(
293                    target: "shape_jit",
294                    "jit-mir compiled top-level code via MirToIR",
295                );
296            }
297            builder.finalize();
298        }
299
300        self.module
301            .define_function(func_id, &mut ctx)
302            .map_err(|e| format!("Failed to define function (strategy): {:?}", e))?;
303
304        self.module.clear_context(&mut ctx);
305
306        Ok(func_id)
307    }
308
309    /// Compute instruction index ranges to skip when compiling the main strategy.
310    ///
311    /// Bytecode layout for programs with user functions:
312    /// ```text
313    /// [0]             Jump → trampoline1     (skip func0 body)
314    /// [entry0 .. t1)  func0 body
315    /// [t1]            Jump → trampoline2     (skip func1 body)
316    /// [entry1 .. t2)  func1 body
317    /// ...
318    /// [main_start ..) main code
319    /// ```
320    ///
321    /// Returns the function body ranges (excluding trampoline jumps between them).
322    pub(super) fn compute_skip_ranges(program: &BytecodeProgram) -> Vec<(usize, usize)> {
323        let mut ranges = Vec::new();
324
325        // Skip function bodies (they are compiled separately).
326        for f in program.functions.iter() {
327            if f.body_length == 0 {
328                continue;
329            }
330            ranges.push((f.entry_point, f.entry_point + f.body_length));
331        }
332
333        ranges
334    }
335
336    // ========================================================================
337    // Simulation Kernel Compilation (Zero-Allocation Hot Path)
338    // ========================================================================
339
340    /// Compile a simulation kernel with the specialized kernel ABI.
341    ///
342    /// The kernel ABI bypasses JITContext to achieve maximum throughput:
343    /// - Direct pointer arithmetic for data access
344    /// - No allocations in the hot path
345    /// - Inlined field access with known offsets
346    ///
347    /// # Arguments
348    /// * `name` - Function name for the compiled kernel
349    /// * `program` - Bytecode program containing the strategy
350    /// * `config` - Kernel configuration with field offset mappings
351    ///
352    /// # Returns
353    /// A function pointer with signature: `fn(usize, *const *const f64, *mut u8) -> i32`
354    ///
355    /// # Generated Code Pattern
356    ///
357    /// For a strategy like:
358    /// ```shape
359    /// let price = candle.close
360    /// if price > state.threshold {
361    ///     state.signal = 1.0
362    /// }
363    /// ```
364    ///
365    /// The kernel generates:
366    /// ```asm
367    /// ; price = candle.close (column 3)
368    /// mov rax, [series_ptrs + 3*8]     ; column pointer
369    /// mov xmm0, [rax + cursor_index*8] ; price value
370    ///
371    /// ; state.threshold (offset 16)
372    /// mov xmm1, [state_ptr + 16]       ; threshold value
373    ///
374    /// ; comparison and store
375    /// ucomisd xmm0, xmm1
376    /// jbe skip
377    /// mov qword [state_ptr + 24], 1.0  ; state.signal
378    /// skip:
379    /// ```
380    #[inline(always)]
381    pub fn compile_simulation_kernel(
382        &mut self,
383        name: &str,
384        program: &BytecodeProgram,
385        config: &SimulationKernelConfig,
386    ) -> Result<SimulationKernelFn, String> {
387        // Kernel ABI signature: fn(cursor_index: usize, series_ptrs: *const *const f64, state_ptr: *mut u8) -> i32
388        let mut sig = self.module.make_signature();
389        sig.params.push(AbiParam::new(types::I64)); // cursor_index
390        sig.params.push(AbiParam::new(types::I64)); // series_ptrs
391        sig.params.push(AbiParam::new(types::I64)); // state_ptr
392        sig.returns.push(AbiParam::new(types::I32)); // result code
393
394        let func_id = self
395            .module
396            .declare_function(name, Linkage::Export, &sig)
397            .map_err(|e| format!("Failed to declare kernel function: {}", e))?;
398
399        let mut ctx = self.module.make_context();
400        ctx.func.signature = sig;
401
402        let mut func_builder_ctx = FunctionBuilderContext::new();
403        {
404            let mut builder = FunctionBuilder::new(&mut ctx.func, &mut func_builder_ctx);
405            let entry_block = builder.create_block();
406            builder.append_block_params_for_function_params(entry_block);
407            builder.switch_to_block(entry_block);
408            builder.seal_block(entry_block);
409
410            // Get kernel parameters
411            let cursor_index = builder.block_params(entry_block)[0];
412            let series_ptrs = builder.block_params(entry_block)[1];
413            let state_ptr = builder.block_params(entry_block)[2];
414
415            // Build kernel-specific IR
416            let result = self.build_kernel_ir(
417                &mut builder,
418                program,
419                config,
420                cursor_index,
421                series_ptrs,
422                state_ptr,
423            )?;
424
425            builder.ins().return_(&[result]);
426            builder.finalize();
427        }
428
429        self.module
430            .define_function(func_id, &mut ctx)
431            .map_err(|e| format!("Failed to define kernel function: {:?}", e))?;
432
433        self.module.clear_context(&mut ctx);
434        self.module
435            .finalize_definitions()
436            .map_err(|e| format!("Failed to finalize kernel: {:?}", e))?;
437
438        let code_ptr = self.module.get_finalized_function(func_id);
439        self.compiled_functions.insert(name.to_string(), code_ptr);
440
441        Ok(unsafe { std::mem::transmute(code_ptr) })
442    }
443
444    /// Build kernel IR using BytecodeToIR in kernel mode.
445    ///
446    /// This compiles bytecode to kernel ABI IR with direct memory access:
447    /// - GetFieldTyped → state_ptr + offset
448    /// - GetDataField → series_ptrs[col][cursor]
449    /// - All locals as Cranelift variables
450    fn build_kernel_ir(
451        &mut self,
452        _builder: &mut FunctionBuilder,
453        _program: &BytecodeProgram,
454        _config: &SimulationKernelConfig,
455        _cursor_index: Value,
456        _series_ptrs: Value,
457        _state_ptr: Value,
458    ) -> Result<Value, String> {
459        Err("Simulation kernel compilation requires v2 runtime migration".to_string())
460    }
461
462    // ========================================================================
463    // Correlated Kernel Compilation (Multi-Series Simulation)
464    // ========================================================================
465
466    /// Compile a correlated (multi-series) simulation kernel.
467    ///
468    /// This extends the simulation kernel to support multiple aligned time series,
469    /// enabling cross-series strategies (e.g., SPY vs VIX, temperature vs pressure).
470    ///
471    /// # Arguments
472    /// * `name` - Function name for the compiled kernel
473    /// * `program` - Bytecode program containing the strategy
474    /// * `config` - Kernel configuration with series mappings
475    ///
476    /// # Returns
477    /// A function pointer with signature:
478    /// `fn(cursor_index: usize, series_ptrs: *const *const f64, table_count: usize, state_ptr: *mut u8) -> i32`
479    ///
480    /// # Generated Code Pattern
481    ///
482    /// For a strategy like:
483    /// ```shape
484    /// let spy_price = context.spy    // series index 0
485    /// let vix_level = context.vix    // series index 1
486    /// if vix_level > 25.0 && state.position == 0 {
487    ///     state.signal = 1.0
488    /// }
489    /// ```
490    ///
491    /// The kernel generates:
492    /// ```asm
493    /// ; spy_price = context.spy (series index 0)
494    /// mov rax, [series_ptrs + 0*8]     ; series 0 pointer
495    /// mov xmm0, [rax + cursor_index*8] ; spy value
496    ///
497    /// ; vix_level = context.vix (series index 1)
498    /// mov rax, [series_ptrs + 1*8]     ; series 1 pointer
499    /// mov xmm1, [rax + cursor_index*8] ; vix value
500    ///
501    /// ; comparison and conditional store
502    /// mov xmm2, [const_25.0]
503    /// ucomisd xmm1, xmm2
504    /// jbe skip
505    /// ; ... check state.position == 0 ...
506    /// mov qword [state_ptr + signal_offset], 1.0
507    /// skip:
508    /// ```
509    #[inline(always)]
510    pub fn compile_correlated_kernel(
511        &mut self,
512        name: &str,
513        program: &BytecodeProgram,
514        config: &SimulationKernelConfig,
515    ) -> Result<CorrelatedKernelFn, String> {
516        // Validate config is for multi-series mode
517        if !config.is_multi_table() {
518            return Err(
519                "compile_correlated_kernel requires multi-series config (use new_multi_table)"
520                    .to_string(),
521            );
522        }
523
524        // Correlated kernel ABI:
525        // fn(cursor_index: usize, series_ptrs: *const *const f64, table_count: usize, state_ptr: *mut u8) -> i32
526        let mut sig = self.module.make_signature();
527        sig.params.push(AbiParam::new(types::I64)); // cursor_index
528        sig.params.push(AbiParam::new(types::I64)); // series_ptrs
529        sig.params.push(AbiParam::new(types::I64)); // table_count
530        sig.params.push(AbiParam::new(types::I64)); // state_ptr
531        sig.returns.push(AbiParam::new(types::I32)); // result code
532
533        let func_id = self
534            .module
535            .declare_function(name, Linkage::Export, &sig)
536            .map_err(|e| format!("Failed to declare correlated kernel function: {}", e))?;
537
538        let mut ctx = self.module.make_context();
539        ctx.func.signature = sig;
540
541        let mut func_builder_ctx = FunctionBuilderContext::new();
542        {
543            let mut builder = FunctionBuilder::new(&mut ctx.func, &mut func_builder_ctx);
544            let entry_block = builder.create_block();
545            builder.append_block_params_for_function_params(entry_block);
546            builder.switch_to_block(entry_block);
547            builder.seal_block(entry_block);
548
549            // Get kernel parameters
550            let cursor_index = builder.block_params(entry_block)[0];
551            let series_ptrs = builder.block_params(entry_block)[1];
552            let _table_count = builder.block_params(entry_block)[2]; // For validation/debugging
553            let state_ptr = builder.block_params(entry_block)[3];
554
555            // Build correlated kernel IR
556            // Note: table_count is known at compile time from config, used for validation
557            let result = self.build_correlated_kernel_ir(
558                &mut builder,
559                program,
560                config,
561                cursor_index,
562                series_ptrs,
563                state_ptr,
564            )?;
565
566            builder.ins().return_(&[result]);
567            builder.finalize();
568        }
569
570        self.module
571            .define_function(func_id, &mut ctx)
572            .map_err(|e| format!("Failed to define correlated kernel function: {:?}", e))?;
573
574        self.module.clear_context(&mut ctx);
575        self.module
576            .finalize_definitions()
577            .map_err(|e| format!("Failed to finalize correlated kernel: {:?}", e))?;
578
579        let code_ptr = self.module.get_finalized_function(func_id);
580        self.compiled_functions.insert(name.to_string(), code_ptr);
581
582        Ok(unsafe { std::mem::transmute(code_ptr) })
583    }
584
585    /// Build correlated kernel IR for multi-series access.
586    ///
587    /// Handles series access via compile-time resolved indices:
588    /// - `context.spy` → `series_ptrs[0][cursor_idx]` (if spy mapped to index 0)
589    /// - `context.vix` → `series_ptrs[1][cursor_idx]` (if vix mapped to index 1)
590    fn build_correlated_kernel_ir(
591        &mut self,
592        builder: &mut FunctionBuilder,
593        program: &BytecodeProgram,
594        config: &SimulationKernelConfig,
595        cursor_index: Value,
596        series_ptrs: Value,
597        state_ptr: Value,
598    ) -> Result<Value, String> {
599        Err("Correlated kernel compilation requires v2 runtime migration".to_string())
600    }
601}