Skip to main content

sigma_proofs/
composition.rs

1//! # Protocol Composition with AND/OR Logic
2//!
3//! This module defines the [`ComposedRelation`] enum, which generalizes the [`CanonicalLinearRelation`]
4//! by enabling compositional logic between multiple proof instances.
5//!
6//! Specifically, it supports:
7//! - Simple atomic proofs (e.g., discrete logarithm, Pedersen commitments)
8//! - Conjunctions (`And`) of multiple sub-protocols
9//! - Disjunctions (`Or`) of multiple sub-protocols
10//! - Thresholds (`Threshold`) over multiple sub-protocols
11//!
12//! ## Example Composition
13//!
14//! ```ignore
15//! And(
16//!    Or(dleq, pedersen_commitment),
17//!    Simple(discrete_logarithm),
18//!    And(pedersen_commitment_dleq, bbs_blind_commitment_computation)
19//! )
20//! ```
21
22use alloc::{vec, vec::Vec};
23use ff::{Field, PrimeField};
24use group::prime::PrimeGroup;
25use itertools::Itertools;
26use sha3::{Digest, Sha3_256};
27use spongefish::{
28    Decoding, Encoding, NargDeserialize, NargSerialize, VerificationError, VerificationResult,
29};
30use subtle::{Choice, ConditionallySelectable, ConstantTimeEq};
31
32use crate::errors::InvalidInstance;
33use crate::traits::ScalarRng;
34use crate::MultiScalarMul;
35use crate::{
36    errors::Error,
37    fiat_shamir::Nizk,
38    linear_relation::{CanonicalLinearRelation, LinearRelation},
39    traits::{SigmaProtocol, SigmaProtocolSimulator},
40};
41
42/// A protocol proving knowledge of a witness for a composition of linear relations.
43///
44/// This implementation generalizes [`CanonicalLinearRelation`] by using AND/OR links.
45///
46/// # Type Parameters
47/// - `G`: A cryptographic group implementing [`group::Group`] and [`group::GroupEncoding`].
48#[derive(Clone)]
49pub enum ComposedRelation<G: PrimeGroup> {
50    Simple(CanonicalLinearRelation<G>),
51    And(Vec<ComposedRelation<G>>),
52    Or(Vec<ComposedRelation<G>>),
53    Threshold(usize, Vec<ComposedRelation<G>>),
54}
55
56impl<G: PrimeGroup + ConstantTimeEq + ConditionallySelectable> ComposedRelation<G> {
57    /// Create a [ComposedRelation] for an AND relation from the given list of relations.
58    pub fn and<T: Into<ComposedRelation<G>>>(witness: impl IntoIterator<Item = T>) -> Self {
59        Self::And(witness.into_iter().map(|x| x.into()).collect())
60    }
61
62    /// Create a [ComposedRelation] for an OR relation from the given list of relations.
63    pub fn or<T: Into<ComposedRelation<G>>>(witness: impl IntoIterator<Item = T>) -> Self {
64        Self::Or(witness.into_iter().map(|x| x.into()).collect())
65    }
66
67    /// Create a [ComposedRelation] for a threshold relation from the given list of relations.
68    pub fn threshold<T: Into<ComposedRelation<G>>>(
69        threshold: usize,
70        witness: impl IntoIterator<Item = T>,
71    ) -> Self {
72        Self::Threshold(threshold, witness.into_iter().map(|x| x.into()).collect())
73    }
74}
75
76impl<G: PrimeGroup> From<CanonicalLinearRelation<G>> for ComposedRelation<G> {
77    fn from(value: CanonicalLinearRelation<G>) -> Self {
78        ComposedRelation::Simple(value)
79    }
80}
81
82impl<G: PrimeGroup + MultiScalarMul> TryFrom<LinearRelation<G>> for ComposedRelation<G> {
83    type Error = InvalidInstance;
84
85    fn try_from(value: LinearRelation<G>) -> Result<Self, Self::Error> {
86        Ok(Self::Simple(CanonicalLinearRelation::try_from(value)?))
87    }
88}
89
90// Structure representing the Commitment type of Protocol as SigmaProtocol
91#[derive(Clone)]
92pub enum ComposedCommitment<G>
93where
94    G: PrimeGroup + ConditionallySelectable + Encoding<[u8]> + NargSerialize + NargDeserialize,
95    G::Scalar:
96        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
97{
98    Simple(Vec<G>),
99    And(Vec<ComposedCommitment<G>>),
100    Or(Vec<ComposedCommitment<G>>),
101    Threshold(Vec<ComposedCommitment<G>>),
102}
103
104impl<G: PrimeGroup> ComposedCommitment<G>
105where
106    G: ConditionallySelectable + Encoding<[u8]> + NargSerialize + NargDeserialize,
107    G::Scalar:
108        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
109{
110    /// Conditionally select between two ComposedCommitment values.
111    /// This function performs constant-time selection of the commitment values.
112    pub fn conditional_select(a: &Self, b: &Self, choice: Choice) -> Self {
113        match (a, b) {
114            (ComposedCommitment::Simple(a_elements), ComposedCommitment::Simple(b_elements)) => {
115                let selected: Vec<G> = a_elements
116                    .iter()
117                    .zip_eq(b_elements.iter())
118                    .map(|(a, b)| G::conditional_select(a, b, choice))
119                    .collect();
120                ComposedCommitment::Simple(selected)
121            }
122            (ComposedCommitment::And(a_commitments), ComposedCommitment::And(b_commitments)) => {
123                let selected: Vec<ComposedCommitment<G>> = a_commitments
124                    .iter()
125                    .zip_eq(b_commitments.iter())
126                    .map(|(a, b)| ComposedCommitment::conditional_select(a, b, choice))
127                    .collect();
128                ComposedCommitment::And(selected)
129            }
130            (ComposedCommitment::Or(a_commitments), ComposedCommitment::Or(b_commitments)) => {
131                let selected: Vec<ComposedCommitment<G>> = a_commitments
132                    .iter()
133                    .zip_eq(b_commitments.iter())
134                    .map(|(a, b)| ComposedCommitment::conditional_select(a, b, choice))
135                    .collect();
136                ComposedCommitment::Or(selected)
137            }
138            (
139                ComposedCommitment::Threshold(a_commitments),
140                ComposedCommitment::Threshold(b_commitments),
141            ) => {
142                let selected: Vec<ComposedCommitment<G>> = a_commitments
143                    .iter()
144                    .zip_eq(b_commitments.iter())
145                    .map(|(a, b)| ComposedCommitment::conditional_select(a, b, choice))
146                    .collect();
147                ComposedCommitment::Threshold(selected)
148            }
149            _ => {
150                unreachable!("Mismatched ComposedCommitment variants in conditional_select");
151            }
152        }
153    }
154}
155
156// Structure representing the ProverState type of Protocol as SigmaProtocol
157pub enum ComposedProverState<G>
158where
159    G: PrimeGroup
160        + ConstantTimeEq
161        + ConditionallySelectable
162        + Encoding<[u8]>
163        + NargSerialize
164        + NargDeserialize
165        + MultiScalarMul,
166    G::Scalar:
167        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
168{
169    Simple(<CanonicalLinearRelation<G> as SigmaProtocol>::ProverState),
170    And(Vec<ComposedProverState<G>>),
171    Or(ComposedOrProverState<G>),
172    Threshold(ComposedThresholdProverState<G>),
173}
174
175pub type ComposedOrProverState<G> = Vec<ComposedOrProverStateEntry<G>>;
176pub struct ComposedOrProverStateEntry<G>(
177    Choice,
178    ComposedProverState<G>,
179    ComposedChallenge<G>,
180    ComposedResponse<G>,
181)
182where
183    G: PrimeGroup
184        + ConstantTimeEq
185        + ConditionallySelectable
186        + Encoding<[u8]>
187        + NargSerialize
188        + NargDeserialize
189        + MultiScalarMul,
190    G::Scalar:
191        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable;
192
193pub type ComposedThresholdProverState<G> = Vec<ComposedThresholdProverStateEntry<G>>;
194pub struct ComposedThresholdProverStateEntry<G>
195where
196    G: PrimeGroup
197        + ConstantTimeEq
198        + ConditionallySelectable
199        + Encoding<[u8]>
200        + NargSerialize
201        + NargDeserialize
202        + MultiScalarMul,
203    G::Scalar:
204        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
205{
206    use_simulator: Choice,
207    prover_state: ComposedProverState<G>,
208    simulated_challenge: ComposedChallenge<G>,
209    simulated_response: ComposedResponse<G>,
210}
211
212// Structure representing the Response type of Protocol as SigmaProtocol
213#[derive(Clone)]
214pub enum ComposedResponse<G>
215where
216    G: PrimeGroup
217        + ConditionallySelectable
218        + Encoding<[u8]>
219        + NargSerialize
220        + NargDeserialize
221        + MultiScalarMul,
222    G::Scalar:
223        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
224{
225    Simple(Vec<<CanonicalLinearRelation<G> as SigmaProtocol>::Response>),
226    And(Vec<ComposedResponse<G>>),
227    Or(Vec<ComposedChallenge<G>>, Vec<ComposedResponse<G>>),
228    Threshold(Vec<ComposedChallenge<G>>, Vec<ComposedResponse<G>>),
229}
230
231const TAG_SIMPLE: u8 = 0;
232const TAG_AND: u8 = 1;
233const TAG_OR: u8 = 2;
234const TAG_THRESHOLD: u8 = 3;
235
236fn read_u32(buf: &mut &[u8]) -> VerificationResult<u32> {
237    if buf.len() < 4 {
238        return Err(VerificationError);
239    }
240    let (head, tail) = buf.split_at(4);
241    *buf = tail;
242    Ok(u32::from_le_bytes(head.try_into().unwrap()))
243}
244
245fn write_len(out: &mut Vec<u8>, len: usize) {
246    out.extend_from_slice(&(len as u32).to_le_bytes());
247}
248
249impl<G> Encoding<[u8]> for ComposedCommitment<G>
250where
251    G: PrimeGroup + ConditionallySelectable + Encoding<[u8]> + NargSerialize + NargDeserialize,
252    G::Scalar:
253        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
254{
255    fn encode(&self) -> impl AsRef<[u8]> {
256        let mut out = Vec::new();
257        match self {
258            ComposedCommitment::Simple(elems) => {
259                out.push(TAG_SIMPLE);
260                write_len(&mut out, elems.len());
261                for elem in elems {
262                    elem.serialize_into_narg(&mut out);
263                }
264            }
265            ComposedCommitment::And(cs) => {
266                out.push(TAG_AND);
267                write_len(&mut out, cs.len());
268                for c in cs {
269                    c.serialize_into_narg(&mut out);
270                }
271            }
272            ComposedCommitment::Or(cs) => {
273                out.push(TAG_OR);
274                write_len(&mut out, cs.len());
275                for c in cs {
276                    c.serialize_into_narg(&mut out);
277                }
278            }
279            ComposedCommitment::Threshold(cs) => {
280                out.push(TAG_THRESHOLD);
281                write_len(&mut out, cs.len());
282                for c in cs {
283                    c.serialize_into_narg(&mut out);
284                }
285            }
286        }
287        out
288    }
289}
290
291impl<G> NargDeserialize for ComposedCommitment<G>
292where
293    G: PrimeGroup + ConditionallySelectable + Encoding<[u8]> + NargSerialize + NargDeserialize,
294    G::Scalar:
295        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
296{
297    fn deserialize_from_narg(buf: &mut &[u8]) -> VerificationResult<Self> {
298        if buf.is_empty() {
299            return Err(VerificationError);
300        }
301        let (tag_bytes, rest) = buf.split_at(1);
302        *buf = rest;
303        match tag_bytes[0] {
304            TAG_SIMPLE => {
305                let len = read_u32(buf)? as usize;
306                let mut elems = Vec::new();
307                for _ in 0..len {
308                    elems.push(G::deserialize_from_narg(buf)?);
309                }
310                Ok(ComposedCommitment::Simple(elems))
311            }
312            TAG_AND => {
313                let len = read_u32(buf)? as usize;
314                let mut entries = Vec::new();
315                for _ in 0..len {
316                    entries.push(ComposedCommitment::deserialize_from_narg(buf)?);
317                }
318                Ok(ComposedCommitment::And(entries))
319            }
320            TAG_OR => {
321                let len = read_u32(buf)? as usize;
322                let mut entries = Vec::new();
323                for _ in 0..len {
324                    entries.push(ComposedCommitment::deserialize_from_narg(buf)?);
325                }
326                Ok(ComposedCommitment::Or(entries))
327            }
328            TAG_THRESHOLD => {
329                let len = read_u32(buf)? as usize;
330                let mut entries = Vec::new();
331                for _ in 0..len {
332                    entries.push(ComposedCommitment::deserialize_from_narg(buf)?);
333                }
334                Ok(ComposedCommitment::Threshold(entries))
335            }
336            _ => Err(VerificationError),
337        }
338    }
339}
340
341impl<G> Encoding<[u8]> for ComposedResponse<G>
342where
343    G: PrimeGroup
344        + ConditionallySelectable
345        + Encoding<[u8]>
346        + NargSerialize
347        + NargDeserialize
348        + MultiScalarMul,
349    G::Scalar:
350        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
351{
352    fn encode(&self) -> impl AsRef<[u8]> {
353        let mut out = Vec::new();
354        match self {
355            ComposedResponse::Simple(responses) => {
356                out.push(TAG_SIMPLE);
357                write_len(&mut out, responses.len());
358                for r in responses {
359                    r.serialize_into_narg(&mut out);
360                }
361            }
362            ComposedResponse::And(entries) => {
363                out.push(TAG_AND);
364                write_len(&mut out, entries.len());
365                for r in entries {
366                    r.serialize_into_narg(&mut out);
367                }
368            }
369            ComposedResponse::Or(challenges, responses) => {
370                out.push(TAG_OR);
371                write_len(&mut out, challenges.len());
372                for c in challenges {
373                    c.serialize_into_narg(&mut out);
374                }
375                write_len(&mut out, responses.len());
376                for r in responses {
377                    r.serialize_into_narg(&mut out);
378                }
379            }
380            ComposedResponse::Threshold(challenges, responses) => {
381                out.push(TAG_THRESHOLD);
382                write_len(&mut out, challenges.len());
383                for c in challenges {
384                    c.serialize_into_narg(&mut out);
385                }
386                write_len(&mut out, responses.len());
387                for r in responses {
388                    r.serialize_into_narg(&mut out);
389                }
390            }
391        }
392        out
393    }
394}
395
396impl<G> NargDeserialize for ComposedResponse<G>
397where
398    G: PrimeGroup
399        + ConditionallySelectable
400        + Encoding<[u8]>
401        + NargSerialize
402        + NargDeserialize
403        + MultiScalarMul,
404    G::Scalar:
405        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
406{
407    fn deserialize_from_narg(buf: &mut &[u8]) -> VerificationResult<Self> {
408        if buf.is_empty() {
409            return Err(VerificationError);
410        }
411        let (tag_bytes, rest) = buf.split_at(1);
412        *buf = rest;
413        match tag_bytes[0] {
414            TAG_SIMPLE => {
415                let len = read_u32(buf)? as usize;
416                let mut elems = Vec::new();
417                for _ in 0..len {
418                    elems.push(G::Scalar::deserialize_from_narg(buf)?);
419                }
420                Ok(ComposedResponse::Simple(elems))
421            }
422            TAG_AND => {
423                let len = read_u32(buf)? as usize;
424                let mut entries = Vec::new();
425                for _ in 0..len {
426                    entries.push(ComposedResponse::deserialize_from_narg(buf)?);
427                }
428                Ok(ComposedResponse::And(entries))
429            }
430            TAG_OR => {
431                let ch_len = read_u32(buf)? as usize;
432                let mut challenges = Vec::new();
433                for _ in 0..ch_len {
434                    challenges.push(G::Scalar::deserialize_from_narg(buf)?);
435                }
436                let resp_len = read_u32(buf)? as usize;
437                let mut responses = Vec::new();
438                for _ in 0..resp_len {
439                    responses.push(ComposedResponse::deserialize_from_narg(buf)?);
440                }
441                Ok(ComposedResponse::Or(challenges, responses))
442            }
443            TAG_THRESHOLD => {
444                let ch_len = read_u32(buf)? as usize;
445                let mut challenges = Vec::new();
446                for _ in 0..ch_len {
447                    challenges.push(G::Scalar::deserialize_from_narg(buf)?);
448                }
449                let resp_len = read_u32(buf)? as usize;
450                let mut responses = Vec::new();
451                for _ in 0..resp_len {
452                    responses.push(ComposedResponse::deserialize_from_narg(buf)?);
453                }
454                Ok(ComposedResponse::Threshold(challenges, responses))
455            }
456            _ => Err(VerificationError),
457        }
458    }
459}
460
461impl<G> ComposedResponse<G>
462where
463    G: PrimeGroup
464        + ConditionallySelectable
465        + Encoding<[u8]>
466        + NargSerialize
467        + NargDeserialize
468        + MultiScalarMul,
469    G::Scalar:
470        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
471{
472    /// Conditionally select between two ComposedResponse values.
473    /// This function performs constant-time selection of the response values.
474    pub fn conditional_select(a: &Self, b: &Self, choice: Choice) -> Self {
475        match (a, b) {
476            (ComposedResponse::Simple(a_scalars), ComposedResponse::Simple(b_scalars)) => {
477                let selected: Vec<G::Scalar> = a_scalars
478                    .iter()
479                    .zip_eq(b_scalars.iter())
480                    .map(|(a, b)| G::Scalar::conditional_select(a, b, choice))
481                    .collect();
482                ComposedResponse::Simple(selected)
483            }
484            (ComposedResponse::And(a_responses), ComposedResponse::And(b_responses)) => {
485                let selected: Vec<ComposedResponse<G>> = a_responses
486                    .iter()
487                    .zip_eq(b_responses.iter())
488                    .map(|(a, b)| ComposedResponse::conditional_select(a, b, choice))
489                    .collect();
490                ComposedResponse::And(selected)
491            }
492            (
493                ComposedResponse::Or(a_challenges, a_responses),
494                ComposedResponse::Or(b_challenges, b_responses),
495            ) => {
496                let selected_challenges: Vec<ComposedChallenge<G>> = a_challenges
497                    .iter()
498                    .zip_eq(b_challenges.iter())
499                    .map(|(a, b)| G::Scalar::conditional_select(a, b, choice))
500                    .collect();
501
502                let selected_responses: Vec<ComposedResponse<G>> = a_responses
503                    .iter()
504                    .zip_eq(b_responses.iter())
505                    .map(|(a, b)| ComposedResponse::conditional_select(a, b, choice))
506                    .collect();
507
508                ComposedResponse::Or(selected_challenges, selected_responses)
509            }
510            (
511                ComposedResponse::Threshold(a_challenges, a_responses),
512                ComposedResponse::Threshold(b_challenges, b_responses),
513            ) => {
514                let selected_challenges: Vec<ComposedChallenge<G>> = a_challenges
515                    .iter()
516                    .zip_eq(b_challenges.iter())
517                    .map(|(a, b)| G::Scalar::conditional_select(a, b, choice))
518                    .collect();
519
520                let selected_responses: Vec<ComposedResponse<G>> = a_responses
521                    .iter()
522                    .zip_eq(b_responses.iter())
523                    .map(|(a, b)| ComposedResponse::conditional_select(a, b, choice))
524                    .collect();
525
526                ComposedResponse::Threshold(selected_challenges, selected_responses)
527            }
528            _ => {
529                unreachable!("Mismatched ComposedResponse variants in conditional_select");
530            }
531        }
532    }
533}
534
535// Structure representing the Witness type of Protocol as SigmaProtocol
536#[derive(Clone)]
537pub enum ComposedWitness<G>
538where
539    G: PrimeGroup + Encoding<[u8]> + NargSerialize + NargDeserialize + MultiScalarMul,
540    G::Scalar: Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]>,
541{
542    Simple(<CanonicalLinearRelation<G> as SigmaProtocol>::Witness),
543    And(Vec<ComposedWitness<G>>),
544    Or(Vec<ComposedWitness<G>>),
545    Threshold(Vec<ComposedWitness<G>>),
546}
547
548impl<G> ComposedWitness<G>
549where
550    G: PrimeGroup + Encoding<[u8]> + NargSerialize + NargDeserialize + MultiScalarMul,
551    G::Scalar: Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]>,
552{
553    /// Create a [ComposedWitness] for an AND relation from the given list of witnesses.
554    pub fn and<T: Into<ComposedWitness<G>>>(witness: impl IntoIterator<Item = T>) -> Self {
555        Self::And(witness.into_iter().map(|x| x.into()).collect())
556    }
557
558    /// Create a [ComposedWitness] for an OR relation from the given list of witnesses.
559    pub fn or<T: Into<ComposedWitness<G>>>(witness: impl IntoIterator<Item = T>) -> Self {
560        Self::Or(witness.into_iter().map(|x| x.into()).collect())
561    }
562
563    /// Create a [ComposedWitness] for a threshold relation from the given list of witnesses.
564    pub fn threshold<T: Into<ComposedWitness<G>>>(witness: impl IntoIterator<Item = T>) -> Self {
565        Self::Threshold(witness.into_iter().map(|x| x.into()).collect())
566    }
567}
568
569impl<G> From<<CanonicalLinearRelation<G> as SigmaProtocol>::Witness> for ComposedWitness<G>
570where
571    G: PrimeGroup + Encoding<[u8]> + NargSerialize + NargDeserialize + MultiScalarMul,
572    G::Scalar:
573        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
574{
575    fn from(value: <CanonicalLinearRelation<G> as SigmaProtocol>::Witness) -> Self {
576        Self::Simple(value)
577    }
578}
579
580type ComposedChallenge<G> = <CanonicalLinearRelation<G> as SigmaProtocol>::Challenge;
581
582fn threshold_x<F: PrimeField>(index: usize) -> F {
583    F::from((index + 1) as u64)
584}
585
586fn poly_mul_linear<F: Field>(coeffs: &[F], constant: F) -> Vec<F> {
587    let mut out = vec![F::ZERO; coeffs.len() + 1];
588    for (i, coeff) in coeffs.iter().enumerate() {
589        out[i] += *coeff * constant;
590        out[i + 1] += *coeff;
591    }
592    out
593}
594
595fn interpolate_polynomial<F: Field>(points: &[Evaluation<F>]) -> Result<Vec<F>, Error> {
596    if points.is_empty() {
597        return Err(Error::InvalidInstanceWitnessPair);
598    }
599
600    let mut coeffs = vec![F::ZERO; points.len()];
601
602    for (i, point_i) in points.iter().enumerate() {
603        let mut basis = vec![F::ONE];
604        let mut denom = F::ONE;
605
606        for (j, point_j) in points.iter().enumerate() {
607            if i == j {
608                continue;
609            }
610            denom *= point_i.x - point_j.x;
611            basis = poly_mul_linear::<F>(&basis, -point_j.x);
612        }
613
614        let denom_inv = denom.invert();
615        if denom_inv.is_none().into() {
616            return Err(Error::InvalidInstanceWitnessPair);
617        }
618        let scale = point_i.y * denom_inv.unwrap_or(F::ZERO);
619        for (coeff, basis_coeff) in coeffs.iter_mut().zip_eq(basis.iter()) {
620            *coeff += *basis_coeff * scale;
621        }
622    }
623
624    Ok(coeffs)
625}
626
627fn evaluate_polynomial<F: Field>(coeffs: &[F], x: F) -> F {
628    coeffs
629        .iter()
630        .rev()
631        .fold(F::ZERO, |acc, coeff| acc * x + coeff)
632}
633
634fn expand_threshold_challenges<F: PrimeField>(
635    threshold: usize,
636    total: usize,
637    challenge: F,
638    compressed_challenges: &[F],
639) -> Result<Vec<F>, Error> {
640    if threshold == 0 || threshold > total {
641        return Err(Error::InvalidInstanceWitnessPair);
642    }
643
644    let degree = total - threshold;
645    if compressed_challenges.len() != degree {
646        return Err(Error::InvalidInstanceWitnessPair);
647    }
648
649    let mut points = Vec::with_capacity(degree + 1);
650    points.push(Evaluation {
651        x: F::ZERO,
652        y: challenge,
653    });
654    for (index, share) in compressed_challenges.iter().enumerate() {
655        points.push(Evaluation {
656            x: threshold_x::<F>(index),
657            y: *share,
658        });
659    }
660
661    let coeffs = interpolate_polynomial::<F>(&points)?;
662    let mut challenges = Vec::with_capacity(total);
663    for index in 0..total {
664        challenges.push(evaluate_polynomial::<F>(&coeffs, threshold_x::<F>(index)));
665    }
666
667    Ok(challenges)
668}
669
670fn count_choices(choices: &[Choice]) -> usize {
671    let mut sum: u32 = 0;
672    for choice in choices {
673        let inc = sum.wrapping_add(1);
674        sum = u32::conditional_select(&sum, &inc, *choice);
675    }
676    sum as usize
677}
678
679#[derive(Clone, Copy)]
680struct Evaluation<T> {
681    x: T,
682    y: T,
683}
684
685impl<T: ConditionallySelectable> ConditionallySelectable for Evaluation<T> {
686    fn conditional_select(a: &Self, b: &Self, choice: Choice) -> Self {
687        Evaluation {
688            x: T::conditional_select(&a.x, &b.x, choice),
689            y: T::conditional_select(&a.y, &b.y, choice),
690        }
691    }
692}
693
694impl<T> From<(T, T)> for Evaluation<T> {
695    fn from(value: (T, T)) -> Self {
696        Evaluation {
697            x: value.0,
698            y: value.1,
699        }
700    }
701}
702
703fn conditional_swap_point<T: ConditionallySelectable>(
704    points: &mut [T],
705    left: usize,
706    right: usize,
707    swap: Choice,
708) {
709    if left == right {
710        return;
711    }
712    if left < right {
713        let (head, tail) = points.split_at_mut(right);
714        T::conditional_swap(&mut head[left], &mut tail[0], swap);
715    } else {
716        let (head, tail) = points.split_at_mut(left);
717        T::conditional_swap(&mut tail[0], &mut head[right], swap);
718    }
719}
720
721fn oroffcompact_points<T: ConditionallySelectable>(
722    points: &mut [T],
723    marks: &[Choice],
724    offset: usize,
725) {
726    let n = points.len();
727    if n <= 1 {
728        return;
729    }
730    debug_assert_eq!(n, marks.len());
731    debug_assert!(n.is_power_of_two());
732
733    let half = n / 2;
734    let mut m = 0usize;
735    for mark in &marks[..half] {
736        m += mark.unwrap_u8() as usize;
737    }
738
739    if n == 2 {
740        let z = Choice::from((offset & 1) as u8);
741        let b = ((!marks[0]) & marks[1]) ^ z;
742        conditional_swap_point(points, 0, 1, b);
743        return;
744    }
745
746    let offset_mod = offset % half;
747    oroffcompact_points(&mut points[..half], &marks[..half], offset_mod);
748    let offset_plus_m_mod = (offset + m) % half;
749    oroffcompact_points(&mut points[half..], &marks[half..], offset_plus_m_mod);
750
751    let s = Choice::from(((offset_mod + m) >= half) as u8) ^ Choice::from((offset >= half) as u8);
752    for i in 0..half {
753        let b = s ^ Choice::from((i >= offset_plus_m_mod) as u8);
754        conditional_swap_point(points, i, i + half, b);
755    }
756}
757
758fn oblivious_compact_points<T: ConditionallySelectable>(points: &mut [T], marks: &[Choice]) {
759    let n = points.len();
760    if n == 0 {
761        return;
762    }
763    debug_assert_eq!(n, marks.len());
764
765    let n1 = 1usize << (usize::BITS as usize - 1 - n.leading_zeros() as usize);
766    let n2 = n - n1;
767    let mut m = 0usize;
768    for mark in &marks[..n2] {
769        m += mark.unwrap_u8() as usize;
770    }
771
772    if n2 > 0 {
773        oblivious_compact_points(&mut points[..n2], &marks[..n2]);
774    }
775    oroffcompact_points(&mut points[n2..], &marks[n2..], (n1 - n2 + m) % n1);
776
777    for i in 0..n2 {
778        let b = Choice::from((i >= m) as u8);
779        conditional_swap_point(points, i, i + n1, b);
780    }
781}
782
783impl<G> ComposedRelation<G>
784where
785    G: PrimeGroup
786        + ConstantTimeEq
787        + ConditionallySelectable
788        + Encoding<[u8]>
789        + NargSerialize
790        + NargDeserialize
791        + MultiScalarMul,
792    G::Scalar:
793        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
794{
795    fn is_witness_valid(&self, witness: &ComposedWitness<G>) -> Choice {
796        match (self, witness) {
797            (ComposedRelation::Simple(instance), ComposedWitness::Simple(witness)) => {
798                instance.is_witness_valid(witness)
799            }
800            (ComposedRelation::And(instances), ComposedWitness::And(witnesses)) => {
801                if instances.len() != witnesses.len() {
802                    return Choice::from(0);
803                }
804                instances
805                    .iter()
806                    .zip_eq(witnesses)
807                    .fold(Choice::from(1), |bit, (instance, witness)| {
808                        bit & instance.is_witness_valid(witness)
809                    })
810            }
811            (ComposedRelation::Or(instances), ComposedWitness::Or(witnesses)) => {
812                if instances.len() != witnesses.len() {
813                    return Choice::from(0);
814                }
815                instances
816                    .iter()
817                    .zip_eq(witnesses)
818                    .fold(Choice::from(0), |bit, (instance, witness)| {
819                        bit | instance.is_witness_valid(witness)
820                    })
821            }
822            (
823                ComposedRelation::Threshold(threshold, instances),
824                ComposedWitness::Threshold(witnesses),
825            ) => {
826                if *threshold == 0 || instances.len() != witnesses.len() {
827                    return Choice::from(0);
828                }
829                let mut count = 0usize;
830                for (instance, witness) in instances.iter().zip_eq(witnesses) {
831                    if instance.is_witness_valid(witness).unwrap_u8() == 1 {
832                        count += 1;
833                    }
834                }
835                Choice::from((count >= *threshold) as u8)
836            }
837            _ => Choice::from(0),
838        }
839    }
840
841    fn prover_commit_simple(
842        protocol: &CanonicalLinearRelation<G>,
843        witness: &<CanonicalLinearRelation<G> as SigmaProtocol>::Witness,
844        rng: &mut impl ScalarRng,
845    ) -> Result<(ComposedCommitment<G>, ComposedProverState<G>), Error> {
846        protocol.prover_commit(witness, rng).map(|(c, s)| {
847            (
848                ComposedCommitment::Simple(c),
849                ComposedProverState::Simple(s),
850            )
851        })
852    }
853
854    fn prover_response_simple(
855        instance: &CanonicalLinearRelation<G>,
856        state: <CanonicalLinearRelation<G> as SigmaProtocol>::ProverState,
857        challenge: &<CanonicalLinearRelation<G> as SigmaProtocol>::Challenge,
858    ) -> Result<ComposedResponse<G>, Error> {
859        instance
860            .prover_response(state, challenge)
861            .map(ComposedResponse::Simple)
862    }
863
864    fn prover_commit_and(
865        protocols: &[ComposedRelation<G>],
866        witnesses: &[ComposedWitness<G>],
867        rng: &mut impl ScalarRng,
868    ) -> Result<(ComposedCommitment<G>, ComposedProverState<G>), Error> {
869        if protocols.len() != witnesses.len() {
870            return Err(Error::InvalidInstanceWitnessPair);
871        }
872
873        let mut commitments = Vec::with_capacity(protocols.len());
874        let mut prover_states = Vec::with_capacity(protocols.len());
875
876        for (p, w) in protocols.iter().zip_eq(witnesses.iter()) {
877            let (mut c, s) = p.prover_commit(w, rng)?;
878            let commitment = c.pop().ok_or(Error::InvalidInstanceWitnessPair)?;
879            if !c.is_empty() {
880                return Err(Error::InvalidInstanceWitnessPair);
881            }
882            commitments.push(commitment);
883            prover_states.push(s);
884        }
885
886        Ok((
887            ComposedCommitment::And(commitments),
888            ComposedProverState::And(prover_states),
889        ))
890    }
891
892    fn prover_response_and(
893        instances: &[ComposedRelation<G>],
894        prover_state: Vec<ComposedProverState<G>>,
895        challenge: &ComposedChallenge<G>,
896    ) -> Result<ComposedResponse<G>, Error> {
897        if instances.len() != prover_state.len() {
898            return Err(Error::InvalidInstanceWitnessPair);
899        }
900
901        let responses: Result<Vec<_>, _> = instances
902            .iter()
903            .zip_eq(prover_state)
904            .map(|(p, s)| {
905                let mut res = p.prover_response(s, challenge)?;
906                res.pop().ok_or(Error::InvalidInstanceWitnessPair)
907            })
908            .collect();
909
910        Ok(ComposedResponse::And(responses?))
911    }
912
913    fn prover_commit_or(
914        instances: &[ComposedRelation<G>],
915        witnesses: &[ComposedWitness<G>],
916        rng: &mut impl ScalarRng,
917    ) -> Result<(ComposedCommitment<G>, ComposedProverState<G>), Error>
918    where
919        G: ConditionallySelectable,
920    {
921        if instances.is_empty() || instances.len() != witnesses.len() {
922            return Err(Error::InvalidInstanceWitnessPair);
923        }
924
925        let mut commitments = Vec::new();
926        let mut prover_states = Vec::new();
927
928        // Selector value set when the first valid witness is found.
929        let mut valid_witness_found = Choice::from(0);
930        for (i, w) in witnesses.iter().enumerate() {
931            let (mut commitment_vec, prover_state) = instances[i].prover_commit(w, rng)?;
932            let commitment = commitment_vec
933                .pop()
934                .ok_or(Error::InvalidInstanceWitnessPair)?;
935            if !commitment_vec.is_empty() {
936                return Err(Error::InvalidInstanceWitnessPair);
937            }
938
939            let (mut simulated_commitment_vec, simulated_challenge, mut simulated_response_vec) =
940                instances[i].simulate_transcript(rng)?;
941            let simulated_commitment = simulated_commitment_vec
942                .pop()
943                .ok_or(Error::InvalidInstanceWitnessPair)?;
944            if !simulated_commitment_vec.is_empty() {
945                return Err(Error::InvalidInstanceWitnessPair);
946            }
947            let simulated_response = simulated_response_vec
948                .pop()
949                .ok_or(Error::InvalidInstanceWitnessPair)?;
950            if !simulated_response_vec.is_empty() {
951                return Err(Error::InvalidInstanceWitnessPair);
952            }
953
954            let valid_witness = instances[i].is_witness_valid(w) & !valid_witness_found;
955            let select_witness = valid_witness;
956
957            let commitment = ComposedCommitment::conditional_select(
958                &simulated_commitment,
959                &commitment,
960                select_witness,
961            );
962
963            commitments.push(commitment);
964            prover_states.push(ComposedOrProverStateEntry(
965                select_witness,
966                prover_state,
967                simulated_challenge,
968                simulated_response,
969            ));
970
971            valid_witness_found |= valid_witness;
972        }
973
974        if valid_witness_found.unwrap_u8() == 0 {
975            Err(Error::InvalidInstanceWitnessPair)
976        } else {
977            Ok((
978                ComposedCommitment::Or(commitments),
979                ComposedProverState::Or(prover_states),
980            ))
981        }
982    }
983
984    fn prover_response_or(
985        instances: &[ComposedRelation<G>],
986        prover_state: ComposedOrProverState<G>,
987        challenge: &ComposedChallenge<G>,
988    ) -> Result<ComposedResponse<G>, Error> {
989        let mut result_challenges = Vec::with_capacity(instances.len());
990        let mut result_responses = Vec::with_capacity(instances.len());
991
992        if instances.is_empty() || instances.len() != prover_state.len() {
993            return Err(Error::InvalidInstanceWitnessPair);
994        }
995
996        let mut witness_challenge = *challenge;
997        for ComposedOrProverStateEntry(
998            valid_witness,
999            _prover_state,
1000            simulated_challenge,
1001            _simulated_response,
1002        ) in &prover_state
1003        {
1004            let c = G::Scalar::conditional_select(
1005                simulated_challenge,
1006                &G::Scalar::ZERO,
1007                *valid_witness,
1008            );
1009            witness_challenge -= c;
1010        }
1011        for (
1012            instance,
1013            ComposedOrProverStateEntry(
1014                valid_witness,
1015                prover_state,
1016                simulated_challenge,
1017                simulated_response,
1018            ),
1019        ) in instances.iter().zip_eq(prover_state)
1020        {
1021            let challenge_i = G::Scalar::conditional_select(
1022                &simulated_challenge,
1023                &witness_challenge,
1024                valid_witness,
1025            );
1026
1027            let mut response_vec = instance.prover_response(prover_state, &challenge_i)?;
1028            let response = response_vec
1029                .pop()
1030                .ok_or(Error::InvalidInstanceWitnessPair)?;
1031            if !response_vec.is_empty() {
1032                return Err(Error::InvalidInstanceWitnessPair);
1033            }
1034            let response =
1035                ComposedResponse::conditional_select(&simulated_response, &response, valid_witness);
1036
1037            result_challenges.push(challenge_i);
1038            result_responses.push(response.clone());
1039        }
1040
1041        result_challenges.pop();
1042        Ok(ComposedResponse::Or(result_challenges, result_responses))
1043    }
1044
1045    fn prover_commit_threshold(
1046        threshold: usize,
1047        instances: &[ComposedRelation<G>],
1048        witnesses: &[ComposedWitness<G>],
1049        rng: &mut impl ScalarRng,
1050    ) -> Result<(ComposedCommitment<G>, ComposedProverState<G>), Error>
1051    where
1052        G: ConditionallySelectable,
1053    {
1054        if instances.len() != witnesses.len() || threshold == 0 || threshold > instances.len() {
1055            return Err(Error::InvalidInstanceWitnessPair);
1056        }
1057        let degree = instances.len() - threshold;
1058
1059        let valid_witnesses = instances
1060            .iter()
1061            .zip_eq(witnesses.iter())
1062            .map(|(x, w)| x.is_witness_valid(w))
1063            .collect::<Vec<Choice>>();
1064
1065        // Degree-(t-1) interpolation can only satisfy t fixed points.
1066        let invalid_count = instances.len() - count_choices(&valid_witnesses);
1067        if invalid_count > degree {
1068            return Err(Error::InvalidInstanceWitnessPair);
1069        }
1070
1071        let mut remaining_seeds = (degree - invalid_count) as u32;
1072        let mut commitments = Vec::with_capacity(instances.len());
1073        let mut prover_states = Vec::with_capacity(instances.len());
1074        for (i, (instance, witness)) in instances.iter().zip_eq(witnesses.iter()).enumerate() {
1075            let (mut commitment_vec, prover_state) = instance.prover_commit(witness, rng)?;
1076            let commitment = commitment_vec
1077                .pop()
1078                .ok_or(Error::InvalidInstanceWitnessPair)?;
1079            if !commitment_vec.is_empty() {
1080                return Err(Error::InvalidInstanceWitnessPair);
1081            }
1082
1083            let (mut simulated_commitments, simulated_challenge, mut simulated_responses) =
1084                instance.simulate_transcript(rng)?;
1085            let simulated_commitment = simulated_commitments
1086                .pop()
1087                .ok_or(Error::InvalidInstanceWitnessPair)?;
1088            if !simulated_commitments.is_empty() {
1089                return Err(Error::InvalidInstanceWitnessPair);
1090            }
1091            let simulated_response = simulated_responses
1092                .pop()
1093                .ok_or(Error::InvalidInstanceWitnessPair)?;
1094            if !simulated_responses.is_empty() {
1095                return Err(Error::InvalidInstanceWitnessPair);
1096            }
1097
1098            let valid_witness = valid_witnesses[i];
1099            let should_seed = valid_witness & Choice::from((remaining_seeds != 0) as u8);
1100            remaining_seeds = remaining_seeds.wrapping_sub(should_seed.unwrap_u8() as u32);
1101            let use_simulator = (!valid_witness) | should_seed;
1102            let commitment = ComposedCommitment::conditional_select(
1103                &commitment,
1104                &simulated_commitment,
1105                use_simulator,
1106            );
1107            commitments.push(commitment);
1108            prover_states.push(ComposedThresholdProverStateEntry {
1109                use_simulator,
1110                prover_state,
1111                simulated_challenge,
1112                simulated_response,
1113            });
1114        }
1115
1116        Ok((
1117            ComposedCommitment::Threshold(commitments),
1118            ComposedProverState::Threshold(prover_states),
1119        ))
1120    }
1121
1122    fn prover_response_threshold(
1123        threshold: usize,
1124        instances: &[ComposedRelation<G>],
1125        prover_states: ComposedThresholdProverState<G>,
1126        challenge: &ComposedChallenge<G>,
1127    ) -> Result<ComposedResponse<G>, Error> {
1128        if threshold == 0 || threshold > instances.len() || instances.len() != prover_states.len() {
1129            return Err(Error::InvalidInstanceWitnessPair);
1130        }
1131        let degree = instances.len() - threshold;
1132
1133        let marks = prover_states
1134            .iter()
1135            .map(|entry| entry.use_simulator)
1136            .collect::<Vec<_>>();
1137        debug_assert_eq!(count_choices(&marks), degree);
1138
1139        let mut points = prover_states
1140            .iter()
1141            .enumerate()
1142            .map(|(i, entry)| Evaluation {
1143                x: threshold_x::<G::Scalar>(i),
1144                y: entry.simulated_challenge,
1145            })
1146            .collect::<Vec<Evaluation<G::Scalar>>>();
1147        oblivious_compact_points(&mut points, &marks);
1148        points.drain(degree..);
1149
1150        let mut full_points = Vec::with_capacity(degree + 1);
1151        full_points.push(Evaluation {
1152            x: G::Scalar::ZERO,
1153            y: *challenge,
1154        });
1155        full_points.extend_from_slice(&points);
1156
1157        let coeffs = interpolate_polynomial::<G::Scalar>(&full_points)?;
1158        let mut compressed_challenges = Vec::with_capacity(degree);
1159        for index in 0..degree {
1160            compressed_challenges.push(evaluate_polynomial::<G::Scalar>(
1161                &coeffs,
1162                threshold_x::<G::Scalar>(index),
1163            ));
1164        }
1165
1166        let expanded_challenges = expand_threshold_challenges::<G::Scalar>(
1167            threshold,
1168            instances.len(),
1169            *challenge,
1170            &compressed_challenges,
1171        )?;
1172
1173        let mut responses = Vec::with_capacity(instances.len());
1174
1175        for (i, (instance, prover_state)) in instances.iter().zip_eq(prover_states).enumerate() {
1176            let poly_challenge = expanded_challenges[i];
1177            let challenge = G::Scalar::conditional_select(
1178                &poly_challenge,
1179                &prover_state.simulated_challenge,
1180                prover_state.use_simulator,
1181            );
1182
1183            let mut response_vec =
1184                instance.prover_response(prover_state.prover_state, &challenge)?;
1185            let response = response_vec
1186                .pop()
1187                .ok_or(Error::InvalidInstanceWitnessPair)?;
1188            if !response_vec.is_empty() {
1189                return Err(Error::InvalidInstanceWitnessPair);
1190            }
1191            let response = ComposedResponse::conditional_select(
1192                &response,
1193                &prover_state.simulated_response,
1194                prover_state.use_simulator,
1195            );
1196
1197            responses.push(response);
1198        }
1199
1200        Ok(ComposedResponse::Threshold(
1201            compressed_challenges,
1202            responses,
1203        ))
1204    }
1205}
1206
1207impl<G> SigmaProtocol for ComposedRelation<G>
1208where
1209    G: PrimeGroup
1210        + ConstantTimeEq
1211        + ConditionallySelectable
1212        + Encoding<[u8]>
1213        + NargSerialize
1214        + NargDeserialize
1215        + MultiScalarMul,
1216    G::Scalar:
1217        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
1218{
1219    type Commitment = ComposedCommitment<G>;
1220    type ProverState = ComposedProverState<G>;
1221    type Response = ComposedResponse<G>;
1222    type Witness = ComposedWitness<G>;
1223    type Challenge = ComposedChallenge<G>;
1224
1225    fn prover_commit(
1226        &self,
1227        witness: &Self::Witness,
1228        rng: &mut impl ScalarRng,
1229    ) -> Result<(Vec<Self::Commitment>, Self::ProverState), Error> {
1230        let (commitment, state) = match (self, witness) {
1231            (ComposedRelation::Simple(p), ComposedWitness::Simple(w)) => {
1232                Self::prover_commit_simple(p, w, rng)
1233            }
1234            (ComposedRelation::And(ps), ComposedWitness::And(ws)) => {
1235                Self::prover_commit_and(ps, ws, rng)
1236            }
1237            (ComposedRelation::Or(ps), ComposedWitness::Or(witnesses)) => {
1238                Self::prover_commit_or(ps, witnesses, rng)
1239            }
1240            (ComposedRelation::Threshold(threshold, ps), ComposedWitness::Threshold(witnesses)) => {
1241                Self::prover_commit_threshold(*threshold, ps, witnesses, rng)
1242            }
1243            _ => Err(Error::InvalidInstanceWitnessPair),
1244        }?;
1245        Ok((vec![commitment], state))
1246    }
1247
1248    fn prover_response(
1249        &self,
1250        state: Self::ProverState,
1251        challenge: &Self::Challenge,
1252    ) -> Result<Vec<Self::Response>, Error> {
1253        let response = match (self, state) {
1254            (ComposedRelation::Simple(instance), ComposedProverState::Simple(state)) => {
1255                Self::prover_response_simple(instance, state, challenge)
1256            }
1257            (ComposedRelation::And(instances), ComposedProverState::And(prover_state)) => {
1258                Self::prover_response_and(instances, prover_state, challenge)
1259            }
1260            (ComposedRelation::Or(instances), ComposedProverState::Or(prover_state)) => {
1261                Self::prover_response_or(instances, prover_state, challenge)
1262            }
1263            (
1264                ComposedRelation::Threshold(threshold, instances),
1265                ComposedProverState::Threshold(prover_state),
1266            ) => Self::prover_response_threshold(*threshold, instances, prover_state, challenge),
1267            _ => Err(Error::InvalidInstanceWitnessPair),
1268        }?;
1269        Ok(vec![response])
1270    }
1271
1272    fn verifier(
1273        &self,
1274        commitment: &[Self::Commitment],
1275        challenge: &Self::Challenge,
1276        response: &[Self::Response],
1277    ) -> Result<(), Error> {
1278        let (commitment, response) = match (commitment.first(), response.first()) {
1279            (Some(c), Some(r)) => (c, r),
1280            _ => return Err(Error::InvalidInstanceWitnessPair),
1281        };
1282
1283        match (self, commitment, response) {
1284            (
1285                ComposedRelation::Simple(p),
1286                ComposedCommitment::Simple(c),
1287                ComposedResponse::Simple(r),
1288            ) => p.verifier(c, challenge, r),
1289            (
1290                ComposedRelation::And(ps),
1291                ComposedCommitment::And(commitments),
1292                ComposedResponse::And(responses),
1293            ) => {
1294                if ps.len() != commitments.len() || commitments.len() != responses.len() {
1295                    return Err(Error::InvalidInstanceWitnessPair);
1296                }
1297                ps.iter()
1298                    .zip_eq(commitments)
1299                    .zip_eq(responses)
1300                    .try_for_each(|((p, c), r)| {
1301                        p.verifier(
1302                            core::slice::from_ref(c),
1303                            challenge,
1304                            core::slice::from_ref(r),
1305                        )
1306                    })
1307            }
1308            (
1309                ComposedRelation::Or(ps),
1310                ComposedCommitment::Or(commitments),
1311                ComposedResponse::Or(challenges, responses),
1312            ) => {
1313                if ps.len() != commitments.len()
1314                    || commitments.len() != responses.len()
1315                    || ps.len().checked_sub(1) != Some(challenges.len())
1316                {
1317                    return Err(Error::InvalidInstanceWitnessPair);
1318                }
1319                let last_challenge = *challenge - challenges.iter().sum::<G::Scalar>();
1320                ps.iter()
1321                    .zip_eq(commitments)
1322                    .zip_eq(challenges.iter().chain(&Some(last_challenge)))
1323                    .zip_eq(responses)
1324                    .try_for_each(|(((p, commitment), challenge), response)| {
1325                        p.verifier(
1326                            core::slice::from_ref(commitment),
1327                            challenge,
1328                            core::slice::from_ref(response),
1329                        )
1330                    })
1331            }
1332            (
1333                ComposedRelation::Threshold(threshold, ps),
1334                ComposedCommitment::Threshold(commitments),
1335                ComposedResponse::Threshold(challenges, responses),
1336            ) => {
1337                if *threshold == 0
1338                    || *threshold > ps.len()
1339                    || commitments.len() != ps.len()
1340                    || challenges.len() != ps.len() - *threshold
1341                    || responses.len() != ps.len()
1342                {
1343                    return Err(Error::InvalidInstanceWitnessPair);
1344                }
1345
1346                let full_challenges = expand_threshold_challenges::<G::Scalar>(
1347                    *threshold,
1348                    ps.len(),
1349                    *challenge,
1350                    challenges,
1351                )?;
1352
1353                ps.iter()
1354                    .zip_eq(commitments)
1355                    .zip_eq(full_challenges.iter())
1356                    .zip_eq(responses)
1357                    .try_for_each(|(((p, commitment), challenge), response)| {
1358                        p.verifier(
1359                            core::slice::from_ref(commitment),
1360                            challenge,
1361                            core::slice::from_ref(response),
1362                        )
1363                    })
1364            }
1365            _ => Err(Error::InvalidInstanceWitnessPair),
1366        }
1367    }
1368
1369    fn commitment_len(&self) -> usize {
1370        1
1371    }
1372
1373    fn response_len(&self) -> usize {
1374        1
1375    }
1376
1377    fn instance_label(&self) -> impl AsRef<[u8]> {
1378        match self {
1379            ComposedRelation::Simple(p) => {
1380                let label = p.instance_label();
1381                label.as_ref().to_vec()
1382            }
1383            ComposedRelation::And(ps) => {
1384                let mut bytes = Vec::new();
1385                for p in ps {
1386                    bytes.extend(p.instance_label().as_ref());
1387                }
1388                bytes
1389            }
1390            ComposedRelation::Or(ps) => {
1391                let mut bytes = Vec::new();
1392                for p in ps {
1393                    bytes.extend(p.instance_label().as_ref());
1394                }
1395                bytes
1396            }
1397            ComposedRelation::Threshold(threshold, ps) => {
1398                let mut bytes = Vec::new();
1399                bytes.extend_from_slice(&((*threshold as u64).to_le_bytes()));
1400                for p in ps {
1401                    bytes.extend(p.instance_label().as_ref());
1402                }
1403                bytes
1404            }
1405        }
1406    }
1407
1408    fn protocol_identifier(&self) -> [u8; 64] {
1409        let mut hasher = Sha3_256::new();
1410
1411        match self {
1412            ComposedRelation::Simple(p) => {
1413                // take the digest of the simple protocol id
1414                hasher.update([0u8; 32]);
1415                hasher.update(p.protocol_identifier());
1416            }
1417            ComposedRelation::And(protocols) => {
1418                hasher.update([1u8; 32]);
1419                for p in protocols {
1420                    hasher.update(p.protocol_identifier().as_ref());
1421                }
1422            }
1423            ComposedRelation::Or(protocols) => {
1424                hasher.update([2u8; 32]);
1425                for p in protocols {
1426                    hasher.update(p.protocol_identifier().as_ref());
1427                }
1428            }
1429            ComposedRelation::Threshold(threshold, protocols) => {
1430                hasher.update([3u8; 32]);
1431                hasher.update(((*threshold as u64).to_le_bytes()).as_ref());
1432                for p in protocols {
1433                    hasher.update(p.protocol_identifier().as_ref());
1434                }
1435            }
1436        }
1437
1438        let mut protocol_id = [0u8; 64];
1439        protocol_id[..32].clone_from_slice(&hasher.finalize());
1440        protocol_id
1441    }
1442}
1443
1444impl<G> SigmaProtocolSimulator for ComposedRelation<G>
1445where
1446    G: PrimeGroup
1447        + ConstantTimeEq
1448        + ConditionallySelectable
1449        + Encoding<[u8]>
1450        + NargSerialize
1451        + NargDeserialize
1452        + MultiScalarMul,
1453    G::Scalar:
1454        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
1455{
1456    fn simulate_commitment(
1457        &self,
1458        challenge: &Self::Challenge,
1459        response: &[Self::Response],
1460    ) -> Result<Vec<Self::Commitment>, Error> {
1461        let response = response.first().ok_or(Error::InvalidInstanceWitnessPair)?;
1462        let commitment = match (self, response) {
1463            (ComposedRelation::Simple(p), ComposedResponse::Simple(r)) => {
1464                ComposedCommitment::Simple(p.simulate_commitment(challenge, r)?)
1465            }
1466            (ComposedRelation::And(ps), ComposedResponse::And(rs)) => {
1467                if ps.len() != rs.len() {
1468                    return Err(Error::InvalidInstanceWitnessPair);
1469                }
1470                let commitments = ps
1471                    .iter()
1472                    .zip_eq(rs)
1473                    .map(|(p, r)| {
1474                        p.simulate_commitment(challenge, core::slice::from_ref(r))
1475                            .and_then(|mut c| c.pop().ok_or(Error::InvalidInstanceWitnessPair))
1476                    })
1477                    .collect::<Result<Vec<_>, _>>()?;
1478                ComposedCommitment::And(commitments)
1479            }
1480            (ComposedRelation::Or(ps), ComposedResponse::Or(challenges, rs)) => {
1481                if rs.len() != ps.len() || ps.len().checked_sub(1) != Some(challenges.len()) {
1482                    return Err(Error::InvalidInstanceWitnessPair);
1483                }
1484                let last_challenge = *challenge - challenges.iter().sum::<G::Scalar>();
1485                let commitments = ps
1486                    .iter()
1487                    .zip_eq(challenges.iter().chain(&Some(last_challenge)))
1488                    .zip_eq(rs)
1489                    .map(|((p, ch), r)| {
1490                        p.simulate_commitment(ch, core::slice::from_ref(r))
1491                            .and_then(|mut c| c.pop().ok_or(Error::InvalidInstanceWitnessPair))
1492                    })
1493                    .collect::<Result<Vec<_>, _>>()?;
1494                ComposedCommitment::Or(commitments)
1495            }
1496            (
1497                ComposedRelation::Threshold(threshold, ps),
1498                ComposedResponse::Threshold(challenges, rs),
1499            ) => {
1500                if rs.len() != ps.len()
1501                    || ps.len() < *threshold
1502                    || challenges.len() != ps.len() - threshold
1503                {
1504                    return Err(Error::InvalidInstanceWitnessPair);
1505                }
1506
1507                let full_challenges = expand_threshold_challenges::<G::Scalar>(
1508                    *threshold,
1509                    ps.len(),
1510                    *challenge,
1511                    challenges,
1512                )?;
1513                let commitments = ps
1514                    .iter()
1515                    .zip_eq(full_challenges.iter())
1516                    .zip_eq(rs)
1517                    .map(|((p, ch), r)| {
1518                        p.simulate_commitment(ch, core::slice::from_ref(r))
1519                            .and_then(|mut c| c.pop().ok_or(Error::InvalidInstanceWitnessPair))
1520                    })
1521                    .collect::<Result<Vec<_>, _>>()?;
1522                ComposedCommitment::Threshold(commitments)
1523            }
1524            _ => return Err(Error::InvalidInstanceWitnessPair),
1525        };
1526
1527        Ok(vec![commitment])
1528    }
1529
1530    fn simulate_response(&self, rng: &mut impl ScalarRng) -> Vec<Self::Response> {
1531        let response = match self {
1532            ComposedRelation::Simple(p) => ComposedResponse::Simple(p.simulate_response(rng)),
1533            ComposedRelation::And(ps) => {
1534                let responses = ps
1535                    .iter()
1536                    .map(|p| {
1537                        let mut r = p.simulate_response(rng);
1538                        r.pop().ok_or(Error::InvalidInstanceWitnessPair)
1539                    })
1540                    .collect::<Result<Vec<_>, _>>()
1541                    .expect("simulate_response invariant");
1542                ComposedResponse::And(responses)
1543            }
1544            ComposedRelation::Or(ps) => {
1545                let challenge_count = ps.len().saturating_sub(1);
1546                let challenges = rng.random_scalars_vec::<G>(challenge_count).to_vec();
1547                let mut responses = Vec::with_capacity(ps.len());
1548                for p in ps.iter() {
1549                    let mut r = p.simulate_response(&mut *rng);
1550                    let resp = r
1551                        .pop()
1552                        .expect("simulate_response should return at least one element");
1553                    responses.push(resp);
1554                }
1555                ComposedResponse::Or(challenges, responses)
1556            }
1557            ComposedRelation::Threshold(threshold, ps) => {
1558                if *threshold == 0 || *threshold > ps.len() {
1559                    return vec![ComposedResponse::Threshold(Vec::new(), Vec::new())];
1560                }
1561
1562                let degree = ps.len() - *threshold;
1563                let compressed_challenges = rng.random_scalars_vec::<G>(degree).to_vec();
1564                let mut responses = Vec::with_capacity(ps.len());
1565                for p in ps.iter() {
1566                    let mut r = p.simulate_response(&mut *rng);
1567                    let response = r
1568                        .pop()
1569                        .expect("simulate_response should return at least one element");
1570                    responses.push(response);
1571                }
1572                ComposedResponse::Threshold(compressed_challenges, responses)
1573            }
1574        };
1575        vec![response]
1576    }
1577
1578    fn simulate_transcript(
1579        &self,
1580        rng: &mut impl ScalarRng,
1581    ) -> Result<(Vec<Self::Commitment>, Self::Challenge, Vec<Self::Response>), Error> {
1582        match self {
1583            ComposedRelation::Simple(p) => {
1584                let (c, ch, r) = p.simulate_transcript(rng)?;
1585                Ok((
1586                    vec![ComposedCommitment::Simple(c)],
1587                    ch,
1588                    vec![ComposedResponse::Simple(r)],
1589                ))
1590            }
1591            ComposedRelation::And(ps) => {
1592                let [challenge] = rng.random_scalars::<G, _>();
1593                let mut responses = Vec::with_capacity(ps.len());
1594                for p in ps.iter() {
1595                    let mut resp = p.simulate_response(&mut *rng);
1596                    let response = resp.pop().ok_or(Error::InvalidInstanceWitnessPair)?;
1597                    if !resp.is_empty() {
1598                        return Err(Error::InvalidInstanceWitnessPair);
1599                    }
1600                    responses.push(response);
1601                }
1602                let commitments = ps
1603                    .iter()
1604                    .enumerate()
1605                    .map(|(i, p)| {
1606                        p.simulate_commitment(&challenge, &[responses[i].clone()])
1607                            .and_then(|mut c| {
1608                                let first = c.pop().ok_or(Error::InvalidInstanceWitnessPair)?;
1609                                if !c.is_empty() {
1610                                    return Err(Error::InvalidInstanceWitnessPair);
1611                                }
1612                                Ok(first)
1613                            })
1614                    })
1615                    .collect::<Result<Vec<_>, Error>>()?;
1616
1617                Ok((
1618                    vec![ComposedCommitment::And(commitments)],
1619                    challenge,
1620                    vec![ComposedResponse::And(responses)],
1621                ))
1622            }
1623            ComposedRelation::Or(ps) => {
1624                let challenge_count = ps
1625                    .len()
1626                    .checked_sub(1)
1627                    .ok_or(Error::InvalidInstanceWitnessPair)?;
1628                let challenges = rng.random_scalars_vec::<G>(challenge_count);
1629                let mut responses = Vec::with_capacity(ps.len());
1630                for p in ps.iter() {
1631                    let mut resp = p.simulate_response(&mut *rng);
1632                    let response = resp.pop().ok_or(Error::InvalidInstanceWitnessPair)?;
1633                    if !resp.is_empty() {
1634                        return Err(Error::InvalidInstanceWitnessPair);
1635                    }
1636                    responses.push(response);
1637                }
1638
1639                let mut commitments = Vec::with_capacity(ps.len());
1640                for i in 0..ps.len() {
1641                    let mut commitment = ps[i].simulate_commitment(
1642                        &if i == challenge_count {
1643                            challenges.iter().fold(G::Scalar::ZERO, |acc, x| acc - x)
1644                        } else {
1645                            challenges[i]
1646                        },
1647                        &[responses[i].clone()],
1648                    )?;
1649                    let commitment = commitment.pop().ok_or(Error::InvalidInstanceWitnessPair)?;
1650                    commitments.push(commitment);
1651                }
1652
1653                Ok((
1654                    vec![ComposedCommitment::Or(commitments)],
1655                    challenges.iter().sum::<G::Scalar>(),
1656                    vec![ComposedResponse::Or(challenges, responses)],
1657                ))
1658            }
1659            ComposedRelation::Threshold(threshold, ps) => {
1660                if *threshold == 0 || *threshold > ps.len() {
1661                    return Err(Error::InvalidInstanceWitnessPair);
1662                }
1663
1664                let degree = ps.len() - *threshold;
1665                let compressed_challenges = rng.random_scalars_vec::<G>(degree);
1666                let mut responses = Vec::with_capacity(ps.len());
1667                for p in ps.iter() {
1668                    let mut resp = p.simulate_response(&mut *rng);
1669                    let response = resp.pop().ok_or(Error::InvalidInstanceWitnessPair)?;
1670                    if !resp.is_empty() {
1671                        return Err(Error::InvalidInstanceWitnessPair);
1672                    }
1673                    responses.push(response);
1674                }
1675
1676                let [challenge] = rng.random_scalars::<G, _>();
1677                let full_challenges = expand_threshold_challenges(
1678                    *threshold,
1679                    ps.len(),
1680                    challenge,
1681                    &compressed_challenges,
1682                )?;
1683                let commitments = ps
1684                    .iter()
1685                    .zip_eq(full_challenges.iter())
1686                    .zip_eq(responses.iter())
1687                    .map(|((p, ch), r)| {
1688                        p.simulate_commitment(ch, core::slice::from_ref(r))
1689                            .and_then(|mut c| {
1690                                let first = c.pop().ok_or(Error::InvalidInstanceWitnessPair)?;
1691                                if !c.is_empty() {
1692                                    return Err(Error::InvalidInstanceWitnessPair);
1693                                }
1694                                Ok(first)
1695                            })
1696                    })
1697                    .collect::<Result<Vec<_>, Error>>()?;
1698                Ok((
1699                    vec![ComposedCommitment::Threshold(commitments)],
1700                    challenge,
1701                    vec![ComposedResponse::Threshold(
1702                        compressed_challenges,
1703                        responses,
1704                    )],
1705                ))
1706            }
1707        }
1708    }
1709}
1710
1711impl<G> ComposedRelation<G>
1712where
1713    G: PrimeGroup
1714        + ConstantTimeEq
1715        + ConditionallySelectable
1716        + Encoding<[u8]>
1717        + NargSerialize
1718        + NargDeserialize
1719        + MultiScalarMul,
1720    G::Scalar:
1721        Encoding<[u8]> + NargSerialize + NargDeserialize + Decoding<[u8]> + ConditionallySelectable,
1722{
1723    /// Convert this Protocol into a non-interactive zero-knowledge proof
1724    /// using the Shake128DuplexSponge codec and a specified session identifier.
1725    ///
1726    /// This method provides a convenient way to create a NIZK from a Protocol
1727    /// without exposing the specific codec type to the API caller.
1728    ///
1729    /// # Parameters
1730    /// - `session_identifier`: Domain separator bytes for the Fiat-Shamir transform
1731    ///
1732    /// # Returns
1733    /// A `Nizk` instance ready for proving and verification
1734    pub fn into_nizk(self, session_identifier: &[u8]) -> Nizk<ComposedRelation<G>> {
1735        Nizk::new(session_identifier, self)
1736    }
1737}