From e9964f9855daa674b76f50886354676e8e6ee738 Mon Sep 17 00:00:00 2001 From: qoosmo Date: Thu, 27 Aug 2026 18:26:15 +0300 Subject: [PATCH 1/4] chore: professionalize multilinear Sumcheck research repository --- .github/workflows/ci.yml | 31 ++++ CONTRIBUTING.md | 22 +++ Cargo.lock | 2 +- Cargo.toml | 13 +- LICENSE-APACHE | 201 ++++++++++++++++++++++++++ LICENSE-MIT | 21 +++ README.md | 235 +++++++++++++++++++++++++++++++ SECURITY.md | 14 ++ benches/sumcheck.rs | 151 ++++++++------------ docs/ALGORITHMS.md | 151 ++++++++++++++++++++ docs/BENCHMARKS.md | 69 +++++++++ examples/basic_sumcheck.rs | 20 +++ src/circuit/bit_reverse_cache.rs | 4 +- src/circuit/canonical.rs | 70 +++++---- src/circuit/lagrange_decomp.rs | 84 +++++------ src/circuit/mod.rs | 8 +- src/circuit/sum_circuit.rs | 158 +++++++++++---------- src/lib.rs | 17 ++- src/poly/canonical.rs | 61 ++++---- src/poly/lagrange.rs | 95 +++++++------ src/poly/mod.rs | 4 - src/poly/traits.rs | 2 +- src/poly/uni.rs | 211 ++++++++++++++++----------- src/sumcheck/mod.rs | 7 - src/sumcheck/proof.rs | 26 ++-- src/sumcheck/prover.rs | 104 ++++++++------ src/sumcheck/verifier.rs | 168 ++++++++++++---------- tests/integration.rs | 176 +++++++++++------------ 28 files changed, 1500 insertions(+), 625 deletions(-) create mode 100644 .github/workflows/ci.yml create mode 100644 CONTRIBUTING.md create mode 100644 LICENSE-APACHE create mode 100644 LICENSE-MIT create mode 100644 README.md create mode 100644 SECURITY.md create mode 100644 docs/ALGORITHMS.md create mode 100644 docs/BENCHMARKS.md create mode 100644 examples/basic_sumcheck.rs diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..53a52fa --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,31 @@ +name: CI + +on: + push: + branches: [main] + pull_request: + +permissions: + contents: read + +jobs: + rust: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - name: Install Rust + uses: dtolnay/rust-toolchain@stable + with: + components: rustfmt, clippy + - name: Cache Cargo + uses: Swatinem/rust-cache@v2 + - name: Formatting + run: cargo fmt --all -- --check + - name: Check all targets + run: cargo check --all-targets --all-features + - name: Tests + run: cargo test --all-features + - name: Clippy + run: cargo clippy --all-targets --all-features + - name: Compile benchmarks + run: cargo bench --bench sumcheck --no-run diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md new file mode 100644 index 0000000..fd76fbe --- /dev/null +++ b/CONTRIBUTING.md @@ -0,0 +1,22 @@ +# Contributing + +Contributions should preserve the repository's research-oriented goals: explicit algebra, reproducible tests, and measurable performance. + +Before opening a pull request, run: + +```bash +cargo fmt --all -- --check +cargo check --all-targets --all-features +cargo test --all-features +cargo clippy --all-targets --all-features +cargo bench --bench sumcheck --no-run +``` + +For algorithmic changes, please include: + +- the mathematical recurrence or invariant being implemented; +- the expected asymptotic and/or field-operation cost; +- correctness tests, ideally including a cross-check against another representation; +- benchmark evidence when performance is part of the claim. + +Avoid performance claims that are not accompanied by enough machine/toolchain information to reproduce them. diff --git a/Cargo.lock b/Cargo.lock index 0d11192..af06f77 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -430,7 +430,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" [[package]] -name = "mlp_pro" +name = "multilinear-sumcheck" version = "0.1.0" dependencies = [ "ark-bn254", diff --git a/Cargo.toml b/Cargo.toml index b4f4868..b61b4cd 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,17 +1,18 @@ [package] -name = "mlp_pro" +name = "multilinear-sumcheck" version = "0.1.0" edition = "2021" authors = ["Ali Mkhida ", "Adil Iguider "] -description = "Multilinear polynomials via tree-based circuit and the Sumcheck protocol" +description = "Basis-aware multilinear polynomial algorithms and the Sumcheck protocol in Rust" license = "MIT OR Apache-2.0" -repository = "https://github.com/algorizk/mlp-pro" +repository = "https://github.com/qoosmo/multilinear-sumcheck" +homepage = "https://github.com/qoosmo/multilinear-sumcheck" +readme = "README.md" keywords = ["cryptography", "sumcheck", "multilinear", "zkp", "zero-knowledge"] categories = ["cryptography", "mathematics"] -readme = "README.md" [lib] -name = "mlp_pro" +name = "multilinear_sumcheck" path = "src/lib.rs" [[bench]] @@ -37,4 +38,4 @@ panic = "abort" [profile.bench] inherits = "release" -debug = true \ No newline at end of file +debug = true diff --git a/LICENSE-APACHE b/LICENSE-APACHE new file mode 100644 index 0000000..4f826bb --- /dev/null +++ b/LICENSE-APACHE @@ -0,0 +1,201 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright 2026 Ali Mkhida and contributors + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/LICENSE-MIT b/LICENSE-MIT new file mode 100644 index 0000000..dc71d54 --- /dev/null +++ b/LICENSE-MIT @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2026 Ali Mkhida and contributors + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/README.md b/README.md new file mode 100644 index 0000000..0abb12c --- /dev/null +++ b/README.md @@ -0,0 +1,235 @@ +# Multilinear Sumcheck + +[![CI](https://github.com/qoosmo/multilinear-sumcheck/actions/workflows/ci.yml/badge.svg)](https://github.com/qoosmo/multilinear-sumcheck/actions/workflows/ci.yml) +[![License: MIT OR Apache-2.0](https://img.shields.io/badge/license-MIT%20OR%20Apache--2.0-blue.svg)](#license) +[![Rust](https://img.shields.io/badge/language-Rust-orange.svg)](https://www.rust-lang.org/) + +**Basis-aware multilinear polynomial algorithms and a research implementation of the Sumcheck protocol in Rust.** + +This repository studies how the representation of a multilinear polynomial changes the structure and cost of evaluation and Sumcheck. It implements the same multilinear object in both the canonical/monomial basis and the Lagrange/evaluation basis over the Boolean hypercube, then exposes the tree decompositions and scalar sum circuits that connect those representations to Sumcheck. + +The goal is not to provide a production SNARK. The goal is to make the algebraic structure explicit, testable, benchmarkable, and easy to inspect. + +## What is implemented + +- **Canonical multilinear polynomials** stored as dense monomial coefficients. +- **Lagrange multilinear polynomials** stored as evaluations over `{0,1}^n`. +- **Linear-time canonical evaluation** by repeated variable folding. +- **Standard, optimized, and Rayon-parallel Lagrange evaluation kernels**. +- **Canonical tree decomposition** by even/odd coefficient splitting. +- **Canonical → Lagrange conversion** through a tree whose leaves are Boolean-hypercube evaluations in bit-reversed order. +- **Bit-reversal caching** used by the decomposition and sum circuits. +- **Boolean sum circuits** for both bases. +- **Canonical and Lagrange Sumcheck provers**. +- **Stateless Sumcheck verifier** with explicit round-consistency and final-oracle checks. +- **End-to-end integration tests** covering completeness, transcript tampering, cross-basis agreement, edge cases, and zero-variable transcripts. +- **Criterion benchmarks** for evaluation, proving, verification, and prover construction. + +## Why the basis matters + +Let `f : F^n -> F` be multilinear and let `N = 2^n`. + +In the canonical basis we store + +```text +f(x_1,...,x_n) = sum_{S subseteq [n]} alpha_S prod_{i in S} x_i. +``` + +In the Lagrange basis we store the evaluation table + +```text +(f(0), f(1), ..., f(N-1)) +``` + +over the Boolean hypercube. + +Both contain the same mathematical information, but they induce different local recurrences and different arithmetic trade-offs. This repository makes those differences explicit rather than hiding them behind a generic polynomial interface. + +## Core architecture + +```text +Canonical coefficients + | + | even/odd tree decomposition + v + canonical q_j tree + | + +--------------------------+ + | | + v v +bit-reversed leaves canonical sum circuit h_j + | | + | a / (a+b) gates | Sumcheck round extraction + v v +Lagrange evaluations Canonical Sumcheck prover + | | + v | +Lagrange sum circuit h_j | + | | + v v +Lagrange Sumcheck prover ----> SumcheckProof + | + v + Verifier + | + v + final oracle evaluation +``` + +## Selected complexity properties + +For `N = 2^n`: + +| Operation | Implementation | Arithmetic / asymptotic cost | +| --- | --- | --- | +| Canonical evaluation | `CanonicalPoly::eval_naive` | `O(N log N)` | +| Canonical evaluation | `CanonicalPoly::eval_circuit` | `N-1` multiplications + `N-1` additions | +| Lagrange evaluation | `eval_standard` | fold `(1-r)u + rv` | +| Lagrange evaluation | `eval_optimized` | fold `u + r(v-u)`, one multiplication per pair | +| Canonical → Lagrange decomposition | `LagrangeDecomp::build` | `n * 2^(n-1)` additions | +| Sum-circuit construction | canonical / Lagrange | `O(N)` | +| Sumcheck proof size | `SumcheckProof` | `2n + 1` field elements | +| Verifier | `Verifier` | `O(n)` round checks + final oracle check | + +The benchmark suite measures implementation-level runtime separately from these field-operation counts. + +## Quick start + +Requirements: + +- stable Rust toolchain +- Cargo + +```bash +cargo test --all-features +cargo run --example basic_sumcheck +cargo bench --bench sumcheck +``` + +## Minimal example + +```rust +use ark_bn254::Fr; +use multilinear_sumcheck::poly::CanonicalPoly; +use multilinear_sumcheck::sumcheck::{CanonicalProver, Verifier}; + +fn main() { + let f = CanonicalPoly::new( + (1u64..=8).map(Fr::from).collect() + ); + let challenges = [Fr::from(3u64), Fr::from(7u64), Fr::from(11u64)]; + + let prover = CanonicalProver::new(&f); + let proof = prover.prove(&challenges); + let oracle_eval = f.eval_circuit(&challenges); + + Verifier::verify(&proof, &challenges, oracle_eval) + .expect("valid Sumcheck proof"); +} +``` + +A complete runnable version is available in [`examples/basic_sumcheck.rs`](examples/basic_sumcheck.rs). + +## Sumcheck boundary + +This repository implements the **algebraic Sumcheck core**. In particular: + +1. the prover constructs one degree-1 round polynomial per variable; +2. the verifier checks Boolean sums across rounds; +3. the verifier checks the last round against an externally supplied evaluation `f(r_1,...,r_n)`. + +The current code intentionally does **not** implement: + +- Fiat-Shamir transcript generation; +- a polynomial commitment scheme; +- Merkle commitments; +- zero-knowledge masking/blinding; +- recursive composition; +- a complete SNARK or STARK; +- production hardening or side-channel guarantees. + +This separation is deliberate: it keeps the repository focused on the multilinear/Sumcheck algebra and makes the trust boundary explicit. + +## Testing + +The integration suite checks: + +- honest canonical proofs verify; +- honest Lagrange proofs verify; +- modified claimed sums are rejected; +- modified intermediate round polynomials are rejected; +- incorrect oracle values are rejected; +- canonical and Lagrange representations agree on the same polynomial; +- proof size scales as `2n+1` field elements; +- zero and constant polynomials behave correctly; +- zero-variable transcripts are handled without panics. + +Run: + +```bash +cargo test --all-features +``` + +## Benchmarks + +Criterion benchmarks cover: + +- standard vs optimized vs parallel multilinear evaluation; +- comparison with `ark-poly::DenseMultilinearExtension`; +- canonical Sumcheck proving and verification; +- Lagrange Sumcheck proving and verification; +- prover construction; +- canonical → Lagrange conversion. + +The benchmark inputs use deterministic randomness and currently include `n = 10, 15, 20`. + +```bash +cargo bench --bench sumcheck +``` + +Criterion writes its HTML report to: + +```text +target/criterion/report/index.html +``` + +See [`docs/BENCHMARKS.md`](docs/BENCHMARKS.md) for reproducibility guidance and a results template. No performance claims are published without the machine/toolchain context needed to reproduce them. + +## Repository layout + +```text +src/ + poly/ canonical and Lagrange multilinear representations + circuit/ decomposition trees, bit reversal, and Boolean sum circuits + sumcheck/ proof transcript, provers, and verifier +benches/ Criterion benchmark suite +tests/ end-to-end integration tests +examples/ minimal runnable examples +docs/ algorithm and benchmark notes +``` + +## Research status + +This is a research-oriented implementation intended for experimentation, validation of algebraic recurrences, operation-count accounting, and benchmarking. APIs may evolve as the underlying research evolves. + +For a more detailed walk-through of the algorithms, see [`docs/ALGORITHMS.md`](docs/ALGORITHMS.md). + +## Security + +This repository is **not production cryptography**. Do not use it to secure funds or sensitive systems without a separate security review and the missing protocol layers described above. + +See [`SECURITY.md`](SECURITY.md) for reporting guidance. + +## Authors + +- Ali Mkhida +- Adil Iguider + +## License + +Licensed under either of: + +- Apache License, Version 2.0 ([`LICENSE-APACHE`](LICENSE-APACHE)); or +- MIT License ([`LICENSE-MIT`](LICENSE-MIT)). + +at your option. diff --git a/SECURITY.md b/SECURITY.md new file mode 100644 index 0000000..ee42494 --- /dev/null +++ b/SECURITY.md @@ -0,0 +1,14 @@ +# Security Policy + +## Research status + +This repository is a research implementation of multilinear polynomial algorithms and the algebraic Sumcheck protocol. It is **not a production proof system** and should not be used to protect funds, credentials, or other sensitive assets without an independent security review and the missing protocol layers described in the README. + +## Reporting a security issue + +Please avoid filing a public issue for a vulnerability that could affect downstream cryptographic use. Send a concise report to the maintainer address listed in `Cargo.toml`, including: + +- affected commit; +- minimal reproduction; +- expected vs observed behavior; +- potential cryptographic impact. diff --git a/benches/sumcheck.rs b/benches/sumcheck.rs index 13a5418..2cc22e8 100644 --- a/benches/sumcheck.rs +++ b/benches/sumcheck.rs @@ -1,10 +1,11 @@ -//! Benchmarks for `mlp_pro`. +//! Benchmarks for `multilinear_sumcheck`. //! //! # Groups //! //! 1. `multilinear_eval` — four evaluation kernels for `LagrangePoly` //! 2. `sumcheck_canonical` — `CanonicalProver::prove` + `Verifier::verify` //! 3. `sumcheck_lagrange` — `LagrangeProver::prove` + `Verifier::verify` +//! 4. `prover_construction` — one-time prover/circuit construction costs //! //! # Running //! @@ -18,14 +19,14 @@ use ark_bn254::Fr; use ark_ff::UniformRand; use ark_poly::DenseMultilinearExtension; use ark_poly::MultilinearExtension; -use ark_std::rand::SeedableRng; use ark_std::rand::rngs::StdRng; +use ark_std::rand::SeedableRng; use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; -use mlp_pro::circuit::LagrangeDecomp; -use mlp_pro::poly::{CanonicalPoly, LagrangePoly}; -use mlp_pro::sumcheck::prover::{CanonicalProver, LagrangeProver}; -use mlp_pro::sumcheck::verifier::Verifier; +use multilinear_sumcheck::circuit::LagrangeDecomp; +use multilinear_sumcheck::poly::{CanonicalPoly, LagrangePoly}; +use multilinear_sumcheck::sumcheck::prover::{CanonicalProver, LagrangeProver}; +use multilinear_sumcheck::sumcheck::verifier::Verifier; // ───────────────────────────────────────────────────────────────────────────── // Input generation @@ -39,18 +40,18 @@ fn make_rng() -> StdRng { /// Random evaluation vector + point of length `n`. fn make_lagrange_inputs(n: usize) -> (Vec, Vec) { let mut rng = make_rng(); - let big_n = 1usize << n; - let evals: Vec = (0..big_n).map(|_| Fr::rand(&mut rng)).collect(); - let point: Vec = (0..n).map(|_| Fr::rand(&mut rng)).collect(); + let big_n = 1usize << n; + let evals: Vec = (0..big_n).map(|_| Fr::rand(&mut rng)).collect(); + let point: Vec = (0..n).map(|_| Fr::rand(&mut rng)).collect(); (evals, point) } /// Random canonical coefficient vector + point of length `n`. fn make_canonical_inputs(n: usize) -> (Vec, Vec) { let mut rng = make_rng(); - let big_n = 1usize << n; + let big_n = 1usize << n; let coeffs: Vec = (0..big_n).map(|_| Fr::rand(&mut rng)).collect(); - let point: Vec = (0..n).map(|_| Fr::rand(&mut rng)).collect(); + let point: Vec = (0..n).map(|_| Fr::rand(&mut rng)).collect(); (coeffs, point) } @@ -67,43 +68,33 @@ fn bench_eval(c: &mut Criterion) { // 1. eval_standard { let poly = LagrangePoly::new(evals.clone()); - group.bench_with_input( - BenchmarkId::new("eval_standard", n), - &n, - |b, _| b.iter(|| poly.eval_standard(&point)), - ); + group.bench_with_input(BenchmarkId::new("eval_standard", n), &n, |b, _| { + b.iter(|| poly.eval_standard(&point)) + }); } // 2. eval_optimized { let poly = LagrangePoly::new(evals.clone()); - group.bench_with_input( - BenchmarkId::new("eval_optimized", n), - &n, - |b, _| b.iter(|| poly.eval_optimized(&point)), - ); + group.bench_with_input(BenchmarkId::new("eval_optimized", n), &n, |b, _| { + b.iter(|| poly.eval_optimized(&point)) + }); } // 3. eval_parallel { let poly = LagrangePoly::new(evals.clone()); - group.bench_with_input( - BenchmarkId::new("eval_parallel", n), - &n, - |b, _| b.iter(|| poly.eval_parallel(&point)), - ); + group.bench_with_input(BenchmarkId::new("eval_parallel", n), &n, |b, _| { + b.iter(|| poly.eval_parallel(&point)) + }); } // 4. ark-poly DenseMultilinearExtension::evaluate { - let ark_poly = DenseMultilinearExtension::from_evaluations_vec( - n, evals.clone(), - ); - group.bench_with_input( - BenchmarkId::new("ark_poly_evaluate", n), - &n, - |b, _| b.iter(|| ark_poly.evaluate(&point)), - ); + let ark_poly = DenseMultilinearExtension::from_evaluations_vec(n, evals.clone()); + group.bench_with_input(BenchmarkId::new("ark_poly_evaluate", n), &n, |b, _| { + b.iter(|| ark_poly.evaluate(&point)) + }); } } @@ -125,38 +116,31 @@ fn bench_sumcheck_canonical(c: &mut Criterion) { let prover = CanonicalProver::new(&f); // Prover: build the full proof transcript. - group.bench_with_input( - BenchmarkId::new("prove", n), - &n, - |b, _| b.iter(|| prover.prove(&challenges)), - ); + group.bench_with_input(BenchmarkId::new("prove", n), &n, |b, _| { + b.iter(|| prover.prove(&challenges)) + }); // Verifier: verify an already-produced proof. // The verifier is O(n) and should be very fast. { - let proof = prover.prove(&challenges); + let proof = prover.prove(&challenges); let oracle_eval = f.eval_circuit(&challenges); - group.bench_with_input( - BenchmarkId::new("verify", n), - &n, - |b, _| b.iter(|| { + group.bench_with_input(BenchmarkId::new("verify", n), &n, |b, _| { + b.iter(|| { Verifier::verify(&proof, &challenges, oracle_eval) .expect("valid proof must verify") - }), - ); + }) + }); } // End-to-end: prove + verify in a single timed block. - group.bench_with_input( - BenchmarkId::new("prove_and_verify", n), - &n, - |b, _| b.iter(|| { - let proof = prover.prove(&challenges); + group.bench_with_input(BenchmarkId::new("prove_and_verify", n), &n, |b, _| { + b.iter(|| { + let proof = prover.prove(&challenges); let oracle_eval = f.eval_circuit(&challenges); - Verifier::verify(&proof, &challenges, oracle_eval) - .expect("valid proof must verify") - }), - ); + Verifier::verify(&proof, &challenges, oracle_eval).expect("valid proof must verify") + }) + }); } group.finish(); @@ -176,37 +160,30 @@ fn bench_sumcheck_lagrange(c: &mut Criterion) { let prover = LagrangeProver::new(&f); // Prover - group.bench_with_input( - BenchmarkId::new("prove", n), - &n, - |b, _| b.iter(|| prover.prove(&challenges)), - ); + group.bench_with_input(BenchmarkId::new("prove", n), &n, |b, _| { + b.iter(|| prover.prove(&challenges)) + }); // Verifier { - let proof = prover.prove(&challenges); + let proof = prover.prove(&challenges); let oracle_eval = f.eval_optimized(&challenges); - group.bench_with_input( - BenchmarkId::new("verify", n), - &n, - |b, _| b.iter(|| { + group.bench_with_input(BenchmarkId::new("verify", n), &n, |b, _| { + b.iter(|| { Verifier::verify(&proof, &challenges, oracle_eval) .expect("valid proof must verify") - }), - ); + }) + }); } // End-to-end - group.bench_with_input( - BenchmarkId::new("prove_and_verify", n), - &n, - |b, _| b.iter(|| { - let proof = prover.prove(&challenges); + group.bench_with_input(BenchmarkId::new("prove_and_verify", n), &n, |b, _| { + b.iter(|| { + let proof = prover.prove(&challenges); let oracle_eval = f.eval_optimized(&challenges); - Verifier::verify(&proof, &challenges, oracle_eval) - .expect("valid proof must verify") - }), - ); + Verifier::verify(&proof, &challenges, oracle_eval).expect("valid proof must verify") + }) + }); } group.finish(); @@ -224,22 +201,18 @@ fn bench_prover_construction(c: &mut Criterion) { for n in [10, 15, 20] { let (coeffs, _) = make_canonical_inputs(n); - let (evals, _) = make_lagrange_inputs(n); + let (evals, _) = make_lagrange_inputs(n); - let canon_f = CanonicalPoly::new(coeffs); + let canon_f = CanonicalPoly::new(coeffs); let lagrange_f = LagrangePoly::new(evals); - group.bench_with_input( - BenchmarkId::new("canonical_new", n), - &n, - |b, _| b.iter(|| CanonicalProver::new(&canon_f)), - ); + group.bench_with_input(BenchmarkId::new("canonical_new", n), &n, |b, _| { + b.iter(|| CanonicalProver::new(&canon_f)) + }); - group.bench_with_input( - BenchmarkId::new("lagrange_new", n), - &n, - |b, _| b.iter(|| LagrangeProver::new(&lagrange_f)), - ); + group.bench_with_input(BenchmarkId::new("lagrange_new", n), &n, |b, _| { + b.iter(|| LagrangeProver::new(&lagrange_f)) + }); // Also measure LagrangeDecomp::build + to_lagrange conversion cost // since that is the typical pipeline starting from a canonical poly. @@ -264,4 +237,4 @@ criterion_group!( bench_sumcheck_lagrange, bench_prover_construction, ); -criterion_main!(benches); \ No newline at end of file +criterion_main!(benches); diff --git a/docs/ALGORITHMS.md b/docs/ALGORITHMS.md new file mode 100644 index 0000000..c37ebc1 --- /dev/null +++ b/docs/ALGORITHMS.md @@ -0,0 +1,151 @@ +# Algorithms + +This note documents the mathematical structure represented by the code. It is intentionally implementation-oriented: notation is kept close to the Rust types and tree indices. + +## 1. Multilinear representations + +For `n` variables, let `N = 2^n`. + +### Canonical basis + +`CanonicalPoly` stores `N` coefficients. Index `j` encodes the monomial by the binary expansion of `j`; a set bit selects the corresponding variable. + +The direct evaluator computes every monomial independently. The circuit evaluator instead folds one variable at a time: + +```text +u, v -> u + r v +``` + +until one field element remains. + +### Lagrange basis + +`LagrangePoly` stores the `N` evaluations on the Boolean hypercube. Evaluation at an arbitrary point uses repeated interpolation folds. + +Standard fold: + +```text +(1-r)u + rv +``` + +Optimized fold: + +```text +u + r(v-u) +``` + +The optimized form removes one multiplication per pair at the cost of an additional addition/subtraction. + +## 2. Tree decomposition + +The canonical decomposition splits a polynomial + +```text +q = a + x_i b +``` + +into the children `a` and `b` by taking even- and odd-indexed coefficients. + +The Lagrange-conversion decomposition uses the children + +```text +p_{2j} = a +p_{2j+1} = a + b +``` + +so that the leaf layer becomes the Boolean evaluation table, up to bit reversal. + +For `N = 2^n`, the full binary tree contains `2N-1` nodes and the total number of field elements stored across all layers is `(n+1)N`. + +The conversion performs exactly + +```text +n * 2^(n-1) +``` + +field additions. + +## 3. Bit-reversal order + +Repeated low-variable-first splitting naturally produces a bit-reversed leaf order. The implementation caches the permutation and applies its inverse when converting the leaf layer into a standard `LagrangePoly` evaluation vector. + +Bit reversal is an involution, so the same table serves both directions. + +## 4. Boolean sum circuits + +For each decomposition node, the code stores the Boolean hypercube sum of the sub-polynomial rooted at that node. + +### Canonical recurrence + +```text +h_j = 2 h_{2j} + h_{2j+1} +``` + +The implementation uses additions (`h_{2j} + h_{2j} + h_{2j+1}`) rather than a general field multiplication by `2`. + +### Lagrange recurrence + +```text +h_j = h_{2j} + h_{2j+1} +``` + +In both cases the root is the claimed Boolean sum + +```text +H(f) = sum_{x in {0,1}^n} f(x). +``` + +## 5. Sumcheck + +For a multilinear polynomial, each Sumcheck message is a degree-1 univariate polynomial + +```text +s_j(X) = a_j + b_j X. +``` + +The proof stores one claimed sum and two field elements per round, giving `2n+1` field elements in total. + +The verifier checks: + +```text +s_1(0) + s_1(1) = claimed_sum +``` + +then, for later rounds, + +```text +s_j(0) + s_j(1) = s_{j-1}(r_{j-1}), +``` + +and finally + +```text +s_n(r_n) = f(r_1,...,r_n). +``` + +The final value is supplied as an oracle evaluation. This repository does not provide the commitment layer that would authenticate that oracle value in a complete proof system. + +## 6. Canonical and Lagrange provers + +Both provers precompute a scalar sum circuit and then extract/fold one tree layer per Sumcheck round. + +The canonical prover uses the canonical recurrence directly. The Lagrange prover uses interpolation folds of the form + +```text +u + r(v-u). +``` + +The integration suite checks that, after converting the same polynomial between bases, both provers agree on the claimed sum, oracle evaluation, and round polynomials. + +## 7. Scope + +The implementation is useful for studying: + +- basis-dependent arithmetic costs; +- tree representations of multilinear evaluation; +- Sumcheck prover structure; +- proof transcript size; +- cross-basis equivalence; +- implementation-level performance. + +It intentionally stops before Fiat-Shamir, polynomial commitments, zero-knowledge masking, recursion, and complete proof-system composition. diff --git a/docs/BENCHMARKS.md b/docs/BENCHMARKS.md new file mode 100644 index 0000000..1c1e702 --- /dev/null +++ b/docs/BENCHMARKS.md @@ -0,0 +1,69 @@ +# Benchmarking + +The benchmark suite uses Criterion and deterministic random inputs so repeated runs compare the same logical instances. + +## Run + +```bash +cargo bench --bench sumcheck +``` + +The HTML report is written to: + +```text +target/criterion/report/index.html +``` + +## Benchmark groups + +1. `multilinear_eval` + - `eval_standard` + - `eval_optimized` + - `eval_parallel` + - `ark_poly_evaluate` +2. `sumcheck_canonical` + - `prove` + - `verify` + - `prove_and_verify` +3. `sumcheck_lagrange` + - `prove` + - `verify` + - `prove_and_verify` +4. `prover_construction` + - canonical prover construction + - Lagrange prover construction + - canonical → Lagrange conversion + +The current benchmark sizes are `n = 10, 15, 20`. + +## Reproducibility record + +Every published benchmark should include: + +```text +Date: +Git commit: +OS: +CPU: +Logical cores: +RAM: +rustc --version: +cargo --version: +Build mode: release / Criterion +``` + +## Results template + +Do not fill this table from memory or from a different machine. Record the actual Criterion estimates from the machine described above. + +| n | ark-poly eval | optimized eval | parallel eval | canonical prove | Lagrange prove | verify | +| ---: | ---: | ---: | ---: | ---: | ---: | ---: | +| 10 | | | | | | | +| 15 | | | | | | | +| 20 | | | | | | | + +## Interpretation + +Runtime results should be kept separate from algebraic operation counts. A method with fewer field multiplications is not automatically faster on every machine: allocation, memory bandwidth, cache behavior, parallel scheduling, and field implementation all matter. + +For that reason the README states exact operation counts only when they follow directly from the algorithm, and leaves machine-dependent performance claims to reproducible benchmark output. diff --git a/examples/basic_sumcheck.rs b/examples/basic_sumcheck.rs new file mode 100644 index 0000000..242f34b --- /dev/null +++ b/examples/basic_sumcheck.rs @@ -0,0 +1,20 @@ +use ark_bn254::Fr; +use multilinear_sumcheck::poly::CanonicalPoly; +use multilinear_sumcheck::sumcheck::{CanonicalProver, Verifier}; + +fn main() { + let f = CanonicalPoly::new((1u64..=8).map(Fr::from).collect()); + let challenges = [Fr::from(3u64), Fr::from(7u64), Fr::from(11u64)]; + + let prover = CanonicalProver::new(&f); + let proof = prover.prove(&challenges); + let oracle_eval = f.eval_circuit(&challenges); + + Verifier::verify(&proof, &challenges, oracle_eval).expect("valid Sumcheck proof must verify"); + + println!( + "verified {}-round Sumcheck proof ({} field elements)", + proof.num_vars(), + proof.size_in_field_elements() + ); +} diff --git a/src/circuit/bit_reverse_cache.rs b/src/circuit/bit_reverse_cache.rs index a4f5038..3cb48d4 100644 --- a/src/circuit/bit_reverse_cache.rs +++ b/src/circuit/bit_reverse_cache.rs @@ -41,7 +41,7 @@ pub fn get_bit_rev_table(n: usize) -> Option<&'static [usize]> { 10 => Some(&BIT_REV_N10), 15 => Some(&BIT_REV_N15), 20 => Some(bit_rev_n20()), - _ => None, + _ => None, } } @@ -145,4 +145,4 @@ mod tests { let cow = get_or_build(8); assert_eq!(cow.as_ref(), build_bit_reverse_table(8).as_slice()); } -} \ No newline at end of file +} diff --git a/src/circuit/canonical.rs b/src/circuit/canonical.rs index 1810a00..95de2c3 100644 --- a/src/circuit/canonical.rs +++ b/src/circuit/canonical.rs @@ -1,7 +1,7 @@ use ark_ff::Field; -use crate::poly::{CanonicalPoly, MlPoly}; use super::bit_reverse_cache; +use crate::poly::{CanonicalPoly, MlPoly}; pub fn bit_reverse(k: usize, n: usize) -> usize { let mut k = k; @@ -27,7 +27,7 @@ pub struct CanonicalDecomp { impl CanonicalDecomp { pub fn build(f: &CanonicalPoly) -> Self { - let n = f.num_vars(); + let n = f.num_vars(); let big_n = f.num_evals(); let mut nodes: Vec>> = vec![None; 2 * big_n - 1]; @@ -35,39 +35,45 @@ impl CanonicalDecomp { for i in 0..n { let layer_start = 1usize << i; - let layer_end = 1usize << (i + 1); + let layer_end = 1usize << (i + 1); for j in layer_start..layer_end { - let qj = nodes[j - 1].take().expect("node must be initialised"); - let beta = qj.coeffs().to_vec(); - let half = beta.len() / 2; + let qj = nodes[j - 1].take().expect("node must be initialised"); + let beta = qj.coeffs().to_vec(); + let half = beta.len() / 2; - let left_coeffs: Vec = (0..half).map(|r| beta[2 * r]).collect(); + let left_coeffs: Vec = (0..half).map(|r| beta[2 * r]).collect(); let right_coeffs: Vec = (0..half).map(|r| beta[2 * r + 1]).collect(); - nodes[j - 1] = Some(qj); + nodes[j - 1] = Some(qj); nodes[2 * j - 1] = Some(CanonicalPoly::new(left_coeffs)); - nodes[2 * j] = Some(CanonicalPoly::new(right_coeffs)); + nodes[2 * j] = Some(CanonicalPoly::new(right_coeffs)); } } let nodes: Vec> = nodes .into_iter() .enumerate() - .map(|(idx, opt)| { - opt.unwrap_or_else(|| panic!("node q_{} was not filled", idx + 1)) - }) + .map(|(idx, opt)| opt.unwrap_or_else(|| panic!("node q_{} was not filled", idx + 1))) .collect(); let bit_rev_table = bit_reverse_cache::get_or_build(n).into_owned(); - Self { n, big_n, nodes, bit_rev_table } + Self { + n, + big_n, + nodes, + bit_rev_table, + } } #[inline] pub fn q(&self, j: usize) -> &CanonicalPoly { - assert!(j >= 1 && j <= 2 * self.big_n - 1, - "q index {j} out of range [1, {}]", 2 * self.big_n - 1); + assert!( + j >= 1 && j <= 2 * self.big_n - 1, + "q index {j} out of range [1, {}]", + 2 * self.big_n - 1 + ); &self.nodes[j - 1] } @@ -76,13 +82,13 @@ impl CanonicalDecomp { } pub fn leaves(&self) -> &[CanonicalPoly] { - &self.nodes[self.big_n - 1 .. 2 * self.big_n - 1] + &self.nodes[self.big_n - 1..2 * self.big_n - 1] } pub fn layer(&self, i: usize) -> &[CanonicalPoly] { assert!(i <= self.n, "layer {i} out of range [0, {}]", self.n); let start = (1usize << i) - 1; - let end = (1usize << (i + 1)) - 1; + let end = (1usize << (i + 1)) - 1; &self.nodes[start..end] } @@ -92,7 +98,7 @@ impl CanonicalDecomp { pub fn leaves_are_bit_reverse_of_root(&self) -> bool { let root_coeffs = self.root().coeffs(); - let leaves = self.leaves(); + let leaves = self.leaves(); (0..self.big_n).all(|k| { let rev_k = self.bit_rev_table[k]; leaves[k].coeffs()[0] == root_coeffs[rev_k] @@ -106,7 +112,9 @@ mod tests { use crate::poly::MlPoly; use ark_bn254::Fr; - fn fr(n: u64) -> Fr { Fr::from(n) } + fn fr(n: u64) -> Fr { + Fr::from(n) + } #[test] fn bit_reverse_n3_cases() { @@ -179,22 +187,22 @@ mod tests { #[test] fn split_matches_paper_example_n3() { - let coeffs = vec![fr(1),fr(2),fr(3),fr(4),fr(5),fr(6),fr(7),fr(8)]; + let coeffs = vec![fr(1), fr(2), fr(3), fr(4), fr(5), fr(6), fr(7), fr(8)]; let f = CanonicalPoly::new(coeffs); let d = CanonicalDecomp::build(&f); - assert_eq!(d.q(2).coeffs(), &[fr(1),fr(3),fr(5),fr(7)]); - assert_eq!(d.q(3).coeffs(), &[fr(2),fr(4),fr(6),fr(8)]); + assert_eq!(d.q(2).coeffs(), &[fr(1), fr(3), fr(5), fr(7)]); + assert_eq!(d.q(3).coeffs(), &[fr(2), fr(4), fr(6), fr(8)]); } #[test] fn second_layer_split_correct_n3() { - let coeffs = vec![fr(1),fr(2),fr(3),fr(4),fr(5),fr(6),fr(7),fr(8)]; + let coeffs = vec![fr(1), fr(2), fr(3), fr(4), fr(5), fr(6), fr(7), fr(8)]; let f = CanonicalPoly::new(coeffs); let d = CanonicalDecomp::build(&f); - assert_eq!(d.q(4).coeffs(), &[fr(1),fr(5)]); - assert_eq!(d.q(5).coeffs(), &[fr(3),fr(7)]); - assert_eq!(d.q(6).coeffs(), &[fr(2),fr(6)]); - assert_eq!(d.q(7).coeffs(), &[fr(4),fr(8)]); + assert_eq!(d.q(4).coeffs(), &[fr(1), fr(5)]); + assert_eq!(d.q(5).coeffs(), &[fr(3), fr(7)]); + assert_eq!(d.q(6).coeffs(), &[fr(2), fr(6)]); + assert_eq!(d.q(7).coeffs(), &[fr(4), fr(8)]); } #[test] @@ -223,14 +231,14 @@ mod tests { let f = CanonicalPoly::new((0..8).map(|i| fr(i as u64 + 1)).collect()); let d = CanonicalDecomp::build(&f); for j in 1..d.big_n { - let qj = d.q(j).coeffs().to_vec(); - let q_left = d.q(2 * j).coeffs().to_vec(); + let qj = d.q(j).coeffs().to_vec(); + let q_left = d.q(2 * j).coeffs().to_vec(); let q_right = d.q(2 * j + 1).coeffs().to_vec(); let m = qj.len(); for r in 0..m / 2 { - assert_eq!(qj[2 * r], q_left[r], "j={j} even r={r}"); + assert_eq!(qj[2 * r], q_left[r], "j={j} even r={r}"); assert_eq!(qj[2 * r + 1], q_right[r], "j={j} odd r={r}"); } } } -} \ No newline at end of file +} diff --git a/src/circuit/lagrange_decomp.rs b/src/circuit/lagrange_decomp.rs index 7e9d26c..4c1f0bc 100644 --- a/src/circuit/lagrange_decomp.rs +++ b/src/circuit/lagrange_decomp.rs @@ -1,7 +1,7 @@ use ark_ff::Field; -use crate::poly::{CanonicalPoly, LagrangePoly, MlPoly}; use super::bit_reverse_cache::get_or_build; +use crate::poly::{CanonicalPoly, LagrangePoly, MlPoly}; /// The full sequence `(p_j)_{1 ≤ j ≤ 2N-1}` from the canonical-to-Lagrange /// circuit decomposition of `f`. @@ -50,7 +50,7 @@ impl LagrangeDecomp { /// # Panics /// Panics if `f.num_evals()` is not a power of two greater than zero. pub fn build(f: &CanonicalPoly) -> Self { - let n = f.num_vars(); + let n = f.num_vars(); let big_n = f.num_evals(); let mut nodes: Vec>> = vec![None; 2 * big_n - 1]; @@ -63,13 +63,13 @@ impl LagrangeDecomp { // At layer i, nodes are at 1-based indices [2^{i-1}, 2^i - 1]. for i in 1..=n { let layer_start = 1usize << (i - 1); // 2^{i-1} (1-based) - let layer_end = 1usize << i; // 2^i (exclusive, 1-based) + let layer_end = 1usize << i; // 2^i (exclusive, 1-based) for j in layer_start..layer_end { // Take p_j out of its slot. - let pj = nodes[j - 1].take().expect("node must be initialised"); + let pj = nodes[j - 1].take().expect("node must be initialised"); let coeffs = pj.coeffs().to_vec(); - let half = coeffs.len() / 2; + let half = coeffs.len() / 2; // Split p_j = a + x_i · b: // a = even-indexed coefficients (p_{2j}) @@ -87,7 +87,7 @@ impl LagrangeDecomp { // Store children. nodes[2 * j - 1] = Some(CanonicalPoly::new(a)); - nodes[2 * j] = Some(CanonicalPoly::new(a_plus_b)); + nodes[2 * j] = Some(CanonicalPoly::new(a_plus_b)); } } @@ -95,14 +95,18 @@ impl LagrangeDecomp { let nodes: Vec> = nodes .into_iter() .enumerate() - .map(|(idx, opt)| { - opt.unwrap_or_else(|| panic!("node p_{} was not filled", idx + 1)) - }) + .map(|(idx, opt)| opt.unwrap_or_else(|| panic!("node p_{} was not filled", idx + 1))) .collect(); let bit_rev_table = get_or_build(n).into_owned(); - Self { n, big_n, nodes, bit_rev_table, addition_count } + Self { + n, + big_n, + nodes, + bit_rev_table, + addition_count, + } } // ── Accessors (paper notation) ──────────────────────────────────────────── @@ -115,7 +119,8 @@ impl LagrangeDecomp { pub fn p(&self, j: usize) -> &CanonicalPoly { assert!( j >= 1 && j <= 2 * self.big_n - 1, - "p index {j} out of range [1, {}]", 2 * self.big_n - 1 + "p index {j} out of range [1, {}]", + 2 * self.big_n - 1 ); &self.nodes[j - 1] } @@ -132,7 +137,7 @@ impl LagrangeDecomp { /// They are in **bit-reverse permutation order**: /// leaf `k` (0-based) holds `f(rev(k))`. pub fn leaves(&self) -> &[CanonicalPoly] { - &self.nodes[self.big_n - 1 .. 2 * self.big_n - 1] + &self.nodes[self.big_n - 1..2 * self.big_n - 1] } /// Layer `i` (1-based, matching the paper): slice of nodes at depth `i`. @@ -144,10 +149,13 @@ impl LagrangeDecomp { /// # Panics /// Panics if `i == 0` or `i > n + 1`. pub fn layer(&self, i: usize) -> &[CanonicalPoly] { - assert!(i >= 1 && i <= self.n + 1, - "layer {i} out of range [1, {}]", self.n + 1); + assert!( + i >= 1 && i <= self.n + 1, + "layer {i} out of range [1, {}]", + self.n + 1 + ); let start = (1usize << (i - 1)) - 1; - let end = (1usize << i) - 1; + let end = (1usize << i) - 1; &self.nodes[start..end] } @@ -164,7 +172,7 @@ impl LagrangeDecomp { let mut evals = vec![F::zero(); self.big_n]; for k in 0..self.big_n { - let rev_k = self.bit_rev_table[k]; + let rev_k = self.bit_rev_table[k]; // leaf k holds f(rev(k)), so f(rev(k)) goes to position rev(k). evals[rev_k] = leaves[k].coeffs()[0]; } @@ -195,7 +203,9 @@ mod tests { use crate::poly::MlPoly; use ark_bn254::Fr; - fn fr(n: u64) -> Fr { Fr::from(n) } + fn fr(n: u64) -> Fr { + Fr::from(n) + } // ── Structural properties ───────────────────────────────────────────────── @@ -255,7 +265,7 @@ mod tests { /// p_3 = a + b = [3,7,11,15] #[test] fn gate_rule_first_layer_n3() { - let coeffs = vec![fr(1),fr(2),fr(3),fr(4),fr(5),fr(6),fr(7),fr(8)]; + let coeffs = vec![fr(1), fr(2), fr(3), fr(4), fr(5), fr(6), fr(7), fr(8)]; let f = CanonicalPoly::new(coeffs); let d = LagrangeDecomp::build(&f); @@ -272,13 +282,13 @@ mod tests { // p_3 = [3,7,11,15]: a=[3,11], b=[7,15] // p_6 = a = [3,11] // p_7 = a+b = [10,26] - let coeffs = vec![fr(1),fr(2),fr(3),fr(4),fr(5),fr(6),fr(7),fr(8)]; + let coeffs = vec![fr(1), fr(2), fr(3), fr(4), fr(5), fr(6), fr(7), fr(8)]; let f = CanonicalPoly::new(coeffs); let d = LagrangeDecomp::build(&f); - assert_eq!(d.p(4).coeffs(), &[fr(1), fr(5)]); - assert_eq!(d.p(5).coeffs(), &[fr(4), fr(12)]); - assert_eq!(d.p(6).coeffs(), &[fr(3), fr(11)]); + assert_eq!(d.p(4).coeffs(), &[fr(1), fr(5)]); + assert_eq!(d.p(5).coeffs(), &[fr(4), fr(12)]); + assert_eq!(d.p(6).coeffs(), &[fr(3), fr(11)]); assert_eq!(d.p(7).coeffs(), &[fr(10), fr(26)]); } @@ -323,9 +333,9 @@ mod tests { let d = LagrangeDecomp::build(&f); let leaves = d.leaves(); // Each leaf is a single-coefficient poly. - assert_eq!(leaves[0].coeffs()[0], fr(1)); // f(0,0) = 1 - assert_eq!(leaves[1].coeffs()[0], fr(4)); // f(0,1) = 4 - assert_eq!(leaves[2].coeffs()[0], fr(3)); // f(1,0) = 3 + assert_eq!(leaves[0].coeffs()[0], fr(1)); // f(0,0) = 1 + assert_eq!(leaves[1].coeffs()[0], fr(4)); // f(0,1) = 4 + assert_eq!(leaves[2].coeffs()[0], fr(3)); // f(1,0) = 3 assert_eq!(leaves[3].coeffs()[0], fr(10)); // f(1,1) = 10 } @@ -345,8 +355,8 @@ mod tests { fn to_lagrange_hypercube_sum_matches_canonical_n3() { // H(f) must be the same whether computed from canonical or Lagrange. let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); - let f = CanonicalPoly::new(coeffs); - let d = LagrangeDecomp::build(&f); + let f = CanonicalPoly::new(coeffs); + let d = LagrangeDecomp::build(&f); let lag = d.to_lagrange(); assert_eq!(f.hypercube_sum(), lag.hypercube_sum()); } @@ -356,28 +366,22 @@ mod tests { // Every entry of the Lagrange eval vector must equal eval_circuit // of the canonical poly at the corresponding Boolean point. let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); - let f = CanonicalPoly::new(coeffs); - let d = LagrangeDecomp::build(&f); + let f = CanonicalPoly::new(coeffs); + let d = LagrangeDecomp::build(&f); let lag = d.to_lagrange(); for b in 0..8usize { - let point: Vec = (0..3) - .map(|k| Fr::from(((b >> k) & 1) as u64)) - .collect(); - assert_eq!( - lag.evals()[b], - f.eval_circuit(&point), - "mismatch at b={b}" - ); + let point: Vec = (0..3).map(|k| Fr::from(((b >> k) & 1) as u64)).collect(); + assert_eq!(lag.evals()[b], f.eval_circuit(&point), "mismatch at b={b}"); } } #[test] fn to_lagrange_n1() { // f = 3 + 7x₁: f(0)=3, f(1)=10 - let f = CanonicalPoly::new(vec![fr(3), fr(7)]); - let d = LagrangeDecomp::build(&f); + let f = CanonicalPoly::new(vec![fr(3), fr(7)]); + let d = LagrangeDecomp::build(&f); let lag = d.to_lagrange(); assert_eq!(lag.evals(), &[fr(3), fr(10)]); } -} \ No newline at end of file +} diff --git a/src/circuit/mod.rs b/src/circuit/mod.rs index 5463c77..cedeb96 100644 --- a/src/circuit/mod.rs +++ b/src/circuit/mod.rs @@ -1,12 +1,8 @@ +pub mod bit_reverse_cache; mod canonical; mod lagrange_decomp; mod sum_circuit; -pub mod bit_reverse_cache; pub use canonical::{bit_reverse, build_bit_reverse_table, CanonicalDecomp}; pub use lagrange_decomp::LagrangeDecomp; -pub use sum_circuit::{ - CanonicalSumCircuit, - LagrangeSumCircuit, - SumCircuit, -}; \ No newline at end of file +pub use sum_circuit::{CanonicalSumCircuit, LagrangeSumCircuit, SumCircuit}; diff --git a/src/circuit/sum_circuit.rs b/src/circuit/sum_circuit.rs index 4a44b4e..aa0a025 100644 --- a/src/circuit/sum_circuit.rs +++ b/src/circuit/sum_circuit.rs @@ -115,11 +115,11 @@ where // Propagate bottom-up from layer n-1 down to layer 0. for i in (0..n).rev() { let layer_start_0based = (1usize << i) - 1; - let layer_size = 1usize << i; + let layer_size = 1usize << i; for t in 0..layer_size { let parent = layer_start_0based + t; - let left = 2 * parent + 1; - let right = 2 * parent + 2; + let left = 2 * parent + 1; + let right = 2 * parent + 2; data[parent] = recurrence(data[left], data[right]); } } @@ -171,23 +171,25 @@ impl CanonicalSumCircuit { /// # Complexity /// `O(N)` — exactly `N − 1` pairs of additions. pub fn build(f: &CanonicalPoly) -> Self { - let n = f.num_vars(); + let n = f.num_vars(); let big_n = f.num_evals(); // For n ∈ {10, 15, 20}: zero-cost borrow of the pre-computed static table. // For other n: single heap allocation, used here and dropped immediately. - let table = get_or_build(n); + let table = get_or_build(n); let coeffs = f.coeffs(); - let leaves: Vec = (0..big_n) - .map(|k| coeffs[table[k]]) - .collect(); + let leaves: Vec = (0..big_n).map(|k| coeffs[table[k]]).collect(); // h_j = h_{2j} + h_{2j} + h_{2j+1} // Replaces 2·h_{2j} + h_{2j+1}: one multiplication + one addition // with two additions — cheaper on BN254 where mul >> add. let data = build_from_leaves(&leaves, |l, r| l + l + r); - Self { num_vars: n, big_n, data } + Self { + num_vars: n, + big_n, + data, + } } /// Build directly from a leaf slice already in bit-reversed order. @@ -199,35 +201,45 @@ impl CanonicalSumCircuit { !leaves.is_empty() && leaves.len().is_power_of_two(), "CanonicalSumCircuit::from_leaves: length must be a power of two" ); - let big_n = leaves.len(); + let big_n = leaves.len(); let num_vars = big_n.trailing_zeros() as usize; - let data = build_from_leaves(leaves, |l, r| l + l + r); - Self { num_vars, big_n, data } + let data = build_from_leaves(leaves, |l, r| l + l + r); + Self { + num_vars, + big_n, + data, + } } } impl SumCircuit for CanonicalSumCircuit { #[inline] - fn num_vars(&self) -> usize { self.num_vars } + fn num_vars(&self) -> usize { + self.num_vars + } #[inline] fn h(&self, j: usize) -> F { assert!( j >= 1 && j <= 2 * self.big_n - 1, - "h index {j} out of range [1, {}]", 2 * self.big_n - 1 + "h index {j} out of range [1, {}]", + 2 * self.big_n - 1 ); self.data[j - 1] } fn leaves(&self) -> &[F] { - &self.data[self.big_n - 1 .. 2 * self.big_n - 1] + &self.data[self.big_n - 1..2 * self.big_n - 1] } fn layer(&self, i: usize) -> &[F] { - assert!(i <= self.num_vars, - "layer {i} out of range [0, {}]", self.num_vars); + assert!( + i <= self.num_vars, + "layer {i} out of range [0, {}]", + self.num_vars + ); let start = (1usize << i) - 1; - let end = (1usize << (i + 1)) - 1; + let end = (1usize << (i + 1)) - 1; &self.data[start..end] } @@ -281,19 +293,21 @@ impl LagrangeSumCircuit { /// # Complexity /// `O(N)` — exactly `N − 1` additions. pub fn build(f: &LagrangePoly) -> Self { - let n = f.num_vars(); + let n = f.num_vars(); let big_n = f.num_evals(); // For n ∈ {10, 15, 20}: zero-cost borrow of the pre-computed static table. // For other n: single heap allocation, used here and dropped immediately. let table = get_or_build(n); let evals = f.evals(); - let leaves: Vec = (0..big_n) - .map(|k| evals[table[k]]) - .collect(); + let leaves: Vec = (0..big_n).map(|k| evals[table[k]]).collect(); let data = build_from_leaves(&leaves, |l, r| l + r); - Self { num_vars: n, big_n, data } + Self { + num_vars: n, + big_n, + data, + } } /// Build directly from a leaf slice already in bit-reversed order. @@ -305,42 +319,50 @@ impl LagrangeSumCircuit { !leaves.is_empty() && leaves.len().is_power_of_two(), "LagrangeSumCircuit::from_leaves: length must be a power of two" ); - let big_n = leaves.len(); + let big_n = leaves.len(); let num_vars = big_n.trailing_zeros() as usize; - let data = build_from_leaves(leaves, |l, r| l + r); - Self { num_vars, big_n, data } + let data = build_from_leaves(leaves, |l, r| l + r); + Self { + num_vars, + big_n, + data, + } } } impl SumCircuit for LagrangeSumCircuit { #[inline] - fn num_vars(&self) -> usize { self.num_vars } + fn num_vars(&self) -> usize { + self.num_vars + } #[inline] fn h(&self, j: usize) -> F { assert!( j >= 1 && j <= 2 * self.big_n - 1, - "h index {j} out of range [1, {}]", 2 * self.big_n - 1 + "h index {j} out of range [1, {}]", + 2 * self.big_n - 1 ); self.data[j - 1] } fn leaves(&self) -> &[F] { - &self.data[self.big_n - 1 .. 2 * self.big_n - 1] + &self.data[self.big_n - 1..2 * self.big_n - 1] } fn layer(&self, i: usize) -> &[F] { - assert!(i <= self.num_vars, - "layer {i} out of range [0, {}]", self.num_vars); + assert!( + i <= self.num_vars, + "layer {i} out of range [0, {}]", + self.num_vars + ); let start = (1usize << i) - 1; - let end = (1usize << (i + 1)) - 1; + let end = (1usize << (i + 1)) - 1; &self.data[start..end] } fn verify_recurrence(&self) -> bool { - (1..self.big_n).all(|j| { - self.data[j - 1] == self.data[2 * j - 1] + self.data[2 * j] - }) + (1..self.big_n).all(|j| self.data[j - 1] == self.data[2 * j - 1] + self.data[2 * j]) } } @@ -354,13 +376,13 @@ mod tests { use crate::poly::{CanonicalPoly, LagrangePoly, MlPoly}; use ark_bn254::Fr; - fn fr(n: u64) -> Fr { Fr::from(n) } + fn fr(n: u64) -> Fr { + Fr::from(n) + } #[test] fn canonical_paper_example_n3() { - let leaves = vec![ - fr(1), fr(4), fr(3), fr(7), fr(2), fr(6), fr(5), fr(8) - ]; + let leaves = vec![fr(1), fr(4), fr(3), fr(7), fr(2), fr(6), fr(5), fr(8)]; let sc = CanonicalSumCircuit::from_leaves(&leaves); assert_eq!(sc.root(), fr(88)); assert_eq!(sc.h(2), fr(25)); @@ -373,16 +395,14 @@ mod tests { #[test] fn canonical_recurrence_holds_paper_example() { - let leaves = vec![ - fr(1), fr(4), fr(3), fr(7), fr(2), fr(6), fr(5), fr(8) - ]; + let leaves = vec![fr(1), fr(4), fr(3), fr(7), fr(2), fr(6), fr(5), fr(8)]; let sc = CanonicalSumCircuit::from_leaves(&leaves); assert!(sc.verify_recurrence()); } #[test] fn canonical_root_equals_hypercube_sum_n2() { - let f = CanonicalPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]); + let f = CanonicalPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]); let sc = CanonicalSumCircuit::build(&f); assert_eq!(sc.root(), f.hypercube_sum()); } @@ -390,7 +410,7 @@ mod tests { #[test] fn canonical_root_equals_hypercube_sum_n3() { let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); - let f = CanonicalPoly::new(coeffs); + let f = CanonicalPoly::new(coeffs); let sc = CanonicalSumCircuit::build(&f); assert_eq!(sc.root(), f.hypercube_sum()); } @@ -398,21 +418,21 @@ mod tests { #[test] fn canonical_root_equals_hypercube_sum_n4() { let coeffs: Vec = (0..16).map(|i| fr(i as u64 * 3 + 1)).collect(); - let f = CanonicalPoly::new(coeffs); + let f = CanonicalPoly::new(coeffs); let sc = CanonicalSumCircuit::build(&f); assert_eq!(sc.root(), f.hypercube_sum()); } #[test] fn canonical_node_count() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); let sc = CanonicalSumCircuit::build(&f); assert_eq!(sc.data.len(), 15); } #[test] fn canonical_layer_sizes() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); let sc = CanonicalSumCircuit::build(&f); assert_eq!(sc.layer(0).len(), 1); assert_eq!(sc.layer(1).len(), 2); @@ -422,7 +442,7 @@ mod tests { #[test] fn canonical_leaves_slice_length() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); let sc = CanonicalSumCircuit::build(&f); assert_eq!(sc.leaves().len(), 8); } @@ -430,14 +450,14 @@ mod tests { #[test] fn canonical_recurrence_holds_n4() { let coeffs: Vec = (0..16).map(|i| fr(i as u64 + 1)).collect(); - let f = CanonicalPoly::new(coeffs); + let f = CanonicalPoly::new(coeffs); let sc = CanonicalSumCircuit::build(&f); assert!(sc.verify_recurrence()); } #[test] fn canonical_n1_edge_case() { - let f = CanonicalPoly::new(vec![fr(3), fr(7)]); + let f = CanonicalPoly::new(vec![fr(3), fr(7)]); let sc = CanonicalSumCircuit::build(&f); assert_eq!(sc.root(), fr(13)); assert_eq!(sc.root(), f.hypercube_sum()); @@ -445,7 +465,7 @@ mod tests { #[test] fn lagrange_root_equals_hypercube_sum_n2() { - let f = LagrangePoly::new(vec![fr(1), fr(3), fr(4), fr(10)]); + let f = LagrangePoly::new(vec![fr(1), fr(3), fr(4), fr(10)]); let sc = LagrangeSumCircuit::build(&f); assert_eq!(sc.root(), f.hypercube_sum()); } @@ -453,7 +473,7 @@ mod tests { #[test] fn lagrange_root_equals_hypercube_sum_n3() { let evals: Vec = (1..=8).map(|i| fr(i)).collect(); - let f = LagrangePoly::new(evals); + let f = LagrangePoly::new(evals); let sc = LagrangeSumCircuit::build(&f); assert_eq!(sc.root(), f.hypercube_sum()); } @@ -461,7 +481,7 @@ mod tests { #[test] fn lagrange_root_equals_hypercube_sum_n4() { let evals: Vec = (0..16).map(|i| fr(i as u64 * 2 + 3)).collect(); - let f = LagrangePoly::new(evals); + let f = LagrangePoly::new(evals); let sc = LagrangeSumCircuit::build(&f); assert_eq!(sc.root(), f.hypercube_sum()); } @@ -469,14 +489,14 @@ mod tests { #[test] fn lagrange_recurrence_holds_n3() { let evals: Vec = (1..=8).map(|i| fr(i)).collect(); - let f = LagrangePoly::new(evals); + let f = LagrangePoly::new(evals); let sc = LagrangeSumCircuit::build(&f); assert!(sc.verify_recurrence()); } #[test] fn lagrange_layer_sizes() { - let f = LagrangePoly::new((0..8).map(|i| fr(i)).collect()); + let f = LagrangePoly::new((0..8).map(|i| fr(i)).collect()); let sc = LagrangeSumCircuit::build(&f); assert_eq!(sc.layer(0).len(), 1); assert_eq!(sc.layer(1).len(), 2); @@ -486,7 +506,7 @@ mod tests { #[test] fn lagrange_n1_edge_case() { - let f = LagrangePoly::new(vec![fr(5), fr(9)]); + let f = LagrangePoly::new(vec![fr(5), fr(9)]); let sc = LagrangeSumCircuit::build(&f); assert_eq!(sc.root(), fr(14)); } @@ -496,9 +516,9 @@ mod tests { use crate::circuit::LagrangeDecomp; let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); let canon = CanonicalPoly::new(coeffs); - let lag = LagrangeDecomp::build(&canon).to_lagrange(); + let lag = LagrangeDecomp::build(&canon).to_lagrange(); let sc_canon = CanonicalSumCircuit::build(&canon); - let sc_lag = LagrangeSumCircuit::build(&lag); + let sc_lag = LagrangeSumCircuit::build(&lag); assert_eq!(sc_canon.root(), sc_lag.root()); assert_eq!(sc_canon.root(), canon.hypercube_sum()); } @@ -508,21 +528,19 @@ mod tests { use crate::circuit::LagrangeDecomp; let coeffs: Vec = (0..16).map(|i| fr(i as u64 * 3 + 1)).collect(); let canon = CanonicalPoly::new(coeffs); - let lag = LagrangeDecomp::build(&canon).to_lagrange(); + let lag = LagrangeDecomp::build(&canon).to_lagrange(); let sc_canon = CanonicalSumCircuit::build(&canon); - let sc_lag = LagrangeSumCircuit::build(&lag); + let sc_lag = LagrangeSumCircuit::build(&lag); assert_eq!(sc_canon.root(), sc_lag.root()); } #[test] fn canonical_from_leaves_matches_build() { let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); - let f = CanonicalPoly::new(coeffs); + let f = CanonicalPoly::new(coeffs); let table = get_or_build(3); - let manual_leaves: Vec = (0..8) - .map(|k| f.coeffs()[table[k]]) - .collect(); - let sc_build = CanonicalSumCircuit::build(&f); + let manual_leaves: Vec = (0..8).map(|k| f.coeffs()[table[k]]).collect(); + let sc_build = CanonicalSumCircuit::build(&f); let sc_from_leaves = CanonicalSumCircuit::from_leaves(&manual_leaves); assert_eq!(sc_build.data, sc_from_leaves.data); } @@ -530,13 +548,11 @@ mod tests { #[test] fn lagrange_from_leaves_matches_build() { let evals: Vec = (1..=8).map(|i| fr(i)).collect(); - let f = LagrangePoly::new(evals); + let f = LagrangePoly::new(evals); let table = get_or_build(3); - let manual_leaves: Vec = (0..8) - .map(|k| f.evals()[table[k]]) - .collect(); - let sc_build = LagrangeSumCircuit::build(&f); + let manual_leaves: Vec = (0..8).map(|k| f.evals()[table[k]]).collect(); + let sc_build = LagrangeSumCircuit::build(&f); let sc_from_leaves = LagrangeSumCircuit::from_leaves(&manual_leaves); assert_eq!(sc_build.data, sc_from_leaves.data); } -} \ No newline at end of file +} diff --git a/src/lib.rs b/src/lib.rs index ab92d5d..9c8454b 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,7 +1,16 @@ -//! mlp_pro — Multilinear polynomials via tree-based circuit and the Sumcheck protocol. +//! Basis-aware multilinear polynomial algorithms and the Sumcheck protocol. //! -//! Built step by step. Each module is activated when its step is complete. +//! The crate provides canonical (coefficient) and Lagrange (Boolean-hypercube +//! evaluation) representations, tree/circuit decompositions, linear-time +//! evaluation kernels, basis-specific Sumcheck provers, and a stateless +//! verifier. +//! +//! This is a research implementation. It is not a complete SNARK/STARK and +//! does not currently include Fiat-Shamir, a polynomial commitment scheme, or +//! zero-knowledge masking. + +#![forbid(unsafe_code)] -pub mod poly; pub mod circuit; -pub mod sumcheck; \ No newline at end of file +pub mod poly; +pub mod sumcheck; diff --git a/src/poly/canonical.rs b/src/poly/canonical.rs index 76101a1..e22afbd 100644 --- a/src/poly/canonical.rs +++ b/src/poly/canonical.rs @@ -102,8 +102,11 @@ impl CanonicalPoly { /// Panics if `r.len() != num_vars`. pub fn eval_naive(&self, r: &[F]) -> F { assert_eq!( - r.len(), self.num_vars, - "eval_naive: expected {} variables, got {}", self.num_vars, r.len() + r.len(), + self.num_vars, + "eval_naive: expected {} variables, got {}", + self.num_vars, + r.len() ); self.coeffs @@ -148,28 +151,31 @@ impl CanonicalPoly { /// # Panics /// Panics if `r.len() != num_vars`. pub fn eval_circuit(&self, r: &[F]) -> F { - assert_eq!( - r.len(), self.num_vars, - "eval_circuit: expected {} variables, got {}", self.num_vars, r.len() - ); - - // Working buffer: starts as a copy of the coefficient vector. - // We fold variable x₁ first (r[0]), then x₂ (r[1]), …, xₙ (r[n-1]). - // This matches the split rule: q₂ holds even-indexed coefficients - // (j₁ = 0) and q₃ holds odd-indexed coefficients (j₁ = 1). - let mut buf = self.coeffs.clone(); - - for k in 0..self.num_vars { - let r_k = r[k]; - let half = buf.len() / 2; - for t in 0..half { - buf[t] = buf[2 * t] + r_k * buf[2 * t + 1]; + assert_eq!( + r.len(), + self.num_vars, + "eval_circuit: expected {} variables, got {}", + self.num_vars, + r.len() + ); + + // Working buffer: starts as a copy of the coefficient vector. + // We fold variable x₁ first (r[0]), then x₂ (r[1]), …, xₙ (r[n-1]). + // This matches the split rule: q₂ holds even-indexed coefficients + // (j₁ = 0) and q₃ holds odd-indexed coefficients (j₁ = 1). + let mut buf = self.coeffs.clone(); + + for k in 0..self.num_vars { + let r_k = r[k]; + let half = buf.len() / 2; + for t in 0..half { + buf[t] = buf[2 * t] + r_k * buf[2 * t + 1]; + } + buf.truncate(half); } - buf.truncate(half); - } - buf[0] -} + buf[0] + } } // ── MlPoly implementation ───────────────────────────────────────────────────── @@ -188,7 +194,7 @@ impl MlPoly for CanonicalPoly { if alpha.is_zero() { acc } else { - let deg = j.count_ones() as usize; + let deg = j.count_ones() as usize; let weight = F::from(1u64 << (n - deg)); acc + alpha * weight } @@ -207,7 +213,9 @@ mod tests { use ark_ff::One; use ark_ff::Zero; - fn fr(n: u64) -> Fr { Fr::from(n) } + fn fr(n: u64) -> Fr { + Fr::from(n) + } // ── CanonicalTerm ───────────────────────────────────────────────────────── @@ -397,7 +405,8 @@ mod tests { assert_eq!( f.eval_naive(&r), f.eval_circuit(&r), - "disagreement at r={:?}", r + "disagreement at r={:?}", + r ); } } @@ -433,4 +442,4 @@ mod tests { let r = [fr(1), fr(2), fr(3), fr(4)]; assert_eq!(f.eval_circuit(&r), Fr::zero()); } -} \ No newline at end of file +} diff --git a/src/poly/lagrange.rs b/src/poly/lagrange.rs index a7145d6..d12e5d5 100644 --- a/src/poly/lagrange.rs +++ b/src/poly/lagrange.rs @@ -73,16 +73,19 @@ impl LagrangePoly { /// Panics if `r.len() != num_vars`. pub fn eval_standard(&self, r: &[F]) -> F { assert_eq!( - r.len(), self.num_vars, - "eval_standard: expected {} variables, got {}", self.num_vars, r.len() + r.len(), + self.num_vars, + "eval_standard: expected {} variables, got {}", + self.num_vars, + r.len() ); let mut buf = self.evals.clone(); for k in 0..self.num_vars { - let r_k = r[k]; + let r_k = r[k]; let one_minus_r = F::one() - r_k; - let half = buf.len() / 2; + let half = buf.len() / 2; for t in 0..half { buf[t] = one_minus_r * buf[2 * t] + r_k * buf[2 * t + 1]; } @@ -104,14 +107,17 @@ impl LagrangePoly { /// Panics if `r.len() != num_vars`. pub fn eval_optimized(&self, r: &[F]) -> F { assert_eq!( - r.len(), self.num_vars, - "eval_optimized: expected {} variables, got {}", self.num_vars, r.len() + r.len(), + self.num_vars, + "eval_optimized: expected {} variables, got {}", + self.num_vars, + r.len() ); let mut buf = self.evals.clone(); for k in 0..self.num_vars { - let r_k = r[k]; + let r_k = r[k]; let half = buf.len() / 2; for t in 0..half { buf[t] = buf[2 * t] + r_k * (buf[2 * t + 1] - buf[2 * t]); @@ -142,30 +148,33 @@ impl LagrangePoly { /// /// # Panics /// Panics if `r.len() != num_vars`. -pub fn eval_parallel(&self, r: &[F]) -> F -where - F: Send + Sync, -{ - assert_eq!( - r.len(), self.num_vars, - "eval_parallel: expected {} variables, got {}", self.num_vars, r.len() - ); - - let mut buf = self.evals.clone(); - // Pre-allocate a reusable buffer — avoids one allocation per layer. - let mut tmp = Vec::with_capacity(buf.len()); - - for k in 0..self.num_vars { - let r_k = r[k]; - tmp.clear(); - buf.par_chunks(2) - .map(|pair| pair[0] + r_k * (pair[1] - pair[0])) - .collect_into_vec(&mut tmp); - std::mem::swap(&mut buf, &mut tmp); - } - - buf[0] -} + pub fn eval_parallel(&self, r: &[F]) -> F + where + F: Send + Sync, + { + assert_eq!( + r.len(), + self.num_vars, + "eval_parallel: expected {} variables, got {}", + self.num_vars, + r.len() + ); + + let mut buf = self.evals.clone(); + // Pre-allocate a reusable buffer — avoids one allocation per layer. + let mut tmp = Vec::with_capacity(buf.len()); + + for k in 0..self.num_vars { + let r_k = r[k]; + tmp.clear(); + buf.par_chunks(2) + .map(|pair| pair[0] + r_k * (pair[1] - pair[0])) + .collect_into_vec(&mut tmp); + std::mem::swap(&mut buf, &mut tmp); + } + + buf[0] + } } // ── MlPoly implementation ───────────────────────────────────────────────────── @@ -191,7 +200,9 @@ mod tests { use ark_bn254::Fr; use ark_ff::Zero; - fn fr(n: u64) -> Fr { Fr::from(n) } + fn fr(n: u64) -> Fr { + Fr::from(n) + } // ── Construction ───────────────────────────────────────────────────────── @@ -346,7 +357,7 @@ mod tests { let s = p.eval_standard(&r); let o = p.eval_optimized(&r); let par = p.eval_parallel(&r); - assert_eq!(s, o, "standard vs optimized at {:?}", r); + assert_eq!(s, o, "standard vs optimized at {:?}", r); assert_eq!(o, par, "optimized vs parallel at {:?}", r); } } @@ -357,8 +368,8 @@ mod tests { fn all_three_agree_at_random_point_n4() { let p = LagrangePoly::new((0..16).map(|i| fr(i as u64 * 3 + 1)).collect()); let r = [fr(2), fr(5), fr(11), fr(7)]; - let s = p.eval_standard(&r); - let o = p.eval_optimized(&r); + let s = p.eval_standard(&r); + let o = p.eval_optimized(&r); let par = p.eval_parallel(&r); assert_eq!(s, o); assert_eq!(o, par); @@ -370,10 +381,10 @@ mod tests { use crate::poly::CanonicalPoly; let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); let canon = CanonicalPoly::new(coeffs); - let lag = LagrangeDecomp::build(&canon).to_lagrange(); - let r = [fr(3), fr(7), fr(13)]; + let lag = LagrangeDecomp::build(&canon).to_lagrange(); + let r = [fr(3), fr(7), fr(13)]; assert_eq!(lag.eval_optimized(&r), canon.eval_circuit(&r)); - assert_eq!(lag.eval_parallel(&r), canon.eval_circuit(&r)); + assert_eq!(lag.eval_parallel(&r), canon.eval_circuit(&r)); } #[test] @@ -382,9 +393,9 @@ mod tests { use crate::poly::CanonicalPoly; let coeffs: Vec = (0..16).map(|i| fr(i as u64 * 3 + 1)).collect(); let canon = CanonicalPoly::new(coeffs); - let lag = LagrangeDecomp::build(&canon).to_lagrange(); - let r = [fr(2), fr(5), fr(11), fr(7)]; + let lag = LagrangeDecomp::build(&canon).to_lagrange(); + let r = [fr(2), fr(5), fr(11), fr(7)]; assert_eq!(lag.eval_optimized(&r), canon.eval_circuit(&r)); - assert_eq!(lag.eval_parallel(&r), canon.eval_circuit(&r)); + assert_eq!(lag.eval_parallel(&r), canon.eval_circuit(&r)); } -} \ No newline at end of file +} diff --git a/src/poly/mod.rs b/src/poly/mod.rs index cb9a964..b407642 100644 --- a/src/poly/mod.rs +++ b/src/poly/mod.rs @@ -19,7 +19,3 @@ pub use canonical::{CanonicalPoly, CanonicalTerm}; pub use lagrange::LagrangePoly; pub use traits::MlPoly; pub use uni::{UniDecomp, UniPoly, UniTerm}; - - - - diff --git a/src/poly/traits.rs b/src/poly/traits.rs index a721747..f791eae 100644 --- a/src/poly/traits.rs +++ b/src/poly/traits.rs @@ -16,4 +16,4 @@ pub trait MlPoly { /// The Boolean hypercube sum `H(f) = Σ_{x ∈ {0,1}^n} f(x)`. fn hypercube_sum(&self) -> F; -} \ No newline at end of file +} diff --git a/src/poly/uni.rs b/src/poly/uni.rs index 9040e0d..be253d2 100644 --- a/src/poly/uni.rs +++ b/src/poly/uni.rs @@ -96,7 +96,10 @@ impl UniPoly { /// # Panics /// Panics if `coeffs` is empty. pub fn new(coeffs: Vec) -> Self { - assert!(!coeffs.is_empty(), "UniPoly: coefficient vector must not be empty"); + assert!( + !coeffs.is_empty(), + "UniPoly: coefficient vector must not be empty" + ); let actual_degree = coeffs .iter() @@ -110,13 +113,16 @@ impl UniPoly { let mut padded = coeffs; padded.resize(padded_len, F::zero()); - Self { coeffs: padded, actual_degree } + Self { + coeffs: padded, + actual_degree, + } } /// The zero polynomial, internally stored with `2^n` coefficients. pub fn zero(n: usize) -> Self { Self { - coeffs: vec![F::zero(); 1 << n], + coeffs: vec![F::zero(); 1 << n], actual_degree: 0, } } @@ -154,7 +160,11 @@ impl UniPoly { /// The coefficient of `Xᵏ`. /// Returns `F::zero()` for `k >= padded_len()`. pub fn coeff(&self, k: usize) -> F { - if k < self.coeffs.len() { self.coeffs[k] } else { F::zero() } + if k < self.coeffs.len() { + self.coeffs[k] + } else { + F::zero() + } } // ── Term views ──────────────────────────────────────────────────────────── @@ -186,8 +196,11 @@ impl UniPoly { .iter() .enumerate() .fold(F::zero(), |acc, (k, &alpha)| { - if alpha.is_zero() { acc } - else { acc + UniTerm::new(alpha, k).eval(r) } + if alpha.is_zero() { + acc + } else { + acc + UniTerm::new(alpha, k).eval(r) + } }) } @@ -232,7 +245,7 @@ impl UniPoly { // At layer k (0-based from bottom), combine with r^{2^k}. for k in 0..n { - let lambda = powers[k]; + let lambda = powers[k]; let new_len = buf.len() / 2; for t in 0..new_len { // buf[t] = buf[2t] + r^{2^k} · buf[2t+1] @@ -267,7 +280,7 @@ impl UniPoly { /// Sumcheck rounds), the bit-reversed layout can be cached. /// Each subsequent evaluation then skips the permutation step entirely. pub fn eval_circuit(&self, r: F) -> F { - let n = self.log_len(); + let n = self.log_len(); let big_n = self.padded_len(); // Power table: powers[k] = r^{2^k} for k = 0, …, n-1. @@ -277,15 +290,13 @@ impl UniPoly { // For n ∈ {10, 15, 20}: zero-cost borrow of the static cached table. // For other n: single heap allocation, used here and dropped. let table = get_or_build(n); - let mut buf: Vec = (0..big_n) - .map(|k| self.coeffs[table[k]]) - .collect(); + let mut buf: Vec = (0..big_n).map(|k| self.coeffs[table[k]]).collect(); // Bottom-up fold. // Layer i from the bottom (i = 0 is the leaf layer): // weight = powers[n-1-i] = r^{2^{n-1-i}} for i in 0..n { - let lambda = powers[n - 1 - i]; + let lambda = powers[n - 1 - i]; let new_len = buf.len() / 2; for t in 0..new_len { buf[t] = buf[2 * t] + lambda * buf[2 * t + 1]; @@ -322,16 +333,14 @@ impl UniPoly { where F: Send + Sync, { - let n = self.log_len(); + let n = self.log_len(); let big_n = self.padded_len(); let powers = power_table(r, n); // For n ∈ {10, 15, 20}: zero-cost borrow of the static cached table. let table = get_or_build(n); - let mut buf: Vec = (0..big_n) - .map(|k| self.coeffs[table[k]]) - .collect(); + let mut buf: Vec = (0..big_n).map(|k| self.coeffs[table[k]]).collect(); let mut tmp = Vec::with_capacity(big_n / 2); @@ -353,7 +362,7 @@ impl UniPoly { /// /// See [`UniDecomp`] for the complete documentation. pub fn decompose(&self) -> UniDecomp { - let n = self.log_len(); + let n = self.log_len(); let big_n = self.padded_len(); let mut nodes: Vec> = vec![Vec::new(); 2 * big_n - 1]; @@ -361,14 +370,14 @@ impl UniPoly { for i in 0..n { let layer_start = 1usize << i; - let layer_end = 1usize << (i + 1); + let layer_end = 1usize << (i + 1); for j in layer_start..layer_end { let parent = nodes[j - 1].clone(); - let half = parent.len() / 2; - let left: Vec = (0..half).map(|r| parent[2 * r]).collect(); + let half = parent.len() / 2; + let left: Vec = (0..half).map(|r| parent[2 * r]).collect(); let right: Vec = (0..half).map(|r| parent[2 * r + 1]).collect(); nodes[2 * j - 1] = left; - nodes[2 * j] = right; + nodes[2 * j] = right; } } @@ -404,25 +413,28 @@ impl UniDecomp { pub fn q(&self, j: usize) -> &[F] { assert!( j >= 1 && j <= 2 * self.big_n - 1, - "q index {j} out of range [1, {}]", 2 * self.big_n - 1 + "q index {j} out of range [1, {}]", + 2 * self.big_n - 1 ); &self.nodes[j - 1] } /// The root `q_1` (full coefficient vector, size `N`). #[inline] - pub fn root(&self) -> &[F] { &self.nodes[0] } + pub fn root(&self) -> &[F] { + &self.nodes[0] + } /// Leaf layer: `q_N, …, q_{2N-1}`, each of size 1. pub fn leaves(&self) -> &[Vec] { - &self.nodes[self.big_n - 1 .. 2 * self.big_n - 1] + &self.nodes[self.big_n - 1..2 * self.big_n - 1] } /// Layer `i` (0-based): slice of node vectors at depth `i`. pub fn layer(&self, i: usize) -> &[Vec] { assert!(i <= self.n, "layer {i} out of range [0, {}]", self.n); let start = (1usize << i) - 1; - let end = (1usize << (i + 1)) - 1; + let end = (1usize << (i + 1)) - 1; &self.nodes[start..end] } @@ -440,9 +452,11 @@ impl UniDecomp { fn field_pow(mut base: F, mut exp: usize) -> F { let mut result = F::one(); while exp > 0 { - if exp & 1 == 1 { result *= base; } - base *= base; - exp >>= 1; + if exp & 1 == 1 { + result *= base; + } + base *= base; + exp >>= 1; } result } @@ -454,12 +468,14 @@ fn field_pow(mut base: F, mut exp: usize) -> F { /// /// Used by `eval_estrin`, `eval_circuit`, and `eval_circuit_parallel`. fn power_table(r: F, n: usize) -> Vec { - if n == 0 { return vec![]; } + if n == 0 { + return vec![]; + } let mut powers = Vec::with_capacity(n); - powers.push(r); // r^{2^0} = r + powers.push(r); // r^{2^0} = r for k in 1..n { let prev = powers[k - 1]; - powers.push(prev * prev); // r^{2^k} = (r^{2^{k-1}})² + powers.push(prev * prev); // r^{2^k} = (r^{2^{k-1}})² } powers } @@ -474,7 +490,9 @@ mod tests { use ark_bn254::Fr; use ark_ff::Zero; - fn fr(n: u64) -> Fr { Fr::from(n) } + fn fr(n: u64) -> Fr { + Fr::from(n) + } // ── UniTerm ─────────────────────────────────────────────────────────────── @@ -592,19 +610,28 @@ mod tests { #[test] fn all_terms_includes_padding() { - assert_eq!(UniPoly::new(vec![fr(1), fr(2), fr(3)]).all_terms().count(), 4); + assert_eq!( + UniPoly::new(vec![fr(1), fr(2), fr(3)]).all_terms().count(), + 4 + ); } // ── eval_naive ──────────────────────────────────────────────────────────── #[test] fn eval_naive_at_zero_returns_constant() { - assert_eq!(UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_naive(fr(0)), fr(1)); + assert_eq!( + UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_naive(fr(0)), + fr(1) + ); } #[test] fn eval_naive_at_one_returns_sum() { - assert_eq!(UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_naive(fr(1)), fr(10)); + assert_eq!( + UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_naive(fr(1)), + fr(10) + ); } #[test] @@ -616,12 +643,18 @@ mod tests { #[test] fn eval_horner_at_zero() { - assert_eq!(UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_horner(fr(0)), fr(1)); + assert_eq!( + UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_horner(fr(0)), + fr(1) + ); } #[test] fn eval_horner_at_one() { - assert_eq!(UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_horner(fr(1)), fr(10)); + assert_eq!( + UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_horner(fr(1)), + fr(10) + ); } #[test] @@ -631,15 +664,15 @@ mod tests { // ── power_table ─────────────────────────────────────────────────────────── -#[test] -fn power_table_n3() { - // r=2: powers[0]=r^1=2, powers[1]=r^2=4, powers[2]=r^4=16 - let pt = power_table(fr(2), 3); - assert_eq!(pt.len(), 3); - assert_eq!(pt[0], fr(2)); // r^1 = 2 - assert_eq!(pt[1], fr(4)); // r^2 = 4 - assert_eq!(pt[2], fr(16)); // r^4 = 16 -} + #[test] + fn power_table_n3() { + // r=2: powers[0]=r^1=2, powers[1]=r^2=4, powers[2]=r^4=16 + let pt = power_table(fr(2), 3); + assert_eq!(pt.len(), 3); + assert_eq!(pt[0], fr(2)); // r^1 = 2 + assert_eq!(pt[1], fr(4)); // r^2 = 4 + assert_eq!(pt[2], fr(16)); // r^4 = 16 + } #[test] fn power_table_n0_is_empty() { @@ -650,7 +683,7 @@ fn power_table_n3() { fn power_table_each_entry_is_square_of_previous() { let pt = power_table(fr(3), 5); for k in 1..5 { - assert_eq!(pt[k], pt[k-1] * pt[k-1]); + assert_eq!(pt[k], pt[k - 1] * pt[k - 1]); } } @@ -658,12 +691,18 @@ fn power_table_n3() { #[test] fn eval_estrin_at_zero() { - assert_eq!(UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_estrin(fr(0)), fr(1)); + assert_eq!( + UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_estrin(fr(0)), + fr(1) + ); } #[test] fn eval_estrin_at_one() { - assert_eq!(UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_estrin(fr(1)), fr(10)); + assert_eq!( + UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_estrin(fr(1)), + fr(10) + ); } #[test] @@ -675,12 +714,18 @@ fn power_table_n3() { #[test] fn eval_circuit_at_zero() { - assert_eq!(UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_circuit(fr(0)), fr(1)); + assert_eq!( + UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_circuit(fr(0)), + fr(1) + ); } #[test] fn eval_circuit_at_one() { - assert_eq!(UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_circuit(fr(1)), fr(10)); + assert_eq!( + UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4)]).eval_circuit(fr(1)), + fr(10) + ); } #[test] @@ -713,14 +758,14 @@ fn power_table_n3() { let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); let p = UniPoly::new(coeffs); for r in [fr(0), fr(1), fr(2), fr(7), fr(100)] { - let naive = p.eval_naive(r); - let horner = p.eval_horner(r); - let estrin = p.eval_estrin(r); - let circuit = p.eval_circuit(r); + let naive = p.eval_naive(r); + let horner = p.eval_horner(r); + let estrin = p.eval_estrin(r); + let circuit = p.eval_circuit(r); let parallel = p.eval_circuit_parallel(r); - assert_eq!(naive, horner, "naive vs horner at r={r}"); - assert_eq!(naive, estrin, "naive vs estrin at r={r}"); - assert_eq!(naive, circuit, "naive vs circuit at r={r}"); + assert_eq!(naive, horner, "naive vs horner at r={r}"); + assert_eq!(naive, estrin, "naive vs estrin at r={r}"); + assert_eq!(naive, circuit, "naive vs circuit at r={r}"); assert_eq!(naive, parallel, "naive vs parallel at r={r}"); } } @@ -730,13 +775,13 @@ fn power_table_n3() { let coeffs: Vec = (0..16).map(|i| fr(i as u64 * 3 + 1)).collect(); let p = UniPoly::new(coeffs); for r in [fr(1), fr(3), fr(11), fr(255)] { - let naive = p.eval_naive(r); + let naive = p.eval_naive(r); let circuit = p.eval_circuit(r); - let estrin = p.eval_estrin(r); - let par = p.eval_circuit_parallel(r); + let estrin = p.eval_estrin(r); + let par = p.eval_circuit_parallel(r); assert_eq!(naive, circuit, "circuit at r={r}"); - assert_eq!(naive, estrin, "estrin at r={r}"); - assert_eq!(naive, par, "parallel at r={r}"); + assert_eq!(naive, estrin, "estrin at r={r}"); + assert_eq!(naive, par, "parallel at r={r}"); } } @@ -746,10 +791,10 @@ fn power_table_n3() { let coeffs: Vec = (1..=5).map(|i| fr(i)).collect(); let p = UniPoly::new(coeffs); for r in [fr(0), fr(1), fr(3), fr(11)] { - let naive = p.eval_naive(r); - let horner = p.eval_horner(r); + let naive = p.eval_naive(r); + let horner = p.eval_horner(r); let circuit = p.eval_circuit(r); - assert_eq!(naive, horner, "at r={r}"); + assert_eq!(naive, horner, "at r={r}"); assert_eq!(naive, circuit, "at r={r}"); } } @@ -757,7 +802,7 @@ fn power_table_n3() { #[test] fn padding_does_not_change_eval_circuit() { // f = 1 + 2X + 3X² (length 3 → padded to 4) - let p_short = UniPoly::new(vec![fr(1), fr(2), fr(3)]); + let p_short = UniPoly::new(vec![fr(1), fr(2), fr(3)]); let p_padded = UniPoly::new(vec![fr(1), fr(2), fr(3), fr(0)]); for r in [fr(0), fr(1), fr(2), fr(7)] { assert_eq!(p_short.eval_circuit(r), p_padded.eval_circuit(r)); @@ -796,28 +841,28 @@ fn power_table_n3() { #[test] fn decompose_first_layer_paper_example_n3() { - let p = UniPoly::new(vec![fr(1),fr(2),fr(3),fr(4),fr(5),fr(6),fr(7),fr(8)]); + let p = UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4), fr(5), fr(6), fr(7), fr(8)]); let d = p.decompose(); - assert_eq!(d.q(2), &[fr(1),fr(3),fr(5),fr(7)]); - assert_eq!(d.q(3), &[fr(2),fr(4),fr(6),fr(8)]); + assert_eq!(d.q(2), &[fr(1), fr(3), fr(5), fr(7)]); + assert_eq!(d.q(3), &[fr(2), fr(4), fr(6), fr(8)]); } #[test] fn decompose_second_layer_paper_example_n3() { - let p = UniPoly::new(vec![fr(1),fr(2),fr(3),fr(4),fr(5),fr(6),fr(7),fr(8)]); + let p = UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4), fr(5), fr(6), fr(7), fr(8)]); let d = p.decompose(); - assert_eq!(d.q(4), &[fr(1),fr(5)]); - assert_eq!(d.q(5), &[fr(3),fr(7)]); - assert_eq!(d.q(6), &[fr(2),fr(6)]); - assert_eq!(d.q(7), &[fr(4),fr(8)]); + assert_eq!(d.q(4), &[fr(1), fr(5)]); + assert_eq!(d.q(5), &[fr(3), fr(7)]); + assert_eq!(d.q(6), &[fr(2), fr(6)]); + assert_eq!(d.q(7), &[fr(4), fr(8)]); } #[test] fn decompose_leaves_paper_example_n3() { - let p = UniPoly::new(vec![fr(1),fr(2),fr(3),fr(4),fr(5),fr(6),fr(7),fr(8)]); + let p = UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4), fr(5), fr(6), fr(7), fr(8)]); let d = p.decompose(); - assert_eq!(d.q(8), &[fr(1)]); - assert_eq!(d.q(9), &[fr(5)]); + assert_eq!(d.q(8), &[fr(1)]); + assert_eq!(d.q(9), &[fr(5)]); assert_eq!(d.q(10), &[fr(3)]); assert_eq!(d.q(11), &[fr(7)]); assert_eq!(d.q(12), &[fr(2)]); @@ -831,12 +876,12 @@ fn power_table_n3() { let p = UniPoly::new((1..=8).map(|i| fr(i)).collect()); let d = p.decompose(); for j in 1..d.big_n { - let qj = d.q(j); - let left = d.q(2 * j); + let qj = d.q(j); + let left = d.q(2 * j); let right = d.q(2 * j + 1); for r in 0..qj.len() / 2 { - assert_eq!(qj[2*r], left[r], "j={j} left r={r}"); - assert_eq!(qj[2*r+1], right[r], "j={j} right r={r}"); + assert_eq!(qj[2 * r], left[r], "j={j} left r={r}"); + assert_eq!(qj[2 * r + 1], right[r], "j={j} right r={r}"); } } } @@ -848,4 +893,4 @@ fn power_table_n3() { assert_eq!(d.q(2), &[fr(1), fr(3)]); assert_eq!(d.q(3), &[fr(2), fr(0)]); } -} \ No newline at end of file +} diff --git a/src/sumcheck/mod.rs b/src/sumcheck/mod.rs index f66f4ac..e35fc1c 100644 --- a/src/sumcheck/mod.rs +++ b/src/sumcheck/mod.rs @@ -1,6 +1,5 @@ //! Sumcheck protocol — proof types, prover, and verifier. - pub mod proof; pub mod prover; pub mod verifier; @@ -8,9 +7,3 @@ pub mod verifier; pub use proof::{RoundPoly, SumcheckProof}; pub use prover::{CanonicalProver, LagrangeProver}; pub use verifier::{Verifier, VerifierError}; - - - - - - diff --git a/src/sumcheck/proof.rs b/src/sumcheck/proof.rs index 7ae5647..8e40551 100644 --- a/src/sumcheck/proof.rs +++ b/src/sumcheck/proof.rs @@ -161,7 +161,10 @@ pub struct SumcheckProof { impl SumcheckProof { /// Construct a proof from a claimed sum and a vector of round polynomials. pub fn new(claimed_sum: F, round_polys: Vec>) -> Self { - Self { claimed_sum, round_polys } + Self { + claimed_sum, + round_polys, + } } /// Number of variables `n` — equals the number of rounds. @@ -181,7 +184,8 @@ impl SumcheckProof { pub fn round_poly(&self, j: usize) -> &RoundPoly { assert!( j >= 1 && j <= self.round_polys.len(), - "round index {j} out of range [1, {}]", self.round_polys.len() + "round index {j} out of range [1, {}]", + self.round_polys.len() ); &self.round_polys[j - 1] } @@ -196,7 +200,9 @@ mod tests { use super::*; use ark_bn254::Fr; - fn fr(n: u64) -> Fr { Fr::from(n) } + fn fr(n: u64) -> Fr { + Fr::from(n) + } // ── RoundPoly construction ──────────────────────────────────────────────── @@ -221,7 +227,7 @@ mod tests { let s1 = fr(13); let s = RoundPoly::from_evaluations(s0, s1); assert_eq!(s.eval_at_zero(), s0); - assert_eq!(s.eval_at_one(), s1); + assert_eq!(s.eval_at_one(), s1); } // ── RoundPoly evaluation ────────────────────────────────────────────────── @@ -287,10 +293,7 @@ mod tests { #[test] fn proof_size_is_2n_plus_1() { - let polys = vec![ - RoundPoly::new(fr(1), fr(2)), - RoundPoly::new(fr(3), fr(4)), - ]; + let polys = vec![RoundPoly::new(fr(1), fr(2)), RoundPoly::new(fr(3), fr(4))]; let proof = SumcheckProof::new(fr(10), polys); // n=2: size = 2*2+1 = 5 assert_eq!(proof.size_in_field_elements(), 5); @@ -344,10 +347,7 @@ mod tests { let r1 = fr(3); // Round 1 check: s_1(0) + s_1(1) = claimed_sum - assert_eq!( - proof.round_poly(1).sum_over_boolean(), - proof.claimed_sum - ); + assert_eq!(proof.round_poly(1).sum_over_boolean(), proof.claimed_sum); // Round 2 check: s_2(0) + s_2(1) = s_1(r_1) assert_eq!( @@ -355,4 +355,4 @@ mod tests { proof.round_poly(1).eval(r1) ); } -} \ No newline at end of file +} diff --git a/src/sumcheck/prover.rs b/src/sumcheck/prover.rs index b05e68b..3f059ac 100644 --- a/src/sumcheck/prover.rs +++ b/src/sumcheck/prover.rs @@ -56,7 +56,9 @@ pub struct CanonicalProver { impl CanonicalProver { /// Build the prover from a canonical polynomial. Cost: `O(N)`. pub fn new(f: &CanonicalPoly) -> Self { - Self { circuit: CanonicalSumCircuit::build(f) } + Self { + circuit: CanonicalSumCircuit::build(f), + } } /// The claimed sum `H(f)` — the first value sent to the verifier. @@ -83,7 +85,7 @@ impl CanonicalProver { // Read layer j: 1-based indices [2^j, 2^{j+1} − 1]. let layer_start = 1usize << j; - let layer_size = 1usize << j; + let layer_size = 1usize << j; let mut h: Vec = (0..layer_size) .map(|t| self.circuit.h(layer_start + t)) .collect(); @@ -93,7 +95,7 @@ impl CanonicalProver { // h'[k] = h[2k] + r · h[2k+2] for k even // h'[k] = h[2k-1] + r · h[2k+1] for k odd for ki in (0..challenges.len()).rev() { - let r = challenges[ki]; + let r = challenges[ki]; let half = h.len() / 2; let mut h_new = Vec::with_capacity(half); for k in 0..half { @@ -118,8 +120,10 @@ impl CanonicalProver { pub fn prove(&self, challenges: &[F]) -> SumcheckProof { let n = self.circuit.num_vars(); assert_eq!( - challenges.len(), n, - "expected {n} challenges, got {}", challenges.len() + challenges.len(), + n, + "expected {n} challenges, got {}", + challenges.len() ); let claimed_sum = self.claimed_sum(); let round_polys = (1..=n) @@ -141,7 +145,9 @@ pub struct LagrangeProver { impl LagrangeProver { /// Build the prover from a Lagrange polynomial. Cost: `O(N)`. pub fn new(f: &LagrangePoly) -> Self { - Self { circuit: LagrangeSumCircuit::build(f) } + Self { + circuit: LagrangeSumCircuit::build(f), + } } /// The claimed sum `H(f)`. @@ -172,7 +178,7 @@ impl LagrangeProver { debug_assert!(j <= self.circuit.num_vars()); let layer_start = 1usize << j; - let layer_size = 1usize << j; + let layer_size = 1usize << j; let mut h: Vec = (0..layer_size) .map(|t| self.circuit.h(layer_start + t)) .collect(); @@ -182,7 +188,7 @@ impl LagrangeProver { // h'[k] = h[2k] + r · (h[2k+2] − h[2k]) for k even // h'[k] = h[2k-1] + r · (h[2k+1] − h[2k-1]) for k odd for ki in (0..challenges.len()).rev() { - let r = challenges[ki]; + let r = challenges[ki]; let half = h.len() / 2; let mut h_new = Vec::with_capacity(half); for k in 0..half { @@ -212,8 +218,10 @@ impl LagrangeProver { pub fn prove(&self, challenges: &[F]) -> SumcheckProof { let n = self.circuit.num_vars(); assert_eq!( - challenges.len(), n, - "expected {n} challenges, got {}", challenges.len() + challenges.len(), + n, + "expected {n} challenges, got {}", + challenges.len() ); let claimed_sum = self.claimed_sum(); let round_polys = (1..=n) @@ -235,7 +243,9 @@ mod tests { use ark_bn254::Fr; use ark_ff::Zero; - fn fr(n: u64) -> Fr { Fr::from(n) } + fn fr(n: u64) -> Fr { + Fr::from(n) + } fn canon(coeffs: &[u64]) -> CanonicalPoly { CanonicalPoly::new(coeffs.iter().map(|&v| fr(v)).collect()) @@ -270,7 +280,7 @@ mod tests { #[test] fn canonical_round_1_is_h2_h3() { let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); - let p = CanonicalProver::new(&f); + let p = CanonicalProver::new(&f); let sc = CanonicalSumCircuit::build(&f); let s1 = p.compute_round_poly(&[]); assert_eq!(s1.a, sc.h(2)); @@ -281,7 +291,10 @@ mod tests { fn canonical_round_1_sums_to_claimed() { let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let p = CanonicalProver::new(&f); - assert_eq!(p.compute_round_poly(&[]).sum_over_boolean(), p.claimed_sum()); + assert_eq!( + p.compute_round_poly(&[]).sum_over_boolean(), + p.claimed_sum() + ); } #[test] @@ -306,21 +319,21 @@ mod tests { #[test] fn canonical_round_consistency_n4() { - let f = canon(&(1u64..=16).collect::>()); - let p = CanonicalProver::new(&f); + let f = canon(&(1u64..=16).collect::>()); + let p = CanonicalProver::new(&f); let ch = [fr(2), fr(5), fr(11), fr(7)]; let mut prev = p.claimed_sum(); for j in 1..=4 { - let sj = p.compute_round_poly(&ch[..j-1]); + let sj = p.compute_round_poly(&ch[..j - 1]); assert_eq!(sj.sum_over_boolean(), prev, "round {j}"); - prev = sj.eval(ch[j-1]); + prev = sj.eval(ch[j - 1]); } } #[test] fn canonical_prove_n3() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); - let p = CanonicalProver::new(&f); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let p = CanonicalProver::new(&f); let ch = [fr(3), fr(7), fr(11)]; let proof = p.prove(&ch); assert_eq!(proof.num_vars(), 3); @@ -328,7 +341,7 @@ mod tests { for j in 1..=3 { let sj = proof.round_poly(j); assert_eq!(sj.sum_over_boolean(), prev, "round {j}"); - prev = sj.eval(ch[j-1]); + prev = sj.eval(ch[j - 1]); } } @@ -347,7 +360,10 @@ mod tests { #[test] fn canonical_zero_poly_claimed_sum_is_zero() { - assert_eq!(CanonicalProver::new(&CanonicalPoly::::zero(3)).claimed_sum(), Fr::zero()); + assert_eq!( + CanonicalProver::new(&CanonicalPoly::::zero(3)).claimed_sum(), + Fr::zero() + ); } // ── LagrangeProver ──────────────────────────────────────────────────────── @@ -371,40 +387,43 @@ mod tests { fn lagrange_round_1_sums_to_claimed() { let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); let p = LagrangeProver::new(&f); - assert_eq!(p.compute_round_poly(&[]).sum_over_boolean(), p.claimed_sum()); + assert_eq!( + p.compute_round_poly(&[]).sum_over_boolean(), + p.claimed_sum() + ); } #[test] fn lagrange_round_consistency_n3() { - let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); - let p = LagrangeProver::new(&f); + let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); + let p = LagrangeProver::new(&f); let ch = [fr(3), fr(7), fr(11)]; let mut prev = p.claimed_sum(); for j in 1..=3 { - let sj = p.compute_round_poly(&ch[..j-1]); + let sj = p.compute_round_poly(&ch[..j - 1]); assert_eq!(sj.sum_over_boolean(), prev, "round {j}"); - prev = sj.eval(ch[j-1]); + prev = sj.eval(ch[j - 1]); } } #[test] fn lagrange_round_consistency_n4() { let evals: Vec = (0..16).map(|i| i * 2 + 1).collect(); - let f = lagrange(&evals); - let p = LagrangeProver::new(&f); + let f = lagrange(&evals); + let p = LagrangeProver::new(&f); let ch = [fr(2), fr(5), fr(11), fr(7)]; let mut prev = p.claimed_sum(); for j in 1..=4 { - let sj = p.compute_round_poly(&ch[..j-1]); + let sj = p.compute_round_poly(&ch[..j - 1]); assert_eq!(sj.sum_over_boolean(), prev, "round {j}"); - prev = sj.eval(ch[j-1]); + prev = sj.eval(ch[j - 1]); } } #[test] fn lagrange_prove_n3() { - let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); - let p = LagrangeProver::new(&f); + let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); + let p = LagrangeProver::new(&f); let ch = [fr(3), fr(7), fr(11)]; let proof = p.prove(&ch); assert_eq!(proof.num_vars(), 3); @@ -412,20 +431,23 @@ mod tests { for j in 1..=3 { let sj = proof.round_poly(j); assert_eq!(sj.sum_over_boolean(), prev, "round {j}"); - prev = sj.eval(ch[j-1]); + prev = sj.eval(ch[j - 1]); } } #[test] fn lagrange_zero_poly_claimed_sum_is_zero() { - assert_eq!(LagrangeProver::new(&LagrangePoly::::zero(3)).claimed_sum(), Fr::zero()); + assert_eq!( + LagrangeProver::new(&LagrangePoly::::zero(3)).claimed_sum(), + Fr::zero() + ); } // ── Canonical and Lagrange agree ────────────────────────────────────────── #[test] fn canonical_and_lagrange_agree_n3() { - let canon_f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let canon_f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let lagrange_f = LagrangeDecomp::build(&canon_f).to_lagrange(); let cp = CanonicalProver::new(&canon_f); let lp = LagrangeProver::new(&lagrange_f); @@ -433,8 +455,8 @@ mod tests { let ch = [fr(3), fr(7), fr(11)]; for j in 1..=3 { assert_eq!( - cp.compute_round_poly(&ch[..j-1]), - lp.compute_round_poly(&ch[..j-1]), + cp.compute_round_poly(&ch[..j - 1]), + lp.compute_round_poly(&ch[..j - 1]), "round {j}" ); } @@ -442,7 +464,7 @@ mod tests { #[test] fn canonical_and_lagrange_agree_n4() { - let canon_f = canon(&(1u64..=16).collect::>()); + let canon_f = canon(&(1u64..=16).collect::>()); let lagrange_f = LagrangeDecomp::build(&canon_f).to_lagrange(); let cp = CanonicalProver::new(&canon_f); let lp = LagrangeProver::new(&lagrange_f); @@ -450,10 +472,10 @@ mod tests { let ch = [fr(2), fr(5), fr(11), fr(7)]; for j in 1..=4 { assert_eq!( - cp.compute_round_poly(&ch[..j-1]), - lp.compute_round_poly(&ch[..j-1]), + cp.compute_round_poly(&ch[..j - 1]), + lp.compute_round_poly(&ch[..j - 1]), "round {j}" ); } } -} \ No newline at end of file +} diff --git a/src/sumcheck/verifier.rs b/src/sumcheck/verifier.rs index e93c0f3..ac2d1c3 100644 --- a/src/sumcheck/verifier.rs +++ b/src/sumcheck/verifier.rs @@ -15,7 +15,7 @@ //! # Usage //! //! ```rust,ignore -//! use mlp_pro::sumcheck::verifier::Verifier; +//! use multilinear_sumcheck::sumcheck::verifier::Verifier; //! //! // Prover side //! let proof = prover.prove(&challenges); @@ -40,29 +40,16 @@ use crate::sumcheck::proof::SumcheckProof; #[derive(Debug, Clone, PartialEq, Eq)] pub enum VerifierError { /// The number of challenges does not match the number of rounds. - ChallengeLengthMismatch { - expected: usize, - got: usize, - }, + ChallengeLengthMismatch { expected: usize, got: usize }, /// Round 1 failed: `s_1(0) + s_1(1) ≠ claimed_sum`. - Round1SumMismatch { - got: F, - expected: F, - }, + Round1SumMismatch { got: F, expected: F }, /// Round `j` (2 ≤ j ≤ n) failed: `s_j(0) + s_j(1) ≠ s_{j-1}(r_{j-1})`. - RoundConsistencyMismatch { - round: usize, - got: F, - expected: F, - }, + RoundConsistencyMismatch { round: usize, got: F, expected: F }, /// Final oracle check failed: `s_n(r_n) ≠ oracle_eval`. - OracleCheckFailed { - got: F, - expected: F, - }, + OracleCheckFailed { got: F, expected: F }, } impl fmt::Display for VerifierError { @@ -76,7 +63,11 @@ impl fmt::Display for VerifierError { f, "round 1 sum check failed: s_1(0)+s_1(1) = {got}, expected {expected}" ), - Self::RoundConsistencyMismatch { round, got, expected } => write!( + Self::RoundConsistencyMismatch { + round, + got, + expected, + } => write!( f, "round {round} consistency check failed: \ s_{round}(0)+s_{round}(1) = {got}, expected {expected}" @@ -117,8 +108,8 @@ impl Verifier { /// /// Does not panic — all error conditions are returned as `Err`. pub fn verify( - proof: &SumcheckProof, - challenges: &[F], + proof: &SumcheckProof, + challenges: &[F], oracle_eval: F, ) -> Result<(), VerifierError> { let n = proof.num_vars(); @@ -127,13 +118,25 @@ impl Verifier { if challenges.len() != n { return Err(VerifierError::ChallengeLengthMismatch { expected: n, - got: challenges.len(), + got: challenges.len(), }); } + // Degenerate zero-variable case: the Sumcheck transcript has no rounds. + // The claimed hypercube sum is simply the oracle evaluation. + if n == 0 { + if proof.claimed_sum != oracle_eval { + return Err(VerifierError::OracleCheckFailed { + got: proof.claimed_sum, + expected: oracle_eval, + }); + } + return Ok(()); + } + // ── Check 1: round 1 ───────────────────────────────────────────────── // s_1(0) + s_1(1) must equal the claimed sum. - let s1 = proof.round_poly(1); + let s1 = proof.round_poly(1); let got = s1.sum_over_boolean(); if got != proof.claimed_sum { return Err(VerifierError::Round1SumMismatch { @@ -145,11 +148,11 @@ impl Verifier { // ── Check 2: round consistency for j = 2 … n ───────────────────────── // s_j(0) + s_j(1) must equal s_{j-1}(r_{j-1}). for j in 2..=n { - let sj = proof.round_poly(j); - let s_prev = proof.round_poly(j - 1); - let r_prev = challenges[j - 2]; // r_{j-1} is challenges[j-2] (0-based) + let sj = proof.round_poly(j); + let s_prev = proof.round_poly(j - 1); + let r_prev = challenges[j - 2]; // r_{j-1} is challenges[j-2] (0-based) let expected = s_prev.eval(r_prev); - let got = sj.sum_over_boolean(); + let got = sj.sum_over_boolean(); if got != expected { return Err(VerifierError::RoundConsistencyMismatch { round: j, @@ -161,12 +164,12 @@ impl Verifier { // ── Check 3: final oracle check ─────────────────────────────────────── // s_n(r_n) must equal f(r_1, …, r_n). - let sn = proof.round_poly(n); - let r_n = challenges[n - 1]; + let sn = proof.round_poly(n); + let r_n = challenges[n - 1]; let sn_at_r = sn.eval(r_n); if sn_at_r != oracle_eval { return Err(VerifierError::OracleCheckFailed { - got: sn_at_r, + got: sn_at_r, expected: oracle_eval, }); } @@ -181,7 +184,7 @@ impl Verifier { /// Useful for testing the prover in isolation before an oracle is /// available. pub fn verify_transcript( - proof: &SumcheckProof, + proof: &SumcheckProof, challenges: &[F], ) -> Result<(), VerifierError> { let n = proof.num_vars(); @@ -189,12 +192,17 @@ impl Verifier { if challenges.len() != n { return Err(VerifierError::ChallengeLengthMismatch { expected: n, - got: challenges.len(), + got: challenges.len(), }); } + // A zero-variable transcript has no round-consistency checks. + if n == 0 { + return Ok(()); + } + // Round 1 - let s1 = proof.round_poly(1); + let s1 = proof.round_poly(1); let got = s1.sum_over_boolean(); if got != proof.claimed_sum { return Err(VerifierError::Round1SumMismatch { @@ -205,11 +213,11 @@ impl Verifier { // Rounds 2..n for j in 2..=n { - let sj = proof.round_poly(j); - let s_prev = proof.round_poly(j - 1); - let r_prev = challenges[j - 2]; + let sj = proof.round_poly(j); + let s_prev = proof.round_poly(j - 1); + let r_prev = challenges[j - 2]; let expected = s_prev.eval(r_prev); - let got = sj.sum_over_boolean(); + let got = sj.sum_over_boolean(); if got != expected { return Err(VerifierError::RoundConsistencyMismatch { round: j, @@ -230,14 +238,16 @@ impl Verifier { #[cfg(test)] mod tests { use super::*; - use crate::poly::{CanonicalPoly, LagrangePoly}; use crate::circuit::LagrangeDecomp; - use crate::sumcheck::prover::{CanonicalProver, LagrangeProver}; + use crate::poly::{CanonicalPoly, LagrangePoly}; use crate::sumcheck::proof::RoundPoly; + use crate::sumcheck::prover::{CanonicalProver, LagrangeProver}; use ark_bn254::Fr; use ark_ff::Zero; - fn fr(n: u64) -> Fr { Fr::from(n) } + fn fr(n: u64) -> Fr { + Fr::from(n) + } fn canon(coeffs: &[u64]) -> CanonicalPoly { CanonicalPoly::new(coeffs.iter().map(|&v| fr(v)).collect()) @@ -253,10 +263,10 @@ mod tests { /// Oracle evaluation is computed from the canonical polynomial directly. #[test] fn canonical_prove_and_verify_n3() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = [fr(3), fr(7), fr(11)]; - let proof = CanonicalProver::new(&f).prove(&ch); + let proof = CanonicalProver::new(&f).prove(&ch); let oracle_eval = f.eval_circuit(&ch); assert!(Verifier::verify(&proof, &ch, oracle_eval).is_ok()); @@ -264,10 +274,10 @@ mod tests { #[test] fn canonical_prove_and_verify_n4() { - let f = canon(&(1u64..=16).collect::>()); + let f = canon(&(1u64..=16).collect::>()); let ch = [fr(2), fr(5), fr(11), fr(7)]; - let proof = CanonicalProver::new(&f).prove(&ch); + let proof = CanonicalProver::new(&f).prove(&ch); let oracle_eval = f.eval_circuit(&ch); assert!(Verifier::verify(&proof, &ch, oracle_eval).is_ok()); @@ -275,10 +285,10 @@ mod tests { #[test] fn canonical_verify_n2() { - let f = canon(&[1, 2, 3, 4]); + let f = canon(&[1, 2, 3, 4]); let ch = [fr(5), fr(9)]; - let proof = CanonicalProver::new(&f).prove(&ch); + let proof = CanonicalProver::new(&f).prove(&ch); let oracle_eval = f.eval_circuit(&ch); assert!(Verifier::verify(&proof, &ch, oracle_eval).is_ok()); @@ -288,10 +298,10 @@ mod tests { #[test] fn lagrange_prove_and_verify_n3() { - let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = [fr(3), fr(7), fr(11)]; - let proof = LagrangeProver::new(&f).prove(&ch); + let proof = LagrangeProver::new(&f).prove(&ch); let oracle_eval = f.eval_optimized(&ch); assert!(Verifier::verify(&proof, &ch, oracle_eval).is_ok()); @@ -300,10 +310,10 @@ mod tests { #[test] fn lagrange_prove_and_verify_n4() { let evals: Vec = (0..16).map(|i| i * 2 + 1).collect(); - let f = lagrange(&evals); + let f = lagrange(&evals); let ch = [fr(2), fr(5), fr(11), fr(7)]; - let proof = LagrangeProver::new(&f).prove(&ch); + let proof = LagrangeProver::new(&f).prove(&ch); let oracle_eval = f.eval_optimized(&ch); assert!(Verifier::verify(&proof, &ch, oracle_eval).is_ok()); @@ -313,24 +323,24 @@ mod tests { #[test] fn canonical_and_lagrange_both_verify_n3() { - let canon_f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let canon_f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let lagrange_f = LagrangeDecomp::build(&canon_f).to_lagrange(); - let ch = [fr(3), fr(7), fr(11)]; + let ch = [fr(3), fr(7), fr(11)]; let cp = CanonicalProver::new(&canon_f); let lp = LagrangeProver::new(&lagrange_f); - let canon_proof = cp.prove(&ch); + let canon_proof = cp.prove(&ch); let lagrange_proof = lp.prove(&ch); - let oracle_canon = canon_f.eval_circuit(&ch); + let oracle_canon = canon_f.eval_circuit(&ch); let oracle_lagrange = lagrange_f.eval_optimized(&ch); // Both oracle evaluations must agree assert_eq!(oracle_canon, oracle_lagrange); // Both proofs must verify - assert!(Verifier::verify(&canon_proof, &ch, oracle_canon).is_ok()); + assert!(Verifier::verify(&canon_proof, &ch, oracle_canon).is_ok()); assert!(Verifier::verify(&lagrange_proof, &ch, oracle_lagrange).is_ok()); } @@ -338,7 +348,7 @@ mod tests { #[test] fn verify_transcript_accepts_valid_proof_n3() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = [fr(3), fr(7), fr(11)]; let proof = CanonicalProver::new(&f).prove(&ch); assert!(Verifier::verify_transcript(&proof, &ch).is_ok()); @@ -348,18 +358,21 @@ mod tests { #[test] fn wrong_challenge_count_returns_error() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = [fr(3), fr(7), fr(11)]; let proof = CanonicalProver::new(&f).prove(&ch); // Too few challenges let err = Verifier::verify(&proof, &[fr(3), fr(7)], fr(0)); - assert!(matches!(err, Err(VerifierError::ChallengeLengthMismatch { .. }))); + assert!(matches!( + err, + Err(VerifierError::ChallengeLengthMismatch { .. }) + )); } #[test] fn tampered_claimed_sum_fails_round1() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = [fr(3), fr(7), fr(11)]; let mut proof = CanonicalProver::new(&f).prove(&ch); @@ -373,7 +386,7 @@ mod tests { #[test] fn tampered_round_poly_fails_consistency() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = [fr(3), fr(7), fr(11)]; let mut proof = CanonicalProver::new(&f).prove(&ch); @@ -390,7 +403,7 @@ mod tests { #[test] fn wrong_oracle_eval_fails_final_check() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = [fr(3), fr(7), fr(11)]; let proof = CanonicalProver::new(&f).prove(&ch); @@ -401,7 +414,7 @@ mod tests { #[test] fn tampered_last_round_poly_fails_oracle() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = [fr(3), fr(7), fr(11)]; let mut proof = CanonicalProver::new(&f).prove(&ch); @@ -412,10 +425,7 @@ mod tests { // Keep s_3(0)+s_3(1) = s_2(r_2) by adjusting both a and b proportionally // Easiest: corrupt a but keep sum_over_boolean intact — impossible cleanly, // so just replace with a zeroed poly and check oracle fails - proof.round_polys[2] = RoundPoly::new( - correct_s3.a + fr(1), - correct_s3.b, - ); + proof.round_polys[2] = RoundPoly::new(correct_s3.a + fr(1), correct_s3.b); let err = Verifier::verify(&proof, &ch, oracle); // Either consistency or oracle fails @@ -426,9 +436,9 @@ mod tests { #[test] fn zero_poly_verifies() { - let f = CanonicalPoly::::zero(3); + let f = CanonicalPoly::::zero(3); let ch = [fr(3), fr(7), fr(11)]; - let proof = CanonicalProver::new(&f).prove(&ch); + let proof = CanonicalProver::new(&f).prove(&ch); let oracle_eval = f.eval_circuit(&ch); assert_eq!(oracle_eval, Fr::zero()); assert!(Verifier::verify(&proof, &ch, oracle_eval).is_ok()); @@ -438,7 +448,7 @@ mod tests { #[test] fn lagrange_tampered_claimed_sum_fails() { - let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = [fr(3), fr(7), fr(11)]; let mut proof = LagrangeProver::new(&f).prove(&ch); proof.claimed_sum = fr(999); @@ -451,7 +461,7 @@ mod tests { #[test] fn lagrange_wrong_oracle_fails() { - let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = [fr(3), fr(7), fr(11)]; let proof = LagrangeProver::new(&f).prove(&ch); assert!(matches!( @@ -459,4 +469,20 @@ mod tests { Err(VerifierError::OracleCheckFailed { .. }) )); } -} \ No newline at end of file + + #[test] + fn zero_variable_proof_verifies() { + let proof = crate::sumcheck::proof::SumcheckProof::::new(fr(7), vec![]); + assert!(Verifier::verify(&proof, &[], fr(7)).is_ok()); + assert!(Verifier::verify_transcript(&proof, &[]).is_ok()); + } + + #[test] + fn zero_variable_wrong_oracle_is_rejected() { + let proof = crate::sumcheck::proof::SumcheckProof::::new(fr(7), vec![]); + assert!(matches!( + Verifier::verify(&proof, &[], fr(8)), + Err(VerifierError::OracleCheckFailed { .. }) + )); + } +} diff --git a/tests/integration.rs b/tests/integration.rs index 6548d05..1da5cca 100644 --- a/tests/integration.rs +++ b/tests/integration.rs @@ -1,4 +1,4 @@ -//! End-to-end integration tests for the `mlp_pro` Sumcheck protocol. +//! End-to-end integration tests for the `multilinear_sumcheck` Sumcheck protocol. //! //! These tests cross all module boundaries and verify the full pipeline: //! @@ -21,20 +21,22 @@ use ark_bn254::Fr; use ark_ff::{UniformRand, Zero}; -use ark_std::rand::SeedableRng; use ark_std::rand::rngs::StdRng; +use ark_std::rand::SeedableRng; -use mlp_pro::circuit::LagrangeDecomp; -use mlp_pro::poly::{CanonicalPoly, LagrangePoly, MlPoly}; -use mlp_pro::sumcheck::proof::RoundPoly; -use mlp_pro::sumcheck::prover::{CanonicalProver, LagrangeProver}; -use mlp_pro::sumcheck::verifier::{Verifier, VerifierError}; +use multilinear_sumcheck::circuit::LagrangeDecomp; +use multilinear_sumcheck::poly::{CanonicalPoly, LagrangePoly, MlPoly}; +use multilinear_sumcheck::sumcheck::proof::RoundPoly; +use multilinear_sumcheck::sumcheck::prover::{CanonicalProver, LagrangeProver}; +use multilinear_sumcheck::sumcheck::verifier::{Verifier, VerifierError}; // ───────────────────────────────────────────────────────────────────────────── // Test helpers // ───────────────────────────────────────────────────────────────────────────── -fn fr(n: u64) -> Fr { Fr::from(n) } +fn fr(n: u64) -> Fr { + Fr::from(n) +} /// Deterministic RNG — same seed across all tests. fn rng() -> StdRng { @@ -79,7 +81,7 @@ fn lagrange(evals: &[u64]) -> LagrangePoly { fn canonical_completeness_n1() { // f = 3 + 7x₁ → H(f) = 3 + 3+7 = 13... wait: H = f(0)+f(1) = 3+10 = 13 // canonical: H = 3·2^1 + 7·2^0 = 6+7 = 13 - let f = canon(&[3, 7]); + let f = canon(&[3, 7]); let ch = vec![fr(5)]; let proof = CanonicalProver::new(&f).prove(&ch); let oracle = f.eval_circuit(&ch); @@ -88,7 +90,7 @@ fn canonical_completeness_n1() { #[test] fn canonical_completeness_n2() { - let f = canon(&[1, 2, 3, 4]); + let f = canon(&[1, 2, 3, 4]); let ch = vec![fr(5), fr(9)]; let proof = CanonicalProver::new(&f).prove(&ch); let oracle = f.eval_circuit(&ch); @@ -97,7 +99,7 @@ fn canonical_completeness_n2() { #[test] fn canonical_completeness_n3() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = vec![fr(3), fr(7), fr(11)]; let proof = CanonicalProver::new(&f).prove(&ch); let oracle = f.eval_circuit(&ch); @@ -106,7 +108,7 @@ fn canonical_completeness_n3() { #[test] fn canonical_completeness_n4() { - let f = canon(&(1u64..=16).collect::>()); + let f = canon(&(1u64..=16).collect::>()); let ch = vec![fr(2), fr(5), fr(11), fr(7)]; let proof = CanonicalProver::new(&f).prove(&ch); let oracle = f.eval_circuit(&ch); @@ -115,7 +117,7 @@ fn canonical_completeness_n4() { #[test] fn canonical_completeness_random_n5() { - let f = random_canonical(5); + let f = random_canonical(5); let ch = random_challenges(5); let proof = CanonicalProver::new(&f).prove(&ch); let oracle = f.eval_circuit(&ch); @@ -125,7 +127,7 @@ fn canonical_completeness_random_n5() { #[test] fn canonical_completeness_random_n10() { // Uses the cached bit-reverse table for n=10. - let f = random_canonical(10); + let f = random_canonical(10); let ch = random_challenges(10); let proof = CanonicalProver::new(&f).prove(&ch); let oracle = f.eval_circuit(&ch); @@ -139,7 +141,7 @@ fn canonical_completeness_random_n10() { #[test] fn lagrange_completeness_n1() { // f(0)=3, f(1)=10 → H = 13 - let f = lagrange(&[3, 10]); + let f = lagrange(&[3, 10]); let ch = vec![fr(5)]; let proof = LagrangeProver::new(&f).prove(&ch); let oracle = f.eval_optimized(&ch); @@ -148,7 +150,7 @@ fn lagrange_completeness_n1() { #[test] fn lagrange_completeness_n2() { - let f = lagrange(&[1, 3, 4, 10]); + let f = lagrange(&[1, 3, 4, 10]); let ch = vec![fr(5), fr(9)]; let proof = LagrangeProver::new(&f).prove(&ch); let oracle = f.eval_optimized(&ch); @@ -157,7 +159,7 @@ fn lagrange_completeness_n2() { #[test] fn lagrange_completeness_n3() { - let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = vec![fr(3), fr(7), fr(11)]; let proof = LagrangeProver::new(&f).prove(&ch); let oracle = f.eval_optimized(&ch); @@ -167,7 +169,7 @@ fn lagrange_completeness_n3() { #[test] fn lagrange_completeness_n4() { let evals: Vec = (0..16).map(|i| i * 2 + 1).collect(); - let f = lagrange(&evals); + let f = lagrange(&evals); let ch = vec![fr(2), fr(5), fr(11), fr(7)]; let proof = LagrangeProver::new(&f).prove(&ch); let oracle = f.eval_optimized(&ch); @@ -176,7 +178,7 @@ fn lagrange_completeness_n4() { #[test] fn lagrange_completeness_random_n5() { - let f = random_lagrange(5); + let f = random_lagrange(5); let ch = random_challenges(5); let proof = LagrangeProver::new(&f).prove(&ch); let oracle = f.eval_optimized(&ch); @@ -185,7 +187,7 @@ fn lagrange_completeness_random_n5() { #[test] fn lagrange_completeness_random_n10() { - let f = random_lagrange(10); + let f = random_lagrange(10); let ch = random_challenges(10); let proof = LagrangeProver::new(&f).prove(&ch); let oracle = f.eval_optimized(&ch); @@ -198,10 +200,10 @@ fn lagrange_completeness_random_n10() { #[test] fn soundness_tampered_claimed_sum_canonical() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = vec![fr(3), fr(7), fr(11)]; let mut proof = CanonicalProver::new(&f).prove(&ch); - let oracle = f.eval_circuit(&ch); + let oracle = f.eval_circuit(&ch); proof.claimed_sum = proof.claimed_sum + fr(1); @@ -213,10 +215,10 @@ fn soundness_tampered_claimed_sum_canonical() { #[test] fn soundness_tampered_claimed_sum_lagrange() { - let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = vec![fr(3), fr(7), fr(11)]; let mut proof = LagrangeProver::new(&f).prove(&ch); - let oracle = f.eval_optimized(&ch); + let oracle = f.eval_optimized(&ch); proof.claimed_sum = fr(0); @@ -228,10 +230,10 @@ fn soundness_tampered_claimed_sum_lagrange() { #[test] fn soundness_tampered_mid_round_poly_canonical() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = vec![fr(3), fr(7), fr(11)]; let mut proof = CanonicalProver::new(&f).prove(&ch); - let oracle = f.eval_circuit(&ch); + let oracle = f.eval_circuit(&ch); // Replace round 2 with a random polynomial. proof.round_polys[1] = RoundPoly::new(fr(999), fr(888)); @@ -244,10 +246,10 @@ fn soundness_tampered_mid_round_poly_canonical() { #[test] fn soundness_tampered_mid_round_poly_lagrange() { - let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = vec![fr(3), fr(7), fr(11)]; let mut proof = LagrangeProver::new(&f).prove(&ch); - let oracle = f.eval_optimized(&ch); + let oracle = f.eval_optimized(&ch); proof.round_polys[1] = RoundPoly::new(fr(777), fr(666)); @@ -259,7 +261,7 @@ fn soundness_tampered_mid_round_poly_lagrange() { #[test] fn soundness_wrong_oracle_canonical() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = vec![fr(3), fr(7), fr(11)]; let proof = CanonicalProver::new(&f).prove(&ch); @@ -272,7 +274,7 @@ fn soundness_wrong_oracle_canonical() { #[test] fn soundness_wrong_oracle_lagrange() { - let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); + let f = lagrange(&[1, 2, 3, 4, 5, 6, 7, 8]); let ch = vec![fr(3), fr(7), fr(11)]; let proof = LagrangeProver::new(&f).prove(&ch); @@ -284,14 +286,17 @@ fn soundness_wrong_oracle_lagrange() { #[test] fn soundness_wrong_challenge_count() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); - let ch = vec![fr(3), fr(7), fr(11)]; + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let ch = vec![fr(3), fr(7), fr(11)]; let proof = CanonicalProver::new(&f).prove(&ch); // Pass only 2 challenges for a 3-variable proof. assert!(matches!( Verifier::verify(&proof, &[fr(3), fr(7)], fr(0)), - Err(VerifierError::ChallengeLengthMismatch { expected: 3, got: 2 }) + Err(VerifierError::ChallengeLengthMismatch { + expected: 3, + got: 2 + }) )); } @@ -299,15 +304,13 @@ fn soundness_wrong_challenge_count() { fn soundness_completely_random_proof_rejected() { // Build a valid proof structure with random garbage values. let mut rng = rng(); - let n = 4; + let n = 4; let round_polys: Vec> = (0..n) .map(|_| RoundPoly::new(Fr::rand(&mut rng), Fr::rand(&mut rng))) .collect(); - let proof = mlp_pro::sumcheck::proof::SumcheckProof::new( - Fr::rand(&mut rng), - round_polys, - ); - let ch = random_challenges(n); + let proof = + multilinear_sumcheck::sumcheck::proof::SumcheckProof::new(Fr::rand(&mut rng), round_polys); + let ch = random_challenges(n); let oracle = Fr::rand(&mut rng); // A random proof almost surely fails — check it is indeed rejected. @@ -324,9 +327,9 @@ fn soundness_completely_random_proof_rejected() { /// - both produce proofs that verify against the same oracle #[test] fn cross_basis_agree_n3() { - let canon_f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let canon_f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); let lagrange_f = LagrangeDecomp::build(&canon_f).to_lagrange(); - let ch = vec![fr(3), fr(7), fr(11)]; + let ch = vec![fr(3), fr(7), fr(11)]; let cp = CanonicalProver::new(&canon_f); let lp = LagrangeProver::new(&lagrange_f); @@ -337,8 +340,8 @@ fn cross_basis_agree_n3() { // Same round polynomials. for j in 1..=3 { assert_eq!( - cp.compute_round_poly(&ch[..j-1]), - lp.compute_round_poly(&ch[..j-1]), + cp.compute_round_poly(&ch[..j - 1]), + lp.compute_round_poly(&ch[..j - 1]), "round {j}" ); } @@ -355,38 +358,30 @@ fn cross_basis_agree_n3() { #[test] fn cross_basis_agree_n4() { - let canon_f = canon(&(1u64..=16).collect::>()); + let canon_f = canon(&(1u64..=16).collect::>()); let lagrange_f = LagrangeDecomp::build(&canon_f).to_lagrange(); - let ch = vec![fr(2), fr(5), fr(11), fr(7)]; + let ch = vec![fr(2), fr(5), fr(11), fr(7)]; let oracle_c = canon_f.eval_circuit(&ch); let oracle_l = lagrange_f.eval_optimized(&ch); assert_eq!(oracle_c, oracle_l); - assert!(Verifier::verify( - &CanonicalProver::new(&canon_f).prove(&ch), &ch, oracle_c - ).is_ok()); - assert!(Verifier::verify( - &LagrangeProver::new(&lagrange_f).prove(&ch), &ch, oracle_l - ).is_ok()); + assert!(Verifier::verify(&CanonicalProver::new(&canon_f).prove(&ch), &ch, oracle_c).is_ok()); + assert!(Verifier::verify(&LagrangeProver::new(&lagrange_f).prove(&ch), &ch, oracle_l).is_ok()); } #[test] fn cross_basis_agree_random_n5() { - let canon_f = random_canonical(5); + let canon_f = random_canonical(5); let lagrange_f = LagrangeDecomp::build(&canon_f).to_lagrange(); - let ch = random_challenges(5); + let ch = random_challenges(5); let oracle_c = canon_f.eval_circuit(&ch); let oracle_l = lagrange_f.eval_optimized(&ch); assert_eq!(oracle_c, oracle_l); - assert!(Verifier::verify( - &CanonicalProver::new(&canon_f).prove(&ch), &ch, oracle_c - ).is_ok()); - assert!(Verifier::verify( - &LagrangeProver::new(&lagrange_f).prove(&ch), &ch, oracle_l - ).is_ok()); + assert!(Verifier::verify(&CanonicalProver::new(&canon_f).prove(&ch), &ch, oracle_c).is_ok()); + assert!(Verifier::verify(&LagrangeProver::new(&lagrange_f).prove(&ch), &ch, oracle_l).is_ok()); } // ───────────────────────────────────────────────────────────────────────────── @@ -395,9 +390,9 @@ fn cross_basis_agree_random_n5() { #[test] fn zero_polynomial_canonical() { - let f = CanonicalPoly::::zero(3); - let ch = vec![fr(1), fr(2), fr(3)]; - let proof = CanonicalProver::new(&f).prove(&ch); + let f = CanonicalPoly::::zero(3); + let ch = vec![fr(1), fr(2), fr(3)]; + let proof = CanonicalProver::new(&f).prove(&ch); let oracle = f.eval_circuit(&ch); assert_eq!(proof.claimed_sum, Fr::zero()); @@ -407,9 +402,9 @@ fn zero_polynomial_canonical() { #[test] fn zero_polynomial_lagrange() { - let f = LagrangePoly::::zero(3); - let ch = vec![fr(1), fr(2), fr(3)]; - let proof = LagrangeProver::new(&f).prove(&ch); + let f = LagrangePoly::::zero(3); + let ch = vec![fr(1), fr(2), fr(3)]; + let proof = LagrangeProver::new(&f).prove(&ch); let oracle = f.eval_optimized(&ch); assert_eq!(proof.claimed_sum, Fr::zero()); @@ -422,9 +417,9 @@ fn constant_polynomial_canonical() { // f = 5 (constant), H(f) = 5 * 2^n = 5 * 8 = 40 for n=3 let mut coeffs = vec![Fr::zero(); 8]; coeffs[0] = fr(5); - let f = CanonicalPoly::new(coeffs); - let ch = vec![fr(2), fr(3), fr(7)]; - let proof = CanonicalProver::new(&f).prove(&ch); + let f = CanonicalPoly::new(coeffs); + let ch = vec![fr(2), fr(3), fr(7)]; + let proof = CanonicalProver::new(&f).prove(&ch); let oracle = f.eval_circuit(&ch); assert_eq!(proof.claimed_sum, fr(40)); @@ -436,40 +431,47 @@ fn constant_polynomial_canonical() { fn proof_size_grows_linearly_with_n() { // Proof has 2n+1 field elements. for n in [1, 2, 3, 4, 5] { - let f = random_canonical(n); - let ch = random_challenges(n); + let f = random_canonical(n); + let ch = random_challenges(n); let proof = CanonicalProver::new(&f).prove(&ch); - assert_eq!(proof.size_in_field_elements(), 2 * n + 1, - "proof size wrong for n={n}"); + assert_eq!( + proof.size_in_field_elements(), + 2 * n + 1, + "proof size wrong for n={n}" + ); } } #[test] fn claimed_sum_matches_hypercube_sum_canonical() { for n in [1, 2, 3, 4] { - let f = random_canonical(n); - let proof = CanonicalProver::new(&f) - .prove(&random_challenges(n)); - assert_eq!(proof.claimed_sum, f.hypercube_sum(), - "claimed sum mismatch for n={n}"); + let f = random_canonical(n); + let proof = CanonicalProver::new(&f).prove(&random_challenges(n)); + assert_eq!( + proof.claimed_sum, + f.hypercube_sum(), + "claimed sum mismatch for n={n}" + ); } } #[test] fn claimed_sum_matches_hypercube_sum_lagrange() { for n in [1, 2, 3, 4] { - let f = random_lagrange(n); - let proof = LagrangeProver::new(&f) - .prove(&random_challenges(n)); - assert_eq!(proof.claimed_sum, f.hypercube_sum(), - "claimed sum mismatch for n={n}"); + let f = random_lagrange(n); + let proof = LagrangeProver::new(&f).prove(&random_challenges(n)); + assert_eq!( + proof.claimed_sum, + f.hypercube_sum(), + "claimed sum mismatch for n={n}" + ); } } #[test] fn verify_transcript_canonical_n3() { - let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); - let ch = vec![fr(3), fr(7), fr(11)]; + let f = canon(&[1, 2, 3, 4, 5, 6, 7, 8]); + let ch = vec![fr(3), fr(7), fr(11)]; let proof = CanonicalProver::new(&f).prove(&ch); // Internal consistency check without oracle. assert!(Verifier::verify_transcript(&proof, &ch).is_ok()); @@ -478,8 +480,8 @@ fn verify_transcript_canonical_n3() { #[test] fn verify_transcript_lagrange_n4() { let evals: Vec = (0..16).map(|i| i * 3 + 1).collect(); - let f = lagrange(&evals); - let ch = vec![fr(2), fr(5), fr(9), fr(13)]; + let f = lagrange(&evals); + let ch = vec![fr(2), fr(5), fr(9), fr(13)]; let proof = LagrangeProver::new(&f).prove(&ch); assert!(Verifier::verify_transcript(&proof, &ch).is_ok()); -} \ No newline at end of file +} From b1e2b8db1b9a7b014286b60386428376d9092b34 Mon Sep 17 00:00:00 2001 From: qoosmo Date: Thu, 27 Aug 2026 18:58:20 +0300 Subject: [PATCH 2/4] chore: enforce warning-free Rust quality gates --- .github/workflows/ci.yml | 2 +- src/circuit/canonical.rs | 14 +++++++------- src/circuit/lagrange_decomp.rs | 24 ++++++++++++------------ src/circuit/sum_circuit.rs | 24 ++++++++++++------------ src/poly/canonical.rs | 7 +++---- src/poly/lagrange.rs | 13 +++++-------- src/poly/uni.rs | 19 +++++++++---------- src/sumcheck/proof.rs | 4 ++-- tests/integration.rs | 2 +- 9 files changed, 52 insertions(+), 57 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 53a52fa..c1a2d16 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -26,6 +26,6 @@ jobs: - name: Tests run: cargo test --all-features - name: Clippy - run: cargo clippy --all-targets --all-features + run: cargo clippy --all-targets --all-features -- -D warnings - name: Compile benchmarks run: cargo bench --bench sumcheck --no-run diff --git a/src/circuit/canonical.rs b/src/circuit/canonical.rs index 95de2c3..c6677c0 100644 --- a/src/circuit/canonical.rs +++ b/src/circuit/canonical.rs @@ -70,7 +70,7 @@ impl CanonicalDecomp { #[inline] pub fn q(&self, j: usize) -> &CanonicalPoly { assert!( - j >= 1 && j <= 2 * self.big_n - 1, + j >= 1 && j < 2 * self.big_n, "q index {j} out of range [1, {}]", 2 * self.big_n - 1 ); @@ -144,14 +144,14 @@ mod tests { #[test] fn node_count_is_2n_minus_1() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let d = CanonicalDecomp::build(&f); assert_eq!(d.nodes.len(), 15); } #[test] fn root_equals_input() { - let coeffs: Vec = (0..8).map(|i| fr(i)).collect(); + let coeffs: Vec = (0..8).map(fr).collect(); let f = CanonicalPoly::new(coeffs.clone()); let d = CanonicalDecomp::build(&f); assert_eq!(d.root().coeffs(), f.coeffs()); @@ -159,14 +159,14 @@ mod tests { #[test] fn total_field_elements_is_n_plus_1_times_n() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let d = CanonicalDecomp::build(&f); assert_eq!(d.total_field_elements(), (d.n + 1) * d.big_n); } #[test] fn layer_sizes_are_correct() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let d = CanonicalDecomp::build(&f); assert_eq!(d.layer(0).len(), 1); assert_eq!(d.layer(1).len(), 2); @@ -176,7 +176,7 @@ mod tests { #[test] fn node_coeff_sizes_decrease_by_layer() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let d = CanonicalDecomp::build(&f); assert_eq!(d.q(1).num_evals(), 8); assert_eq!(d.q(2).num_evals(), 4); @@ -207,7 +207,7 @@ mod tests { #[test] fn leaves_are_bit_reverse_permutation_of_root_n3() { - let f = CanonicalPoly::new((1..=8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((1..=8).map(fr).collect()); let d = CanonicalDecomp::build(&f); assert!(d.leaves_are_bit_reverse_of_root()); } diff --git a/src/circuit/lagrange_decomp.rs b/src/circuit/lagrange_decomp.rs index 4c1f0bc..8a008fe 100644 --- a/src/circuit/lagrange_decomp.rs +++ b/src/circuit/lagrange_decomp.rs @@ -118,7 +118,7 @@ impl LagrangeDecomp { #[inline] pub fn p(&self, j: usize) -> &CanonicalPoly { assert!( - j >= 1 && j <= 2 * self.big_n - 1, + j >= 1 && j < 2 * self.big_n, "p index {j} out of range [1, {}]", 2 * self.big_n - 1 ); @@ -171,10 +171,10 @@ impl LagrangeDecomp { let leaves = self.leaves(); let mut evals = vec![F::zero(); self.big_n]; - for k in 0..self.big_n { + for (k, leaf) in leaves.iter().enumerate().take(self.big_n) { let rev_k = self.bit_rev_table[k]; // leaf k holds f(rev(k)), so f(rev(k)) goes to position rev(k). - evals[rev_k] = leaves[k].coeffs()[0]; + evals[rev_k] = leaf.coeffs()[0]; } LagrangePoly::new(evals) @@ -211,14 +211,14 @@ mod tests { #[test] fn node_count_is_2n_minus_1() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let d = LagrangeDecomp::build(&f); assert_eq!(d.nodes.len(), 15); } #[test] fn root_equals_input() { - let coeffs: Vec = (0..8).map(|i| fr(i)).collect(); + let coeffs: Vec = (0..8).map(fr).collect(); let f = CanonicalPoly::new(coeffs.clone()); let d = LagrangeDecomp::build(&f); assert_eq!(d.root().coeffs(), f.coeffs()); @@ -226,14 +226,14 @@ mod tests { #[test] fn total_field_elements_is_n_plus_1_times_n() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let d = LagrangeDecomp::build(&f); assert_eq!(d.total_field_elements(), (d.n + 1) * d.big_n); } #[test] fn layer_sizes_are_correct() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let d = LagrangeDecomp::build(&f); assert_eq!(d.layer(1).len(), 1); assert_eq!(d.layer(2).len(), 2); @@ -243,7 +243,7 @@ mod tests { #[test] fn node_coeff_sizes_decrease_by_layer() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let d = LagrangeDecomp::build(&f); assert_eq!(d.p(1).num_evals(), 8); assert_eq!(d.p(2).num_evals(), 4); @@ -297,7 +297,7 @@ mod tests { #[test] fn addition_count_matches_theory_n3() { // n · 2^{n-1} = 3 · 4 = 12 - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let d = LagrangeDecomp::build(&f); assert!(d.addition_count_is_correct()); assert_eq!(d.addition_count, 12); @@ -306,7 +306,7 @@ mod tests { #[test] fn addition_count_matches_theory_n4() { // n · 2^{n-1} = 4 · 8 = 32 - let f = CanonicalPoly::new((0..16).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..16).map(fr).collect()); let d = LagrangeDecomp::build(&f); assert!(d.addition_count_is_correct()); assert_eq!(d.addition_count, 32); @@ -354,7 +354,7 @@ mod tests { #[test] fn to_lagrange_hypercube_sum_matches_canonical_n3() { // H(f) must be the same whether computed from canonical or Lagrange. - let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); + let coeffs: Vec = (1..=8).map(fr).collect(); let f = CanonicalPoly::new(coeffs); let d = LagrangeDecomp::build(&f); let lag = d.to_lagrange(); @@ -365,7 +365,7 @@ mod tests { fn to_lagrange_evals_match_canonical_eval_circuit_n3() { // Every entry of the Lagrange eval vector must equal eval_circuit // of the canonical poly at the corresponding Boolean point. - let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); + let coeffs: Vec = (1..=8).map(fr).collect(); let f = CanonicalPoly::new(coeffs); let d = LagrangeDecomp::build(&f); let lag = d.to_lagrange(); diff --git a/src/circuit/sum_circuit.rs b/src/circuit/sum_circuit.rs index aa0a025..ec2d9ca 100644 --- a/src/circuit/sum_circuit.rs +++ b/src/circuit/sum_circuit.rs @@ -221,7 +221,7 @@ impl SumCircuit for CanonicalSumCircuit { #[inline] fn h(&self, j: usize) -> F { assert!( - j >= 1 && j <= 2 * self.big_n - 1, + j >= 1 && j < 2 * self.big_n, "h index {j} out of range [1, {}]", 2 * self.big_n - 1 ); @@ -339,7 +339,7 @@ impl SumCircuit for LagrangeSumCircuit { #[inline] fn h(&self, j: usize) -> F { assert!( - j >= 1 && j <= 2 * self.big_n - 1, + j >= 1 && j < 2 * self.big_n, "h index {j} out of range [1, {}]", 2 * self.big_n - 1 ); @@ -409,7 +409,7 @@ mod tests { #[test] fn canonical_root_equals_hypercube_sum_n3() { - let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); + let coeffs: Vec = (1..=8).map(fr).collect(); let f = CanonicalPoly::new(coeffs); let sc = CanonicalSumCircuit::build(&f); assert_eq!(sc.root(), f.hypercube_sum()); @@ -425,14 +425,14 @@ mod tests { #[test] fn canonical_node_count() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let sc = CanonicalSumCircuit::build(&f); assert_eq!(sc.data.len(), 15); } #[test] fn canonical_layer_sizes() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let sc = CanonicalSumCircuit::build(&f); assert_eq!(sc.layer(0).len(), 1); assert_eq!(sc.layer(1).len(), 2); @@ -442,7 +442,7 @@ mod tests { #[test] fn canonical_leaves_slice_length() { - let f = CanonicalPoly::new((0..8).map(|i| fr(i)).collect()); + let f = CanonicalPoly::new((0..8).map(fr).collect()); let sc = CanonicalSumCircuit::build(&f); assert_eq!(sc.leaves().len(), 8); } @@ -472,7 +472,7 @@ mod tests { #[test] fn lagrange_root_equals_hypercube_sum_n3() { - let evals: Vec = (1..=8).map(|i| fr(i)).collect(); + let evals: Vec = (1..=8).map(fr).collect(); let f = LagrangePoly::new(evals); let sc = LagrangeSumCircuit::build(&f); assert_eq!(sc.root(), f.hypercube_sum()); @@ -488,7 +488,7 @@ mod tests { #[test] fn lagrange_recurrence_holds_n3() { - let evals: Vec = (1..=8).map(|i| fr(i)).collect(); + let evals: Vec = (1..=8).map(fr).collect(); let f = LagrangePoly::new(evals); let sc = LagrangeSumCircuit::build(&f); assert!(sc.verify_recurrence()); @@ -496,7 +496,7 @@ mod tests { #[test] fn lagrange_layer_sizes() { - let f = LagrangePoly::new((0..8).map(|i| fr(i)).collect()); + let f = LagrangePoly::new((0..8).map(fr).collect()); let sc = LagrangeSumCircuit::build(&f); assert_eq!(sc.layer(0).len(), 1); assert_eq!(sc.layer(1).len(), 2); @@ -514,7 +514,7 @@ mod tests { #[test] fn canonical_and_lagrange_agree_on_root_n3() { use crate::circuit::LagrangeDecomp; - let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); + let coeffs: Vec = (1..=8).map(fr).collect(); let canon = CanonicalPoly::new(coeffs); let lag = LagrangeDecomp::build(&canon).to_lagrange(); let sc_canon = CanonicalSumCircuit::build(&canon); @@ -536,7 +536,7 @@ mod tests { #[test] fn canonical_from_leaves_matches_build() { - let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); + let coeffs: Vec = (1..=8).map(fr).collect(); let f = CanonicalPoly::new(coeffs); let table = get_or_build(3); let manual_leaves: Vec = (0..8).map(|k| f.coeffs()[table[k]]).collect(); @@ -547,7 +547,7 @@ mod tests { #[test] fn lagrange_from_leaves_matches_build() { - let evals: Vec = (1..=8).map(|i| fr(i)).collect(); + let evals: Vec = (1..=8).map(fr).collect(); let f = LagrangePoly::new(evals); let table = get_or_build(3); let manual_leaves: Vec = (0..8).map(|k| f.evals()[table[k]]).collect(); diff --git a/src/poly/canonical.rs b/src/poly/canonical.rs index e22afbd..0e46b56 100644 --- a/src/poly/canonical.rs +++ b/src/poly/canonical.rs @@ -165,8 +165,7 @@ impl CanonicalPoly { // (j₁ = 0) and q₃ holds odd-indexed coefficients (j₁ = 1). let mut buf = self.coeffs.clone(); - for k in 0..self.num_vars { - let r_k = r[k]; + for &r_k in r.iter().take(self.num_vars) { let half = buf.len() / 2; for t in 0..half { buf[t] = buf[2 * t] + r_k * buf[2 * t + 1]; @@ -396,7 +395,7 @@ mod tests { #[test] fn naive_and_circuit_agree_at_boolean_points_n3() { // f = α₀ + α₁x₁ + … + α₇x₁x₂x₃ - let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); + let coeffs: Vec = (1..=8).map(fr).collect(); let f = CanonicalPoly::new(coeffs); for b0 in [fr(0), fr(1)] { for b1 in [fr(0), fr(1)] { @@ -415,7 +414,7 @@ mod tests { #[test] fn naive_and_circuit_agree_at_random_point_n3() { - let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); + let coeffs: Vec = (1..=8).map(fr).collect(); let f = CanonicalPoly::new(coeffs); // Use a fixed non-boolean point to test the general case. let r = [fr(2), fr(5), fr(11)]; diff --git a/src/poly/lagrange.rs b/src/poly/lagrange.rs index d12e5d5..2d9b7a2 100644 --- a/src/poly/lagrange.rs +++ b/src/poly/lagrange.rs @@ -82,8 +82,7 @@ impl LagrangePoly { let mut buf = self.evals.clone(); - for k in 0..self.num_vars { - let r_k = r[k]; + for &r_k in r.iter().take(self.num_vars) { let one_minus_r = F::one() - r_k; let half = buf.len() / 2; for t in 0..half { @@ -116,8 +115,7 @@ impl LagrangePoly { let mut buf = self.evals.clone(); - for k in 0..self.num_vars { - let r_k = r[k]; + for &r_k in r.iter().take(self.num_vars) { let half = buf.len() / 2; for t in 0..half { buf[t] = buf[2 * t] + r_k * (buf[2 * t + 1] - buf[2 * t]); @@ -164,8 +162,7 @@ impl LagrangePoly { // Pre-allocate a reusable buffer — avoids one allocation per layer. let mut tmp = Vec::with_capacity(buf.len()); - for k in 0..self.num_vars { - let r_k = r[k]; + for &r_k in r.iter().take(self.num_vars) { tmp.clear(); buf.par_chunks(2) .map(|pair| pair[0] + r_k * (pair[1] - pair[0])) @@ -349,7 +346,7 @@ mod tests { #[test] fn all_three_agree_at_boolean_points_n3() { - let p = LagrangePoly::new((1..=8).map(|i| fr(i)).collect()); + let p = LagrangePoly::new((1..=8).map(fr).collect()); for b0 in [fr(0), fr(1)] { for b1 in [fr(0), fr(1)] { for b2 in [fr(0), fr(1)] { @@ -379,7 +376,7 @@ mod tests { fn lagrange_eval_agrees_with_canonical_eval_circuit_n3() { use crate::circuit::LagrangeDecomp; use crate::poly::CanonicalPoly; - let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); + let coeffs: Vec = (1..=8).map(fr).collect(); let canon = CanonicalPoly::new(coeffs); let lag = LagrangeDecomp::build(&canon).to_lagrange(); let r = [fr(3), fr(7), fr(13)]; diff --git a/src/poly/uni.rs b/src/poly/uni.rs index be253d2..c5c31a6 100644 --- a/src/poly/uni.rs +++ b/src/poly/uni.rs @@ -244,8 +244,7 @@ impl UniPoly { let mut buf = self.coeffs.clone(); // At layer k (0-based from bottom), combine with r^{2^k}. - for k in 0..n { - let lambda = powers[k]; + for &lambda in powers.iter().take(n) { let new_len = buf.len() / 2; for t in 0..new_len { // buf[t] = buf[2t] + r^{2^k} · buf[2t+1] @@ -412,7 +411,7 @@ impl UniDecomp { #[inline] pub fn q(&self, j: usize) -> &[F] { assert!( - j >= 1 && j <= 2 * self.big_n - 1, + j >= 1 && j < 2 * self.big_n, "q index {j} out of range [1, {}]", 2 * self.big_n - 1 ); @@ -755,7 +754,7 @@ mod tests { #[test] fn all_methods_agree_n3() { - let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); + let coeffs: Vec = (1..=8).map(fr).collect(); let p = UniPoly::new(coeffs); for r in [fr(0), fr(1), fr(2), fr(7), fr(100)] { let naive = p.eval_naive(r); @@ -788,7 +787,7 @@ mod tests { #[test] fn all_methods_agree_non_power_of_two_input() { // length 5 → padded to 8 - let coeffs: Vec = (1..=5).map(|i| fr(i)).collect(); + let coeffs: Vec = (1..=5).map(fr).collect(); let p = UniPoly::new(coeffs); for r in [fr(0), fr(1), fr(3), fr(11)] { let naive = p.eval_naive(r); @@ -813,20 +812,20 @@ mod tests { #[test] fn decompose_node_count() { - let d = UniPoly::new((0..8).map(|i| fr(i)).collect()).decompose(); + let d = UniPoly::new((0..8).map(fr).collect()).decompose(); assert_eq!(d.nodes.len(), 15); } #[test] fn decompose_root_equals_input() { - let coeffs: Vec = (1..=8).map(|i| fr(i)).collect(); + let coeffs: Vec = (1..=8).map(fr).collect(); let p = UniPoly::new(coeffs.clone()); assert_eq!(p.decompose().root(), coeffs.as_slice()); } #[test] fn decompose_layer_sizes() { - let d = UniPoly::new((0..8).map(|i| fr(i)).collect()).decompose(); + let d = UniPoly::new((0..8).map(fr).collect()).decompose(); assert_eq!(d.layer(0).len(), 1); assert_eq!(d.layer(1).len(), 2); assert_eq!(d.layer(2).len(), 4); @@ -835,7 +834,7 @@ mod tests { #[test] fn decompose_total_field_elements() { - let d = UniPoly::new((0..8).map(|i| fr(i)).collect()).decompose(); + let d = UniPoly::new((0..8).map(fr).collect()).decompose(); assert_eq!(d.total_field_elements(), (d.n + 1) * d.big_n); } @@ -873,7 +872,7 @@ mod tests { #[test] fn decompose_gate_relation_holds_n3() { - let p = UniPoly::new((1..=8).map(|i| fr(i)).collect()); + let p = UniPoly::new((1..=8).map(fr).collect()); let d = p.decompose(); for j in 1..d.big_n { let qj = d.q(j); diff --git a/src/sumcheck/proof.rs b/src/sumcheck/proof.rs index 8e40551..2f9b742 100644 --- a/src/sumcheck/proof.rs +++ b/src/sumcheck/proof.rs @@ -3,9 +3,9 @@ //! This module defines the two data types that constitute a Sumcheck proof: //! //! - [`RoundPoly`] — the degree-1 polynomial `s_j(X_j) = a + b·X_j` -//! sent by the prover at each round. +//! sent by the prover at each round. //! - [`SumcheckProof`] — the complete transcript: claimed sum `h` plus -//! the `n` round polynomials `s_1, …, s_n`. +//! the `n` round polynomials `s_1, …, s_n`. //! //! # Protocol recap //! diff --git a/tests/integration.rs b/tests/integration.rs index 1da5cca..93d0a94 100644 --- a/tests/integration.rs +++ b/tests/integration.rs @@ -205,7 +205,7 @@ fn soundness_tampered_claimed_sum_canonical() { let mut proof = CanonicalProver::new(&f).prove(&ch); let oracle = f.eval_circuit(&ch); - proof.claimed_sum = proof.claimed_sum + fr(1); + proof.claimed_sum += fr(1); assert!(matches!( Verifier::verify(&proof, &ch, oracle), From 8b40dcb1635193d8aa9b530ef610cec722b78e2e Mon Sep 17 00:00:00 2001 From: qoosmo Date: Thu, 27 Aug 2026 19:04:21 +0300 Subject: [PATCH 3/4] docs: make research presentation self-contained --- CONTRIBUTING.md | 2 +- src/circuit/lagrange_decomp.rs | 7 ++++--- src/circuit/sum_circuit.rs | 6 +++--- src/poly/lagrange.rs | 12 +++++++----- src/sumcheck/proof.rs | 14 +++++++------- src/sumcheck/prover.rs | 17 ++++++++--------- 6 files changed, 30 insertions(+), 28 deletions(-) diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index fd76fbe..9baf489 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -8,7 +8,7 @@ Before opening a pull request, run: cargo fmt --all -- --check cargo check --all-targets --all-features cargo test --all-features -cargo clippy --all-targets --all-features +cargo clippy --all-targets --all-features -- -D warnings cargo bench --bench sumcheck --no-run ``` diff --git a/src/circuit/lagrange_decomp.rs b/src/circuit/lagrange_decomp.rs index 8a008fe..ad72c83 100644 --- a/src/circuit/lagrange_decomp.rs +++ b/src/circuit/lagrange_decomp.rs @@ -24,7 +24,8 @@ use crate::poly::{CanonicalPoly, LagrangePoly, MlPoly}; /// /// # Cost /// -/// Exactly `n · 2^{n-1}` field additions — proved in §A of the paper. +/// Exactly `n · 2^{n-1}` field additions. The implementation records +/// this count and the test suite checks it for multiple dimensions. pub struct LagrangeDecomp { /// Number of variables `n`. pub n: usize, @@ -33,7 +34,7 @@ pub struct LagrangeDecomp { pub big_n: usize, /// The sequence `(p_j)_{1 ≤ j ≤ 2N-1}`, stored 0-based. - /// Paper index `j` → `nodes[j - 1]`. + /// 1-based tree index `j` → `nodes[j - 1]`. pub nodes: Vec>, /// Cached bit-reverse permutation table of length `N`. @@ -59,7 +60,7 @@ impl LagrangeDecomp { // p_1 = f (root) nodes[0] = Some(f.clone()); - // Layer i = 1, …, n (1-based to match the paper) + // Layer i = 1, …, n in the 1-based tree notation. // At layer i, nodes are at 1-based indices [2^{i-1}, 2^i - 1]. for i in 1..=n { let layer_start = 1usize << (i - 1); // 2^{i-1} (1-based) diff --git a/src/circuit/sum_circuit.rs b/src/circuit/sum_circuit.rs index ec2d9ca..1c17ee4 100644 --- a/src/circuit/sum_circuit.rs +++ b/src/circuit/sum_circuit.rs @@ -16,8 +16,8 @@ //! //! # Indexing convention //! -//! 1-based throughout, matching the paper exactly. -//! Internal storage is 0-based: paper index `j` → `data[j-1]`. +//! The mathematical tree notation is 1-based. +//! Internal storage is 0-based: tree index `j` → `data[j-1]`. //! //! ```text //! layer 0 : j = 1 (root = H(f)) @@ -58,7 +58,7 @@ pub trait SumCircuit { 1 << self.num_vars() } - /// Get `h_j` using the **1-based paper index**. + /// Get `h_j` using the **1-based tree index**. /// /// # Panics /// Panics if `j == 0` or `j > 2N − 1`. diff --git a/src/poly/lagrange.rs b/src/poly/lagrange.rs index 2d9b7a2..6e5fbbd 100644 --- a/src/poly/lagrange.rs +++ b/src/poly/lagrange.rs @@ -100,7 +100,9 @@ impl LagrangePoly { /// Cost per pair: 1 multiplication, 2 additions. /// /// Saves one multiplication per pair vs `eval_standard`. - /// Paper reports ~25–35% improvement over `eval_standard`. + /// + /// Runtime impact is hardware- and field-dependent; use the Criterion + /// benchmark suite to compare the kernels on the target machine. /// /// # Panics /// Panics if `r.len() != num_vars`. @@ -139,10 +141,10 @@ impl LagrangePoly { /// /// # When this is faster /// - /// The parallel version is faster than `eval_optimized` only when - /// the per-layer work is large enough to amortise thread synchronisation. - /// In practice this means `n ≥ 18` on most machines. - /// At `n < 16` the sequential version is typically faster. + /// The parallel version can outperform `eval_optimized` only when + /// the per-layer work is large enough to amortise thread scheduling and + /// synchronisation. The crossover point is machine- and field-dependent; + /// benchmark both kernels on the target workload. /// /// # Panics /// Panics if `r.len() != num_vars`. diff --git a/src/sumcheck/proof.rs b/src/sumcheck/proof.rs index 2f9b742..0806df9 100644 --- a/src/sumcheck/proof.rs +++ b/src/sumcheck/proof.rs @@ -23,10 +23,10 @@ //! The verifier checks `s_j(0) + s_j(1) = s_{j-1}(r_{j-1})` and then //! samples a fresh challenge `r_j`. //! -//! From the paper (§5): +//! In the tree notation used by this crate: //! - **Canonical basis:** `s_j(X_j) = h_2' + h_3' · X_j` -//! - **Lagrange basis:** `s_j(X_j) = h_2' + h_3' · X_j` (same form, -//! different recurrence for `h_2'` and `h_3'`) +//! - **Lagrange basis:** the same affine form is reconstructed from +//! `s_j(0)` and `s_j(1)`, using the Lagrange recurrence. //! //! In both cases the round polynomial is degree-1 and is fully described //! by two field elements. @@ -42,12 +42,12 @@ use std::fmt; /// /// This is the message the prover sends in each round of the Sumcheck /// protocol. Since `f` is multilinear (degree at most 1 in each -/// variable), every round polynomial has degree exactly 1. +/// variable), every round polynomial has degree at most 1. /// -/// # Paper notation +/// # Tree notation /// -/// The paper writes `s_j(X_j) = h_2' + h_3' · X_j`. -/// Here `a = h_2'` and `b = h_3'`. +/// In the canonical tree representation, +/// `s_j(X_j) = h_2' + h_3' · X_j`, so `a = h_2'` and `b = h_3'`. /// /// # Encoding /// diff --git a/src/sumcheck/prover.rs b/src/sumcheck/prover.rs index 3f059ac..197474b 100644 --- a/src/sumcheck/prover.rs +++ b/src/sumcheck/prover.rs @@ -8,7 +8,7 @@ //! the challenges received so far, and return a [`RoundPoly`]. //! 3. Assemble the full [`SumcheckProof`] transcript. //! -//! # Round polynomial algorithm (canonical basis, §5.4) +//! # Round polynomial algorithm (canonical basis) //! //! At round `j` with challenges `r_1, …, r_{j-1}` already received: //! @@ -20,7 +20,7 @@ //! ``` //! 3. Return `s_j(X_j) = h[0] + h[1] · X_j`. //! -//! # Round polynomial algorithm (Lagrange basis, §5.5) +//! # Round polynomial algorithm (Lagrange basis) //! //! Same fold rule but with the optimized formula: //! ```text @@ -30,13 +30,12 @@ //! Returns `s_j` via `from_evaluations(h[0], h[1])` since in the Lagrange //! basis `s_j(X) = h[0]·(1−X) + h[1]·X`. //! -//! # Complexity (from the paper, Table 1) +//! # Complexity //! -//! | Algorithm | Multiplications | Additions | -//! |---------------------|----------------|-----------| -//! | LinearTimeSC | `2N` | `3N` | -//! | **CanonicalProver** | **`2N`** | **`2N`** | -//! | **LagrangeProver** | **`2N`** | **`4N`** | +//! Prover construction is linear in the evaluation-table size `N = 2^n`. +//! Round extraction folds stored tree layers rather than re-evaluating the +//! original multilinear polynomial. Basis-dependent arithmetic costs and +//! runtime measurements are documented separately in the repository. use ark_ff::Field; @@ -91,7 +90,7 @@ impl CanonicalProver { .collect(); // Fold with r_{j-1}, r_{j-2}, …, r_1 (most recent first). - // Paper §5.4 fold rule: + // Canonical fold rule: // h'[k] = h[2k] + r · h[2k+2] for k even // h'[k] = h[2k-1] + r · h[2k+1] for k odd for ki in (0..challenges.len()).rev() { From 85a9f407cc2350a4c5c70863a1fa3ad2bc7bd2ea Mon Sep 17 00:00:00 2001 From: qoosmo Date: Thu, 27 Aug 2026 19:08:09 +0300 Subject: [PATCH 4/4] docs: remove orphaned research references --- src/circuit/canonical.rs | 2 +- src/circuit/lagrange_decomp.rs | 8 ++++---- src/circuit/sum_circuit.rs | 4 ++-- src/poly/canonical.rs | 10 ++++------ src/poly/uni.rs | 10 +++++----- src/sumcheck/proof.rs | 4 ++-- 6 files changed, 18 insertions(+), 20 deletions(-) diff --git a/src/circuit/canonical.rs b/src/circuit/canonical.rs index c6677c0..ea0b7ed 100644 --- a/src/circuit/canonical.rs +++ b/src/circuit/canonical.rs @@ -186,7 +186,7 @@ mod tests { } #[test] - fn split_matches_paper_example_n3() { + fn split_matches_reference_example_n3() { let coeffs = vec![fr(1), fr(2), fr(3), fr(4), fr(5), fr(6), fr(7), fr(8)]; let f = CanonicalPoly::new(coeffs); let d = CanonicalDecomp::build(&f); diff --git a/src/circuit/lagrange_decomp.rs b/src/circuit/lagrange_decomp.rs index ad72c83..7a181bb 100644 --- a/src/circuit/lagrange_decomp.rs +++ b/src/circuit/lagrange_decomp.rs @@ -110,9 +110,9 @@ impl LagrangeDecomp { } } - // ── Accessors (paper notation) ──────────────────────────────────────────── + // ── Accessors (1-based tree notation) ──────────────────────────────────────────── - /// Return a reference to `p_j` using the **1-based paper index**. + /// Return a reference to `p_j` using the **1-based tree index**. /// /// # Panics /// Panics if `j == 0` or `j > 2N - 1`. @@ -141,7 +141,7 @@ impl LagrangeDecomp { &self.nodes[self.big_n - 1..2 * self.big_n - 1] } - /// Layer `i` (1-based, matching the paper): slice of nodes at depth `i`. + /// Layer `i` (1-based tree convention): slice of nodes at depth `i`. /// /// - Layer `1` : `[p_1]` (the root) /// - Layer `i` : `p_{2^{i-1}}, …, p_{2^i - 1}` @@ -255,7 +255,7 @@ mod tests { // ── Gate rule: p_{2j} = a, p_{2j+1} = a + b ────────────────────────────── - /// Paper example: f = α₀ + α₁x₁ + α₂x₂ + α₄x₃ + α₃x₁x₂ + α₅x₁x₃ + α₆x₂x₃ + α₇x₁x₂x₃ + /// Reference example: f = α₀ + α₁x₁ + α₂x₂ + α₄x₃ + α₃x₁x₂ + α₅x₁x₃ + α₆x₂x₃ + α₇x₁x₂x₃ /// coeffs = [1, 2, 3, 4, 5, 6, 7, 8] /// /// p_1 = f = a + x₁·b where: diff --git a/src/circuit/sum_circuit.rs b/src/circuit/sum_circuit.rs index 1c17ee4..3edfde0 100644 --- a/src/circuit/sum_circuit.rs +++ b/src/circuit/sum_circuit.rs @@ -381,7 +381,7 @@ mod tests { } #[test] - fn canonical_paper_example_n3() { + fn canonical_reference_example_n3() { let leaves = vec![fr(1), fr(4), fr(3), fr(7), fr(2), fr(6), fr(5), fr(8)]; let sc = CanonicalSumCircuit::from_leaves(&leaves); assert_eq!(sc.root(), fr(88)); @@ -394,7 +394,7 @@ mod tests { } #[test] - fn canonical_recurrence_holds_paper_example() { + fn canonical_recurrence_holds_reference_example() { let leaves = vec![fr(1), fr(4), fr(3), fr(7), fr(2), fr(6), fr(5), fr(8)]; let sc = CanonicalSumCircuit::from_leaves(&leaves); assert!(sc.verify_recurrence()); diff --git a/src/poly/canonical.rs b/src/poly/canonical.rs index 0e46b56..a77a632 100644 --- a/src/poly/canonical.rs +++ b/src/poly/canonical.rs @@ -130,8 +130,7 @@ impl CanonicalPoly { /// Circuit-based evaluation at `r = (r₁, …, rₙ)`. /// - /// Implements the bottom-up traversal of the `(q_j)` tree from §3.3 - /// of the paper. + /// Implements a bottom-up traversal of the `(q_j)` decomposition tree. /// /// # Algorithm /// @@ -431,10 +430,9 @@ mod tests { // ── eval_circuit buffer size after each fold ────────────────────────────── - /// The circuit evaluation must return the same value whether we fold - /// variables in forward or reverse order — but our implementation - /// always folds x_n first (right-to-left), matching the paper's - /// canonical order x₁ → x₂ → … → xₙ for the decomposition tree. + /// The circuit evaluator folds challenges in variable order + /// `x₁, x₂, …, xₙ`, matching the even/odd coefficient decomposition used + /// by the canonical tree. #[test] fn eval_circuit_zero_poly_is_zero() { let f = CanonicalPoly::::zero(4); diff --git a/src/poly/uni.rs b/src/poly/uni.rs index c5c31a6..dd1ac27 100644 --- a/src/poly/uni.rs +++ b/src/poly/uni.rs @@ -399,12 +399,12 @@ pub struct UniDecomp { pub n: usize, /// `N = 2^n`. pub big_n: usize, - /// Flat storage. Paper index `j` (1-based) → `nodes[j-1]`. + /// Flat storage. Tree index `j` (1-based) → `nodes[j-1]`. pub nodes: Vec>, } impl UniDecomp { - /// `q_j` using the **1-based paper index**. + /// `q_j` using the **1-based tree index**. /// /// # Panics /// Panics if `j == 0` or `j > 2N - 1`. @@ -839,7 +839,7 @@ mod tests { } #[test] - fn decompose_first_layer_paper_example_n3() { + fn decompose_first_layer_reference_example_n3() { let p = UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4), fr(5), fr(6), fr(7), fr(8)]); let d = p.decompose(); assert_eq!(d.q(2), &[fr(1), fr(3), fr(5), fr(7)]); @@ -847,7 +847,7 @@ mod tests { } #[test] - fn decompose_second_layer_paper_example_n3() { + fn decompose_second_layer_reference_example_n3() { let p = UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4), fr(5), fr(6), fr(7), fr(8)]); let d = p.decompose(); assert_eq!(d.q(4), &[fr(1), fr(5)]); @@ -857,7 +857,7 @@ mod tests { } #[test] - fn decompose_leaves_paper_example_n3() { + fn decompose_leaves_reference_example_n3() { let p = UniPoly::new(vec![fr(1), fr(2), fr(3), fr(4), fr(5), fr(6), fr(7), fr(8)]); let d = p.decompose(); assert_eq!(d.q(8), &[fr(1)]); diff --git a/src/sumcheck/proof.rs b/src/sumcheck/proof.rs index 0806df9..3abde82 100644 --- a/src/sumcheck/proof.rs +++ b/src/sumcheck/proof.rs @@ -154,7 +154,7 @@ pub struct SumcheckProof { /// The `n` round polynomials `s_1, …, s_n`. /// - /// `round_polys[j-1]` = `s_j` (0-based storage, 1-based paper index). + /// `round_polys[j-1]` = `s_j` (0-based storage, 1-based round index). pub round_polys: Vec>, } @@ -177,7 +177,7 @@ impl SumcheckProof { 1 + 2 * self.round_polys.len() } - /// The round polynomial `s_j` using the **1-based paper index**. + /// The round polynomial `s_j` using the **1-based round index**. /// /// # Panics /// Panics if `j == 0` or `j > n`.