use crate::{PlanError, PlanErrorCode};
use blazingly_json::Value;
use std::{collections::BTreeMap, io};
const MAX_SAFE_INTEGER: u64 = 9_007_199_254_740_991;
const MAX_SAFE_INTEGER_F64: f64 = 9_007_199_254_740_991.0;
const MAX_SAFE_INTEGER_I64: i64 = 9_007_199_254_740_991;
#[derive(Clone, Copy, Debug)]
pub(crate) struct JsonLimits {
pub bytes: usize,
pub nodes: usize,
pub depth: usize,
pub key_bytes: usize,
}
#[derive(Debug)]
pub(crate) struct JsonBudget {
limits: JsonLimits,
nodes: usize,
bytes: ByteCounter,
}
impl JsonBudget {
pub(crate) const fn new(limits: JsonLimits) -> Self {
Self {
limits,
nodes: 0,
bytes: ByteCounter::new(limits.bytes),
}
}
pub(crate) fn visit_map(&mut self, values: &BTreeMap<String, Value>) -> Result<(), PlanError> {
if values.is_empty() {
return Ok(());
}
if values.keys().any(|key| key.len() > self.limits.key_bytes) {
return Err(too_large(format!(
"extension key exceeds the {}-byte limit",
self.limits.key_bytes
)));
}
for value in values.values() {
self.visit_value(value)?;
}
if let Err(error) = blazingly_json::to_writer(&mut self.bytes, values) {
if self.bytes.exceeded {
return Err(too_large(format!(
"extension JSON exceeds the combined {}-byte limit",
self.limits.bytes
)));
}
return Err(PlanError::new(
PlanErrorCode::JsonEncoding,
format!("could not encode extension JSON: {error}"),
));
}
Ok(())
}
fn visit_value(&mut self, root: &Value) -> Result<(), PlanError> {
if let Value::Null | Value::Bool(_) | Value::Number(_) | Value::String(_) = root {
return self.admit_node(root, 1);
}
let mut stack = vec![(root, 1_usize)];
while let Some((value, depth)) = stack.pop() {
self.admit_node(value, depth)?;
match value {
Value::Array(values) => self.push_children(&mut stack, values.iter(), depth)?,
Value::Object(values) => {
self.push_children(&mut stack, values.values(), depth)?;
}
Value::Null | Value::Bool(_) | Value::Number(_) | Value::String(_) => {}
}
}
Ok(())
}
fn admit_node(&mut self, value: &Value, depth: usize) -> Result<(), PlanError> {
if self.nodes >= self.limits.nodes {
return Err(too_large(format!(
"extension JSON has more than {} values",
self.limits.nodes
)));
}
if depth > self.limits.depth {
return Err(too_large(format!(
"extension JSON exceeds depth {}",
self.limits.depth
)));
}
self.nodes += 1;
if let Value::Number(number) = value {
validate_number(number)?;
}
Ok(())
}
fn push_children<'a, I>(
&self,
stack: &mut Vec<(&'a Value, usize)>,
values: I,
depth: usize,
) -> Result<(), PlanError>
where
I: ExactSizeIterator<Item = &'a Value>,
{
if values.len() == 0 {
return Ok(());
}
if depth >= self.limits.depth
|| values.len()
> self
.limits
.nodes
.saturating_sub(self.nodes)
.saturating_sub(stack.len())
{
return Err(too_large("extension JSON exceeds its node or depth budget"));
}
stack.extend(values.map(|value| (value, depth + 1)));
Ok(())
}
}
fn validate_number(number: &blazingly_json::Number) -> Result<(), PlanError> {
if number.is_f64() {
let value = number.as_f64().expect("float number");
if value == 0.0 && value.is_sign_negative() {
return Err(unsafe_number("negative zero is not fingerprint-safe"));
}
if value.fract() == 0.0 && value.abs() > MAX_SAFE_INTEGER_F64 {
return Err(unsafe_number("integer exceeds the IEEE-754 safe range"));
}
return Ok(());
}
if let Some(value) = number.as_u64() {
if value > MAX_SAFE_INTEGER {
return Err(unsafe_number("integer exceeds the IEEE-754 safe range"));
}
return Ok(());
}
if number
.as_i64()
.is_some_and(|value| value < -MAX_SAFE_INTEGER_I64)
{
return Err(unsafe_number("integer exceeds the IEEE-754 safe range"));
}
Ok(())
}
#[derive(Debug)]
struct ByteCounter {
total: usize,
max: usize,
exceeded: bool,
}
impl ByteCounter {
const fn new(max: usize) -> Self {
Self {
total: 0,
max,
exceeded: false,
}
}
}
impl io::Write for ByteCounter {
fn write(&mut self, buffer: &[u8]) -> io::Result<usize> {
let Some(next) = self.total.checked_add(buffer.len()) else {
self.exceeded = true;
return Err(io::Error::other("extension byte count overflow"));
};
if next > self.max {
self.exceeded = true;
return Err(io::Error::other("extension byte budget exceeded"));
}
self.total = next;
Ok(buffer.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
fn too_large(message: impl Into<String>) -> PlanError {
PlanError::new(PlanErrorCode::EvidenceTooLarge, message)
}
fn unsafe_number(message: impl Into<String>) -> PlanError {
PlanError::new(PlanErrorCode::UnsafeNumber, message)
}
#[cfg(test)]
mod tests {
use super::{JsonBudget, JsonLimits};
use blazingly_json::{Number, Value};
use std::collections::BTreeMap;
#[test]
fn rejects_negative_zero_and_unsafe_integers() {
for number in [
Number::from_f64(-0.0).unwrap(),
Number::from(9_007_199_254_740_992_u64),
] {
let values = BTreeMap::from([("number".to_owned(), Value::Number(number))]);
let mut budget = JsonBudget::new(JsonLimits {
bytes: 100,
nodes: 10,
depth: 2,
key_bytes: 100,
});
assert!(budget.visit_map(&values).is_err());
}
}
}