use std::env;
use serde::{Deserialize, Serialize};
use sysinfo::System;
const MAX_SHARD_SIZE: usize = 1 << 22;
const MAX_SHARD_BATCH_SIZE: usize = 8;
const DEFAULT_TRACE_GEN_WORKERS: usize = 1;
const DEFAULT_CHECKPOINTS_CHANNEL_CAPACITY: usize = 128;
const DEFAULT_RECORDS_AND_TRACES_CHANNEL_CAPACITY: usize = 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct SP1ProverOpts {
pub core_opts: SP1CoreOpts,
pub recursion_opts: SP1CoreOpts,
}
impl Default for SP1ProverOpts {
fn default() -> Self {
Self { core_opts: SP1CoreOpts::default(), recursion_opts: SP1CoreOpts::recursion() }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct SP1CoreOpts {
pub shard_size: usize,
pub shard_batch_size: usize,
pub split_opts: SplitOpts,
pub reconstruct_commitments: bool,
pub trace_gen_workers: usize,
pub checkpoints_channel_capacity: usize,
pub records_and_traces_channel_capacity: usize,
}
#[allow(clippy::cast_precision_loss)]
fn shard_size(total_available_mem: u64) -> usize {
let log_shard_size = match total_available_mem {
0..=14 => 18,
m => (((m as f64).log2() * 0.619) + 17.2).floor() as usize,
};
std::cmp::min(1 << log_shard_size, MAX_SHARD_SIZE)
}
fn shard_batch_size(total_available_mem: u64) -> usize {
match total_available_mem {
0..=16 => 1,
17..=48 => 2,
256.. => MAX_SHARD_BATCH_SIZE,
_ => 4,
}
}
impl Default for SP1CoreOpts {
fn default() -> Self {
let split_threshold = env::var("SPLIT_THRESHOLD")
.map(|s| s.parse::<usize>().unwrap_or(DEFERRED_SPLIT_THRESHOLD))
.unwrap_or(DEFERRED_SPLIT_THRESHOLD);
let sys = System::new_all();
let total_available_mem = sys.total_memory() / (1024 * 1024 * 1024);
let default_shard_size = shard_size(total_available_mem);
let default_shard_batch_size = shard_batch_size(total_available_mem);
Self {
shard_size: env::var("SHARD_SIZE").map_or_else(
|_| default_shard_size,
|s| s.parse::<usize>().unwrap_or(default_shard_size),
),
shard_batch_size: env::var("SHARD_BATCH_SIZE").map_or_else(
|_| default_shard_batch_size,
|s| s.parse::<usize>().unwrap_or(default_shard_batch_size),
),
split_opts: SplitOpts::new(split_threshold),
reconstruct_commitments: true,
trace_gen_workers: env::var("TRACE_GEN_WORKERS").map_or_else(
|_| DEFAULT_TRACE_GEN_WORKERS,
|s| s.parse::<usize>().unwrap_or(DEFAULT_TRACE_GEN_WORKERS),
),
checkpoints_channel_capacity: env::var("CHECKPOINTS_CHANNEL_CAPACITY").map_or_else(
|_| DEFAULT_CHECKPOINTS_CHANNEL_CAPACITY,
|s| s.parse::<usize>().unwrap_or(DEFAULT_CHECKPOINTS_CHANNEL_CAPACITY),
),
records_and_traces_channel_capacity: env::var("RECORDS_AND_TRACES_CHANNEL_CAPACITY")
.map_or_else(
|_| DEFAULT_RECORDS_AND_TRACES_CHANNEL_CAPACITY,
|s| s.parse::<usize>().unwrap_or(DEFAULT_RECORDS_AND_TRACES_CHANNEL_CAPACITY),
),
}
}
}
impl SP1CoreOpts {
#[must_use]
pub fn recursion() -> Self {
let mut opts = Self::default();
opts.reconstruct_commitments = false;
opts.shard_size = MAX_SHARD_SIZE;
opts
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct SplitOpts {
pub deferred: usize,
pub keccak: usize,
pub sha_extend: usize,
pub sha_compress: usize,
pub memory: usize,
}
impl SplitOpts {
#[must_use]
pub fn new(deferred_shift_threshold: usize) -> Self {
Self {
deferred: deferred_shift_threshold,
keccak: deferred_shift_threshold / 24,
sha_extend: deferred_shift_threshold / 48,
sha_compress: deferred_shift_threshold / 80,
memory: deferred_shift_threshold * 4,
}
}
}
pub const DEFERRED_SPLIT_THRESHOLD: usize = 1 << 19;