openvm_continuations/
commit_bytes.rs1use std::{array::from_fn, fmt};
2
3use num_bigint::BigUint;
4use openvm_stark_backend::codec::{Decode, Encode};
5use openvm_stark_sdk::config::baby_bear_poseidon2::{DIGEST_SIZE, F};
6use openvm_verify_stark_host::pvs::VkCommit;
7use p3_field::{PrimeCharacteristicRing, PrimeField32};
8use serde::{Deserialize, Deserializer, Serialize, Serializer};
9
10pub const COMMIT_NUM_BYTES: usize = 32;
11
12#[derive(Copy, Clone, Debug, PartialEq, Eq, Encode, Decode)]
17pub struct CommitBytes([u8; COMMIT_NUM_BYTES]);
18
19impl CommitBytes {
20 pub fn new(bytes: [u8; COMMIT_NUM_BYTES]) -> Self {
21 assert!(
22 u32_digest_to_bytes(&bytes_to_u32_digest(&bytes)) == bytes,
23 "non-canonical CommitBytes for BabyBear digest"
24 );
25 Self(bytes)
26 }
27
28 pub fn as_slice(&self) -> &[u8; COMMIT_NUM_BYTES] {
29 &self.0
30 }
31
32 pub fn to_field_le_bytes(self) -> [u8; COMMIT_NUM_BYTES] {
33 let digest: [F; DIGEST_SIZE] = self.into();
34 let mut bytes = [0u8; COMMIT_NUM_BYTES];
35 for (i, limb) in digest.into_iter().enumerate() {
36 bytes[4 * i..4 * (i + 1)].copy_from_slice(&limb.to_unique_u32().to_le_bytes());
37 }
38 bytes
39 }
40}
41
42#[derive(Copy, Clone, Debug, PartialEq, Eq)]
43pub struct VkCommitBytes {
44 pub cached_commit: CommitBytes,
45 pub vk_pre_hash: CommitBytes,
46}
47
48impl<F: PrimeCharacteristicRing> From<VkCommitBytes> for VkCommit<F> {
49 fn from(value: VkCommitBytes) -> Self {
50 VkCommit {
51 cached_commit: value.cached_commit.into(),
52 vk_pre_hash: value.vk_pre_hash.into(),
53 }
54 }
55}
56
57impl From<[F; DIGEST_SIZE]> for CommitBytes {
58 fn from(value: [F; DIGEST_SIZE]) -> Self {
59 Self::from(value.map(|x| x.as_canonical_u32()))
60 }
61}
62
63impl From<[u32; DIGEST_SIZE]> for CommitBytes {
64 fn from(value: [u32; DIGEST_SIZE]) -> Self {
65 assert!(
66 value.iter().all(|&digit| digit < F::ORDER_U32),
67 "non-canonical BabyBear digest limb"
68 );
69 Self(u32_digest_to_bytes(&value))
70 }
71}
72
73impl<F: PrimeCharacteristicRing> From<CommitBytes> for [F; DIGEST_SIZE] {
74 fn from(value: CommitBytes) -> Self {
75 assert!(
76 u32_digest_to_bytes(&bytes_to_u32_digest(&value.0)) == value.0,
77 "non-canonical CommitBytes for BabyBear digest"
78 );
79 bytes_to_u32_digest(&value.0).map(F::from_u32)
80 }
81}
82
83fn bytes_to_biguint(bytes: &[u8; COMMIT_NUM_BYTES]) -> BigUint {
84 let mut bigint = BigUint::ZERO;
85 for byte in bytes.iter() {
86 bigint <<= 8;
87 bigint += BigUint::from(*byte);
88 }
89 bigint
90}
91
92fn biguint_to_u32_digest(mut bigint: BigUint) -> [u32; DIGEST_SIZE] {
93 let order = F::ORDER_U32;
94 from_fn(|_| {
95 let bigint_digit = bigint.clone() % order;
96 let digit = if bigint_digit == BigUint::ZERO {
97 0u32
98 } else {
99 bigint_digit.to_u32_digits()[0]
100 };
101 bigint /= order;
102 digit
103 })
104}
105
106fn u32_digest_to_biguint(digest: &[u32; DIGEST_SIZE]) -> BigUint {
107 let mut bigint = BigUint::ZERO;
108 let mut base = BigUint::from(1u32);
109 let order = BigUint::from(F::ORDER_U32);
110 for digit in digest {
111 bigint += &base * BigUint::from(*digit);
112 base *= ℴ
113 }
114 bigint
115}
116
117fn bytes_to_u32_digest(bytes: &[u8; COMMIT_NUM_BYTES]) -> [u32; DIGEST_SIZE] {
118 biguint_to_u32_digest(bytes_to_biguint(bytes))
119}
120
121fn u32_digest_to_bytes(digest: &[u32; DIGEST_SIZE]) -> [u8; COMMIT_NUM_BYTES] {
122 let mut ret = [0u8; COMMIT_NUM_BYTES];
123 let bytes = u32_digest_to_biguint(digest).to_bytes_be();
124 let start = COMMIT_NUM_BYTES - bytes.len();
125 ret[start..].copy_from_slice(&bytes);
126 ret
127}
128
129impl fmt::Display for CommitBytes {
130 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
131 write!(f, "0x{}", hex::encode(self.as_slice()))
132 }
133}
134
135impl Serialize for CommitBytes {
136 fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
137 self.to_string().serialize(serializer)
138 }
139}
140
141impl<'de> Deserialize<'de> for CommitBytes {
142 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
143 let hex_str = String::deserialize(deserializer)?;
144 let hex_str = hex_str.strip_prefix("0x").unwrap_or(&hex_str);
145 let bytes: [u8; COMMIT_NUM_BYTES] = hex::decode(hex_str)
146 .map_err(serde::de::Error::custom)?
147 .try_into()
148 .map_err(|_| serde::de::Error::custom("expected 32 bytes"))?;
149 Ok(CommitBytes::new(bytes))
150 }
151}
152
153#[cfg(feature = "root-prover")]
154mod bn254 {
155 use p3_bn254::Bn254;
156 use p3_field::PrimeField;
157
158 use super::*;
159
160 impl From<Bn254> for CommitBytes {
161 fn from(value: Bn254) -> Self {
162 Self::new(bn254_to_bytes(value))
163 }
164 }
165
166 impl From<[Bn254; 1]> for CommitBytes {
167 fn from(value: [Bn254; 1]) -> Self {
168 CommitBytes::from(value[0])
169 }
170 }
171
172 impl From<CommitBytes> for Bn254 {
173 fn from(value: CommitBytes) -> Self {
174 bytes_to_bn254(&value.0)
175 }
176 }
177
178 fn bytes_to_bn254(bytes: &[u8; COMMIT_NUM_BYTES]) -> Bn254 {
179 let order = Bn254::from_u32(1 << 8);
180 let mut ret = Bn254::ZERO;
181 let mut base = Bn254::ONE;
182 for byte in bytes.iter().rev() {
183 ret += base * Bn254::from_u8(*byte);
184 base *= order;
185 }
186 ret
187 }
188
189 fn bn254_to_bytes(bn254: Bn254) -> [u8; COMMIT_NUM_BYTES] {
190 let mut ret = [0u8; COMMIT_NUM_BYTES];
191 let bytes = bn254.as_canonical_biguint().to_bytes_be();
192 let start = COMMIT_NUM_BYTES - bytes.len();
193 ret[start..].copy_from_slice(&bytes);
194 ret
195 }
196}