Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
341 changes: 341 additions & 0 deletions docs/protocol-projection-v1.md

Large diffs are not rendered by default.

5 changes: 4 additions & 1 deletion src/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,10 @@ authors = [
]
description = "A Python library for zero-knowledge proof generation and verification"
readme = "../readme.md" # Update path to point to root README
requires-python = ">=3.8"
# Floor matches actual usage: PEP 604 unions evaluated at import time in
# lora_contributor_mpi require 3.10+, and pathlib.Path.is_relative_to
# requires 3.9+. CI exercises 3.11.
requires-python = ">=3.10"
classifiers = [
"Programming Language :: Python :: 3",
"License :: Other/Proprietary License",
Expand Down
294 changes: 281 additions & 13 deletions src/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ use halo2_proofs::{
pasta::{vesta, EqAffine, Fp},
plonk::{
create_proof, keygen_pk, keygen_vk, verify_proof, Advice, Circuit, Column,
ConstraintSystem, Error, Instance, Selector, SingleVerifier,
ConstraintSystem, Error, Instance, ProvingKey, Selector, SingleVerifier, VerifyingKey,
},
poly::commitment::Params,
poly::Rotation,
Expand All @@ -22,14 +22,21 @@ use num_traits::{One, Signed, Zero};
use rand_core::OsRng;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use std::collections::HashMap;
use std::convert::TryInto;
use std::sync::{Arc, Mutex, OnceLock};

const ADAPTER_COMMITMENT_DOMAIN: u64 = 0x5a4b4c4f5241; // "ZKLORA"
const ADAPTER_COMMITMENT_VERSION: u64 = 1;
// Must match proof_contract.SCHEMA_VERSION; it is hashed into adapter commitments.
const ARTIFACT_SCHEMA_VERSION: u64 = 2;
const FIELD_SAFE_BITS: usize = 250;
const POSEIDON_PAIR_ROWS: usize = 96;
// Caps for the legacy backend: artifacts beyond these shapes are rejected before
// any keygen work so a hostile statement cannot stall the verifier.
const MAX_LEGACY_K: u32 = 24;
const MAX_LEGACY_DIM: usize = 16_384;
const MAX_LEGACY_RANK: usize = 1_024;

#[derive(Debug, thiserror::Error)]
pub enum NativeError {
Expand Down Expand Up @@ -608,6 +615,21 @@ impl LoraCircuit {
));
}
let in_dim = self.in_dim();
if in_dim > MAX_LEGACY_DIM || self.out_dim() > MAX_LEGACY_DIM {
return Err(NativeError::InvalidDimensions(format!(
"legacy artifact exceeds verification caps: dims {}x{} beyond {}",
in_dim,
self.out_dim(),
MAX_LEGACY_DIM
)));
}
if self.rank() > MAX_LEGACY_RANK {
return Err(NativeError::InvalidDimensions(format!(
"legacy artifact exceeds verification caps: rank {} beyond {}",
self.rank(),
MAX_LEGACY_RANK
)));
}
for row in &self.a {
if row.len() != in_dim {
return Err(NativeError::InvalidDimensions(
Expand Down Expand Up @@ -1197,6 +1219,180 @@ fn k_for(circuit: &LoraCircuit) -> u32 {
rows.trailing_zeros().max(8)
}

/// Cache key covering everything the circuit structure (and therefore the
/// params/keys) depends on: the in-circuit constants are derived solely from
/// dims, fixed-point widths, and scaling; witness values never enter keygen.
type LegacyShapeKey = (u32, usize, usize, usize, u32, u32, u32, i64, i64);

/// Both caches are bounded: the cache keys span every attacker-influenceable
/// statement field, so an unbounded map fed varying shapes is a slow memory
/// DoS. Entries near k = MAX_LEGACY_K are GB-scale, hence the small caps.
const MAX_LEGACY_KEY_CACHE_ENTRIES: usize = 4;
const MAX_LEGACY_PARAMS_CACHE_ENTRIES: usize = 4;

/// Minimal LRU map: a HashMap with a monotonically increasing use stamp per
/// entry; inserting beyond capacity evicts the least recently used entry.
struct BoundedLru<K, V> {
map: HashMap<K, (u64, V)>,
counter: u64,
capacity: usize,
}

impl<K: std::hash::Hash + Eq + Clone, V: Clone> BoundedLru<K, V> {
fn new(capacity: usize) -> Self {
BoundedLru {
map: HashMap::new(),
counter: 0,
capacity: capacity.max(1),
}
}

fn get(&mut self, key: &K) -> Option<V> {
self.counter += 1;
let stamp = self.counter;
self.map.get_mut(key).map(|slot| {
slot.0 = stamp;
slot.1.clone()
})
}

fn insert(&mut self, key: K, value: V) {
self.counter += 1;
if !self.map.contains_key(&key) && self.map.len() >= self.capacity {
if let Some(oldest) = self
.map
.iter()
.min_by_key(|(_, (stamp, _))| *stamp)
.map(|(k, _)| k.clone())
{
self.map.remove(&oldest);
}
}
self.map.insert(key, (self.counter, value));
}

#[cfg(test)]
fn len(&self) -> usize {
self.map.len()
}

#[cfg(test)]
fn contains_key(&self, key: &K) -> bool {
self.map.contains_key(key)
}
}

/// Verification only ever needs the verifying key; the proving key is built
/// (and cached) lazily the first time a shape is actually proven, so a
/// verifier never pays keygen_pk for hostile or one-off shapes.
enum LegacyKeys {
VerifyOnly {
params: Arc<Params<EqAffine>>,
vk: VerifyingKey<EqAffine>,
},
Prover {
params: Arc<Params<EqAffine>>,
pk: ProvingKey<EqAffine>,
},
}

impl LegacyKeys {
fn params(&self) -> &Params<EqAffine> {
match self {
LegacyKeys::VerifyOnly { params, .. } => params,
LegacyKeys::Prover { params, .. } => params,
}
}

fn vk(&self) -> &VerifyingKey<EqAffine> {
match self {
LegacyKeys::VerifyOnly { vk, .. } => vk,
LegacyKeys::Prover { pk, .. } => pk.get_vk(),
}
}

fn pk(&self) -> Option<&ProvingKey<EqAffine>> {
match self {
LegacyKeys::VerifyOnly { .. } => None,
LegacyKeys::Prover { pk, .. } => Some(pk),
}
}
}

static LEGACY_KEY_CACHE: OnceLock<Mutex<BoundedLru<LegacyShapeKey, Arc<LegacyKeys>>>> =
OnceLock::new();
static LEGACY_PARAMS_CACHE: OnceLock<Mutex<BoundedLru<u32, Arc<Params<EqAffine>>>>> =
OnceLock::new();

/// Cached values are immutable once inserted (Arc'd keys/params plus LRU
/// bookkeeping), so a panic in another thread cannot leave them torn;
/// recover from poisoning instead of propagating panics through PyO3.
fn lock_recovering<T>(mutex: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
mutex
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}

fn legacy_cache() -> &'static Mutex<BoundedLru<LegacyShapeKey, Arc<LegacyKeys>>> {
LEGACY_KEY_CACHE.get_or_init(|| Mutex::new(BoundedLru::new(MAX_LEGACY_KEY_CACHE_ENTRIES)))
}

fn legacy_params_for(k: u32) -> Arc<Params<EqAffine>> {
let cache = LEGACY_PARAMS_CACHE
.get_or_init(|| Mutex::new(BoundedLru::new(MAX_LEGACY_PARAMS_CACHE_ENTRIES)));
if let Some(found) = lock_recovering(cache).get(&k) {
return found;
}
// Built outside the lock: Params::new at large k takes seconds and two
// racing builders are deterministic, so last-write-wins is harmless.
let params = Arc::new(Params::<EqAffine>::new(k));
lock_recovering(cache).insert(k, params.clone());
params
}

fn legacy_shape_key(circuit: &LoraCircuit, k: u32) -> LegacyShapeKey {
(
k,
circuit.in_dim(),
circuit.rank(),
circuit.out_dim(),
circuit.fixed_point.scale_bits,
circuit.fixed_point.value_bits,
circuit.fixed_point.intermediate_bits,
circuit.scaling_num,
circuit.scaling_den,
)
}

fn legacy_keys_for(circuit: &LoraCircuit, need_pk: bool) -> Result<Arc<LegacyKeys>, NativeError> {
let k = k_for(circuit);
if k > MAX_LEGACY_K {
return Err(NativeError::InvalidDimensions(format!(
"legacy artifact exceeds verification caps: k {k} beyond {MAX_LEGACY_K}"
)));
}
let key = legacy_shape_key(circuit, k);
if let Some(found) = lock_recovering(legacy_cache()).get(&key) {
if !need_pk || found.pk().is_some() {
return Ok(found);
}
}
// Keygen runs outside the lock so concurrent callers on other shapes are
// not serialized behind it; duplicated keygen on the same shape is
// deterministic and last-write-wins.
let params = legacy_params_for(k);
let empty = circuit.without_witnesses();
let vk = keygen_vk(&params, &empty).map_err(|e| NativeError::Halo2(e.to_string()))?;
let entry = if need_pk {
let pk = keygen_pk(&params, vk, &empty).map_err(|e| NativeError::Halo2(e.to_string()))?;
Arc::new(LegacyKeys::Prover { params, pk })
} else {
Arc::new(LegacyKeys::VerifyOnly { params, vk })
};
lock_recovering(legacy_cache()).insert(key, entry.clone());
Ok(entry)
}

fn circuit_from_json(statement_json: &str, witness_json: &str) -> Result<LoraCircuit, NativeError> {
let statement: NativeStatement =
serde_json::from_str(statement_json).map_err(|e| NativeError::Json(e.to_string()))?;
Expand Down Expand Up @@ -1236,16 +1432,14 @@ fn default_scaling_den() -> i64 {

pub fn prove_bytes(statement_json: &str, witness_json: &str) -> Result<Vec<u8>, NativeError> {
let circuit = circuit_from_json(statement_json, witness_json)?;
let k = k_for(&circuit);
let params: Params<EqAffine> = Params::new(k);
let vk = keygen_vk(&params, &circuit).map_err(|e| NativeError::Halo2(e.to_string()))?;
let pk = keygen_pk(&params, vk, &circuit).map_err(|e| NativeError::Halo2(e.to_string()))?;
let keys = legacy_keys_for(&circuit, true)?;
let pk = keys.pk().expect("prover cache entry carries a proving key");
let instances = public_inputs(&circuit)?;
let instance_refs: Vec<&[Fp]> = vec![instances.as_slice()];
let mut transcript = Blake2bWrite::<_, vesta::Affine, Challenge255<_>>::init(vec![]);
create_proof(
&params,
&pk,
keys.params(),
pk,
&[circuit],
&[instance_refs.as_slice()],
&mut OsRng,
Expand All @@ -1266,16 +1460,14 @@ pub fn verify_bytes(statement_json: &str, proof: &[u8]) -> Result<bool, NativeEr
statement_json,
&serde_json::to_string(&witness_shape).map_err(|e| NativeError::Json(e.to_string()))?,
)?;
let k = k_for(&circuit);
let params: Params<EqAffine> = Params::new(k);
let vk = keygen_vk(&params, &circuit).map_err(|e| NativeError::Halo2(e.to_string()))?;
let keys = legacy_keys_for(&circuit, false)?;
let instances = public_inputs(&circuit)?;
let instance_refs: Vec<&[Fp]> = vec![instances.as_slice()];
let mut transcript = Blake2bRead::<_, vesta::Affine, Challenge255<_>>::init(proof);
let result = verify_proof(
&params,
&vk,
SingleVerifier::new(&params),
keys.params(),
keys.vk(),
SingleVerifier::new(keys.params()),
&[instance_refs.as_slice()],
&mut transcript,
);
Expand Down Expand Up @@ -1405,6 +1597,82 @@ mod tests {
assert_ne!(first, adapter_commitment_for_input(&changed).unwrap());
}

#[test]
fn legacy_key_cache_reuses_and_upgrades_entries() {
let circuit = minimal_circuit();
let verify_only = legacy_keys_for(&circuit, false).unwrap();
assert!(verify_only.pk().is_none());
let cached = legacy_keys_for(&circuit, false).unwrap();
assert!(Arc::ptr_eq(&verify_only, &cached));

// A prover call on the same shape upgrades the entry in place...
let prover = legacy_keys_for(&circuit, true).unwrap();
assert!(prover.pk().is_some());
// ...and both later verifiers and provers share the upgraded entry.
let reused_verify = legacy_keys_for(&circuit, false).unwrap();
assert!(Arc::ptr_eq(&prover, &reused_verify));
let reused_prover = legacy_keys_for(&circuit, true).unwrap();
assert!(Arc::ptr_eq(&prover, &reused_prover));

let key = legacy_shape_key(&circuit, k_for(&circuit));
assert!(lock_recovering(legacy_cache()).contains_key(&key));
}

#[test]
fn bounded_lru_evicts_least_recently_used() {
let mut lru: BoundedLru<u32, u32> = BoundedLru::new(2);
lru.insert(1, 10);
lru.insert(2, 20);
assert_eq!(lru.get(&1), Some(10)); // touch 1 so 2 becomes the oldest
lru.insert(3, 30);
assert_eq!(lru.len(), 2);
assert!(lru.contains_key(&1));
assert!(!lru.contains_key(&2));
assert!(lru.contains_key(&3));

// Re-inserting an existing key must not evict anything.
lru.insert(1, 11);
assert_eq!(lru.len(), 2);
assert_eq!(lru.get(&1), Some(11));
assert!(lru.contains_key(&3));
}

#[test]
fn legacy_caps_reject_oversized_dimensions() {
let fixed_point = FixedPointConfig {
scale_bits: 1,
value_bits: 8,
intermediate_bits: 16,
};
let wide = LoraCircuit {
a: vec![vec![0; MAX_LEGACY_DIM + 1]],
b: vec![vec![0]],
x: vec![0; MAX_LEGACY_DIM + 1],
delta: vec![0],
fixed_point: fixed_point.clone(),
scaling_num: 1,
scaling_den: 1,
adapter_commitment: "0".to_string(),
statement_digest: "22".repeat(32),
};
let err = wide.validate().unwrap_err();
assert!(err.to_string().contains("exceeds verification caps"));

let deep = LoraCircuit {
a: vec![vec![0]; MAX_LEGACY_RANK + 1],
b: vec![vec![0; MAX_LEGACY_RANK + 1]],
x: vec![0],
delta: vec![0],
fixed_point,
scaling_num: 1,
scaling_den: 1,
adapter_commitment: "0".to_string(),
statement_digest: "22".repeat(32),
};
let err = deep.validate().unwrap_err();
assert!(err.to_string().contains("exceeds verification caps"));
}

#[test]
fn mock_prover_accepts_valid_lora_relation() {
let circuit = valid_circuit();
Expand Down
1 change: 1 addition & 0 deletions src/zklora/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
"generate_proofs": ("zklora.zk_proof_generator", "generate_proofs"),
"adapter_manifest_entry": ("zklora.proof_contract", "adapter_manifest_entry"),
"write_adapter_manifest": ("zklora.proof_contract", "write_adapter_manifest"),
"expand_statement_rows": ("zklora.proof_v3", "expand_statement_rows"),
"commit_activations": ("zklora.polynomial_commit", "commit_activations"),
"verify_commitment": ("zklora.polynomial_commit", "verify_commitment"),
}
Expand Down
Loading
Loading