openvm_continuations/
commit_bytes.rs

1use 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/// Wrapper for the canonical big-endian byte representation of a BabyBear digest interpreted as
13/// an unsigned integer in base `F::MODULUS`. Each commit can be converted to a Bn254 using the
14/// trivial identification as natural numbers or into a `u32` digest by decomposing the big integer
15/// base-`F::MODULUS`.
16#[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 *= &order;
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}