native_range_check.rs 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435
  1. use group::ff::{Field, PrimeFieldBits};
  2. use halo2_proofs::{
  3. circuit::{AssignedCell, Chip, Layouter, Region, Value},
  4. pasta::pallas,
  5. plonk,
  6. plonk::{Advice, Column, ConstraintSystem, Selector, TableColumn},
  7. poly::Rotation,
  8. };
  9. #[derive(Clone, Debug)]
  10. pub struct NativeRangeCheckConfig<
  11. const WINDOW_SIZE: usize,
  12. const NUM_BITS: usize,
  13. const NUM_WINDOWS: usize,
  14. > {
  15. pub z: Column<Advice>,
  16. pub s_rc: Selector,
  17. pub k_values_table: TableColumn,
  18. }
  19. #[derive(Clone, Debug)]
  20. pub struct NativeRangeCheckChip<
  21. const WINDOW_SIZE: usize,
  22. const NUM_BITS: usize,
  23. const NUM_WINDOWS: usize,
  24. > {
  25. config: NativeRangeCheckConfig<WINDOW_SIZE, NUM_BITS, NUM_WINDOWS>,
  26. }
  27. impl<const WINDOW_SIZE: usize, const NUM_BITS: usize, const NUM_WINDOWS: usize> Chip<pallas::Base>
  28. for NativeRangeCheckChip<WINDOW_SIZE, NUM_BITS, NUM_WINDOWS>
  29. {
  30. type Config = NativeRangeCheckConfig<WINDOW_SIZE, NUM_BITS, NUM_WINDOWS>;
  31. type Loaded = ();
  32. fn config(&self) -> &Self::Config {
  33. &self.config
  34. }
  35. fn loaded(&self) -> &Self::Loaded {
  36. &()
  37. }
  38. }
  39. impl<const WINDOW_SIZE: usize, const NUM_BITS: usize, const NUM_WINDOWS: usize>
  40. NativeRangeCheckChip<WINDOW_SIZE, NUM_BITS, NUM_WINDOWS>
  41. {
  42. pub fn construct(config: NativeRangeCheckConfig<WINDOW_SIZE, NUM_BITS, NUM_WINDOWS>) -> Self {
  43. Self { config }
  44. }
  45. pub fn configure(
  46. meta: &mut ConstraintSystem<pallas::Base>,
  47. z: Column<Advice>,
  48. k_values_table: TableColumn,
  49. ) -> NativeRangeCheckConfig<WINDOW_SIZE, NUM_BITS, NUM_WINDOWS> {
  50. // Enable permutation on z column
  51. meta.enable_equality(z);
  52. let s_rc = meta.complex_selector();
  53. meta.lookup(|meta| {
  54. let s_rc = meta.query_selector(s_rc);
  55. let z_curr = meta.query_advice(z, Rotation::cur());
  56. let z_next = meta.query_advice(z, Rotation::next());
  57. // z_next = (z_curr - k_i) / 2^K
  58. // => k_i = z_curr - (z_next * 2^K)
  59. vec![(s_rc * (z_curr - z_next * pallas::Base::from(1 << WINDOW_SIZE)), k_values_table)]
  60. });
  61. NativeRangeCheckConfig { z, s_rc, k_values_table }
  62. }
  63. /// `k_values_table` should be reused across different chips
  64. /// which is why we don't limit it to a specific instance.
  65. pub fn load_k_table(
  66. layouter: &mut impl Layouter<pallas::Base>,
  67. k_values_table: TableColumn,
  68. ) -> Result<(), plonk::Error> {
  69. layouter.assign_table(
  70. || format!("{} window table", WINDOW_SIZE),
  71. |mut table| {
  72. for index in 0..(1 << WINDOW_SIZE) {
  73. table.assign_cell(
  74. || format!("{} window assign", WINDOW_SIZE),
  75. k_values_table,
  76. index,
  77. || Value::known(pallas::Base::from(index as u64)),
  78. )?;
  79. }
  80. Ok(())
  81. },
  82. )
  83. }
  84. fn decompose_value(value: &pallas::Base) -> Vec<[bool; WINDOW_SIZE]> {
  85. let padding = (WINDOW_SIZE - NUM_BITS % WINDOW_SIZE) % WINDOW_SIZE;
  86. let bits: Vec<bool> = value
  87. .to_le_bits()
  88. .into_iter()
  89. .take(NUM_BITS)
  90. .chain(std::iter::repeat(false).take(padding))
  91. .collect();
  92. assert_eq!(bits.len(), NUM_BITS + padding);
  93. bits.chunks_exact(WINDOW_SIZE)
  94. .map(|x| {
  95. let mut chunks = [false; WINDOW_SIZE];
  96. chunks.copy_from_slice(x);
  97. chunks
  98. })
  99. .collect()
  100. }
  101. pub fn decompose(
  102. &self,
  103. region: &mut Region<'_, pallas::Base>,
  104. z_0: AssignedCell<pallas::Base, pallas::Base>,
  105. offset: usize,
  106. strict: bool,
  107. ) -> Result<(), plonk::Error> {
  108. assert!(WINDOW_SIZE * NUM_WINDOWS < NUM_BITS + WINDOW_SIZE);
  109. // Enable selectors
  110. for index in 0..NUM_WINDOWS {
  111. self.config.s_rc.enable(region, index + offset)?;
  112. }
  113. let mut z_values: Vec<AssignedCell<pallas::Base, pallas::Base>> = vec![z_0.clone()];
  114. let mut z = z_0.clone();
  115. let decomposed_chunks = z_0.value().map(Self::decompose_value).transpose_vec(NUM_WINDOWS);
  116. let two_pow_k_inverse =
  117. Value::known(pallas::Base::from(1 << WINDOW_SIZE as u64).invert().unwrap());
  118. for (i, chunk) in decomposed_chunks.iter().enumerate() {
  119. let z_next = {
  120. let z_curr = z.value().copied();
  121. let chunk_value = chunk.map(|c| {
  122. pallas::Base::from(c.iter().rev().fold(0, |acc, c| (acc << 1) + *c as u64))
  123. });
  124. // z_next = (z_curr - k_i) / 2^K
  125. let z_next = (z_curr - chunk_value) * two_pow_k_inverse;
  126. region.assign_advice(
  127. || format!("z_{}", i + offset + 1),
  128. self.config.z,
  129. i + offset + 1,
  130. || z_next,
  131. )?
  132. };
  133. z_values.push(z_next.clone());
  134. z = z_next.clone();
  135. }
  136. assert!(z_values.len() == NUM_WINDOWS + 1);
  137. if strict {
  138. // Constrain the remaining bits to be zero
  139. region.constrain_constant(z_values.last().unwrap().cell(), pallas::Base::zero())?;
  140. }
  141. Ok(())
  142. }
  143. pub fn witness_range_check(
  144. &self,
  145. mut layouter: impl Layouter<pallas::Base>,
  146. value: Value<pallas::Base>,
  147. strict: bool,
  148. ) -> Result<(), plonk::Error> {
  149. layouter.assign_region(
  150. || format!("witness {}-bit native range check", NUM_BITS),
  151. |mut region: Region<'_, pallas::Base>| {
  152. let z_0 = region.assign_advice(|| "z_0", self.config.z, 0, || value)?;
  153. self.decompose(&mut region, z_0, 0, strict)?;
  154. Ok(())
  155. },
  156. )
  157. }
  158. pub fn copy_range_check(
  159. &self,
  160. mut layouter: impl Layouter<pallas::Base>,
  161. value: AssignedCell<pallas::Base, pallas::Base>,
  162. strict: bool,
  163. ) -> Result<(), plonk::Error> {
  164. layouter.assign_region(
  165. || format!("copy {}-bit native range check", NUM_BITS),
  166. |mut region: Region<'_, pallas::Base>| {
  167. let z_0 = value.copy_advice(|| "z_0", &mut region, self.config.z, 0)?;
  168. self.decompose(&mut region, z_0, 0, strict)?;
  169. Ok(())
  170. },
  171. )
  172. }
  173. }
  174. #[cfg(test)]
  175. mod tests {
  176. use super::*;
  177. use crate::zk::assign_free_advice;
  178. use group::ff::PrimeField;
  179. use halo2_proofs::{
  180. circuit::floor_planner,
  181. dev::{CircuitLayout, MockProver},
  182. plonk::Circuit,
  183. };
  184. use pasta_curves::arithmetic::FieldExt;
  185. macro_rules! test_circuit {
  186. ($window_size:expr, $num_bits:expr, $num_windows:expr) => {
  187. #[derive(Default)]
  188. struct RangeCheckCircuit {
  189. a: Value<pallas::Base>,
  190. }
  191. impl Circuit<pallas::Base> for RangeCheckCircuit {
  192. type Config =
  193. (NativeRangeCheckConfig<$window_size, $num_bits, $num_windows>, Column<Advice>);
  194. type FloorPlanner = floor_planner::V1;
  195. fn without_witnesses(&self) -> Self {
  196. Self::default()
  197. }
  198. fn configure(meta: &mut ConstraintSystem<pallas::Base>) -> Self::Config {
  199. let w = meta.advice_column();
  200. meta.enable_equality(w);
  201. let z = meta.advice_column();
  202. let table_column = meta.lookup_table_column();
  203. let constants = meta.fixed_column();
  204. meta.enable_constant(constants);
  205. (
  206. NativeRangeCheckChip::<$window_size, $num_bits, $num_windows>::configure(
  207. meta,
  208. z,
  209. table_column,
  210. ),
  211. w,
  212. )
  213. }
  214. fn synthesize(
  215. &self,
  216. config: Self::Config,
  217. mut layouter: impl Layouter<pallas::Base>,
  218. ) -> Result<(), plonk::Error> {
  219. let rangecheck_chip =
  220. NativeRangeCheckChip::<$window_size, $num_bits, $num_windows>::construct(
  221. config.0.clone(),
  222. );
  223. NativeRangeCheckChip::<$window_size, $num_bits, $num_windows>::load_k_table(
  224. &mut layouter,
  225. config.0.k_values_table,
  226. )?;
  227. let a = assign_free_advice(layouter.namespace(|| "load a"), config.1, self.a)?;
  228. rangecheck_chip.copy_range_check(
  229. layouter.namespace(|| "copy a and range check"),
  230. a,
  231. true,
  232. )?;
  233. rangecheck_chip.witness_range_check(
  234. layouter.namespace(|| "witness a and range check"),
  235. self.a,
  236. true,
  237. )?;
  238. Ok(())
  239. }
  240. }
  241. };
  242. }
  243. // cargo test --release --all-features --lib native_range_check -- --nocapture
  244. #[test]
  245. fn native_range_check_64() {
  246. test_circuit!(3, 64, 22);
  247. let k = 6;
  248. let valid_values = vec![
  249. pallas::Base::zero(),
  250. pallas::Base::one(),
  251. pallas::Base::from(u64::MAX),
  252. pallas::Base::from(rand::random::<u64>()),
  253. ];
  254. let invalid_values = vec![
  255. -pallas::Base::one(),
  256. pallas::Base::from_u128(u64::MAX as u128 + 1),
  257. -pallas::Base::from_u128(u64::MAX as u128 + 1),
  258. pallas::Base::from_u128(rand::random::<u128>()),
  259. // The following two are valid
  260. // 2 = -28948022309329048855892746252171976963363056481941560715954676764349967630335
  261. //-pallas::Base::from_str_vartime(
  262. // "28948022309329048855892746252171976963363056481941560715954676764349967630335",
  263. //)
  264. //.unwrap(),
  265. // 1 = -28948022309329048855892746252171976963363056481941560715954676764349967630336
  266. //-pallas::Base::from_str_vartime(
  267. // "28948022309329048855892746252171976963363056481941560715954676764349967630336",
  268. //)
  269. //.unwrap(),
  270. ];
  271. use plotters::prelude::*;
  272. let circuit = RangeCheckCircuit { a: Value::known(pallas::Base::one()) };
  273. let root =
  274. BitMapBackend::new("target/native_range_check_64_circuit_layout.png", (3840, 2160))
  275. .into_drawing_area();
  276. root.fill(&WHITE).unwrap();
  277. let root =
  278. root.titled("64-bit Native Range Check Circuit Layout", ("sans-serif", 60)).unwrap();
  279. CircuitLayout::default().render(k, &circuit, &root).unwrap();
  280. for i in valid_values {
  281. println!("64-bit (valid) range check for {:?}", i);
  282. let circuit = RangeCheckCircuit { a: Value::known(i) };
  283. let prover = MockProver::run(k, &circuit, vec![]).unwrap();
  284. prover.assert_satisfied();
  285. println!("Constraints satisfied");
  286. }
  287. for i in invalid_values {
  288. println!("64-bit (invalid) range check for {:?}", i);
  289. let circuit = RangeCheckCircuit { a: Value::known(i) };
  290. let prover = MockProver::run(k, &circuit, vec![]).unwrap();
  291. assert!(prover.verify().is_err());
  292. }
  293. }
  294. #[test]
  295. fn native_range_check_128() {
  296. test_circuit!(3, 128, 43);
  297. let k = 7;
  298. let valid_values = vec![
  299. pallas::Base::zero(),
  300. pallas::Base::one(),
  301. pallas::Base::from_u128(u128::MAX),
  302. pallas::Base::from_u128(rand::random::<u128>()),
  303. ];
  304. let invalid_values = vec![
  305. -pallas::Base::one(),
  306. pallas::Base::from_u128(u128::MAX) + pallas::Base::one(),
  307. -pallas::Base::from_u128(u128::MAX) + pallas::Base::one(),
  308. -pallas::Base::from_u128(u128::MAX),
  309. ];
  310. use plotters::prelude::*;
  311. let circuit = RangeCheckCircuit { a: Value::known(pallas::Base::one()) };
  312. let root =
  313. BitMapBackend::new("target/native_range_check_128_circuit_layout.png", (3840, 2160))
  314. .into_drawing_area();
  315. root.fill(&WHITE).unwrap();
  316. let root =
  317. root.titled("128-bit Native Range Check Circuit Layout", ("sans-serif", 60)).unwrap();
  318. CircuitLayout::default().render(k, &circuit, &root).unwrap();
  319. for i in valid_values {
  320. println!("128-bit (valid) range check for {:?}", i);
  321. let circuit = RangeCheckCircuit { a: Value::known(i) };
  322. let prover = MockProver::run(k, &circuit, vec![]).unwrap();
  323. prover.assert_satisfied();
  324. println!("Constraints satisfied");
  325. }
  326. for i in invalid_values {
  327. println!("128-bit (invalid) range check for {:?}", i);
  328. let circuit = RangeCheckCircuit { a: Value::known(i) };
  329. let prover = MockProver::run(k, &circuit, vec![]).unwrap();
  330. assert!(prover.verify().is_err());
  331. }
  332. }
  333. #[test]
  334. fn native_range_check_253() {
  335. test_circuit!(3, 253, 85);
  336. let k = 8;
  337. let valid_values = vec![
  338. pallas::Base::zero(),
  339. pallas::Base::one(),
  340. // 2^253 - 1
  341. pallas::Base::from_str_vartime(
  342. "14474011154664524427946373126085988481658748083205070504932198000989141204991",
  343. )
  344. .unwrap(),
  345. // 2^253 / 2
  346. pallas::Base::from_str_vartime(
  347. "7237005577332262213973186563042994240829374041602535252466099000494570602496",
  348. )
  349. .unwrap(),
  350. ];
  351. let invalid_values = vec![
  352. -pallas::Base::one(),
  353. // p - 1
  354. pallas::Base::from_str_vartime(
  355. "28948022309329048855892746252171976963363056481941560715954676764349967630336",
  356. )
  357. .unwrap(),
  358. ];
  359. use plotters::prelude::*;
  360. let circuit = RangeCheckCircuit { a: Value::known(pallas::Base::one()) };
  361. let root =
  362. BitMapBackend::new("target/native_range_check_253_circuit_layout.png", (3840, 2160))
  363. .into_drawing_area();
  364. root.fill(&WHITE).unwrap();
  365. let root =
  366. root.titled("253-bit Native Range Check Circuit Layout", ("sans-serif", 60)).unwrap();
  367. CircuitLayout::default().render(k, &circuit, &root).unwrap();
  368. for i in valid_values {
  369. println!("253-bit (valid) range check for {:?}", i);
  370. let circuit = RangeCheckCircuit { a: Value::known(i) };
  371. let prover = MockProver::run(k, &circuit, vec![]).unwrap();
  372. prover.assert_satisfied();
  373. println!("Constraints satisfied");
  374. }
  375. for i in invalid_values {
  376. println!("253-bit (invalid) range check for {:?}", i);
  377. let circuit = RangeCheckCircuit { a: Value::known(i) };
  378. let prover = MockProver::run(k, &circuit, vec![]).unwrap();
  379. assert!(prover.verify().is_err());
  380. }
  381. }
  382. }