1use std::fmt;
4
5use rangeset::iter::{FromRangeIterator, IntoRangeIterator};
6use serde::{Deserialize, Serialize};
7
8use crate::{
9 hash::HashAlgId,
10 transcript::{
11 Direction, RangeSet, Transcript,
12 hash::{PlaintextHash, PlaintextHashSecret},
13 },
14};
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
18#[non_exhaustive]
19pub enum TranscriptCommitmentKind {
20 Hash {
22 alg: HashAlgId,
24 },
25}
26
27impl fmt::Display for TranscriptCommitmentKind {
28 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
29 match self {
30 Self::Hash { alg } => write!(f, "hash ({alg})"),
31 }
32 }
33}
34
35#[derive(Debug, Clone, Serialize, Deserialize)]
37#[non_exhaustive]
38pub enum TranscriptCommitment {
39 Hash(PlaintextHash),
41}
42
43#[derive(Debug, Clone, Serialize, Deserialize)]
45#[non_exhaustive]
46pub enum TranscriptSecret {
47 Hash(PlaintextHashSecret),
49}
50
51#[derive(Debug, Clone, Serialize, Deserialize)]
53pub struct TranscriptCommitConfig {
54 commits: Vec<((Direction, RangeSet<usize>), TranscriptCommitmentKind)>,
55}
56
57impl TranscriptCommitConfig {
58 pub fn builder(transcript: &Transcript) -> TranscriptCommitConfigBuilder<'_> {
60 TranscriptCommitConfigBuilder::new(transcript)
61 }
62
63 pub fn has_hash(&self) -> bool {
65 self.commits
66 .iter()
67 .any(|(_, kind)| matches!(kind, TranscriptCommitmentKind::Hash { .. }))
68 }
69
70 pub fn iter_hash(&self) -> impl Iterator<Item = (&(Direction, RangeSet<usize>), &HashAlgId)> {
74 self.commits.iter().map(|(idx, kind)| match kind {
75 TranscriptCommitmentKind::Hash { alg } => (idx, alg),
76 })
77 }
78
79 pub fn to_request(&self) -> TranscriptCommitRequest {
81 TranscriptCommitRequest {
82 hash: self
83 .iter_hash()
84 .map(|((dir, idx), alg)| (*dir, idx.clone(), *alg))
85 .collect(),
86 }
87 }
88}
89
90#[derive(Debug)]
92pub struct TranscriptCommitConfigBuilder<'a> {
93 transcript: &'a Transcript,
94 default_kind: TranscriptCommitmentKind,
95 commits: Vec<((Direction, RangeSet<usize>), TranscriptCommitmentKind)>,
96}
97
98impl<'a> TranscriptCommitConfigBuilder<'a> {
99 pub fn new(transcript: &'a Transcript) -> Self {
101 Self {
102 transcript,
103 default_kind: TranscriptCommitmentKind::Hash {
104 alg: HashAlgId::BLAKE3,
105 },
106 commits: Vec::default(),
107 }
108 }
109
110 pub fn default_kind(&mut self, default_kind: TranscriptCommitmentKind) -> &mut Self {
112 self.default_kind = default_kind;
113 self
114 }
115
116 pub fn commit_with_kind(
124 &mut self,
125 ranges: impl IntoRangeIterator<usize>,
126 direction: Direction,
127 kind: TranscriptCommitmentKind,
128 ) -> Result<&mut Self, TranscriptCommitConfigBuilderError> {
129 self.commit_with_kind_inner(RangeSet::from_range_iter(ranges), direction, kind)
130 }
131
132 fn commit_with_kind_inner(
133 &mut self,
134 idx: RangeSet<usize>,
135 direction: Direction,
136 kind: TranscriptCommitmentKind,
137 ) -> Result<&mut Self, TranscriptCommitConfigBuilderError> {
138 if idx.end().unwrap_or(0) > self.transcript.len_of_direction(direction) {
139 return Err(TranscriptCommitConfigBuilderError::new(
140 ErrorKind::Index,
141 format!(
142 "range is out of bounds of the transcript ({}): {} > {}",
143 direction,
144 idx.end().unwrap_or(0),
145 self.transcript.len_of_direction(direction)
146 ),
147 ));
148 }
149
150 let commit = ((direction, idx), kind);
151 if !self.commits.contains(&commit) {
152 self.commits.push(commit);
153 }
154
155 Ok(self)
156 }
157
158 pub fn commit(
165 &mut self,
166 ranges: impl IntoRangeIterator<usize>,
167 direction: Direction,
168 ) -> Result<&mut Self, TranscriptCommitConfigBuilderError> {
169 self.commit_with_kind_inner(
170 RangeSet::from_range_iter(ranges),
171 direction,
172 self.default_kind,
173 )
174 }
175
176 pub fn commit_sent(
182 &mut self,
183 ranges: impl IntoRangeIterator<usize>,
184 ) -> Result<&mut Self, TranscriptCommitConfigBuilderError> {
185 self.commit_with_kind_inner(
186 RangeSet::from_range_iter(ranges),
187 Direction::Sent,
188 self.default_kind,
189 )
190 }
191
192 pub fn commit_recv(
198 &mut self,
199 ranges: impl IntoRangeIterator<usize>,
200 ) -> Result<&mut Self, TranscriptCommitConfigBuilderError> {
201 self.commit_with_kind_inner(
202 RangeSet::from_range_iter(ranges),
203 Direction::Received,
204 self.default_kind,
205 )
206 }
207
208 pub fn build(self) -> Result<TranscriptCommitConfig, TranscriptCommitConfigBuilderError> {
210 Ok(TranscriptCommitConfig {
211 commits: self.commits,
212 })
213 }
214}
215
216#[derive(Debug, thiserror::Error)]
218pub struct TranscriptCommitConfigBuilderError {
219 kind: ErrorKind,
220 source: Option<Box<dyn std::error::Error + Send + Sync>>,
221}
222
223impl TranscriptCommitConfigBuilderError {
224 fn new<E>(kind: ErrorKind, source: E) -> Self
225 where
226 E: Into<Box<dyn std::error::Error + Send + Sync>>,
227 {
228 Self {
229 kind,
230 source: Some(source.into()),
231 }
232 }
233}
234
235#[derive(Debug)]
236enum ErrorKind {
237 Index,
238}
239
240impl fmt::Display for TranscriptCommitConfigBuilderError {
241 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
242 match self.kind {
243 ErrorKind::Index => f.write_str("index error")?,
244 }
245
246 if let Some(source) = &self.source {
247 write!(f, " caused by: {source}")?;
248 }
249
250 Ok(())
251 }
252}
253
254#[derive(Debug, Clone, Serialize, Deserialize)]
256pub struct TranscriptCommitRequest {
257 hash: Vec<(Direction, RangeSet<usize>, HashAlgId)>,
258}
259
260impl TranscriptCommitRequest {
261 pub fn has_hash(&self) -> bool {
263 !self.hash.is_empty()
264 }
265
266 pub fn iter_hash(&self) -> impl Iterator<Item = &(Direction, RangeSet<usize>, HashAlgId)> {
268 self.hash.iter()
269 }
270}
271
272#[cfg(test)]
273mod tests {
274 use super::*;
275
276 #[test]
277 fn test_range_out_of_bounds() {
278 let transcript = Transcript::new(
279 [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11],
280 [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11],
281 );
282 let mut builder = TranscriptCommitConfigBuilder::new(&transcript);
283
284 assert!(builder.commit_sent(&(10..15)).is_err());
285 assert!(builder.commit_recv(&(10..15)).is_err());
286 }
287
288 #[test]
289 fn test_commitment_order_matches_insertion_order() {
290 let transcript = Transcript::new([0; 12], [0; 12]);
291 let mut builder = TranscriptCommitConfigBuilder::new(&transcript);
292
293 builder.commit_recv(&(8..10)).unwrap();
294 builder.commit_sent(&(1..3)).unwrap();
295 builder.commit_recv(&(4..6)).unwrap();
296 builder.commit_recv(&(8..10)).unwrap();
297
298 let config = builder.build().unwrap();
299 let commits = config
300 .iter_hash()
301 .map(|((direction, idx), alg)| (*direction, idx.clone(), *alg))
302 .collect::<Vec<_>>();
303
304 assert_eq!(
305 commits,
306 vec![
307 (
308 Direction::Received,
309 RangeSet::from(8..10),
310 HashAlgId::BLAKE3
311 ),
312 (Direction::Sent, RangeSet::from(1..3), HashAlgId::BLAKE3),
313 (Direction::Received, RangeSet::from(4..6), HashAlgId::BLAKE3),
314 ]
315 );
316 }
317}