use std::collections::HashMap;
use cranelift_jit::JITModule;
use crate::ast::PolydatNode;
use crate::kernel::ProvMask;
#[derive(Clone)]
pub struct JitCode(std::sync::Arc<FinalizedModule>);
struct FinalizedModule {
#[allow(dead_code)]
module: JITModule,
#[allow(dead_code)]
kits: Vec<super::codegen::SlotKitRef>,
fallible: bool,
}
#[derive(Clone, Default)]
pub(crate) struct ScratchPlan {
pub(crate) elems: Vec<crate::ast::ScratchElem>,
pub(crate) refs: Vec<(usize, usize)>,
}
unsafe impl Send for FinalizedModule {}
unsafe impl Sync for FinalizedModule {}
impl JitCode {
pub(crate) fn new(
module: JITModule,
kits: Vec<super::codegen::SlotKitRef>,
fallible: bool,
) -> Self {
JitCode(std::sync::Arc::new(FinalizedModule {
module,
kits,
fallible,
}))
}
pub(crate) fn fallible(&self) -> bool {
self.0.fallible
}
}
pub type JitParts = (super::codegen::NativeFn, JitCode);
pub(super) struct JitCore {
pub(super) buffer: Vec<u64>,
pub(super) coord_count: usize,
pub(super) output_map: HashMap<String, usize>,
pub(super) guard_slots: Vec<bool>,
pub(super) output_types: HashMap<String, crate::ast::PortType>,
pub(super) externs: crate::compile::externs::Externs,
pub(super) traversals: std::sync::Arc<[crate::dsl::traversal::Traversal]>,
pub(super) _module: JitCode,
pub(super) fallible: bool,
pub(super) _nodes: std::sync::Arc<Vec<Box<dyn PolydatNode>>>,
pub(super) drive: crate::compile::Drive,
pub(super) sites: std::sync::Arc<crate::compile::Attribution>,
pub(super) tracker: usize,
pub(super) scratch: Vec<crate::ast::ScratchBuf>,
pub(super) ref_scratch: Vec<(usize, usize)>,
pub(super) volatile_steps: Vec<usize>,
}
impl Clone for JitCore {
fn clone(&self) -> Self {
let mut core = JitCore {
buffer: self.buffer.clone(),
coord_count: self.coord_count,
output_map: self.output_map.clone(),
guard_slots: self.guard_slots.clone(),
output_types: self.output_types.clone(),
externs: self.externs.clone(),
traversals: self.traversals.clone(),
_module: self._module.clone(),
fallible: self.fallible,
_nodes: self._nodes.clone(),
drive: self.drive.clone(),
sites: self.sites.clone(),
tracker: self.tracker,
scratch: self.scratch.clone(),
ref_scratch: self.ref_scratch.clone(),
volatile_steps: self.volatile_steps.clone(),
};
for &(slot, idx) in &core.ref_scratch {
let (p, l) = core.scratch[idx].ptr_len();
core.buffer[slot] = p;
core.buffer[slot + 1] = l;
}
core.externs.seed(&mut core.buffer, None);
core
}
}
impl JitCore {
pub(super) fn slot_value(&self, slot: usize, ty: crate::ast::PortType) -> crate::ast::Value {
crate::compile::marshal::decode_output(&self.buffer, slot, ty)
}
pub(super) fn plan(&self) -> crate::EnginePlan {
crate::EnginePlan {
native_segments: 1,
..Default::default()
}
}
pub(super) fn invalidate_all(&mut self) {
self.drive.stale = true;
}
pub(super) fn new(
total_slots: usize,
coord_count: usize,
output_map: HashMap<String, usize>,
code: JitCode,
nodes: Vec<Box<dyn PolydatNode>>,
scratch: ScratchPlan,
volatile_steps: Vec<usize>,
) -> Self {
Self {
buffer: vec![0u64; total_slots + 1],
coord_count,
output_map,
guard_slots: Vec::new(),
output_types: HashMap::new(),
externs: crate::compile::externs::Externs::default(),
traversals: Vec::new().into(),
fallible: code.fallible(),
_module: code,
_nodes: std::sync::Arc::new(nodes),
drive: crate::compile::Drive::default(),
sites: std::sync::Arc::default(),
tracker: total_slots,
scratch: scratch
.elems
.iter()
.map(|e| crate::ast::ScratchBuf::new(*e))
.collect(),
ref_scratch: scratch.refs,
volatile_steps,
}
}
#[inline]
fn has_volatile(&self) -> bool {
!self.volatile_steps.is_empty()
}
#[cfg(debug_assertions)]
fn validate_refs(&self) {
for &(slot, idx) in &self.ref_scratch {
let (p, l) = self.scratch[idx].ptr_len();
assert!(
self.buffer[slot] == p && self.buffer[slot + 1] == l,
"S9 ref-validator: slot pair ({slot}, {}) = ({:#x}, {}) does not match \
scratch[{idx}] = ({p:#x}, {l})",
slot + 1,
self.buffer[slot],
self.buffer[slot + 1],
);
}
}
pub(super) fn set_externs(&mut self, externs: crate::compile::externs::Externs) {
externs.seed(&mut self.buffer, None);
self.externs = externs;
}
fn set_extern(&mut self, name: &str, value: crate::ast::Value) -> Result<usize, String> {
Ok(self.externs.set(name, value, &mut self.buffer)?.0)
}
fn set_extern_at(&mut self, index: usize, value: crate::ast::Value) -> Result<usize, String> {
Ok(self.externs.set_at(index, value, &mut self.buffer)?.0)
}
fn attach_cell(&mut self, name: &str, cell: crate::kernel::SharedCell) -> Result<(), String> {
self.externs.attach_cell(name, cell)?;
self.drive.stale = true;
Ok(())
}
#[inline]
pub(super) fn run(&mut self, native: impl FnOnce()) {
if self.externs.cells_dirty() {
self.externs.refresh_cells(&mut self.buffer);
}
if let Some((name, ty)) = self.externs.first_unset() {
panic!(
"extern '{name}' ({ty}) has no value: it has no default, so set it with \
set_input before the first run (native code cannot carry `None`; \
docs/design/engine_parity.md, A12)"
);
}
if !self.fallible {
native();
} else {
self.buffer[self.tracker] = u64::MAX;
let capture = crate::kernel::engines::EvalPanicCaptureGuard::arm();
let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
super::codegen::invoke_with_catch(native)
}));
drop(capture);
if let Err(payload) = outcome {
let step = self.buffer[self.tracker] as usize;
let sites = std::sync::Arc::clone(&self.sites);
sites.reraise(payload, step, &self.buffer, None);
}
}
#[cfg(debug_assertions)]
self.validate_refs();
}
}
macro_rules! jit_accessors {
() => {
pub fn coord_count(&self) -> usize {
self.core.coord_count
}
pub fn resolve_output(&self, name: &str) -> Option<usize> {
self.core.output_map.get(name).copied()
}
#[inline]
pub fn get(&self, name: &str) -> u64 {
self.get_slot(self.core.output_map[name])
}
#[inline]
pub fn get_slot(&self, slot: usize) -> u64 {
if self.core.guard_slots.get(slot).copied().unwrap_or(false) {
panic!(
"slot {slot} is Ref2-colored; a raw u64 read would leak an interior \
address. Use get_value to decode it."
);
}
self.core.buffer[slot]
}
pub fn get_value(&self, name: &str) -> crate::ast::Value {
let slot = self.core.output_map[name];
let ty = self
.core
.output_types
.get(name)
.copied()
.unwrap_or(crate::ast::PortType::U64);
crate::compile::marshal::decode_output(&self.core.buffer, slot, ty)
}
pub(crate) fn set_slot_info(
&mut self,
guard_slots: Vec<bool>,
output_types: HashMap<String, crate::ast::PortType>,
) {
self.core.guard_slots = guard_slots;
self.core.output_types = output_types;
}
pub(crate) fn set_attribution(
&mut self,
sites: std::sync::Arc<crate::compile::Attribution>,
) {
self.core.sites = sites;
}
pub fn set_input(&mut self, name: &str, value: crate::ast::Value) -> Result<(), String> {
let slot = self.core.set_extern(name, value)?;
self.mark_input_changed(slot);
Ok(())
}
pub fn set_input_at(
&mut self,
index: usize,
value: crate::ast::Value,
) -> Result<(), String> {
let slot = self.core.set_extern_at(index, value)?;
self.mark_input_changed(slot);
Ok(())
}
pub fn externs(&self) -> Vec<(&str, crate::ast::PortType)> {
self.core.externs.names()
}
fn mark_all_dirty(&mut self) {
for i in 0..self.core.coord_count {
self.mark_input_changed(i);
}
}
fn pull_value(&mut self, name: &str) -> crate::ast::Value {
if self.core.drive.stale || self.core.externs.cells_dirty() {
self.eval_pending();
self.core.drive.stale = false;
}
self.get_value(name)
}
fn pull_value_at(&mut self, index: usize) -> crate::ast::Value {
let name = self
.core
.externs
.output_names()
.get(index)
.cloned()
.unwrap_or_else(|| panic!("no output at index {index}"));
self.pull_value(&name)
}
fn eval_pending(&mut self) {
let coords = std::mem::take(&mut self.core.drive.coords);
self.eval(&coords);
self.core.drive.coords = coords;
}
pub fn cursor_schemas(&self) -> &[crate::iteration::source::SourceSchema] {
self.core.externs.cursor_schemas()
}
pub fn set_cursor(
&mut self,
name: &str,
partition: &crate::iteration::cursor_partition::Partition,
) -> Result<(), String> {
for (slot, value) in self.core.externs.cursor_writes(name, partition)? {
self.set_input(&slot, value)?;
}
Ok(())
}
};
}
#[derive(Clone)]
#[doc(hidden)]
pub struct JitKernelRaw {
pub(super) core: JitCore,
pub(super) code_fn: super::codegen::NativeFn,
}
impl JitKernelRaw {
#[inline]
pub fn eval(&mut self, coords: &[u64]) {
for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
if self.core.buffer[i] != c {
self.core.buffer[i] = c;
}
}
let code_fn = self.code_fn;
let buf_ptr_const = self.core.buffer.as_ptr();
let buf_ptr_mut = self.core.buffer.as_mut_ptr();
let sc = self.core.scratch.as_mut_ptr();
self.core.run(move || unsafe {
(code_fn)(buf_ptr_const, buf_ptr_mut, sc);
});
}
#[inline]
pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
self.eval(coords);
self.core.buffer[slot]
}
pub fn into_parts(self) -> JitParts {
(self.code_fn, self.core._module)
}
fn mark_input_changed(&mut self, _slot: usize) {}
jit_accessors!();
}
#[derive(Clone)]
#[doc(hidden)]
pub struct JitKernelPush {
pub(super) core: JitCore,
pub(super) code_fn_prov: super::codegen::NativeProvFn,
pub(super) node_clean: Vec<u8>,
pub(super) input_dependents: Vec<Vec<usize>>,
}
impl JitKernelPush {
#[inline]
fn set_inputs(&mut self, coords: &[u64]) {
for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
if self.core.buffer[i] != c {
self.core.buffer[i] = c;
self.mark_input_changed(i);
}
}
for &step_idx in &self.core.volatile_steps {
self.node_clean[step_idx] = 0;
}
}
fn mark_input_changed(&mut self, slot: usize) {
if slot < self.input_dependents.len() {
for &step_idx in &self.input_dependents[slot] {
self.node_clean[step_idx] = 0;
}
}
for &step_idx in &self.core.volatile_steps {
self.node_clean[step_idx] = 0;
}
}
#[inline]
pub fn eval(&mut self, coords: &[u64]) {
self.set_inputs(coords);
let code_fn = self.code_fn_prov;
let buf_const = self.core.buffer.as_ptr();
let buf_mut = self.core.buffer.as_mut_ptr();
let sc = self.core.scratch.as_mut_ptr();
let clean_mut = self.node_clean.as_mut_ptr();
self.core.run(move || unsafe {
(code_fn)(buf_const, buf_mut, sc, clean_mut);
});
}
#[inline]
pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
self.eval(coords);
self.core.buffer[slot]
}
jit_accessors!();
}
#[derive(Clone)]
#[doc(hidden)]
pub struct JitKernelPull {
pub(super) core: JitCore,
pub(super) code_fn: super::codegen::NativeFn,
pub(super) slot_provenance: Vec<ProvMask>,
pub(super) changed_mask: ProvMask,
pub(super) force_run: bool,
}
impl JitKernelPull {
#[inline]
fn set_inputs(&mut self, coords: &[u64]) {
self.changed_mask.clear();
for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
if self.core.buffer[i] != c {
self.core.buffer[i] = c;
self.changed_mask.set(i);
}
}
if self.core.has_volatile() {
self.force_run = true;
}
}
fn mark_input_changed(&mut self, _slot: usize) {
self.force_run = true;
}
#[inline]
pub fn eval(&mut self, coords: &[u64]) {
self.set_inputs(coords);
self.force_run = false;
let code_fn = self.code_fn;
let buf_const = self.core.buffer.as_ptr();
let buf_mut = self.core.buffer.as_mut_ptr();
let sc = self.core.scratch.as_mut_ptr();
self.core.run(move || unsafe {
(code_fn)(buf_const, buf_mut, sc);
});
}
#[inline]
pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
self.set_inputs(coords);
if !self.force_run
&& slot < self.slot_provenance.len()
&& !self.slot_provenance[slot].intersects(&self.changed_mask)
{
return self.core.buffer[slot];
}
self.force_run = false;
let code_fn = self.code_fn;
let buf_const = self.core.buffer.as_ptr();
let buf_mut = self.core.buffer.as_mut_ptr();
let sc = self.core.scratch.as_mut_ptr();
self.core.run(move || unsafe {
(code_fn)(buf_const, buf_mut, sc);
});
self.core.buffer[slot]
}
jit_accessors!();
}
#[derive(Clone)]
#[doc(hidden)]
pub struct JitKernelPushPull {
pub(super) core: JitCore,
pub(super) code_fn_prov: super::codegen::NativeProvFn,
pub(super) node_clean: Vec<u8>,
pub(super) input_dependents: Vec<Vec<usize>>,
pub(super) slot_provenance: Vec<ProvMask>,
pub(super) changed_mask: ProvMask,
pub(super) force_run: bool,
}
impl JitKernelPushPull {
#[inline]
fn set_inputs(&mut self, coords: &[u64]) {
self.changed_mask.clear();
for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
if self.core.buffer[i] != c {
self.core.buffer[i] = c;
self.changed_mask.set(i);
if i < self.input_dependents.len() {
for &step_idx in &self.input_dependents[i] {
self.node_clean[step_idx] = 0;
}
}
}
}
if self.core.has_volatile() {
for &step_idx in &self.core.volatile_steps {
self.node_clean[step_idx] = 0;
}
self.force_run = true;
}
}
fn mark_input_changed(&mut self, slot: usize) {
if slot < self.input_dependents.len() {
for &step_idx in &self.input_dependents[slot] {
self.node_clean[step_idx] = 0;
}
}
for &step_idx in &self.core.volatile_steps {
self.node_clean[step_idx] = 0;
}
self.force_run = true;
}
#[inline]
pub fn eval(&mut self, coords: &[u64]) {
self.set_inputs(coords);
self.force_run = false;
let code_fn = self.code_fn_prov;
let buf_const = self.core.buffer.as_ptr();
let buf_mut = self.core.buffer.as_mut_ptr();
let sc = self.core.scratch.as_mut_ptr();
let clean_mut = self.node_clean.as_mut_ptr();
self.core.run(move || unsafe {
(code_fn)(buf_const, buf_mut, sc, clean_mut);
});
}
#[inline]
pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
self.set_inputs(coords);
if !self.force_run
&& slot < self.slot_provenance.len()
&& !self.slot_provenance[slot].intersects(&self.changed_mask)
{
return self.core.buffer[slot];
}
self.force_run = false;
let code_fn = self.code_fn_prov;
let buf_const = self.core.buffer.as_ptr();
let buf_mut = self.core.buffer.as_mut_ptr();
let sc = self.core.scratch.as_mut_ptr();
let clean_mut = self.node_clean.as_mut_ptr();
self.core.run(move || unsafe {
(code_fn)(buf_const, buf_mut, sc, clean_mut);
});
self.core.buffer[slot]
}
jit_accessors!();
}
use crate::compile::select::{Engine, Provenance};
crate::compile::impl_kernel_trait!(JitKernelRaw, Engine::Native(Provenance::Raw));
crate::compile::impl_kernel_trait!(JitKernelPush, Engine::Native(Provenance::Push));
crate::compile::impl_kernel_trait!(JitKernelPull, Engine::Native(Provenance::Pull));
crate::compile::impl_kernel_trait!(JitKernelPushPull, Engine::Native(Provenance::PushPull));