Skip to main content

arete_hash/
canonical.rs

1use serde::de::{self, Deserialize, Deserializer, Error as _, MapAccess, SeqAccess, Visitor};
2use serde::Serialize;
3use serde_json::{Map, Value};
4use std::collections::HashSet;
5use std::fmt;
6
7use crate::{
8    hash_canonical_payload, push_framed_bytes, require_profile, CanonicalizationProfile, HashError,
9    HashId, Kind,
10};
11
12const MAX_SAFE_INTEGER: i64 = 9_007_199_254_740_991;
13const MAX_SAFE_INTEGER_DIGITS: &str = "9007199254740991";
14const DUPLICATE_KEY_MARKER: &str = "__ARETE_DUPLICATE_KEY__:";
15const UNSAFE_INTEGER_MARKER: &str = "__ARETE_UNSAFE_INTEGER__:";
16const NON_FINITE_MARKER: &str = "__ARETE_NON_FINITE_NUMBER__";
17
18pub fn hash_raw_bytes<K: Kind>(bytes: &[u8]) -> Result<HashId<K>, HashError> {
19    require_profile::<K>(CanonicalizationProfile::RawBytesV1)?;
20    Ok(hash_canonical_payload(bytes))
21}
22
23/// Parse JSON bytes without first converting through a JavaScript number.
24///
25/// This rejects malformed UTF-8, duplicate object keys, non-finite/out-of-range
26/// numbers, and integer tokens outside JavaScript's inclusive safe range.
27pub fn parse_json_bytes_strict(bytes: &[u8]) -> Result<Value, HashError> {
28    reject_unsafe_integer_tokens(bytes)?;
29    let mut deserializer = serde_json::Deserializer::from_slice(bytes);
30    let value = StrictValue::deserialize(&mut deserializer)
31        .map_err(|error| classify_json_error(error.to_string()))?;
32    deserializer
33        .end()
34        .map_err(|error| classify_json_error(error.to_string()))?;
35    Ok(value.0)
36}
37
38/// RFC 8785 treats every number as an IEEE-754 double, but the Arete profile
39/// rejects integer *tokens* outside the inclusive safe range before they can
40/// be converted through a double. serde_json silently falls back to `f64` for
41/// integer literals beyond `u64`/`i64`, so token-level validation must happen
42/// lexically, before deserialization.
43fn reject_unsafe_integer_tokens(bytes: &[u8]) -> Result<(), HashError> {
44    let mut index = 0;
45    while index < bytes.len() {
46        match bytes[index] {
47            b'"' => {
48                index = skip_string_token(bytes, index);
49            }
50            b'-' | b'0'..=b'9' => {
51                if let Some((token_end, integer_digits)) = scan_number_token(bytes, index) {
52                    if let Some(digits) = integer_digits {
53                        let magnitude = digits.trim_start_matches('-');
54                        let unsafe_integer = magnitude.len() > MAX_SAFE_INTEGER_DIGITS.len()
55                            || (magnitude.len() == MAX_SAFE_INTEGER_DIGITS.len()
56                                && magnitude > MAX_SAFE_INTEGER_DIGITS);
57                        if unsafe_integer {
58                            let token =
59                                std::string::String::from_utf8_lossy(&bytes[index..token_end])
60                                    .into_owned();
61                            return Err(HashError::UnsafeJsonInteger(token));
62                        }
63                    }
64                    index = token_end;
65                } else {
66                    index += 1;
67                }
68            }
69            _ => index += 1,
70        }
71    }
72    Ok(())
73}
74
75/// Skip a string token starting at the opening quote. Escaped code units never
76/// contain a raw `"` or `\` byte, so skipping `\` plus one byte is exact for
77/// scanning; malformed escapes are rejected later by the real parser.
78fn skip_string_token(bytes: &[u8], start: usize) -> usize {
79    let mut index = start + 1;
80    while index < bytes.len() {
81        match bytes[index] {
82            b'\\' => index += 2,
83            b'"' => return index + 1,
84            _ => index += 1,
85        }
86    }
87    index
88}
89
90/// Scan a JSON number token starting at `start`. Returns the token end and,
91/// for integer-form tokens (no fraction or exponent), the digit text. Returns
92/// `None` when the bytes do not form a complete, well-delimited JSON number;
93/// the real parser reports those as syntax errors.
94fn scan_number_token(bytes: &[u8], start: usize) -> Option<(usize, Option<String>)> {
95    let mut index = start;
96    if bytes.get(index) == Some(&b'-') {
97        index += 1;
98    }
99    match bytes.get(index) {
100        Some(b'0') => {
101            index += 1;
102            if bytes.get(index).is_some_and(|byte| byte.is_ascii_digit()) {
103                return None;
104            }
105        }
106        Some(b'1'..=b'9') => {
107            while bytes.get(index).is_some_and(|byte| byte.is_ascii_digit()) {
108                index += 1;
109            }
110        }
111        _ => return None,
112    }
113    let integer_digits = std::str::from_utf8(&bytes[start..index]).ok()?.to_string();
114
115    let mut is_integer = true;
116    if bytes.get(index) == Some(&b'.') {
117        is_integer = false;
118        index += 1;
119        let fraction_start = index;
120        while bytes.get(index).is_some_and(|byte| byte.is_ascii_digit()) {
121            index += 1;
122        }
123        if index == fraction_start {
124            return None;
125        }
126    }
127    if matches!(bytes.get(index), Some(b'e') | Some(b'E')) {
128        is_integer = false;
129        index += 1;
130        if matches!(bytes.get(index), Some(b'+') | Some(b'-')) {
131            index += 1;
132        }
133        let exponent_start = index;
134        while bytes.get(index).is_some_and(|byte| byte.is_ascii_digit()) {
135            index += 1;
136        }
137        if index == exponent_start {
138            return None;
139        }
140    }
141
142    let complete = match bytes.get(index) {
143        None => true,
144        Some(byte) => matches!(byte, b',' | b']' | b'}' | b' ' | b'\t' | b'\n' | b'\r'),
145    };
146    if !complete {
147        return None;
148    }
149    Some((index, is_integer.then_some(integer_digits)))
150}
151
152pub fn canonicalize_json_bytes(bytes: &[u8]) -> Result<Vec<u8>, HashError> {
153    let value = parse_json_bytes_strict(bytes)?;
154    canonicalize_json_value(&value)
155}
156
157pub fn canonicalize_jcs<T: Serialize>(value: &T) -> Result<Vec<u8>, HashError> {
158    // serde_json converts non-finite floats to null when building a Value. Run
159    // the RFC 8785 serializer first so NaN and infinities fail before that
160    // lossy conversion, then strictly parse its output to validate integer
161    // bounds without canonicalizing a second time.
162    let canonical = serde_json_canonicalizer::to_vec(value).map_err(|error| {
163        let message = error.to_string();
164        if message.contains("NaN") || message.contains("Infinity") || message.contains("finite") {
165            HashError::NonFiniteNumber
166        } else {
167            HashError::Serialization(message)
168        }
169    })?;
170    parse_json_bytes_strict(&canonical)?;
171    Ok(canonical)
172}
173
174pub fn canonicalize_json_value(value: &Value) -> Result<Vec<u8>, HashError> {
175    validate_json_value(value)?;
176    serde_json_canonicalizer::to_vec(value)
177        .map_err(|error| HashError::Serialization(error.to_string()))
178}
179
180pub fn hash_json_bytes<K: Kind>(bytes: &[u8]) -> Result<HashId<K>, HashError> {
181    require_profile::<K>(CanonicalizationProfile::AreteJcsV1)?;
182    let payload = canonicalize_json_bytes(bytes)?;
183    Ok(hash_canonical_payload(&payload))
184}
185
186pub fn hash_jcs<K: Kind, T: Serialize>(value: &T) -> Result<HashId<K>, HashError> {
187    require_profile::<K>(CanonicalizationProfile::AreteJcsV1)?;
188    let payload = canonicalize_jcs(value)?;
189    Ok(hash_canonical_payload(&payload))
190}
191
192#[derive(Debug, Clone, Copy)]
193pub struct TupleField<'a> {
194    pub label: &'a str,
195    pub value: &'a [u8],
196}
197
198impl<'a> TupleField<'a> {
199    pub const fn new(label: &'a str, value: &'a [u8]) -> Self {
200        Self { label, value }
201    }
202}
203
204pub fn framed_tuple_payload(fields: &[TupleField<'_>]) -> Result<Vec<u8>, HashError> {
205    let mut labels = HashSet::with_capacity(fields.len());
206    let mut payload = Vec::new();
207    payload.extend_from_slice(&(fields.len() as u64).to_be_bytes());
208    for field in fields {
209        if !labels.insert(field.label) {
210            return Err(HashError::DuplicateTupleLabel(field.label.to_string()));
211        }
212        push_framed_bytes(&mut payload, field.label.as_bytes());
213        push_framed_bytes(&mut payload, field.value);
214    }
215    Ok(payload)
216}
217
218pub fn hash_framed_tuple<K: Kind>(fields: &[TupleField<'_>]) -> Result<HashId<K>, HashError> {
219    require_profile::<K>(CanonicalizationProfile::FramedTupleV1)?;
220    let payload = framed_tuple_payload(fields)?;
221    Ok(hash_canonical_payload(&payload))
222}
223
224fn validate_json_value(value: &Value) -> Result<(), HashError> {
225    match value {
226        Value::Null | Value::Bool(_) | Value::String(_) => Ok(()),
227        Value::Array(values) => values.iter().try_for_each(validate_json_value),
228        Value::Object(values) => values.values().try_for_each(validate_json_value),
229        Value::Number(number) => {
230            if let Some(integer) = number.as_i64() {
231                if !(-MAX_SAFE_INTEGER..=MAX_SAFE_INTEGER).contains(&integer) {
232                    return Err(HashError::UnsafeJsonInteger(integer.to_string()));
233                }
234            } else if let Some(integer) = number.as_u64() {
235                if integer > MAX_SAFE_INTEGER as u64 {
236                    return Err(HashError::UnsafeJsonInteger(integer.to_string()));
237                }
238            } else if !number.as_f64().is_some_and(f64::is_finite) {
239                return Err(HashError::NonFiniteNumber);
240            }
241            Ok(())
242        }
243    }
244}
245
246fn classify_json_error(message: String) -> HashError {
247    if let Some(value) = marker_value(&message, DUPLICATE_KEY_MARKER) {
248        return HashError::DuplicateJsonKey(value);
249    }
250    if let Some(value) = marker_value(&message, UNSAFE_INTEGER_MARKER) {
251        return HashError::UnsafeJsonInteger(value);
252    }
253    if marker_value(&message, NON_FINITE_MARKER).is_some() {
254        return HashError::NonFiniteNumber;
255    }
256    if message.contains("number out of range") {
257        return HashError::NonFiniteNumber;
258    }
259    HashError::InvalidJson(message)
260}
261
262fn marker_value(message: &str, marker: &str) -> Option<String> {
263    let rest = message.split_once(marker)?.1;
264    Some(rest.split(" at line ").next().unwrap_or(rest).to_string())
265}
266
267struct StrictValue(Value);
268
269impl<'de> Deserialize<'de> for StrictValue {
270    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
271    where
272        D: Deserializer<'de>,
273    {
274        deserializer.deserialize_any(StrictValueVisitor)
275    }
276}
277
278struct StrictValueVisitor;
279
280impl<'de> Visitor<'de> for StrictValueVisitor {
281    type Value = StrictValue;
282
283    fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
284        formatter.write_str("a strict JSON value")
285    }
286
287    fn visit_unit<E>(self) -> Result<Self::Value, E>
288    where
289        E: de::Error,
290    {
291        Ok(StrictValue(Value::Null))
292    }
293
294    fn visit_bool<E>(self, value: bool) -> Result<Self::Value, E>
295    where
296        E: de::Error,
297    {
298        Ok(StrictValue(Value::Bool(value)))
299    }
300
301    fn visit_i64<E>(self, value: i64) -> Result<Self::Value, E>
302    where
303        E: de::Error,
304    {
305        if !(-MAX_SAFE_INTEGER..=MAX_SAFE_INTEGER).contains(&value) {
306            return Err(E::custom(format!("{UNSAFE_INTEGER_MARKER}{value}")));
307        }
308        Ok(StrictValue(Value::Number(value.into())))
309    }
310
311    fn visit_u64<E>(self, value: u64) -> Result<Self::Value, E>
312    where
313        E: de::Error,
314    {
315        if value > MAX_SAFE_INTEGER as u64 {
316            return Err(E::custom(format!("{UNSAFE_INTEGER_MARKER}{value}")));
317        }
318        Ok(StrictValue(Value::Number(value.into())))
319    }
320
321    fn visit_f64<E>(self, value: f64) -> Result<Self::Value, E>
322    where
323        E: de::Error,
324    {
325        if !value.is_finite() {
326            return Err(E::custom(format!("{NON_FINITE_MARKER}number")));
327        }
328        let number = serde_json::Number::from_f64(value)
329            .ok_or_else(|| E::custom(format!("{NON_FINITE_MARKER}number")))?;
330        Ok(StrictValue(Value::Number(number)))
331    }
332
333    fn visit_str<E>(self, value: &str) -> Result<Self::Value, E>
334    where
335        E: de::Error,
336    {
337        Ok(StrictValue(Value::String(value.to_string())))
338    }
339
340    fn visit_string<E>(self, value: String) -> Result<Self::Value, E>
341    where
342        E: de::Error,
343    {
344        Ok(StrictValue(Value::String(value)))
345    }
346
347    fn visit_seq<A>(self, mut sequence: A) -> Result<Self::Value, A::Error>
348    where
349        A: SeqAccess<'de>,
350    {
351        let mut values = Vec::new();
352        while let Some(value) = sequence.next_element::<StrictValue>()? {
353            values.push(value.0);
354        }
355        Ok(StrictValue(Value::Array(values)))
356    }
357
358    fn visit_map<A>(self, mut object: A) -> Result<Self::Value, A::Error>
359    where
360        A: MapAccess<'de>,
361    {
362        let mut values = Map::new();
363        while let Some(key) = object.next_key::<String>()? {
364            if values.contains_key(&key) {
365                return Err(A::Error::custom(format!("{DUPLICATE_KEY_MARKER}{key}")));
366            }
367            let value = object.next_value::<StrictValue>()?;
368            values.insert(key, value.0);
369        }
370        Ok(StrictValue(Value::Object(values)))
371    }
372}