1use std::collections::HashMap;
12
13use cranelift_jit::JITModule;
14
15use crate::ast::PolydatNode;
16use crate::kernel::ProvMask;
17
18#[derive(Clone)]
24pub struct JitCode(std::sync::Arc<FinalizedModule>);
25
26struct FinalizedModule {
29 #[allow(dead_code)]
30 module: JITModule,
31 #[allow(dead_code)]
32 kits: Vec<super::codegen::SlotKitRef>,
33 fallible: bool,
34}
35
36#[derive(Clone, Default)]
40pub(crate) struct ScratchPlan {
41 pub(crate) elems: Vec<crate::ast::ScratchElem>,
42 pub(crate) refs: Vec<(usize, usize)>,
43}
44
45unsafe impl Send for FinalizedModule {}
48unsafe impl Sync for FinalizedModule {}
49
50impl JitCode {
51 pub(crate) fn new(
52 module: JITModule,
53 kits: Vec<super::codegen::SlotKitRef>,
54 fallible: bool,
55 ) -> Self {
56 JitCode(std::sync::Arc::new(FinalizedModule {
57 module,
58 kits,
59 fallible,
60 }))
61 }
62
63 pub(crate) fn fallible(&self) -> bool {
70 self.0.fallible
71 }
72}
73
74pub type JitParts = (super::codegen::NativeFn, JitCode);
76
77pub(super) struct JitCore {
83 pub(super) buffer: Vec<u64>,
84 pub(super) coord_count: usize,
85 pub(super) output_map: HashMap<String, usize>,
86 pub(super) guard_slots: Vec<bool>,
90 pub(super) output_types: HashMap<String, crate::ast::PortType>,
92 pub(super) externs: crate::compile::externs::Externs,
94 pub(super) traversals: std::sync::Arc<[crate::dsl::traversal::Traversal]>,
97 pub(super) _module: JitCode,
98 pub(super) fallible: bool,
100 pub(super) _nodes: std::sync::Arc<Vec<Box<dyn PolydatNode>>>,
101 pub(super) drive: crate::compile::Drive,
104 pub(super) sites: std::sync::Arc<crate::compile::Attribution>,
106 pub(super) tracker: usize,
109 pub(super) scratch: Vec<crate::ast::ScratchBuf>,
112 pub(super) ref_scratch: Vec<(usize, usize)>,
115 pub(super) volatile_steps: Vec<usize>,
121}
122
123impl Clone for JitCore {
124 fn clone(&self) -> Self {
125 let mut core = JitCore {
126 buffer: self.buffer.clone(),
127 coord_count: self.coord_count,
128 output_map: self.output_map.clone(),
129 guard_slots: self.guard_slots.clone(),
130 output_types: self.output_types.clone(),
131 externs: self.externs.clone(),
132 traversals: self.traversals.clone(),
133 _module: self._module.clone(),
134 fallible: self.fallible,
135 _nodes: self._nodes.clone(),
136 drive: self.drive.clone(),
137 sites: self.sites.clone(),
138 tracker: self.tracker,
139 scratch: self.scratch.clone(),
140 ref_scratch: self.ref_scratch.clone(),
141 volatile_steps: self.volatile_steps.clone(),
142 };
143 for &(slot, idx) in &core.ref_scratch {
146 let (p, l) = core.scratch[idx].ptr_len();
147 core.buffer[slot] = p;
148 core.buffer[slot + 1] = l;
149 }
150 core.externs.seed(&mut core.buffer, None);
151 core
152 }
153}
154
155impl JitCore {
156 pub(super) fn slot_value(&self, slot: usize, ty: crate::ast::PortType) -> crate::ast::Value {
158 crate::compile::marshal::decode_output(&self.buffer, slot, ty)
159 }
160
161 pub(super) fn plan(&self) -> crate::EnginePlan {
163 crate::EnginePlan {
164 native_segments: 1,
165 ..Default::default()
166 }
167 }
168
169 pub(super) fn invalidate_all(&mut self) {
171 self.drive.stale = true;
172 }
173
174 pub(super) fn new(
175 total_slots: usize,
176 coord_count: usize,
177 output_map: HashMap<String, usize>,
178 code: JitCode,
179 nodes: Vec<Box<dyn PolydatNode>>,
180 scratch: ScratchPlan,
181 volatile_steps: Vec<usize>,
182 ) -> Self {
183 Self {
184 buffer: vec![0u64; total_slots + 1],
185 coord_count,
186 output_map,
187 guard_slots: Vec::new(),
188 output_types: HashMap::new(),
189 externs: crate::compile::externs::Externs::default(),
190 traversals: Vec::new().into(),
191 fallible: code.fallible(),
192 _module: code,
193 _nodes: std::sync::Arc::new(nodes),
194 drive: crate::compile::Drive::default(),
195 sites: std::sync::Arc::default(),
196 tracker: total_slots,
197 scratch: scratch
198 .elems
199 .iter()
200 .map(|e| crate::ast::ScratchBuf::new(*e))
201 .collect(),
202 ref_scratch: scratch.refs,
203 volatile_steps,
204 }
205 }
206
207 #[inline]
210 fn has_volatile(&self) -> bool {
211 !self.volatile_steps.is_empty()
212 }
213
214 #[cfg(debug_assertions)]
217 fn validate_refs(&self) {
218 for &(slot, idx) in &self.ref_scratch {
219 let (p, l) = self.scratch[idx].ptr_len();
220 assert!(
221 self.buffer[slot] == p && self.buffer[slot + 1] == l,
222 "S9 ref-validator: slot pair ({slot}, {}) = ({:#x}, {}) does not match \
223 scratch[{idx}] = ({p:#x}, {l})",
224 slot + 1,
225 self.buffer[slot],
226 self.buffer[slot + 1],
227 );
228 }
229 }
230
231 pub(super) fn set_externs(&mut self, externs: crate::compile::externs::Externs) {
233 externs.seed(&mut self.buffer, None);
234 self.externs = externs;
235 }
236
237 fn set_extern(&mut self, name: &str, value: crate::ast::Value) -> Result<usize, String> {
239 Ok(self.externs.set(name, value, &mut self.buffer)?.0)
240 }
241
242 fn set_extern_at(&mut self, index: usize, value: crate::ast::Value) -> Result<usize, String> {
244 Ok(self.externs.set_at(index, value, &mut self.buffer)?.0)
245 }
246
247 fn attach_cell(&mut self, name: &str, cell: crate::kernel::SharedCell) -> Result<(), String> {
250 self.externs.attach_cell(name, cell)?;
251 self.drive.stale = true;
252 Ok(())
253 }
254
255 #[inline]
259 pub(super) fn run(&mut self, native: impl FnOnce()) {
260 if self.externs.cells_dirty() {
261 self.externs.refresh_cells(&mut self.buffer);
262 }
263 if let Some((name, ty)) = self.externs.first_unset() {
264 panic!(
265 "extern '{name}' ({ty}) has no value: it has no default, so set it with \
266 set_input before the first run (native code cannot carry `None`; \
267 docs/design/engine_parity.md, A12)"
268 );
269 }
270 if !self.fallible {
276 native();
277 } else {
278 self.buffer[self.tracker] = u64::MAX;
279 let capture = crate::kernel::engines::EvalPanicCaptureGuard::arm();
280 let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
281 super::codegen::invoke_with_catch(native)
282 }));
283 drop(capture);
284 if let Err(payload) = outcome {
285 let step = self.buffer[self.tracker] as usize;
286 let sites = std::sync::Arc::clone(&self.sites);
287 sites.reraise(payload, step, &self.buffer, None);
288 }
289 }
290 #[cfg(debug_assertions)]
291 self.validate_refs();
292 }
293}
294
295macro_rules! jit_accessors {
296 () => {
297 pub fn coord_count(&self) -> usize {
299 self.core.coord_count
300 }
301
302 pub fn resolve_output(&self, name: &str) -> Option<usize> {
304 self.core.output_map.get(name).copied()
305 }
306
307 #[inline]
309 pub fn get(&self, name: &str) -> u64 {
310 self.get_slot(self.core.output_map[name])
311 }
312
313 #[inline]
317 pub fn get_slot(&self, slot: usize) -> u64 {
318 if self.core.guard_slots.get(slot).copied().unwrap_or(false) {
319 panic!(
320 "slot {slot} is Ref2-colored; a raw u64 read would leak an interior \
321 address. Use get_value to decode it."
322 );
323 }
324 self.core.buffer[slot]
325 }
326
327 pub fn get_value(&self, name: &str) -> crate::ast::Value {
331 let slot = self.core.output_map[name];
332 let ty = self
333 .core
334 .output_types
335 .get(name)
336 .copied()
337 .unwrap_or(crate::ast::PortType::U64);
338 crate::compile::marshal::decode_output(&self.core.buffer, slot, ty)
339 }
340
341 pub(crate) fn set_slot_info(
344 &mut self,
345 guard_slots: Vec<bool>,
346 output_types: HashMap<String, crate::ast::PortType>,
347 ) {
348 self.core.guard_slots = guard_slots;
349 self.core.output_types = output_types;
350 }
351
352 pub(crate) fn set_attribution(
354 &mut self,
355 sites: std::sync::Arc<crate::compile::Attribution>,
356 ) {
357 self.core.sites = sites;
358 }
359
360 pub fn set_input(&mut self, name: &str, value: crate::ast::Value) -> Result<(), String> {
366 let slot = self.core.set_extern(name, value)?;
367 self.mark_input_changed(slot);
368 Ok(())
369 }
370
371 pub fn set_input_at(
373 &mut self,
374 index: usize,
375 value: crate::ast::Value,
376 ) -> Result<(), String> {
377 let slot = self.core.set_extern_at(index, value)?;
378 self.mark_input_changed(slot);
379 Ok(())
380 }
381
382 pub fn externs(&self) -> Vec<(&str, crate::ast::PortType)> {
384 self.core.externs.names()
385 }
386
387 fn mark_all_dirty(&mut self) {
391 for i in 0..self.core.coord_count {
392 self.mark_input_changed(i);
393 }
394 }
395
396 fn pull_value(&mut self, name: &str) -> crate::ast::Value {
400 if self.core.drive.stale || self.core.externs.cells_dirty() {
402 self.eval_pending();
403 self.core.drive.stale = false;
404 }
405 self.get_value(name)
406 }
407
408 fn pull_value_at(&mut self, index: usize) -> crate::ast::Value {
411 let name = self
412 .core
413 .externs
414 .output_names()
415 .get(index)
416 .cloned()
417 .unwrap_or_else(|| panic!("no output at index {index}"));
418 self.pull_value(&name)
419 }
420
421 fn eval_pending(&mut self) {
423 let coords = std::mem::take(&mut self.core.drive.coords);
424 self.eval(&coords);
425 self.core.drive.coords = coords;
426 }
427
428 pub fn cursor_schemas(&self) -> &[crate::iteration::source::SourceSchema] {
432 self.core.externs.cursor_schemas()
433 }
434
435 pub fn set_cursor(
439 &mut self,
440 name: &str,
441 partition: &crate::iteration::cursor_partition::Partition,
442 ) -> Result<(), String> {
443 for (slot, value) in self.core.externs.cursor_writes(name, partition)? {
444 self.set_input(&slot, value)?;
445 }
446 Ok(())
447 }
448 };
449}
450
451#[derive(Clone)]
455#[doc(hidden)]
456pub struct JitKernelRaw {
457 pub(super) core: JitCore,
458 pub(super) code_fn: super::codegen::NativeFn,
459}
460
461impl JitKernelRaw {
462 #[inline]
470 pub fn eval(&mut self, coords: &[u64]) {
471 for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
475 if self.core.buffer[i] != c {
476 self.core.buffer[i] = c;
477 }
478 }
479 let code_fn = self.code_fn;
480 let buf_ptr_const = self.core.buffer.as_ptr();
481 let buf_ptr_mut = self.core.buffer.as_mut_ptr();
482 let sc = self.core.scratch.as_mut_ptr();
483 self.core.run(move || unsafe {
484 (code_fn)(buf_ptr_const, buf_ptr_mut, sc);
485 });
486 }
487
488 #[inline]
490 pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
491 self.eval(coords);
492 self.core.buffer[slot]
493 }
494
495 pub fn into_parts(self) -> JitParts {
498 (self.code_fn, self.core._module)
499 }
500
501 fn mark_input_changed(&mut self, _slot: usize) {}
503
504 jit_accessors!();
505}
506
507#[derive(Clone)]
511#[doc(hidden)]
512pub struct JitKernelPush {
513 pub(super) core: JitCore,
514 pub(super) code_fn_prov: super::codegen::NativeProvFn,
515 pub(super) node_clean: Vec<u8>,
516 pub(super) input_dependents: Vec<Vec<usize>>,
517}
518
519impl JitKernelPush {
520 #[inline]
521 fn set_inputs(&mut self, coords: &[u64]) {
522 for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
523 if self.core.buffer[i] != c {
524 self.core.buffer[i] = c;
525 self.mark_input_changed(i);
526 }
527 }
528 for &step_idx in &self.core.volatile_steps {
530 self.node_clean[step_idx] = 0;
531 }
532 }
533
534 fn mark_input_changed(&mut self, slot: usize) {
537 if slot < self.input_dependents.len() {
538 for &step_idx in &self.input_dependents[slot] {
539 self.node_clean[step_idx] = 0;
540 }
541 }
542 for &step_idx in &self.core.volatile_steps {
543 self.node_clean[step_idx] = 0;
544 }
545 }
546
547 #[inline]
549 pub fn eval(&mut self, coords: &[u64]) {
550 self.set_inputs(coords);
551 let code_fn = self.code_fn_prov;
552 let buf_const = self.core.buffer.as_ptr();
553 let buf_mut = self.core.buffer.as_mut_ptr();
554 let sc = self.core.scratch.as_mut_ptr();
555 let clean_mut = self.node_clean.as_mut_ptr();
556 self.core.run(move || unsafe {
557 (code_fn)(buf_const, buf_mut, sc, clean_mut);
558 });
559 }
560
561 #[inline]
563 pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
564 self.eval(coords);
565 self.core.buffer[slot]
566 }
567
568 jit_accessors!();
569}
570
571#[derive(Clone)]
576#[doc(hidden)]
577pub struct JitKernelPull {
578 pub(super) core: JitCore,
579 pub(super) code_fn: super::codegen::NativeFn,
580 pub(super) slot_provenance: Vec<ProvMask>,
581 pub(super) changed_mask: ProvMask,
582 pub(super) force_run: bool,
585}
586
587impl JitKernelPull {
588 #[inline]
589 fn set_inputs(&mut self, coords: &[u64]) {
590 self.changed_mask.clear();
591 for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
592 if self.core.buffer[i] != c {
593 self.core.buffer[i] = c;
594 self.changed_mask.set(i);
595 }
596 }
597 if self.core.has_volatile() {
600 self.force_run = true;
601 }
602 }
603
604 fn mark_input_changed(&mut self, _slot: usize) {
607 self.force_run = true;
608 }
609
610 #[inline]
612 pub fn eval(&mut self, coords: &[u64]) {
613 self.set_inputs(coords);
614 self.force_run = false;
615 let code_fn = self.code_fn;
616 let buf_const = self.core.buffer.as_ptr();
617 let buf_mut = self.core.buffer.as_mut_ptr();
618 let sc = self.core.scratch.as_mut_ptr();
619 self.core.run(move || unsafe {
620 (code_fn)(buf_const, buf_mut, sc);
621 });
622 }
623
624 #[inline]
627 pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
628 self.set_inputs(coords);
629 if !self.force_run
630 && slot < self.slot_provenance.len()
631 && !self.slot_provenance[slot].intersects(&self.changed_mask)
632 {
633 return self.core.buffer[slot];
634 }
635 self.force_run = false;
636 let code_fn = self.code_fn;
637 let buf_const = self.core.buffer.as_ptr();
638 let buf_mut = self.core.buffer.as_mut_ptr();
639 let sc = self.core.scratch.as_mut_ptr();
640 self.core.run(move || unsafe {
641 (code_fn)(buf_const, buf_mut, sc);
642 });
643 self.core.buffer[slot]
644 }
645
646 jit_accessors!();
647}
648
649#[derive(Clone)]
653#[doc(hidden)]
654pub struct JitKernelPushPull {
655 pub(super) core: JitCore,
656 pub(super) code_fn_prov: super::codegen::NativeProvFn,
657 pub(super) node_clean: Vec<u8>,
658 pub(super) input_dependents: Vec<Vec<usize>>,
659 pub(super) slot_provenance: Vec<ProvMask>,
660 pub(super) changed_mask: ProvMask,
661 pub(super) force_run: bool,
664}
665
666impl JitKernelPushPull {
667 #[inline]
668 fn set_inputs(&mut self, coords: &[u64]) {
669 self.changed_mask.clear();
670 for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
671 if self.core.buffer[i] != c {
672 self.core.buffer[i] = c;
673 self.changed_mask.set(i);
674 if i < self.input_dependents.len() {
675 for &step_idx in &self.input_dependents[i] {
676 self.node_clean[step_idx] = 0;
677 }
678 }
679 }
680 }
681 if self.core.has_volatile() {
684 for &step_idx in &self.core.volatile_steps {
685 self.node_clean[step_idx] = 0;
686 }
687 self.force_run = true;
688 }
689 }
690
691 fn mark_input_changed(&mut self, slot: usize) {
694 if slot < self.input_dependents.len() {
695 for &step_idx in &self.input_dependents[slot] {
696 self.node_clean[step_idx] = 0;
697 }
698 }
699 for &step_idx in &self.core.volatile_steps {
700 self.node_clean[step_idx] = 0;
701 }
702 self.force_run = true;
703 }
704
705 #[inline]
707 pub fn eval(&mut self, coords: &[u64]) {
708 self.set_inputs(coords);
709 self.force_run = false;
710 let code_fn = self.code_fn_prov;
711 let buf_const = self.core.buffer.as_ptr();
712 let buf_mut = self.core.buffer.as_mut_ptr();
713 let sc = self.core.scratch.as_mut_ptr();
714 let clean_mut = self.node_clean.as_mut_ptr();
715 self.core.run(move || unsafe {
716 (code_fn)(buf_const, buf_mut, sc, clean_mut);
717 });
718 }
719
720 #[inline]
723 pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
724 self.set_inputs(coords);
725 if !self.force_run
726 && slot < self.slot_provenance.len()
727 && !self.slot_provenance[slot].intersects(&self.changed_mask)
728 {
729 return self.core.buffer[slot];
730 }
731 self.force_run = false;
732 let code_fn = self.code_fn_prov;
733 let buf_const = self.core.buffer.as_ptr();
734 let buf_mut = self.core.buffer.as_mut_ptr();
735 let sc = self.core.scratch.as_mut_ptr();
736 let clean_mut = self.node_clean.as_mut_ptr();
737 self.core.run(move || unsafe {
738 (code_fn)(buf_const, buf_mut, sc, clean_mut);
739 });
740 self.core.buffer[slot]
741 }
742
743 jit_accessors!();
744}
745
746use crate::compile::select::{Engine, Provenance};
749
750crate::compile::impl_kernel_trait!(JitKernelRaw, Engine::Native(Provenance::Raw));
751crate::compile::impl_kernel_trait!(JitKernelPush, Engine::Native(Provenance::Push));
752crate::compile::impl_kernel_trait!(JitKernelPull, Engine::Native(Provenance::Pull));
753crate::compile::impl_kernel_trait!(JitKernelPushPull, Engine::Native(Provenance::PushPull));