1mod algebra;
9mod ranks;
10use ranks::generalized_rank;
11use std::fmt;
12
13use sha2::{Digest, Sha256};
14
15use crate::field::{MODULUS_LIMIT, is_prime};
16use crate::{Error, Result};
17
18#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20#[non_exhaustive]
21pub struct ZigzagLimits {
22 pub max_nodes: usize,
24 pub max_total_dimension: usize,
26 pub max_map_terms: usize,
28 pub max_rank_work: usize,
30}
31
32impl Default for ZigzagLimits {
33 fn default() -> Self {
34 Self {
35 max_nodes: 2_049,
36 max_total_dimension: 100_000,
37 max_map_terms: 20_000_000,
38 max_rank_work: 100_000_000,
39 }
40 }
41}
42
43#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
45pub enum ZigzagDirection {
46 Forward,
48 Backward,
50}
51
52#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
54pub struct ZigzagTerm {
55 pub target: usize,
57 pub coefficient: u32,
59}
60
61#[derive(Debug, Clone, PartialEq, Eq)]
63pub struct ZigzagMap {
64 direction: ZigzagDirection,
65 columns: Vec<Vec<ZigzagTerm>>,
66}
67
68impl ZigzagMap {
69 pub fn new(direction: ZigzagDirection, columns: Vec<Vec<ZigzagTerm>>) -> Self {
71 Self { direction, columns }
72 }
73
74 pub fn direction(&self) -> ZigzagDirection {
76 self.direction
77 }
78
79 pub fn columns(&self) -> &[Vec<ZigzagTerm>] {
81 &self.columns
82 }
83}
84
85#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
87pub struct ZigzagModuleId([u8; 32]);
88
89impl ZigzagModuleId {
90 pub fn as_bytes(&self) -> &[u8; 32] {
92 &self.0
93 }
94}
95
96impl fmt::Display for ZigzagModuleId {
97 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
98 write_hex(formatter, &self.0)
99 }
100}
101
102#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
104pub struct ZigzagIntervalId([u8; 32]);
105
106impl ZigzagIntervalId {
107 pub fn as_bytes(&self) -> &[u8; 32] {
109 &self.0
110 }
111}
112
113impl fmt::Display for ZigzagIntervalId {
114 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
115 write_hex(formatter, &self.0)
116 }
117}
118
119#[derive(Debug, Clone, PartialEq, Eq)]
121pub struct ZigzagInterval {
122 pub id: ZigzagIntervalId,
124 pub start: usize,
126 pub end: usize,
128 pub multiplicity: usize,
130}
131
132#[derive(Debug, Clone, PartialEq, Eq)]
134pub struct ZigzagBarcode {
135 pub module: ZigzagModuleId,
137 pub dimensions: Vec<usize>,
139 pub generalized_ranks: Vec<usize>,
141 pub intervals: Vec<ZigzagInterval>,
143}
144
145impl ZigzagBarcode {
146 pub fn rank(&self, start: usize, end: usize) -> Option<usize> {
148 (start <= end && end < self.dimensions.len())
149 .then(|| self.generalized_ranks[start * self.dimensions.len() + end])
150 }
151}
152
153#[derive(Debug, Clone, PartialEq, Eq)]
155pub struct ZigzagModule {
156 id: ZigzagModuleId,
157 modulus: u32,
158 dimensions: Vec<usize>,
159 maps: Vec<ZigzagMap>,
160 limits: ZigzagLimits,
161}
162
163impl ZigzagModule {
164 pub fn new(
166 modulus: u32,
167 dimensions: Vec<usize>,
168 maps: Vec<ZigzagMap>,
169 limits: ZigzagLimits,
170 ) -> Result<Self> {
171 validate_module_shape(modulus, &dimensions, &maps, limits)?;
172 let total_dimension = total_dimension(&dimensions)?;
173 if total_dimension > limits.max_total_dimension {
174 return Err(Error::InvalidInput(format!(
175 "zigzag total dimension exceeds the limit {}",
176 limits.max_total_dimension
177 )));
178 }
179 let mut terms = 0usize;
180 for (position, map) in maps.iter().enumerate() {
181 terms = terms
182 .checked_add(validate_map(position, map, &dimensions, modulus)?)
183 .ok_or_else(|| Error::InvalidInput("zigzag map term count overflows".into()))?;
184 }
185 if terms > limits.max_map_terms {
186 return Err(Error::InvalidInput(format!(
187 "zigzag map term count exceeds the limit {}",
188 limits.max_map_terms
189 )));
190 }
191 let id = module_id(modulus, &dimensions, &maps);
192 Ok(Self {
193 id,
194 modulus,
195 dimensions,
196 maps,
197 limits,
198 })
199 }
200
201 pub fn id(&self) -> ZigzagModuleId {
203 self.id
204 }
205
206 pub fn modulus(&self) -> u32 {
208 self.modulus
209 }
210
211 pub fn dimensions(&self) -> &[usize] {
213 &self.dimensions
214 }
215
216 pub fn maps(&self) -> &[ZigzagMap] {
218 &self.maps
219 }
220
221 pub fn decompose(&self) -> Result<ZigzagBarcode> {
223 let generalized_ranks = self.compute_generalized_ranks()?;
224 let intervals =
225 interval_multiplicities(self.id, &generalized_ranks, self.dimensions.len())?;
226 Ok(ZigzagBarcode {
227 module: self.id,
228 dimensions: self.dimensions.clone(),
229 generalized_ranks,
230 intervals,
231 })
232 }
233
234 fn compute_generalized_ranks(&self) -> Result<Vec<usize>> {
235 let count = self.dimensions.len();
236 let mut ranks = vec![0usize; count * count];
237 let mut work = 0usize;
238 for start in (0..count).rev() {
239 for end in start..count {
240 work = work
241 .checked_add(self.rank_work(start, end)?)
242 .ok_or_else(rank_work_overflow)?;
243 if work > self.limits.max_rank_work {
244 return Err(Error::InvalidInput(format!(
245 "zigzag generalized-rank work exceeds the limit {}",
246 self.limits.max_rank_work
247 )));
248 }
249 ranks[start * count + end] = generalized_rank(self, start, end)?;
250 }
251 }
252 Ok(ranks)
253 }
254
255 fn rank_work(&self, start: usize, end: usize) -> Result<usize> {
256 let ambient = self.dimensions[start..=end].iter().sum::<usize>();
257 let arrows = (start..end).try_fold(0usize, |sum, position| {
258 let (source, target) =
259 map_shape(&self.dimensions, position, self.maps[position].direction);
260 sum.checked_add(source)
261 .and_then(|value| value.checked_add(target))
262 .ok_or_else(rank_work_overflow)
263 })?;
264 ambient.checked_add(arrows).ok_or_else(rank_work_overflow)
265 }
266}
267
268fn validate_module_shape(
269 modulus: u32,
270 dimensions: &[usize],
271 maps: &[ZigzagMap],
272 limits: ZigzagLimits,
273) -> Result<()> {
274 if !is_prime(modulus as u64) || u64::from(modulus) >= MODULUS_LIMIT {
275 return Err(Error::InvalidInput(
276 "zigzag modulus must be a supported prime".into(),
277 ));
278 }
279 if dimensions.is_empty() || dimensions.len() > limits.max_nodes {
280 return Err(Error::InvalidInput(format!(
281 "zigzag node count must be in 1..={}",
282 limits.max_nodes
283 )));
284 }
285 if maps.len() + 1 != dimensions.len() {
286 return Err(Error::InvalidInput(
287 "zigzag requires one map between each adjacent node".into(),
288 ));
289 }
290 Ok(())
291}
292
293fn total_dimension(dimensions: &[usize]) -> Result<usize> {
294 dimensions.iter().try_fold(0usize, |sum, value| {
295 sum.checked_add(*value)
296 .ok_or_else(|| Error::InvalidInput("zigzag total dimension overflows".into()))
297 })
298}
299
300fn validate_map(
301 position: usize,
302 map: &ZigzagMap,
303 dimensions: &[usize],
304 modulus: u32,
305) -> Result<usize> {
306 let (source, target) = map_shape(dimensions, position, map.direction);
307 if map.columns.len() != source {
308 return Err(Error::InvalidInput(format!(
309 "zigzag map {position} has {} columns but its source dimension is {source}",
310 map.columns.len()
311 )));
312 }
313 let mut terms = 0usize;
314 for column in &map.columns {
315 validate_column(position, column, target, modulus)?;
316 terms = terms
317 .checked_add(column.len())
318 .ok_or_else(|| Error::InvalidInput("zigzag map term count overflows".into()))?;
319 }
320 Ok(terms)
321}
322
323fn validate_column(
324 position: usize,
325 column: &[ZigzagTerm],
326 target: usize,
327 modulus: u32,
328) -> Result<()> {
329 let mut previous = None;
330 for term in column {
331 if term.target >= target
332 || term.coefficient == 0
333 || term.coefficient >= modulus
334 || previous.is_some_and(|value| value >= term.target)
335 {
336 return Err(Error::InvalidInput(format!(
337 "zigzag map {position} has a noncanonical term"
338 )));
339 }
340 previous = Some(term.target);
341 }
342 Ok(())
343}
344
345fn rank_work_overflow() -> Error {
346 Error::InvalidInput("zigzag rank work overflows".into())
347}
348
349fn interval_multiplicities(
350 module: ZigzagModuleId,
351 ranks: &[usize],
352 count: usize,
353) -> Result<Vec<ZigzagInterval>> {
354 let mut intervals = Vec::new();
355 for start in 0..count {
356 for end in start..count {
357 if let Some(interval) = interval_multiplicity(module, ranks, count, start, end)? {
358 intervals.push(interval);
359 }
360 }
361 }
362 Ok(intervals)
363}
364
365fn interval_multiplicity(
366 module: ZigzagModuleId,
367 ranks: &[usize],
368 count: usize,
369 start: usize,
370 end: usize,
371) -> Result<Option<ZigzagInterval>> {
372 let rank = |left: usize, right: usize| -> i128 { ranks[left * count + right] as i128 };
373 let mut multiplicity = rank(start, end);
374 if start > 0 {
375 multiplicity -= rank(start - 1, end);
376 }
377 if end + 1 < count {
378 multiplicity -= rank(start, end + 1);
379 }
380 if start > 0 && end + 1 < count {
381 multiplicity += rank(start - 1, end + 1);
382 }
383 if multiplicity < 0 {
384 return Err(Error::InvalidInput(
385 "zigzag generalized ranks violate interval decomposability".into(),
386 ));
387 }
388 if multiplicity == 0 {
389 return Ok(None);
390 }
391 let multiplicity = usize::try_from(multiplicity)
392 .map_err(|_| Error::InvalidInput("zigzag interval multiplicity overflows".into()))?;
393 Ok(Some(ZigzagInterval {
394 id: interval_id(module, start, end),
395 start,
396 end,
397 multiplicity,
398 }))
399}
400
401fn map_shape(dimensions: &[usize], position: usize, direction: ZigzagDirection) -> (usize, usize) {
402 match direction {
403 ZigzagDirection::Forward => (dimensions[position], dimensions[position + 1]),
404 ZigzagDirection::Backward => (dimensions[position + 1], dimensions[position]),
405 }
406}
407
408fn module_id(modulus: u32, dimensions: &[usize], maps: &[ZigzagMap]) -> ZigzagModuleId {
409 let mut hash = Sha256::new();
410 hash.update(b"holos-zigzag-module-v1");
411 hash.update(modulus.to_be_bytes());
412 hash.update((dimensions.len() as u64).to_be_bytes());
413 for dimension in dimensions {
414 hash.update((*dimension as u64).to_be_bytes());
415 }
416 for map in maps {
417 hash.update([match map.direction {
418 ZigzagDirection::Forward => 1,
419 ZigzagDirection::Backward => 2,
420 }]);
421 hash.update((map.columns.len() as u64).to_be_bytes());
422 for column in &map.columns {
423 hash.update((column.len() as u64).to_be_bytes());
424 for term in column {
425 hash.update((term.target as u64).to_be_bytes());
426 hash.update(term.coefficient.to_be_bytes());
427 }
428 }
429 }
430 ZigzagModuleId(hash.finalize().into())
431}
432
433fn interval_id(module: ZigzagModuleId, start: usize, end: usize) -> ZigzagIntervalId {
434 let mut hash = Sha256::new();
435 hash.update(b"holos-zigzag-interval-v1");
436 hash.update(module.as_bytes());
437 hash.update((start as u64).to_be_bytes());
438 hash.update((end as u64).to_be_bytes());
439 ZigzagIntervalId(hash.finalize().into())
440}
441
442fn write_hex(formatter: &mut fmt::Formatter<'_>, bytes: &[u8; 32]) -> fmt::Result {
443 for byte in bytes {
444 write!(formatter, "{byte:02x}")?;
445 }
446 Ok(())
447}
448
449#[cfg(test)]
450mod tests;