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
23pub 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
38fn 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
75fn 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
90fn 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 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}