Skip to main content

ferrin_schema/
provider_schema.rs

1//! Provider schemas with reference-compatible object parsing.
2//!
3//! Object stripping and defaults follow the Zod input schemas used by Vercel AI SDK
4//! (Apache-2.0, Copyright 2023 Vercel, Inc.); implemented independently in Rust.
5
6use std::io;
7use std::io::Write;
8
9use serde_json::Map;
10use serde_json::Value;
11
12use crate::Schema;
13use crate::TypeValidationError;
14use crate::json::DEFAULT_MAX_BYTES;
15use crate::json::DEFAULT_MAX_DEPTH;
16
17#[derive(Debug, thiserror::Error)]
18enum ProviderSchemaError {
19    #[error("provider schema input exceeds the JSON resource limits")]
20    Limit,
21    #[error("provider schema contains an unsupported reference")]
22    Reference,
23}
24
25impl Schema<Value> {
26    /// Creates a provider-input schema that applies defaults and strips undeclared object fields.
27    ///
28    /// Use JSON Schema generated for Zod's input mode. Objects with explicit
29    /// `additionalProperties` retain that policy: records preserve arbitrary keys,
30    /// and `false` rejects unknown keys. The ordinary JSON Schema constructor keeps
31    /// its validation-only behavior. Parsing uses the default JSON depth/byte limits.
32    ///
33    /// # Examples
34    ///
35    /// ```
36    /// use ferrin_schema::Schema;
37    /// use serde_json::json;
38    /// let schema = Schema::from_provider_json_schema(json!({
39    ///     "type": "object", "properties": {"content": {"type": "array", "default": []}}
40    /// }));
41    /// assert_eq!(schema.validate(json!({"ignored": 1}))?, json!({"content": []}));
42    /// # Ok::<(), ferrin_schema::TypeValidationError>(())
43    /// ```
44    #[must_use]
45    pub fn from_provider_json_schema(schema: Value) -> Self {
46        let validator = Self::from_json_schema(schema.clone());
47        Self::with_json_schema_and_validator(schema.clone(), move |value| {
48            check_limits(&value).map_err(|error| TypeValidationError::new(value.clone(), error))?;
49            let normalized = normalize(&schema, &schema, value.clone(), 0)
50                .map_err(|error| TypeValidationError::new(value, error))?;
51            check_limits(&normalized)
52                .map_err(|error| TypeValidationError::new(normalized.clone(), error))?;
53            validator.validate(normalized)
54        })
55    }
56}
57
58fn check_limits(value: &Value) -> Result<(), ProviderSchemaError> {
59    let mut stack = vec![(value, 0)];
60    while let Some((value, depth)) = stack.pop() {
61        if depth > DEFAULT_MAX_DEPTH {
62            return Err(ProviderSchemaError::Limit);
63        }
64        match value {
65            Value::Array(items) => stack.extend(items.iter().map(|value| (value, depth + 1))),
66            Value::Object(items) => stack.extend(items.values().map(|value| (value, depth + 1))),
67            _ => {}
68        }
69    }
70    serde_json::to_writer(ByteBudget(0), value).map_err(|_| ProviderSchemaError::Limit)
71}
72
73struct ByteBudget(usize);
74impl Write for ByteBudget {
75    fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
76        self.0 = self.0.saturating_add(bytes.len());
77        if self.0 > DEFAULT_MAX_BYTES {
78            return Err(io::Error::other("JSON size limit exceeded"));
79        }
80        Ok(bytes.len())
81    }
82    fn flush(&mut self) -> io::Result<()> {
83        Ok(())
84    }
85}
86
87fn with_definitions(root: &Value, schema: &Value) -> Value {
88    let mut result = schema.clone();
89    if let Some(object) = result.as_object_mut() {
90        for key in ["$defs", "definitions"] {
91            if let Some(definitions) = root.get(key) {
92                object.entry(key).or_insert_with(|| definitions.clone());
93            }
94        }
95    }
96    result
97}
98
99fn normalize(
100    root: &Value,
101    schema: &Value,
102    mut value: Value,
103    depth: usize,
104) -> Result<Value, ProviderSchemaError> {
105    if depth > DEFAULT_MAX_DEPTH {
106        return Err(ProviderSchemaError::Limit);
107    }
108    if let Some(reference) = schema.get("$ref").and_then(Value::as_str) {
109        let target = reference
110            .strip_prefix('#')
111            .and_then(|path| root.pointer(path))
112            .ok_or(ProviderSchemaError::Reference)?;
113        return normalize(root, target, value, depth + 1);
114    }
115    for key in ["anyOf", "oneOf"] {
116        if let Some(branches) = schema.get(key).and_then(Value::as_array) {
117            for branch in branches {
118                let candidate = normalize(root, branch, value.clone(), depth + 1)?;
119                if Schema::from_json_schema(with_definitions(root, branch))
120                    .validate(candidate.clone())
121                    .is_ok()
122                {
123                    value = candidate;
124                    break;
125                }
126            }
127        }
128    }
129    if let Some(branches) = schema.get("allOf").and_then(Value::as_array) {
130        let mut merged = Map::new();
131        let mut all_objects = true;
132        for branch in branches {
133            match normalize(root, branch, value.clone(), depth + 1)? {
134                Value::Object(object) => merged.extend(object),
135                _ => all_objects = false,
136            }
137        }
138        if all_objects {
139            value = Value::Object(merged);
140        }
141    }
142    if let Some(properties) = schema.get("properties").and_then(Value::as_object)
143        && let Value::Object(object) = &mut value
144    {
145        for (key, field_schema) in properties {
146            if let Some(field) = object.remove(key) {
147                object.insert(
148                    key.clone(),
149                    normalize(root, field_schema, field, depth + 1)?,
150                );
151            } else if let Some(default) = field_schema.get("default") {
152                object.insert(
153                    key.clone(),
154                    normalize(root, field_schema, default.clone(), depth + 1)?,
155                );
156            }
157        }
158        match schema.get("additionalProperties") {
159            None => object.retain(|key, _| properties.contains_key(key)),
160            Some(Value::Object(additional)) => {
161                let additional = Value::Object(additional.clone());
162                for (_, field) in object
163                    .iter_mut()
164                    .filter(|(key, _)| !properties.contains_key(*key))
165                {
166                    *field = normalize(root, &additional, field.take(), depth + 1)?;
167                }
168            }
169            _ => {}
170        }
171    } else if let (Some(additional), Value::Object(object)) = (
172        schema.get("additionalProperties").filter(|v| v.is_object()),
173        &mut value,
174    ) {
175        for field in object.values_mut() {
176            *field = normalize(root, additional, field.take(), depth + 1)?;
177        }
178    }
179    if let (Some(items), Value::Array(values)) = (schema.get("items"), &mut value) {
180        for (index, field) in values.iter_mut().enumerate() {
181            let item_schema = if let Some(tuple) = items.as_array() {
182                tuple.get(index)
183            } else {
184                Some(items)
185            };
186            if let Some(item_schema) = item_schema {
187                *field = normalize(root, item_schema, field.take(), depth + 1)?;
188            }
189        }
190    }
191    Ok(value)
192}