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(super) struct JitCore {
pub(super) engine: crate::compile::select::Engine,
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>,
pub(super) cones: ConePlan,
pub(super) unit_clean: Vec<u8>,
pub(super) volatile_units: Vec<usize>,
pub(super) entry: super::codegen::NativeDispatchFn,
pub(super) outputs_at: Vec<(usize, crate::ast::PortType)>,
}
#[derive(Clone, Default)]
pub(super) struct ConePlan {
inputs: std::sync::Arc<[Box<[usize]>]>,
producer: std::sync::Arc<[usize]>,
unit_of: std::sync::Arc<[u32]>,
members: std::sync::Arc<[Box<[usize]>]>,
all: std::sync::Arc<[u32]>,
by_slot: Vec<Option<std::sync::Arc<[u32]>>>,
reads_by_slot: Vec<Option<std::sync::Arc<[usize]>>>,
}
impl ConePlan {
pub(super) fn new(
steps: &[(super::codegen::JitOp, Vec<usize>, Vec<usize>)],
slots: usize,
units: &crate::compile::fusion_units::UnitPlan,
outputs: impl IntoIterator<Item = usize>,
) -> Self {
let mut producer = vec![usize::MAX; slots + 1];
for (i, (_, _, outs)) in steps.iter().enumerate() {
for &s in outs {
if s < producer.len() {
producer[s] = i;
}
}
}
let mut plan = ConePlan {
inputs: steps
.iter()
.map(|(_, ins, _)| ins.clone().into_boxed_slice())
.collect(),
producer: producer.into(),
unit_of: units.unit_of.iter().map(|&u| u as u32).collect(),
members: units
.units
.iter()
.map(|m| m.clone().into_boxed_slice())
.collect(),
all: (0..units.units.len() as u32).collect(),
by_slot: vec![None; slots + 1],
reads_by_slot: vec![None; slots + 1],
};
for slot in outputs {
plan.of(slot);
}
plan
}
pub(super) fn unit_count(&self) -> usize {
self.all.len()
}
pub(super) fn unit_of(&self, step: usize) -> usize {
self.unit_of[step] as usize
}
fn unit_of_slot(&self, slot: usize) -> Option<usize> {
let step = self
.producer
.get(slot)
.copied()
.filter(|&p| p != usize::MAX)?;
Some(self.unit_of[step] as usize)
}
pub(super) fn all(&self) -> &[u32] {
&self.all
}
#[inline]
pub(super) fn of(&mut self, slot: usize) -> &[u32] {
if self.by_slot.get(slot).is_none_or(|c| c.is_none()) {
self.find(slot);
}
self.by_slot[slot].as_deref().unwrap_or(&[])
}
fn reads_of(&mut self, slot: usize) -> &[usize] {
if self.reads_by_slot.get(slot).is_none_or(|r| r.is_none()) {
self.find(slot);
}
self.reads_by_slot[slot].as_deref().unwrap_or(&[])
}
#[cold]
fn find(&mut self, slot: usize) {
if slot >= self.by_slot.len() {
self.by_slot.resize(slot + 1, None);
self.reads_by_slot.resize(slot + 1, None);
}
{
let mut unit_seen = vec![false; self.members.len()];
let producer_of = |s: usize| self.producer.get(s).copied().filter(|&p| p != usize::MAX);
let mut stack: Vec<usize> = producer_of(slot).into_iter().collect();
let mut units: Vec<u32> = Vec::new();
let mut reads: Vec<usize> = Vec::new();
if stack.is_empty() {
reads.push(slot);
}
while let Some(step) = stack.pop() {
let unit = self.unit_of[step];
if unit_seen[unit as usize] {
continue;
}
unit_seen[unit as usize] = true;
units.push(unit);
for &m in self.members[unit as usize].iter() {
for &s in self.inputs[m].iter() {
match producer_of(s) {
Some(p) => stack.push(p),
None => reads.push(s),
}
}
}
}
units.sort_unstable();
units.dedup();
reads.sort_unstable();
reads.dedup();
self.by_slot[slot] = Some(units.into());
self.reads_by_slot[slot] = Some(reads.into());
}
}
}
impl Clone for JitCore {
fn clone(&self) -> Self {
let mut core = JitCore {
engine: self.engine,
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(),
cones: self.cones.clone(),
unit_clean: self.unit_clean.clone(),
volatile_units: self.volatile_units.clone(),
entry: self.entry,
outputs_at: self.outputs_at.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 {
fn dirty_input(&mut self, _slot: usize) {}
fn program_identity(&self) -> usize {
std::sync::Arc::as_ptr(&self._nodes) as *const () as usize
}
fn output_cell_for(&self, name: &str) -> Option<crate::kernel::SharedCell> {
let slot = *self.output_map.get(name)?;
let uncomputed = self
.cones
.unit_of_slot(slot)
.is_some_and(|u| self.unit_clean[u] == 0);
let initial = if uncomputed {
crate::ast::Value::None
} else {
let ty = self
.output_types
.get(name)
.copied()
.unwrap_or(crate::ast::PortType::U64);
self.slot_value(slot, ty)
};
Some(self.externs.output_cell(slot, initial))
}
#[cold]
#[inline(never)]
pub(super) fn publish_slot(&self, slot: usize, value: &crate::ast::Value) {
if let Some(cell) = self.externs.published_output(slot) {
cell.publish(value.clone());
}
}
fn ref_entry(&self, slot: usize) -> &crate::ast::ScratchBuf {
match self.ref_scratch.iter().find(|(s, _)| *s == slot) {
Some(&(_, idx)) => &self.scratch[idx],
None if self.guard_slots.get(slot).copied().unwrap_or(false) => panic!(
"slot {slot} is a Ref pair owned by the CALLER (a kernel \
input) — read it on the caller side"
),
None => panic!("slot {slot} is not a Ref2-colored slot"),
}
}
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;
}
#[allow(clippy::too_many_arguments)]
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>,
entry: super::codegen::NativeDispatchFn,
cones: ConePlan,
) -> Self {
let mut volatile_units: Vec<usize> =
volatile_steps.iter().map(|&s| cones.unit_of(s)).collect();
volatile_units.sort_unstable();
volatile_units.dedup();
let unit_count = cones.unit_count();
let mut core = Self {
engine: crate::compile::select::Engine::PureNative(
crate::compile::select::Provenance::PushPull,
),
buffer: vec![0u64; total_slots + 1],
coord_count,
output_map,
guard_slots: Vec::new(),
output_types: HashMap::new(),
externs: crate::compile::externs::Externs::coordinates_only(coord_count),
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,
cones,
unit_clean: vec![0u8; unit_count],
volatile_units,
entry,
outputs_at: Vec::new(),
};
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
}
#[inline]
pub(super) fn run_units(&mut self, slot: Option<usize>) {
if self.externs.any_unset_read() {
self.refuse_unset_read(slot);
}
let units: &[u32] = match slot {
Some(s) => self.cones.of(s),
None => self.cones.all(),
};
let list = units.as_ptr();
let len = units.len() as u64;
let entry = self.entry;
let buf_const = self.buffer.as_ptr();
let buf_mut = self.buffer.as_mut_ptr();
let sc = self.scratch.as_mut_ptr();
let clean = self.unit_clean.as_mut_ptr();
self.run(move || unsafe {
(entry)(buf_const, buf_mut, sc, list, len, clean);
});
}
pub(super) fn dirty_all_units(&mut self) {
self.unit_clean.fill(0);
}
pub(super) fn dirty_volatile_units(&mut self) {
for &u in &self.volatile_units {
self.unit_clean[u] = 0;
}
}
fn output_at(&mut self, index: usize) -> (usize, crate::ast::PortType) {
if self.outputs_at.is_empty() {
self.outputs_at = self
.externs
.output_names()
.iter()
.map(|n| {
let slot = self.output_map[n];
let ty = self
.output_types
.get(n)
.copied()
.unwrap_or(crate::ast::PortType::U64);
(slot, ty)
})
.collect();
}
*self
.outputs_at
.get(index)
.unwrap_or_else(|| panic!("no output at index {index}"))
}
#[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, crate::kernel::WriteError> {
Ok(self.externs.set(name, value, &mut self.buffer)?.0)
}
fn set_extern_at(
&mut self,
index: usize,
value: crate::ast::Value,
) -> Result<usize, crate::kernel::WriteError> {
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(())
}
#[cold]
#[inline(never)]
pub(super) fn take_refreshed(&mut self) -> Vec<usize> {
self.externs.refresh_cells(&mut self.buffer);
self.externs.take_changed()
}
#[inline]
fn run(&mut self, native: impl FnOnce()) {
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();
}
#[cold]
#[inline(never)]
fn refuse_unset_read(&mut self, slot: Option<usize>) {
let unset = match slot {
None => self.externs.first_unset(|_| true),
Some(s) => {
let reads = self.cones.reads_of(s);
self.externs
.first_unset(|x| reads.binary_search(&x).is_ok())
}
};
let Some((name, ty)) = unset else {
return;
};
panic!(
"extern '{name}' ({ty}) has no value and this evaluation reads it, on \
the pure native tier, which cannot carry a `None`: every step is native \
code and there is no closure to propagate one through. Either it was \
declared without a default and never set, or a host cleared it after \
the build. Set it with set_input before pulling, or run this program on \
`native`, which answers a cleared extern with `None` as the interpreter \
does (docs/design/engines.md §3.3)"
);
}
pub(super) fn fold_constants(
&mut self,
folded: &[(super::codegen::JitOp, Vec<usize>, Vec<usize>)],
origin: &[usize],
total_slots: usize,
) -> Result<(), crate::KernelError> {
if folded.is_empty() {
return Ok(());
}
let (code_fn, code) = super::codegen::compile_jit_entry(folded, Some(total_slots))
.map_err(|reason| crate::KernelError::ConstantFold { reason })?;
let buf_ptr_const = self.buffer.as_ptr();
let buf_ptr_mut = self.buffer.as_mut_ptr();
let sc = self.scratch.as_mut_ptr();
let native = move || unsafe {
(code_fn)(buf_ptr_const, buf_ptr_mut, sc);
};
if !code.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 step = origin.get(step).copied().unwrap_or(step);
let sites = std::sync::Arc::clone(&self.sites);
return Err(crate::KernelError::ConstantFold {
reason: sites.describe(payload, step, &self.buffer, None),
});
}
}
drop(code);
Ok(())
}
}
macro_rules! jit_accessors {
() => {
crate::compile::ref_readers!();
pub fn coord_count(&self) -> usize {
self.core.externs.coordinate_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(crate) fn fold_constants(
&mut self,
folded: &[(super::codegen::JitOp, Vec<usize>, Vec<usize>)],
origin: &[usize],
total_slots: usize,
) -> Result<(), crate::KernelError> {
self.core.fold_constants(folded, origin, total_slots)
}
pub fn set_input(
&mut self,
name: &str,
value: crate::ast::Value,
) -> Result<(), crate::kernel::WriteError> {
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<(), crate::kernel::WriteError> {
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 {
let slot = self.core.output_map[name];
let ty = self
.core
.output_types
.get(name)
.copied()
.unwrap_or(crate::ast::PortType::U64);
self.pull_publishing(slot, ty)
}
fn pull_value_at(&mut self, index: usize) -> crate::ast::Value {
let (slot, ty) = self.core.output_at(index);
self.pull_publishing(slot, ty)
}
#[inline]
fn pull_publishing(&mut self, slot: usize, ty: crate::ast::PortType) -> crate::ast::Value {
if self.core.externs.broadcasts() {
let value = self.pull_slot(slot, ty);
self.core.publish_slot(slot, &value);
return value;
}
self.pull_slot(slot, ty)
}
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<(), crate::kernel::WriteError> {
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,
}
impl JitKernelRaw {
fn pull_slot(&mut self, slot: usize, ty: crate::ast::PortType) -> crate::ast::Value {
if self.core.drive.stale || self.core.externs.cells_dirty() {
let coords = std::mem::take(&mut self.core.drive.coords);
self.write_coords(&coords);
self.core.drive.coords = coords;
self.core.drive.stale = false;
self.refresh_cells();
self.core.dirty_all_units();
} else if self.core.has_volatile() {
self.core.dirty_volatile_units();
}
self.core.run_units(Some(slot));
self.core.slot_value(slot, ty)
}
#[inline]
fn write_coords(&mut self, coords: &[u64]) {
for (i, &c) in coords
.iter()
.enumerate()
.take(self.core.externs.coordinate_slots())
{
if self.core.buffer[i] != c {
self.core.buffer[i] = c;
}
}
}
#[inline]
pub fn eval(&mut self, coords: &[u64]) {
self.write_coords(coords);
self.refresh_cells();
self.core.dirty_all_units();
self.core.run_units(None);
}
#[inline]
fn refresh_cells(&mut self) {
if self.core.externs.cells_dirty() {
let changed = self.core.take_refreshed();
self.core.externs.return_changed(changed);
}
}
#[inline]
pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
self.eval(coords);
self.core.buffer[slot]
}
fn mark_input_changed(&mut self, _slot: usize) {
self.core.dirty_all_units();
}
jit_accessors!();
}
#[derive(Clone)]
#[doc(hidden)]
pub struct JitKernelPushPull {
pub(super) core: JitCore,
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.externs.coordinate_slots())
{
if self.core.buffer[i] != c {
self.core.buffer[i] = c;
self.changed_mask.set(i);
self.dirty_dependents(i);
}
}
if self.core.has_volatile() {
self.core.dirty_volatile_units();
self.force_run = true;
}
}
#[inline]
fn dirty_dependents(&mut self, slot: usize) {
if let Some(units) = self.input_dependents.get(slot) {
for &u in units {
self.core.unit_clean[u] = 0;
}
}
}
fn mark_input_changed(&mut self, slot: usize) {
self.dirty_dependents(slot);
self.core.dirty_volatile_units();
self.force_run = true;
}
fn pull_slot(&mut self, slot: usize, ty: crate::ast::PortType) -> crate::ast::Value {
if self.core.drive.stale {
let coords = std::mem::take(&mut self.core.drive.coords);
self.set_inputs(&coords);
self.core.drive.coords = coords;
self.core.drive.stale = false;
}
self.refresh_cells();
if self.core.has_volatile() {
self.core.dirty_volatile_units();
}
self.core.run_units(Some(slot));
self.core.slot_value(slot, ty)
}
#[inline]
pub fn eval(&mut self, coords: &[u64]) {
self.set_inputs(coords);
self.refresh_cells();
self.force_run = false;
self.core.run_units(None);
}
#[inline]
fn refresh_cells(&mut self) {
if self.core.externs.cells_dirty() {
self.dirty_refreshed();
}
}
#[cold]
#[inline(never)]
fn dirty_refreshed(&mut self) {
let changed = self.core.take_refreshed();
for &slot in &changed {
self.mark_input_changed(slot);
}
self.core.externs.return_changed(changed);
}
#[inline]
pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
self.set_inputs(coords);
self.refresh_cells();
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;
self.core.run_units(Some(slot));
self.core.buffer[slot]
}
jit_accessors!();
}
crate::compile::impl_kernel_trait!(JitKernelRaw);
crate::compile::impl_kernel_trait!(JitKernelPushPull);
crate::compile::impl_slot_kernel!(JitKernelRaw);
crate::compile::impl_slot_kernel!(JitKernelPushPull);