use std::{
any::{type_name, Any},
iter::{self, zip},
sync::Arc,
};
use getset::{CopyGetters, Getters};
use openvm_circuit_primitives::{
var_range::{SharedVariableRangeCheckerChip, VariableRangeCheckerAir},
AnyChip, Chip, ColumnsAir,
};
use openvm_cpu_backend::CpuBackend;
use openvm_instructions::{PhantomDiscriminant, VmOpcode};
use openvm_stark_backend::{
interaction::BusIndex,
keygen::{types::MultiStarkProvingKey, MultiStarkKeygenBuilder},
prover::{AirProvingContext, MatrixDimensions, ProverBackend, ProvingContext},
AirRef, AnyAir, StarkEngine, StarkProtocolConfig, Val,
};
use rustc_hash::FxHashMap;
use tracing::info_span;
use super::{GenerationError, PhantomSubExecutor, SystemConfig};
use crate::{
arch::Arena,
system::{
memory::{BOUNDARY_AIR_OFFSET, MERKLE_AIR_OFFSET},
phantom::PhantomExecutor,
SystemAirInventory, SystemChipComplex, SystemRecords,
},
};
pub const PROGRAM_AIR_ID: usize = 0;
pub const PROGRAM_CACHED_TRACE_INDEX: usize = 0;
pub const CONNECTOR_AIR_ID: usize = 1;
pub const MEMORY_AIRS_START_IDX: usize = 2;
pub const BOUNDARY_AIR_ID: usize = MEMORY_AIRS_START_IDX + BOUNDARY_AIR_OFFSET;
pub const MERKLE_AIR_ID: usize = MEMORY_AIRS_START_IDX + MERKLE_AIR_OFFSET;
pub type ExecutorId = u32;
pub trait AnyAirWithColumns<SC: StarkProtocolConfig>: AnyAir<SC> + ColumnsAir {}
impl<SC, T> AnyAirWithColumns<SC> for T
where
SC: StarkProtocolConfig,
T: AnyAir<SC> + ColumnsAir,
{
}
pub type AirRefWithColumns<SC> = Arc<dyn AnyAirWithColumns<SC>>;
pub trait VmExecutionExtension<F> {
type Executor: AnyEnum;
fn extend_execution(
&self,
inventory: &mut ExecutorInventoryBuilder<F, Self::Executor>,
) -> Result<(), ExecutorInventoryError>;
}
pub trait VmCircuitExtension<SC: StarkProtocolConfig> {
fn extend_circuit(&self, inventory: &mut AirInventory<SC>) -> Result<(), AirInventoryError>;
}
pub trait VmProverExtension<E, RA, EXT>
where
E: StarkEngine,
EXT: VmExecutionExtension<Val<E::SC>> + VmCircuitExtension<E::SC>,
{
fn extend_prover(
&self,
extension: &EXT,
inventory: &mut ChipInventory<E::SC, RA, E::PB>,
) -> Result<(), ChipInventoryError>;
}
pub struct ExecutorInventory<E> {
config: SystemConfig,
pub instruction_lookup: FxHashMap<VmOpcode, ExecutorId>,
pub executors: Vec<E>,
ext_start: Vec<usize>,
}
pub struct ExecutorInventoryBuilder<'a, F, E> {
old_executors: Vec<&'a dyn AnyEnum>,
new_inventory: ExecutorInventory<E>,
phantom_executors: FxHashMap<PhantomDiscriminant, Arc<dyn PhantomSubExecutor<F>>>,
}
#[derive(Clone, Getters, CopyGetters)]
pub struct AirInventory<SC: StarkProtocolConfig> {
#[get = "pub"]
config: SystemConfig,
#[get = "pub"]
system: SystemAirInventory,
#[get = "pub"]
ext_airs: Vec<AirRefWithColumns<SC>>,
ext_start: Vec<usize>,
bus_idx_mgr: BusIndexManager,
}
#[derive(Clone, Copy, Debug, Default)]
pub struct BusIndexManager {
bus_idx_max: BusIndex,
}
#[derive(Getters)]
pub struct ChipInventory<SC, RA, PB>
where
SC: StarkProtocolConfig,
PB: ProverBackend,
{
#[get = "pub"]
airs: AirInventory<SC>,
#[get = "pub"]
chips: Vec<Box<dyn AnyChip<RA, PB>>>,
cur_num_exts: usize,
pub executor_idx_to_insertion_idx: Vec<usize>,
}
#[derive(Getters)]
pub struct VmChipComplex<SC, RA, PB, SCC>
where
SC: StarkProtocolConfig,
PB: ProverBackend,
{
pub system: SCC,
pub inventory: ChipInventory<SC, RA, PB>,
}
impl<E> ExecutorInventory<E> {
#[allow(clippy::new_without_default)]
pub fn new(config: SystemConfig) -> Self {
Self {
config,
instruction_lookup: Default::default(),
executors: Default::default(),
ext_start: vec![0],
}
}
pub fn add_executor(
&mut self,
executor: impl Into<E>,
opcodes: impl IntoIterator<Item = VmOpcode>,
) -> Result<(), ExecutorInventoryError> {
let opcodes: Vec<_> = opcodes.into_iter().collect();
for opcode in &opcodes {
if let Some(id) = self.instruction_lookup.get(opcode) {
return Err(ExecutorInventoryError::ExecutorExists {
opcode: *opcode,
id: *id,
});
}
}
let id = self.executors.len();
self.executors.push(executor.into());
for opcode in opcodes {
self.instruction_lookup
.insert(opcode, id.try_into().unwrap());
}
Ok(())
}
pub fn extend<F, E3, EXT>(
self,
other: &EXT,
) -> Result<ExecutorInventory<E3>, ExecutorInventoryError>
where
F: 'static,
E: Into<E3> + AnyEnum,
E3: AnyEnum,
EXT: VmExecutionExtension<F>,
EXT::Executor: Into<E3>,
{
let mut builder: ExecutorInventoryBuilder<F, EXT::Executor> = self.builder();
other.extend_execution(&mut builder)?;
let other_inventory = builder.new_inventory;
let other_phantom_executors = builder.phantom_executors;
let mut inventory_ext = self.transmute();
inventory_ext.append(other_inventory.transmute())?;
let phantom_chip: &mut PhantomExecutor<F> = inventory_ext
.find_executor_mut()
.next()
.expect("system always has phantom chip");
let phantom_executors = &mut phantom_chip.phantom_executors;
for (discriminant, sub_executor) in other_phantom_executors {
if phantom_executors
.insert(discriminant, sub_executor)
.is_some()
{
return Err(ExecutorInventoryError::PhantomSubExecutorExists { discriminant });
}
}
Ok(inventory_ext)
}
pub fn builder<F, E2>(&self) -> ExecutorInventoryBuilder<'_, F, E2>
where
F: 'static,
E: AnyEnum,
{
let old_executors = self.executors.iter().map(|e| e as &dyn AnyEnum).collect();
ExecutorInventoryBuilder {
old_executors,
new_inventory: ExecutorInventory::new(self.config.clone()),
phantom_executors: Default::default(),
}
}
pub fn transmute<E2>(self) -> ExecutorInventory<E2>
where
E: Into<E2>,
{
ExecutorInventory {
config: self.config,
instruction_lookup: self.instruction_lookup,
executors: self.executors.into_iter().map(|e| e.into()).collect(),
ext_start: self.ext_start,
}
}
fn append(&mut self, mut other: ExecutorInventory<E>) -> Result<(), ExecutorInventoryError> {
let num_executors = self.executors.len();
for (opcode, mut id) in other.instruction_lookup.into_iter() {
id = id.checked_add(num_executors.try_into().unwrap()).unwrap();
if let Some(old_id) = self.instruction_lookup.insert(opcode, id) {
return Err(ExecutorInventoryError::ExecutorExists { opcode, id: old_id });
}
}
for id in &mut other.ext_start {
*id = id.checked_add(num_executors).unwrap();
}
self.executors.append(&mut other.executors);
self.ext_start.append(&mut other.ext_start);
Ok(())
}
pub fn get_executor(&self, opcode: VmOpcode) -> Option<&E> {
let id = self.instruction_lookup.get(&opcode)?;
self.executors.get(*id as usize)
}
pub fn get_mut_executor(&mut self, opcode: &VmOpcode) -> Option<&mut E> {
let id = self.instruction_lookup.get(opcode)?;
self.executors.get_mut(*id as usize)
}
pub fn executors(&self) -> &[E] {
&self.executors
}
pub fn find_executor<EX: 'static>(&self) -> impl Iterator<Item = &'_ EX>
where
E: AnyEnum,
{
self.executors
.iter()
.filter_map(|e| e.as_any_kind().downcast_ref())
}
pub fn find_executor_mut<EX: 'static>(&mut self) -> impl Iterator<Item = &'_ mut EX>
where
E: AnyEnum,
{
self.executors
.iter_mut()
.filter_map(|e| e.as_any_kind_mut().downcast_mut())
}
pub fn config(&self) -> &SystemConfig {
&self.config
}
}
impl<F, E> ExecutorInventoryBuilder<'_, F, E> {
pub fn add_executor(
&mut self,
executor: impl Into<E>,
opcodes: impl IntoIterator<Item = VmOpcode>,
) -> Result<(), ExecutorInventoryError> {
self.new_inventory.add_executor(executor, opcodes)
}
pub fn add_phantom_sub_executor<PE>(
&mut self,
phantom_sub: PE,
discriminant: PhantomDiscriminant,
) -> Result<(), ExecutorInventoryError>
where
E: AnyEnum,
F: 'static,
PE: PhantomSubExecutor<F> + 'static,
{
let existing = self
.phantom_executors
.insert(discriminant, Arc::new(phantom_sub));
if existing.is_some() {
return Err(ExecutorInventoryError::PhantomSubExecutorExists { discriminant });
}
Ok(())
}
pub fn find_executor<EX: 'static>(&self) -> impl Iterator<Item = &'_ EX>
where
E: AnyEnum,
{
self.old_executors
.iter()
.filter_map(|e| e.as_any_kind().downcast_ref())
}
pub fn pointer_max_bits(&self) -> usize {
self.new_inventory.config().memory_config.pointer_max_bits
}
}
impl<SC: StarkProtocolConfig> AirInventory<SC> {
pub(crate) fn new(
config: SystemConfig,
system: SystemAirInventory,
bus_idx_mgr: BusIndexManager,
) -> Self {
Self {
config,
system,
ext_start: Vec::new(),
ext_airs: Vec::new(),
bus_idx_mgr,
}
}
pub fn start_new_extension(&mut self) {
self.ext_start.push(self.ext_airs.len());
}
pub fn new_bus_idx(&mut self) -> BusIndex {
self.bus_idx_mgr.new_bus_idx()
}
pub fn find_air<A: 'static>(&self) -> impl Iterator<Item = &'_ A> {
self.ext_airs
.iter()
.filter_map(|air| air.as_any().downcast_ref())
}
pub fn add_air<A: AnyAirWithColumns<SC> + 'static>(&mut self, air: A) {
self.add_air_ref(Arc::new(air));
}
pub fn add_air_ref(&mut self, air: AirRefWithColumns<SC>) {
self.ext_airs.push(air);
}
pub fn range_checker(&self) -> &VariableRangeCheckerAir {
self.find_air()
.next()
.expect("system always has range checker AIR")
}
pub fn into_airs(self) -> impl Iterator<Item = AirRefWithColumns<SC>> {
self.system
.into_airs()
.into_iter()
.chain(self.ext_airs.into_iter().rev())
}
pub fn keygen(self, config: &SC) -> MultiStarkProvingKey<SC> {
let system_config = self.config.clone();
let mut keygen_builder = MultiStarkKeygenBuilder::new(config.clone());
for (air_id, air) in self.into_airs().enumerate() {
if system_config.is_required_air_id(air_id) {
keygen_builder.add_required_air(air as AirRef<_>);
} else {
keygen_builder.add_air(air as AirRef<_>);
}
}
keygen_builder.generate_pk().unwrap()
}
pub fn num_airs(&self) -> usize {
self.config.num_airs() + self.ext_airs.len()
}
pub fn pointer_max_bits(&self) -> usize {
self.config.memory_config.pointer_max_bits
}
}
impl BusIndexManager {
pub fn new() -> Self {
Self { bus_idx_max: 0 }
}
pub fn new_bus_idx(&mut self) -> BusIndex {
let idx = self.bus_idx_max;
self.bus_idx_max = self.bus_idx_max.checked_add(1).unwrap();
idx
}
}
impl<SC, RA, PB> ChipInventory<SC, RA, PB>
where
SC: StarkProtocolConfig,
PB: ProverBackend,
{
pub fn new(airs: AirInventory<SC>) -> Self {
Self {
airs,
chips: Vec::new(),
cur_num_exts: 0,
executor_idx_to_insertion_idx: Vec::new(),
}
}
pub fn config(&self) -> &SystemConfig {
&self.airs.config
}
pub fn start_new_extension(&mut self) -> Result<(), ChipInventoryError> {
if self.cur_num_exts >= self.airs.ext_start.len() {
return Err(ChipInventoryError::MissingCircuitExtension(
self.airs.ext_start.len(),
));
}
if self.chips.len() != self.airs.ext_start[self.cur_num_exts] {
return Err(ChipInventoryError::MissingChip {
actual: self.chips.len(),
expected: self.airs.ext_start[self.cur_num_exts],
});
}
self.cur_num_exts += 1;
Ok(())
}
pub fn next_air<A: 'static>(&self) -> Result<&A, ChipInventoryError> {
let cur_idx = self.chips.len();
self.airs
.ext_airs
.get(cur_idx)
.and_then(|air| air.as_any().downcast_ref())
.ok_or_else(|| ChipInventoryError::AirNotFound {
name: type_name::<A>().to_string(),
})
}
pub fn find_chip<C: 'static>(&self) -> impl Iterator<Item = &'_ C> {
self.chips.iter().filter_map(|c| c.as_any().downcast_ref())
}
pub fn add_periphery_chip<C: Chip<RA, PB> + 'static>(&mut self, chip: C) {
self.chips.push(Box::new(chip));
}
pub fn add_executor_chip<C: Chip<RA, PB> + 'static>(&mut self, chip: C) {
tracing::debug!("add_executor_chip: {}", type_name::<C>());
self.executor_idx_to_insertion_idx.push(self.chips.len());
self.chips.push(Box::new(chip));
}
pub fn executor_idx_to_air_idx(&self) -> Vec<usize> {
let num_airs = self.airs.num_airs();
assert_eq!(
num_airs,
self.config().num_airs() + self.chips.len(),
"Number of chips does not match number of AIRs"
);
self.executor_idx_to_insertion_idx
.iter()
.map(|insertion_idx| {
num_airs
.checked_sub(insertion_idx.checked_add(1).unwrap())
.unwrap_or_else(|| {
panic!(
"Attempt to subtract num_airs={num_airs} by {}",
insertion_idx + 1
)
})
})
.collect()
}
pub fn timestamp_max_bits(&self) -> usize {
self.airs.config().memory_config.timestamp_max_bits
}
pub fn constant_trace_heights(&self) -> Vec<Option<usize>> {
let num_system = self.airs.config().num_airs();
let mut heights = vec![None; num_system];
heights.extend(
self.chips
.iter()
.rev()
.map(|chip| chip.constant_trace_height()),
);
heights
}
}
impl<SC, RA> ChipInventory<SC, RA, CpuBackend<SC>>
where
SC: StarkProtocolConfig,
{
pub fn range_checker(&self) -> Result<&SharedVariableRangeCheckerChip, ChipInventoryError> {
self.find_chip::<SharedVariableRangeCheckerChip>()
.next()
.ok_or_else(|| ChipInventoryError::ChipNotFound {
name: "VariableRangeCheckerChip".to_string(),
})
}
}
#[derive(thiserror::Error, Debug)]
pub enum ExecutorInventoryError {
#[error("Opcode {opcode} already owned by executor id {id}")]
ExecutorExists { opcode: VmOpcode, id: ExecutorId },
#[error("Phantom discriminant {} already has sub-executor", .discriminant.0)]
PhantomSubExecutorExists { discriminant: PhantomDiscriminant },
}
#[derive(thiserror::Error, Debug)]
pub enum AirInventoryError {
#[error("AIR {name} not found")]
AirNotFound { name: String },
}
#[derive(thiserror::Error, Debug)]
pub enum ChipInventoryError {
#[error("Air {name} not found")]
AirNotFound { name: String },
#[error("Chip {name} not found")]
ChipNotFound { name: String },
#[error("Adding prover extension without execution extension. Number of execution extensions is {0}")]
MissingExecutionExtension(usize),
#[error(
"Adding prover extension without circuit extension. Number of circuit extensions is {0}"
)]
MissingCircuitExtension(usize),
#[error("Missing chip. Number of chips is {actual}, expected number is {expected}")]
MissingChip { actual: usize, expected: usize },
#[error("Missing executor chip. Number of executors with associated chips is {actual}, expected number is {expected}")]
MissingExecutor { actual: usize, expected: usize },
}
impl<SC, RA, PB, SCC> VmChipComplex<SC, RA, PB, SCC>
where
SC: StarkProtocolConfig,
RA: Arena,
PB: ProverBackend,
SCC: SystemChipComplex<RA, PB>,
{
pub fn system_config(&self) -> &SystemConfig {
self.inventory.config()
}
pub(crate) fn generate_proving_ctx(
&mut self,
system_records: SystemRecords<PB::Val>,
record_arenas: Vec<RA>,
) -> Result<ProvingContext<PB>, GenerationError> {
let num_sys_airs = self.system_config().num_airs();
let num_airs = num_sys_airs + self.inventory.chips.len();
if num_airs != record_arenas.len() {
return Err(GenerationError::UnexpectedNumArenas {
actual: record_arenas.len(),
expected: num_airs,
});
}
let mut _record_arenas = record_arenas;
let record_arenas = _record_arenas.split_off(num_sys_airs);
let sys_record_arenas = _record_arenas;
let ctx_without_empties: Vec<(usize, AirProvingContext<_>)> = iter::empty()
.chain(info_span!("system_trace_gen").in_scope(|| {
self.system
.generate_proving_ctx(system_records, sys_record_arenas)
}))
.chain(
zip(self.inventory.chips.iter().enumerate().rev(), record_arenas).map(
|((insertion_idx, chip), records)| {
let _span = (!records.is_empty()).then(|| {
let air_name = self.inventory.airs.ext_airs[insertion_idx].name();
info_span!("single_trace_gen", air = air_name).entered()
});
#[cfg(feature = "metrics")]
if let Some(allocated_bytes) = (!records.is_empty())
.then(|| records.allocated_bytes())
.flatten()
{
let air_name = self.inventory.airs.ext_airs[insertion_idx].name();
let labels = [
("air_name", air_name.to_string()),
("air_id", (num_sys_airs + insertion_idx).to_string()),
];
metrics::counter!("trace_gen.record_arena_bytes", &labels)
.absolute(allocated_bytes as u64);
}
chip.generate_proving_ctx(records)
},
),
)
.enumerate()
.filter(|(_air_id, ctx)| ctx.common_main.height() > 0)
.collect();
Ok(ProvingContext::new(ctx_without_empties))
}
}
impl<F, EXT: VmExecutionExtension<F>> VmExecutionExtension<F> for Option<EXT> {
type Executor = EXT::Executor;
fn extend_execution(
&self,
inventory: &mut ExecutorInventoryBuilder<F, Self::Executor>,
) -> Result<(), ExecutorInventoryError> {
if let Some(extension) = self {
extension.extend_execution(inventory)
} else {
Ok(())
}
}
}
impl<SC: StarkProtocolConfig, EXT: VmCircuitExtension<SC>> VmCircuitExtension<SC> for Option<EXT> {
fn extend_circuit(&self, inventory: &mut AirInventory<SC>) -> Result<(), AirInventoryError> {
if let Some(extension) = self {
extension.extend_circuit(inventory)
} else {
Ok(())
}
}
}
pub trait AnyEnum {
fn as_any_kind(&self) -> &dyn Any;
fn as_any_kind_mut(&mut self) -> &mut dyn Any;
}
impl AnyEnum for () {
fn as_any_kind(&self) -> &dyn Any {
self
}
fn as_any_kind_mut(&mut self) -> &mut dyn Any {
self
}
}
#[cfg(test)]
mod tests {
use openvm_circuit_derive::AnyEnum;
use openvm_stark_sdk::config::baby_bear_poseidon2::BabyBearPoseidon2Config;
use super::*;
use crate::arch::VmCircuitConfig;
#[allow(dead_code)]
#[derive(Copy, Clone)]
enum EnumA {
A(u8),
B(u32),
}
enum EnumB {
C(u64),
D(EnumA),
}
#[derive(AnyEnum)]
enum EnumC {
C(u64),
#[any_enum]
D(EnumA),
}
impl AnyEnum for EnumA {
fn as_any_kind(&self) -> &dyn Any {
match self {
EnumA::A(a) => a,
EnumA::B(b) => b,
}
}
fn as_any_kind_mut(&mut self) -> &mut dyn Any {
match self {
EnumA::A(a) => a,
EnumA::B(b) => b,
}
}
}
impl AnyEnum for EnumB {
fn as_any_kind(&self) -> &dyn Any {
match self {
EnumB::C(c) => c,
EnumB::D(d) => d.as_any_kind(),
}
}
fn as_any_kind_mut(&mut self) -> &mut dyn Any {
match self {
EnumB::C(c) => c,
EnumB::D(d) => d.as_any_kind_mut(),
}
}
}
#[test]
fn test_any_enum_downcast() {
let a = EnumA::A(1);
assert_eq!(a.as_any_kind().downcast_ref::<u8>(), Some(&1));
let b = EnumB::D(a);
assert!(b.as_any_kind().downcast_ref::<u64>().is_none());
assert!(b.as_any_kind().downcast_ref::<EnumA>().is_none());
assert_eq!(b.as_any_kind().downcast_ref::<u8>(), Some(&1));
let c = EnumB::C(3);
assert_eq!(c.as_any_kind().downcast_ref::<u64>(), Some(&3));
let d = EnumC::D(a);
assert!(d.as_any_kind().downcast_ref::<u64>().is_none());
assert!(d.as_any_kind().downcast_ref::<EnumA>().is_none());
assert_eq!(d.as_any_kind().downcast_ref::<u8>(), Some(&1));
let e = EnumC::C(3);
assert_eq!(e.as_any_kind().downcast_ref::<u64>(), Some(&3));
}
#[test]
fn test_system_bus_indices() {
let config = SystemConfig::default();
let inventory: AirInventory<BabyBearPoseidon2Config> = config.create_airs().unwrap();
let system = inventory.system();
let port = system.port();
assert_eq!(port.execution_bus.index(), 0);
assert_eq!(port.memory_bridge.memory_bus().index(), 1);
assert_eq!(port.program_bus.index(), 2);
assert_eq!(port.memory_bridge.range_bus().index(), 3);
assert_eq!(system.memory.interface.boundary.merkle_bus.index, 4);
assert_eq!(system.memory.interface.boundary.compression_bus.index, 5);
}
}