use std::{num::ParseIntError, str::ParseBoolError, sync::Arc};
use num_bigint::{BigInt, ParseBigIntError, TryFromBigIntError};
use thiserror::Error;
#[derive(Error, Debug)]
pub enum Error {
#[error("Failure while parsing field element")]
FieldParsingError,
#[error("Not enough constants")]
NotEnoughConstants,
#[error("IO cell iterator was exhausted")]
NotEnoughIOCells,
#[error("Parse failure")]
IntParse(#[from] ParseIntError),
#[error("Parse failure")]
BoolParse(#[from] ParseBoolError),
#[error("Parse failure")]
BigUintParse(#[from] ParseBigIntError),
#[error("Synthesis error")]
Plonk(Arc<dyn std::error::Error>),
#[error("Error")]
StrError(&'static str),
#[error(transparent)]
IntCast(#[from] std::num::TryFromIntError),
#[error(transparent)]
BigIntCast(#[from] TryFromBigIntError<BigInt>),
#[error("{header}Was expecting {expected} elements but got {actual}")]
UnexpectedElements {
header: String,
expected: usize,
actual: usize,
},
}
impl From<&'static str> for Error {
fn from(value: &'static str) -> Self {
Self::StrError(value)
}
}
unsafe impl Send for Error {}
unsafe impl Sync for Error {}
#[macro_export]
macro_rules! expect_elements {
(($($tokens:tt)+) , $($fmt:tt)+) => {
$crate::__expect_elements_parse! { [] [$($tokens)+] format!($($fmt)+) }
};
($($tokens:tt)+) => {
$crate::__expect_elements_parse! { [] [$($tokens)+] String::new() }
};
}
#[macro_export]
#[doc(hidden)]
macro_rules! __expect_elements_parse {
([$($lhs:tt)+] [== $($rhs:tt)+] $msg:expr) => {
$crate::__expect_elements_finish! {
($($lhs)*)
(==)
($($rhs)*)
($msg)
}
};
([$($lhs:tt)+] [!= $($rhs:tt)+] $msg:expr) => {
$crate::__expect_elements_finish! {
($($lhs)*)
(!=)
($($rhs)*)
($msg)
}
};
([$($lhs:tt)+] [< $($rhs:tt)+] $msg:expr) => {
$crate::__expect_elements_finish! {
($($lhs)*)
(<)
($($rhs)*)
($msg)
}
};
([$($lhs:tt)+] [<= $($rhs:tt)+] $msg:expr) => {
$crate::__expect_elements_finish! {
($($lhs)*)
(<=)
($($rhs)*)
($msg)
}
};
([$($lhs:tt)+] [> $($rhs:tt)+] $msg:expr) => {
$crate::__expect_elements_finish! {
($($lhs)*)
(>)
($($rhs)*)
($msg)
}
};
([$($lhs:tt)+] [>= $($rhs:tt)+] $msg:expr) => {
$crate::__expect_elements_finish! {
($($lhs)*)
(>=)
($($rhs)*)
($msg)
}
};
([$($lhs:tt)*] [$next:tt $($rest:tt)*] $msg:expr) => {
$crate::__expect_elements_parse! { [$($lhs)* $next] [$($rest)*] $msg }
};
([$($lhs:tt)*] [] $msg:expr) => {
compile_error!("expected a comparison expression such as `lhs == rhs`");
};
}
#[doc(hidden)]
#[macro_export]
macro_rules! __expect_elements_finish {
(($($lhs:tt)+) (==) ($($rhs:tt)+) ($msg:expr)) => {{
let lhs_val: usize = $($lhs)*;
let rhs_val: usize = $($rhs)*;
$crate::error::__expect_elements_impl($msg, lhs_val, rhs_val, lhs_val == rhs_val)?;
}};
(($($lhs:tt)+) (!=) ($($rhs:tt)+) ($msg:expr)) => {{
let lhs_val: usize = $($lhs)*;
let rhs_val: usize = $($rhs)*;
$crate::error::__expect_elements_impl($msg, lhs_val, rhs_val, lhs_val != rhs_val)?;
}};
(($($lhs:tt)+) (<) ($($rhs:tt)+) ($msg:expr)) => {{
let lhs_val: usize = $($lhs)*;
let rhs_val: usize = $($rhs)*;
$crate::error::__expect_elements_impl($msg, lhs_val, rhs_val, lhs_val < rhs_val)?;
}};
(($($lhs:tt)+) (<=) ($($rhs:tt)+) ($msg:expr)) => {{
let lhs_val: usize = $($lhs)*;
let rhs_val: usize = $($rhs)*;
$crate::error::__expect_elements_impl($msg, lhs_val, rhs_val, lhs_val <= rhs_val)?;
}};
(($($lhs:tt)+) (>) ($($rhs:tt)+) ($msg:expr)) => {{
let lhs_val: usize = $($lhs)*;
let rhs_val: usize = $($rhs)*;
$crate::error::__expect_elements_impl($msg, lhs_val, rhs_val, lhs_val > rhs_val)?;
}};
(($($lhs:tt)+) (>=) ($($rhs:tt)+) ($msg:expr)) => {{
let lhs_val: usize = $($lhs)*;
let rhs_val: usize = $($rhs)*;
$crate::error::__expect_elements_impl($msg, lhs_val, rhs_val, lhs_val >= rhs_val)?;
}};
}
#[doc(hidden)]
#[inline]
pub fn __expect_elements_impl(
header: String,
expected: usize,
actual: usize,
passed: bool,
) -> Result<(), Error> {
if passed {
Ok(())
} else {
Err(Error::UnexpectedElements {
header,
expected,
actual,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
#[derive(Debug)]
enum Cmp {
Eq,
Ne,
Lt,
Le,
Gt,
Ge,
}
#[rstest]
#[case(Cmp::Eq, 1, 1)]
#[case(Cmp::Ne, 1, 2)]
#[case(Cmp::Lt, 1, 2)]
#[case(Cmp::Le, 1, 1)]
#[case(Cmp::Gt, 2, 1)]
#[case(Cmp::Ge, 1, 1)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Eq, 1, 2)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Ne, 1, 1)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Lt, 1, 0)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Le, 1, 0)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Gt, 2, 3)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Ge, 1, 3)]
fn expect_elements_test(#[case] cmp: Cmp, #[case] expected: usize, #[case] actual: usize) {
fn do_test(cmp: Cmp, expected: usize, actual: usize) -> Result<(), Error> {
match cmp {
Cmp::Eq => expect_elements!((expected == actual), "unexpected elements error"),
Cmp::Ne => expect_elements!((expected != actual), "unexpected elements error"),
Cmp::Lt => expect_elements!((expected < actual), "unexpected elements error"),
Cmp::Le => expect_elements!((expected <= actual), "unexpected elements error"),
Cmp::Gt => expect_elements!((expected > actual), "unexpected elements error"),
Cmp::Ge => expect_elements!((expected >= actual), "unexpected elements error"),
}
Ok(())
}
eprintln!("cmp = {cmp:?}, expected = {expected}, actual = {actual}");
do_test(cmp, expected, actual).unwrap();
}
#[rstest]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Eq)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Ne)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Lt)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Le)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Gt)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Ge)]
fn expect_elements_complex_expr_rhs(#[case] cmp: Cmp) {
fn do_test(cmp: Cmp) -> Result<(), Error> {
let v = vec![1, 2, 3];
match cmp {
Cmp::Eq => expect_elements!((2 == v.len()), "unexpected elements error"),
Cmp::Ne => expect_elements!((3 != v.len()), "unexpected elements error"),
Cmp::Lt => expect_elements!((4 < v.len()), "unexpected elements error"),
Cmp::Le => expect_elements!((4 <= v.len()), "unexpected elements error"),
Cmp::Gt => expect_elements!((2 > v.len()), "unexpected elements error"),
Cmp::Ge => expect_elements!((2 >= v.len()), "unexpected elements error"),
};
Ok(())
}
do_test(cmp).unwrap();
}
#[rstest]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Eq)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Ne)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Lt)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Le)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Gt)]
#[should_panic(expected = "unexpected elements error")]
#[case(Cmp::Ge)]
fn expect_elements_complex_expr_lhs(#[case] cmp: Cmp) {
fn do_test(cmp: Cmp) -> Result<(), Error> {
let v = vec![1, 2, 3];
match cmp {
Cmp::Eq => expect_elements!((v.len() == 2), "unexpected elements error"),
Cmp::Ne => expect_elements!((v.len() != 3), "unexpected elements error"),
Cmp::Lt => expect_elements!((v.len() < 2), "unexpected elements error"),
Cmp::Le => expect_elements!((v.len() <= 2), "unexpected elements error"),
Cmp::Gt => expect_elements!((v.len() > 4), "unexpected elements error"),
Cmp::Ge => expect_elements!((v.len() >= 4), "unexpected elements error"),
};
Ok(())
}
do_test(cmp).unwrap();
}
}