use crate::{ProofOptions, TraceInfo, TraceLayout};
use math::StarkField;
use utils::{
collections::Vec, string::ToString, ByteReader, ByteWriter, Deserializable,
DeserializationError, Serializable,
};
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct Context {
trace_layout: TraceLayout,
trace_length: usize,
trace_meta: Vec<u8>,
field_modulus_bytes: Vec<u8>,
options: ProofOptions,
}
impl Context {
pub fn new<B: StarkField>(trace_info: &TraceInfo, options: ProofOptions) -> Self {
Context {
trace_layout: trace_info.layout().clone(),
trace_length: trace_info.length(),
trace_meta: trace_info.meta().to_vec(),
field_modulus_bytes: B::get_modulus_le_bytes(),
options,
}
}
pub fn trace_layout(&self) -> &TraceLayout {
&self.trace_layout
}
pub fn trace_length(&self) -> usize {
self.trace_length
}
pub fn get_trace_info(&self) -> TraceInfo {
TraceInfo::new_multi_segment(
self.trace_layout.clone(),
self.trace_length(),
self.trace_meta.clone(),
)
}
pub fn lde_domain_size(&self) -> usize {
self.trace_length() * self.options.blowup_factor()
}
pub fn field_modulus_bytes(&self) -> &[u8] {
&self.field_modulus_bytes
}
pub fn num_modulus_bits(&self) -> u32 {
let mut num_bits = self.field_modulus_bytes.len() as u32 * 8;
for &byte in self.field_modulus_bytes.iter().rev() {
if byte != 0 {
num_bits -= byte.leading_zeros();
return num_bits;
}
num_bits -= 8;
}
0
}
pub fn options(&self) -> &ProofOptions {
&self.options
}
}
impl Serializable for Context {
fn write_into<W: ByteWriter>(&self, target: &mut W) {
self.trace_layout.write_into(target);
target.write_u8(math::log2(self.trace_length) as u8); target.write_u16(self.trace_meta.len() as u16);
target.write_u8_slice(&self.trace_meta);
assert!(self.field_modulus_bytes.len() < u8::MAX as usize);
target.write_u8(self.field_modulus_bytes.len() as u8);
target.write_u8_slice(&self.field_modulus_bytes);
self.options.write_into(target);
}
}
impl Deserializable for Context {
fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
let trace_layout = TraceLayout::read_from(source)?;
let trace_length = source.read_u8()?;
if trace_length < math::log2(TraceInfo::MIN_TRACE_LENGTH) as u8 {
return Err(DeserializationError::InvalidValue(format!(
"trace length cannot be smaller than 2^{}, but was 2^{}",
math::log2(TraceInfo::MIN_TRACE_LENGTH),
trace_length
)));
}
let trace_length = 2_usize.pow(trace_length as u32);
let num_meta_bytes = source.read_u16()? as usize;
let trace_meta = if num_meta_bytes != 0 {
source.read_u8_vec(num_meta_bytes)?
} else {
vec![]
};
let num_modulus_bytes = source.read_u8()? as usize;
if num_modulus_bytes == 0 {
return Err(DeserializationError::InvalidValue(
"field modulus cannot be an empty value".to_string(),
));
}
let field_modulus_bytes = source.read_u8_vec(num_modulus_bytes)?;
let options = ProofOptions::read_from(source)?;
Ok(Context {
trace_layout,
trace_length,
trace_meta,
field_modulus_bytes,
options,
})
}
}