Federated learning on MNIST with secure aggregation (Bonawitz-style pairwise masking), a continuous model-poisoning attack that exploits the privacy that secure aggregation buys, and a zero-knowledge norm-bound defence that restores robustness without breaking the privacy guarantee.
The whole stack is built from primitives in pure PyTorch + Python, so each layer is inspectable as a few hundred lines of code. The ZKP stage ships two interchangeable backends — a from-scratch Python NIZK (Pedersen + Schnorr-OR + Fiat–Shamir) used by default, and an optional Rust
- arkworks Groth16 SNARK
extension (
03_zkp_defence/zkp_rust/) selected at runtime by settingUSE_RUST_ZKP=true.
┌─────────────────────────────────────────────────────────────────┐
│ Stage 1 — Federated Averaging │
│ Clients run local SGD, send deltas Δᵢ │
│ Server averages weighted by dataset size │
│ Result: 0.9904 test accuracy in 30 rounds │
└─────────────────────────────────────────────────────────────────┘
│
▼ (server sees individual Δᵢ — privacy leak)
┌─────────────────────────────────────────────────────────────────┐
│ Stage 2 — Secure Aggregation + Poisoning Demonstration │
│ Pairwise masks sᵢⱼ = −sⱼᵢ hide every individual Δᵢ │
│ Privacy gained ✓ │
│ Robustness lost ✗ — gradient-reversal attack at ρ=0.3, α=3 │
│ Result: clean 0.9905 vs attack 0.098 (sustained 87–92 pp) │
└─────────────────────────────────────────────────────────────────┘
│
▼ (need a defence that doesn't see Δᵢ)
┌─────────────────────────────────────────────────────────────────┐
│ Stage 3 — Zero-Knowledge Norm-Bound Defence │
│ Each client proves ‖Δᵢ‖₂ ≤ B without revealing Δᵢ │
│ Built from Pedersen commitments + Schnorr-OR + Fiat–Shamir │
│ Server filters proofs that fail; aggregation continues clean │
│ Result: 0.098 → 0.9522 recovery; +27.7 % time, +0.23 % comm │
└─────────────────────────────────────────────────────────────────┘
| Stage | Folder | Headline output |
|---|---|---|
| 1. FedAvg baseline | 01_fedavg/ |
results.csv — 30-round accuracy curve, final 0.9904 |
| 2. Secure aggregation + poisoning | 02_secure_aggregation/ |
results.csv — clean vs. attacked accuracy, 89 pp gap sustained |
| 3. ZKP defence + overhead | 03_zkp_defence/ |
results.csv — round-by-round time + comm; report_4_stats.json — summary |
| Writeups | reports/ |
secure_aggregation_tradeoff.pdf, zkp_defence_overhead.pdf |
client.py:train_local()snapshotsw_G(t), runsEepochs of local SGD with momentum 0.9 and returns the deltaΔᵢ(t+1) = wᵢ(t+1) − w_G(t).server.py:aggregate()computes the weighted averageΔ_agg = Σᵢ (|Dᵢ| / |D|) · Δᵢand applies the global stepw_G(t+1) = w_G(t) + η · Δ_aggwithη = 1.main.pydrives 30 rounds across 20 clients on an IID MNIST partition.
secure_aggregation.pyimplements pairwise additive maskssᵢⱼ = −sⱼᵢ. Each client transmits its weighted update plus the sum of pairwise masks; the server sums every message and the masks telescope to zero, leaving only the aggregate learnable.attack.py:AccuracyDegradationAttackflips and amplifies the local delta:Δ̃ᵢ = −α · Δᵢ.main.pyruns clean (ρ = 0) and attacked (ρ = 0.3, α = 3) FedAvg back-to-back on the same seed and writes both accuracy series side-by- side into one CSV.
With FedAvg's (1 − ρ) − ρ · α = −0.20 effective learning rate, the
aggregate update pushes the global model backwards every round; the
attacked accuracy is pinned at ≈ 0.10 (random) while the clean baseline
converges past 0.99 — illustrating exactly why secure aggregation
complicates robustness.
A non-interactive zero-knowledge range proof for the scaled integer norm
n* = ⌊‖Δᵢ‖² · SCALE⌋ (SCALE = 10⁴). The construction composes four
textbook primitives:
- Group setup. RFC 5114 1024-bit MODP safe prime
p; subgroup of orderq = (p − 1) / 2.g = 2is a generator;h = g^seed mod pis a second generator with unknown DL baseg. - Pedersen commitment.
C(x; r) = g^x · h^r mod p. Perfectly hiding, computationally binding under DLP. - Schnorr-OR bit proof. A non-interactive proof that
Copens to either 0 or 1, built with the Cramer–Damgård–Schoenmakers OR composition: one real Schnorr branch + one simulated branch, glued by Fiat–Shamir under SHA-256. - Bit-decomposition range proof. Decompose
n*intokbits and commit to each one; pick the randomness so the bit commitments combine to the norm commitment viaC = Π Cᵢ^(2^i). The verifier dictatesk = ⌈log₂(B² · SCALE + 1)⌉from its policy, so the prover cannot enlarge the range.
A valid proof certifies n* ∈ [0, 2^k − 1], i.e.
‖Δᵢ‖₂ ≤ √((2^k − 1) / SCALE) ≈ B. Soundness reduces to DLP in the
subgroup, zero-knowledge to Pedersen hiding plus the Schnorr-OR
simulator, non-interactivity to Fiat–Shamir in the random-oracle
model.
server.py:filter_updates() rejects clients whose proofs do not verify;
the survivors are aggregated under the same secure-aggregation channel.
Caveat (honest-prover scope). The prover computes n* from its own
Δᵢ and refuses to attach a proof when n* > 2^k − 1. Binding n* to
the actual Δᵢ would require a SNARK over an arithmetic circuit for the
norm computation (e.g. Groth16 on bn254 R1CS), which would change the
implementation effort by an order of magnitude and is out of scope for
a CPU-only laptop build target.
The same ‖Δᵢ‖₂ ≤ B predicate can be discharged with a real Groth16
zk-SNARK over BN254. 03_zkp_defence/zkp_rust/ ships a PyO3 Rust
extension that wraps an arkworks-based Groth16 circuit; the vendored
arkworks dependency (03_zkp_defence/groth16/) is included so the build
is self-contained.
# Build the Rust extension (requires Rust + maturin)
pip install maturin
cd 03_zkp_defence/zkp_rust && maturin develop
# Switch the client and server to the Rust prover/verifier at runtime
cd .. && USE_RUST_ZKP=true python main.pyzkp.py exposes ZKPProverRust / ZKPVerifierRust, which the
client.py and server.py files prefer when USE_RUST_ZKP=true and
fall back to the pure-Python NIZK otherwise. See
03_zkp_defence/zkp_rust/BUILD.md for the full build walk-through.
| Metric | Value | Source |
|---|---|---|
| FedAvg final test accuracy | 0.9904 | 01_fedavg/results.csv |
| Clean / attack accuracy gap (ρ=0.3, α=3) | 87–92 pp sustained over 30 rounds | 02_secure_aggregation/results.csv |
| Accuracy recovery from ZKP filter (α=5, B=4) | 0.098 → 0.9522 | 03_zkp_defence/report_4_stats.json |
| Per-round time overhead (10 rounds, 20 clients) | +27.7 % (58.2 → 74.3 s/round) | 03_zkp_defence/results.csv |
| Per-round communication overhead | +0.23 % (proofs ≈ 15 KB / client) | 03_zkp_defence/results.csv |
# 1. Create the conda environment (Python 3.10 + PyTorch 2.0+ + numpy)
conda env create -f environment.yml
conda activate fl-secure-aggregation-zkp
# 2. Reproduce each stage (CPU is enough; tuned defaults reproduce the numbers above)
cd 01_fedavg && python main.py # ~25 min on CPU
cd ../02_secure_aggregation && python main.py # ~50 min on CPU (two FedAvg runs)
cd ../03_zkp_defence && python main.py # ~25 min on CPU
# 3. (Optional) regenerate the PDF writeups
cd ../reports
python build_report2.py
python build_report4.pyPass --help to any main.py for the full CLI surface (number of clients
and rounds, malicious ratio, attack strength, ZKP bound B, norm type,
etc.).
.
├── README.md
├── LICENSE # MIT
├── .gitignore
├── environment.yml # conda env: Python 3.10 + PyTorch 2.0+ + numpy
├── requirements.txt # pip alternative
│
├── 01_fedavg/ # Stage 1 — vanilla FedAvg baseline
│ ├── client.py / server.py / main.py
│ ├── model.py # 2-conv CNN
│ ├── data_utils.py # MNIST loader + IID/non-IID splitter
│ └── results.csv # 30-round accuracy curve
│
├── 02_secure_aggregation/ # Stage 2 — pairwise-mask SA + poisoning demo
│ ├── secure_aggregation.py # SecureAggregator (pairwise masks)
│ ├── attack.py # gradient-reversal model poisoning
│ ├── client.py / server.py / main.py
│ ├── model.py / data_utils.py
│ └── results.csv # clean vs. attacked, per round
│
├── 03_zkp_defence/ # Stage 3 — ZKP norm-bound defence + overhead
│ ├── zkp.py # Pedersen + Schnorr-OR NIZK (default backend)
│ ├── server.py # ZKP-filtering SA server
│ ├── client.py # client with ZKP prover
│ ├── attack.py / secure_aggregation.py
│ ├── main.py # no-ZKP vs ZKP comparison + overhead
│ ├── model.py / data_utils.py
│ ├── results.csv # round / time_no_zkp / time_zkp / comm_no_zkp / comm_zkp
│ ├── report_4_stats.json # summary stats (consumed by the report builder)
│ ├── zkp_rust/ # optional Rust + arkworks Groth16 backend (PyO3)
│ │ ├── Cargo.toml / src/ # maturin develop → loaded when USE_RUST_ZKP=true
│ │ └── BUILD.md # build instructions + troubleshooting
│ └── groth16/ # vendored arkworks ark-groth16 dependency
│
└── reports/
├── build_report2.py
├── build_report4.py
├── secure_aggregation_tradeoff.pdf # 1-page privacy↔robustness writeup
└── zkp_defence_overhead.pdf # 1-page ZKP design + overhead writeup
MIT — see LICENSE.