binius_prover/protocols/shift/prove.rs
1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use binius_compute::Allocator;
5use binius_core::word::Word;
6use binius_field::{BinaryField, PackedField};
7use binius_ip_prover::channel::IPProverChannel;
8use binius_math::{BinarySubspace, univariate::EvaluationDomain};
9
10use super::{
11 SegmentWords,
12 claims::OperandClaims,
13 key_collection::KeyCollection,
14 phase_1::prove_phase_1,
15 phase_2::{ShiftOutput, prove_phase_2},
16 shift_ind::{ShiftChallengePoint, ShiftIndSumcheck},
17};
18
19/// Proves the shift protocol reduction, collapsing every operation's claims into one.
20///
21/// The result is a single multilinear evaluation claim on the witness.
22/// It is reached in five prover phases.
23/// A shifted value index names two shifts applied in sequence.
24/// The reduction peels them off from the output end inward:
25///
26/// 1. bind the outer shift slot, then the inner one, then the bit position within a word;
27/// 2. bind the bit index of the intermediate word, where the two shift indicators meet;
28/// 3. bind the output bit index the reduction's first factor attaches to;
29/// 4. reduce what is left to a witness evaluation, against the constraint-matrix multilinear.
30///
31/// # Arguments
32///
33/// - `key_collection`: the prover's key collection for the constraint system.
34/// - `public_words`: the constants followed by the inout values, as the circuit declares them.
35/// - `hidden_words`: the private values, as the circuit declares them.
36/// - `claims`: the operand evaluation claims, all at one constraint point.
37/// - `domain_subspace`: the univariate evaluation domain.
38/// - `channel`: the prover channel the interactive rounds run over.
39/// - `alloc`: the allocator the intermediate buffers are drawn from.
40///
41/// # Returns
42///
43/// The final challenges with the witness evaluation.
44/// Also the wiring multilinear's evaluation, for the caller to send.
45pub fn prove<F, P, Channel, A>(
46 key_collection: &KeyCollection,
47 public_words: &[Word],
48 hidden_words: &[Word],
49 claims: OperandClaims<F>,
50 domain_subspace: &BinarySubspace<F>,
51 channel: &mut Channel,
52 alloc: &A,
53) -> ShiftOutput<F>
54where
55 F: BinaryField,
56 P: PackedField<Scalar = F>,
57 Channel: IPProverChannel<F>,
58 A: Allocator,
59{
60 // The segments are passed as the circuit declares them, at whatever length that is.
61 // Neither phase needs them padded.
62 let words = SegmentWords {
63 public: public_words,
64 hidden: hidden_words,
65 };
66
67 // One batching coefficient per operation, folded into its operand weights, and one expansion
68 // of the constraint point shared by every operation.
69 // SOUNDNESS: this must draw in the same order the verifier draws in.
70 let prepared = {
71 let _scope = tracing::debug_span!("Expand tensor queries").entered();
72 claims.prepare(channel)
73 };
74
75 // The weights the reduction's first factor carries, one per bit position.
76 // Phase 1 and phase 3 both need them, so they are computed once here.
77 let oblong_weights = domain_subspace.lagrange_evals_buffer(prepared.r_zhat_prime);
78
79 // Phase 1: bind the shift variant, the shift amount, and the bit position.
80 let phase_1_output = prove_phase_1::<_, P, _, _>(
81 key_collection,
82 words,
83 &prepared,
84 oblong_weights.as_ref(),
85 channel,
86 alloc,
87 );
88
89 // Phases 2 and 3 bind the two bit indices the shift indicators chain through.
90 // Phase 2 takes the intermediate word's, phase 3 the reduction's first-factor output bit.
91 //
92 // Phase 2 runs against phase 1's leftover weights, carrying its evaluation as a constant.
93 let inner = ShiftIndSumcheck::<P, _>::new(
94 alloc,
95 &phase_1_output.psi,
96 &ShiftChallengePoint::new(&phase_1_output.r_j, &phase_1_output.inner),
97 phase_1_output.g_eval,
98 );
99 debug_assert_eq!(inner.beta(), phase_1_output.gamma);
100 let inner_output = inner.prove(channel, alloc);
101
102 // Phase 3 runs against the reduction's first-factor weights, carrying what phase 2 fixed.
103 // Its own weights evaluate to a factor the verifier recomputes independently.
104 // So no division is needed between phases.
105 let outer = ShiftIndSumcheck::<P, _>::new(
106 alloc,
107 oblong_weights.as_ref(),
108 &ShiftChallengePoint::new(&inner_output.point, &phase_1_output.outer),
109 inner_output.ind_eval * phase_1_output.g_eval,
110 );
111 debug_assert_eq!(outer.beta(), inner_output.eval);
112 let outer_output = outer.prove(channel, alloc);
113
114 // Phase 4 reduces to the final challenges and witness evaluation.
115 // It runs against the constraint-matrix multilinear, scaled by the three factors above.
116 prove_phase_2::<_, P, _, _>(
117 key_collection,
118 words,
119 &prepared,
120 phase_1_output,
121 outer_output.weights_eval * outer_output.ind_eval * inner_output.ind_eval,
122 outer_output.eval,
123 channel,
124 alloc,
125 )
126}