use std::cell::RefCell;
use getset::CopyGetters;
use crate::{
gates::{
circuit::CircuitBuilderStage,
flex_gate::{BasicGateConfig, ThreadBreakPoints},
},
utils::halo2::{raw_assign_advice, raw_constrain_equal},
utils::ScalarField,
virtual_region::copy_constraints::{CopyConstraintManager, SharedCopyConstraintManager},
Context, ContextCell,
};
use crate::{
halo2_proofs::circuit::{Region, Value},
virtual_region::manager::VirtualRegionManager,
};
#[derive(Clone, Debug, Default, CopyGetters)]
pub struct SinglePhaseCoreManager<F: ScalarField> {
pub threads: Vec<Context<F>>,
pub copy_manager: SharedCopyConstraintManager<F>,
#[getset(get_copy = "pub")]
witness_gen_only: bool,
#[getset(get_copy = "pub")]
pub(crate) use_unknown: bool,
#[getset(get_copy = "pub", set)]
pub(crate) phase: usize,
pub break_points: RefCell<Option<ThreadBreakPoints>>,
}
impl<F: ScalarField> SinglePhaseCoreManager<F> {
pub fn new(witness_gen_only: bool, copy_manager: SharedCopyConstraintManager<F>) -> Self {
Self {
threads: vec![],
witness_gen_only,
use_unknown: false,
phase: 0,
copy_manager,
..Default::default()
}
}
pub fn in_phase(self, phase: usize) -> Self {
Self { phase, ..self }
}
pub fn from_stage(
stage: CircuitBuilderStage,
copy_manager: SharedCopyConstraintManager<F>,
) -> Self {
Self::new(stage.witness_gen_only(), copy_manager)
.unknown(stage == CircuitBuilderStage::Keygen)
}
pub fn unknown(self, use_unknown: bool) -> Self {
Self { use_unknown, ..self }
}
pub fn set_copy_manager(&mut self, copy_manager: SharedCopyConstraintManager<F>) {
self.copy_manager = copy_manager.clone();
for ctx in &mut self.threads {
ctx.copy_manager = copy_manager.clone();
}
}
pub fn use_copy_manager(mut self, copy_manager: SharedCopyConstraintManager<F>) -> Self {
self.set_copy_manager(copy_manager);
self
}
pub fn clear(&mut self) {
self.threads = vec![];
self.copy_manager.lock().unwrap().clear();
}
pub fn main(&mut self) -> &mut Context<F> {
if self.threads.is_empty() {
self.new_thread()
} else {
self.threads.last_mut().unwrap()
}
}
pub fn thread_count(&self) -> usize {
self.threads.len()
}
pub fn type_of(&self) -> &'static str {
match self.phase {
0 => "halo2-base:SinglePhaseCoreManager:FirstPhase",
1 => "halo2-base:SinglePhaseCoreManager:SecondPhase",
2 => "halo2-base:SinglePhaseCoreManager:ThirdPhase",
_ => panic!("Unsupported phase"),
}
}
pub fn new_context(&self, context_id: usize) -> Context<F> {
Context::new(
self.witness_gen_only,
self.phase,
self.type_of(),
context_id,
self.copy_manager.clone(),
)
}
pub fn new_thread(&mut self) -> &mut Context<F> {
let context_id = self.thread_count();
self.threads.push(self.new_context(context_id));
self.threads.last_mut().unwrap()
}
pub fn total_advice(&self) -> usize {
self.threads.iter().map(|ctx| ctx.advice.len()).sum::<usize>()
}
}
impl<F: ScalarField> VirtualRegionManager<F> for SinglePhaseCoreManager<F> {
type Config = (Vec<BasicGateConfig<F>>, usize); type Assignment = ();
fn assign_raw(&self, (config, usable_rows): &Self::Config, region: &mut Region<F>) {
if self.witness_gen_only {
let binding = self.break_points.borrow();
let break_points = binding.as_ref().expect("break points not set");
assign_witnesses(&self.threads, config, region, break_points);
} else {
let mut copy_manager = self.copy_manager.lock().unwrap();
let break_points = assign_with_constraints::<F, 4>(
&self.threads,
config,
region,
&mut copy_manager,
*usable_rows,
self.use_unknown,
);
let mut bp = self.break_points.borrow_mut();
if let Some(bp) = bp.as_ref() {
assert_eq!(bp, &break_points, "break points don't match");
} else {
*bp = Some(break_points);
}
}
}
}
pub fn assign_with_constraints<F: ScalarField, const ROTATIONS: usize>(
threads: &[Context<F>],
basic_gates: &[BasicGateConfig<F>],
region: &mut Region<F>,
copy_manager: &mut CopyConstraintManager<F>,
max_rows: usize,
use_unknown: bool,
) -> ThreadBreakPoints {
let mut break_points = vec![];
let mut gate_index = 0;
let mut row_offset = 0;
for ctx in threads {
if ctx.advice.is_empty() {
continue;
}
let mut basic_gate = basic_gates
.get(gate_index)
.unwrap_or_else(|| panic!("NOT ENOUGH ADVICE COLUMNS. Perhaps blinding factors were not taken into account. The max non-poisoned rows is {max_rows}"));
assert_eq!(ctx.selector.len(), ctx.advice.len());
for (i, (advice, &q)) in ctx.advice.iter().zip(ctx.selector.iter()).enumerate() {
let column = basic_gate.value;
let value = if use_unknown { Value::unknown() } else { Value::known(advice) };
let cell = region.assign_advice(column, row_offset, value).cell();
if let Some(old_cell) = copy_manager
.assigned_advices
.insert(ContextCell::new(ctx.type_id, ctx.context_id, i), cell)
{
assert!(
old_cell.row_offset == cell.row_offset && old_cell.column == cell.column,
"Trying to overwrite virtual cell with a different raw cell"
);
}
if (q && row_offset + ROTATIONS > max_rows) || row_offset >= max_rows - 1 {
break_points.push(row_offset);
row_offset = 0;
gate_index += 1;
if ROTATIONS > 1 && i + 2 >= ROTATIONS {
for delta in 1..ROTATIONS - 1 {
assert!(
!ctx.selector[i - delta],
"We do not support overlaps with delta = {delta}"
);
}
}
basic_gate = basic_gates
.get(gate_index)
.unwrap_or_else(|| panic!("NOT ENOUGH ADVICE COLUMNS. Perhaps blinding factors were not taken into account. The max non-poisoned rows is {max_rows}"));
let column = basic_gate.value;
let ncell = region.assign_advice(column, row_offset, value);
raw_constrain_equal(region, ncell.cell(), cell);
}
if q {
basic_gate
.q_enable
.enable(region, row_offset)
.expect("enable selector should not fail");
}
row_offset += 1;
}
}
break_points
}
pub fn assign_witnesses<F: ScalarField>(
threads: &[Context<F>],
basic_gates: &[BasicGateConfig<F>],
region: &mut Region<F>,
break_points: &ThreadBreakPoints,
) {
if basic_gates.is_empty() {
assert_eq!(
threads.iter().map(|ctx| ctx.advice.len()).sum::<usize>(),
0,
"Trying to assign threads in a phase with no columns"
);
return;
}
let mut break_points = break_points.clone().into_iter();
let mut break_point = break_points.next();
let mut gate_index = 0;
let mut column = basic_gates[gate_index].value;
let mut row_offset = 0;
for ctx in threads {
for advice in &ctx.advice {
raw_assign_advice(region, column, row_offset, Value::known(advice));
if break_point == Some(row_offset) {
break_point = break_points.next();
row_offset = 0;
gate_index += 1;
column = basic_gates[gate_index].value;
raw_assign_advice(region, column, row_offset, Value::known(advice));
}
row_offset += 1;
}
}
}