use std::collections::HashMap;
use std::fs::File;
use std::io::{self, BufReader, BufWriter, Read, Write};
use anyhow::{bail, Result};
use crate::io::bin_file_writer::BinFileWriter;
pub fn reorder_plonk_pols_for_pilout(
fixed_pols: &[pil2_stark_recurser::plonk2pil::FixedPol],
symbols: &[pil2_pilout::pilout::Symbol],
air_group_id: u32,
air_id: u32,
) -> Vec<Vec<u64>> {
let pol_map: HashMap<(String, usize), &Vec<u64>> =
fixed_pols.iter().map(|fp| ((fp.name.clone(), fp.index), &fp.values)).collect();
let mut fixed_syms: Vec<&pil2_pilout::pilout::Symbol> = symbols
.iter()
.filter(|s| {
s.r#type == 1 && s.air_group_id == Some(air_group_id)
&& s.air_id == Some(air_id)
})
.collect();
fixed_syms.sort_by_key(|s| s.id);
let mut result = Vec::new();
for sym in fixed_syms {
if sym.lengths.is_empty() {
continue;
}
let total: usize = sym.lengths.iter().map(|&l| l as usize).product();
for idx in 0..total {
if let Some(vals) = pol_map.get(&(sym.name.clone(), idx)) {
result.push((*vals).clone());
}
}
}
result
}
#[derive(Debug, Clone)]
pub struct FixedPolInfo {
pub lengths: Vec<u32>,
pub values: Vec<u64>,
}
pub fn read_fixed_pols_bin(
fixed_info: &mut HashMap<String, HashMap<String, Vec<FixedPolInfo>>>,
bin_filename: &str,
) -> Result<()> {
let file = File::open(bin_filename)?;
let mut reader = BufReader::new(file);
let mut magic = [0u8; 4];
reader.read_exact(&mut magic)?;
if &magic != b"cnst" {
bail!("Invalid magic in fixed pols file: expected 'cnst'");
}
let _version = read_u32_le(&mut reader)?;
let _n_sections = read_u32_le(&mut reader)?;
let _section_id = read_u32_le(&mut reader)?;
let _section_size = read_u64_le(&mut reader)?;
let airgroup_name = read_string(&mut reader)?;
let air_name = read_string(&mut reader)?;
let n = read_u64_le(&mut reader)?;
let n_fixed_pols = read_u32_le(&mut reader)?;
let mut pols_info: HashMap<String, Vec<FixedPolInfo>> = HashMap::new();
for _ in 0..n_fixed_pols {
let name = read_string(&mut reader)?;
let n_lengths = read_u32_le(&mut reader)?;
let mut lengths = Vec::with_capacity(n_lengths as usize);
for _ in 0..n_lengths {
lengths.push(read_u32_le(&mut reader)?);
}
let mut values = Vec::with_capacity(n as usize);
let mut buf = vec![0u8; n as usize * 8];
reader.read_exact(&mut buf)?;
for i in 0..n as usize {
let val = u64::from_le_bytes([
buf[i * 8],
buf[i * 8 + 1],
buf[i * 8 + 2],
buf[i * 8 + 3],
buf[i * 8 + 4],
buf[i * 8 + 5],
buf[i * 8 + 6],
buf[i * 8 + 7],
]);
values.push(val);
}
pols_info.entry(name).or_default().push(FixedPolInfo { lengths, values });
}
let key = format!("{}_{}", airgroup_name, air_name);
fixed_info.insert(key, pols_info);
Ok(())
}
pub fn write_fixed_pols_bin(
bin_filename: &str,
airgroup_name: &str,
air_name: &str,
n: u64,
fixed_info: &[(String, Vec<u32>, Vec<u64>)],
) -> Result<()> {
let mut writer = BinFileWriter::new(bin_filename, "cnst", 1, 1)?;
writer.start_write_section(1)?;
writer.write_string(airgroup_name)?;
writer.write_string(air_name)?;
writer.write_u64(n)?;
writer.write_u32(fixed_info.len() as u32)?;
for (name, lengths, values) in fixed_info {
writer.write_string(name)?;
writer.write_u32(lengths.len() as u32)?;
for &len in lengths {
writer.write_u32(len)?;
}
let mut buf = vec![0u8; values.len() * 8];
for (i, &v) in values.iter().enumerate() {
buf[i * 8..i * 8 + 8].copy_from_slice(&v.to_le_bytes());
}
writer.write_bytes(&buf)?;
}
writer.end_write_section()?;
writer.close()
}
pub fn write_const_file(path: &str, air: &pil2_pilout::pilout::Air, plonk_pol_values: &[Vec<u64>]) -> Result<()> {
let n_rows = air.num_rows.unwrap_or(0) as usize;
let n_rows = if n_rows > 0 {
n_rows
} else if let Some(first) = plonk_pol_values.first() {
first.len()
} else {
0
};
let n_constants = air.fixed_cols.len();
let mut flat_buffer = vec![0u64; n_rows * n_constants];
let mut plonk_idx = 0usize;
for (col_idx, fixed_col) in air.fixed_cols.iter().enumerate() {
if fixed_col.values.is_empty() {
if let Some(vals) = plonk_pol_values.get(plonk_idx) {
for (row, &val) in vals.iter().enumerate() {
if row < n_rows {
flat_buffer[row * n_constants + col_idx] = val;
}
}
plonk_idx += 1;
}
} else {
for (row, val_bytes) in fixed_col.values.iter().enumerate() {
if row >= n_rows {
break;
}
flat_buffer[row * n_constants + col_idx] = bytes_to_u64_be(val_bytes);
}
}
}
write_fixed_cols_raw(path, &flat_buffer)
}
pub fn write_fixed_cols_raw(path: &str, buffer: &[u64]) -> Result<()> {
let file = File::create(path)?;
let mut writer = BufWriter::new(file);
for &val in buffer {
writer.write_all(&val.to_le_bytes())?;
}
writer.flush()?;
Ok(())
}
fn read_u32_le(reader: &mut impl Read) -> io::Result<u32> {
let mut buf = [0u8; 4];
reader.read_exact(&mut buf)?;
Ok(u32::from_le_bytes(buf))
}
fn read_u64_le(reader: &mut impl Read) -> io::Result<u64> {
let mut buf = [0u8; 8];
reader.read_exact(&mut buf)?;
Ok(u64::from_le_bytes(buf))
}
fn read_string(reader: &mut impl Read) -> io::Result<String> {
let mut buf = Vec::new();
let mut byte = [0u8; 1];
loop {
reader.read_exact(&mut byte)?;
if byte[0] == 0 {
break;
}
buf.push(byte[0]);
}
String::from_utf8(buf).map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))
}
fn bytes_to_u64_be(bytes: &[u8]) -> u64 {
let mut val = 0u64;
for &b in bytes.iter().take(8) {
val = (val << 8) | (b as u64);
}
val
}