use serde_json::{Map, Value};
use crate::SchemaDialect;
const STRICT_STRING_FORMATS: &[&str] = &[
"date-time",
"time",
"date",
"duration",
"email",
"hostname",
"ipv4",
"ipv6",
"uuid",
];
pub fn schema_for_dialect(schema: &Value, dialect: SchemaDialect) -> Value {
let mut schema = schema.clone();
if let Value::Object(root) = &mut schema {
root.remove("$schema");
}
match dialect {
SchemaDialect::Generic => schema = inline_refs(&schema),
SchemaDialect::OpenAiStrict => strict(&mut schema),
}
schema
}
fn inline_refs(schema: &Value) -> Value {
let mut kept = Vec::new();
let mut inlined = expand(schema, schema, &mut Vec::new(), &mut kept);
if let Value::Object(root) = &mut inlined {
for key in ["$defs", "definitions"] {
if let Some(Value::Object(defs)) = root.get_mut(key) {
defs.retain(|name, _| kept.contains(name));
if defs.is_empty() {
root.remove(key);
}
}
}
}
inlined
}
fn expand(node: &Value, root: &Value, stack: &mut Vec<String>, kept: &mut Vec<String>) -> Value {
match node {
Value::Object(map) => {
if let Some(Value::String(reference)) = map.get("$ref") {
let name = reference.rsplit('/').next().unwrap_or_default().to_string();
let target = reference.strip_prefix('#').and_then(|p| root.pointer(p));
match target {
Some(target) if !stack.contains(&name) => {
stack.push(name);
let mut expanded = expand(target, root, stack, kept);
stack.pop();
if let Value::Object(expanded_map) = &mut expanded {
for (key, value) in map.iter().filter(|(k, _)| *k != "$ref") {
expanded_map.insert(key.clone(), expand(value, root, stack, kept));
}
}
return expanded;
}
_ => {
if !kept.contains(&name) {
kept.push(name.clone());
if let Some(target) = target {
stack.push(name);
let _ = expand(target, root, stack, kept);
stack.pop();
}
}
return node.clone();
}
}
}
let mut out = serde_json::Map::new();
for (key, value) in map {
let value = if key == "$defs" || key == "definitions" {
expand_defs(value, root, kept)
} else {
expand(value, root, stack, kept)
};
out.insert(key.clone(), value);
}
Value::Object(out)
}
Value::Array(items) => {
Value::Array(items.iter().map(|v| expand(v, root, stack, kept)).collect())
}
other => other.clone(),
}
}
fn expand_defs(defs: &Value, root: &Value, kept: &mut Vec<String>) -> Value {
let Value::Object(defs) = defs else {
return defs.clone();
};
let mut out = serde_json::Map::new();
for (name, def) in defs {
let mut stack = vec![name.clone()];
out.insert(name.clone(), expand(def, root, &mut stack, kept));
}
Value::Object(out)
}
fn strict(schema: &mut Value) {
let Value::Object(map) = schema else { return };
if let Some(one_of) = map.remove("oneOf") {
map.insert("anyOf".into(), one_of);
}
if let Some(value) = map.remove("const") {
map.insert("enum".into(), Value::Array(vec![value]));
}
let keep_format = map.get("format").and_then(Value::as_str).is_some_and(|f| {
map.get("type") == Some(&Value::String("string".into()))
&& STRICT_STRING_FORMATS.contains(&f)
});
if !keep_format {
map.remove("format");
}
if map.get("type") == Some(&Value::String("object".into())) || map.contains_key("properties") {
make_object_strict(map);
}
for (key, child) in map.iter_mut() {
match key.as_str() {
"properties" | "$defs" | "definitions" => {
if let Value::Object(children) = child {
children.values_mut().for_each(strict);
}
}
"items" | "additionalProperties" => strict(child),
"anyOf" | "allOf" => {
if let Value::Array(branches) = child {
branches.iter_mut().for_each(strict);
}
}
_ => {}
}
}
}
fn make_object_strict(map: &mut Map<String, Value>) {
map.insert("additionalProperties".into(), Value::Bool(false));
let required: Vec<String> = map
.get("required")
.and_then(Value::as_array)
.map(|r| {
r.iter()
.filter_map(Value::as_str)
.map(String::from)
.collect()
})
.unwrap_or_default();
let Some(Value::Object(properties)) = map.get_mut("properties") else {
map.insert("properties".into(), Value::Object(Map::new()));
map.insert("required".into(), Value::Array(Vec::new()));
return;
};
for (name, property) in properties.iter_mut() {
if !required.contains(name) {
make_nullable(property);
}
}
let all: Vec<Value> = properties.keys().cloned().map(Value::String).collect();
map.insert("required".into(), Value::Array(all));
}
fn make_nullable(property: &mut Value) {
if accepts_null(property) {
return;
}
if let Value::Object(map) = property {
if let Some(Value::String(t)) = map.get("type") {
if !map.contains_key("enum") {
let t = t.clone();
map.insert("type".into(), serde_json::json!([t, "null"]));
return;
}
}
}
let original = std::mem::take(property);
*property = serde_json::json!({ "anyOf": [original, { "type": "null" }] });
}
fn accepts_null(schema: &Value) -> bool {
let is_null = |t: &Value| t.as_str() == Some("null");
match schema.get("type") {
Some(Value::String(t)) if t == "null" => return true,
Some(Value::Array(ts)) if ts.iter().any(is_null) => return true,
_ => {}
}
["anyOf", "oneOf"].iter().any(|k| {
schema
.get(*k)
.and_then(Value::as_array)
.is_some_and(|bs| bs.iter().any(accepts_null))
})
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn generic_drops_meta_schema() {
let schema = json!({"$schema": "x", "type": "object", "oneOf": []});
assert_eq!(
schema_for_dialect(&schema, SchemaDialect::Generic),
json!({"type": "object", "oneOf": []})
);
}
#[test]
fn generic_inlines_refs_and_keeps_field_descriptions() {
let schema = json!({
"type": "object",
"properties": {
"urgency": {"$ref": "#/$defs/Urgency", "description": "How urgent"},
"items": {"type": "array", "items": {"$ref": "#/$defs/Item"}}
},
"$defs": {
"Urgency": {"type": "string", "enum": ["Low", "High"]},
"Item": {"type": "object", "properties": {"u": {"$ref": "#/$defs/Urgency"}}}
}
});
let generic = schema_for_dialect(&schema, SchemaDialect::Generic);
assert_eq!(
generic,
json!({
"type": "object",
"properties": {
"urgency": {"type": "string", "enum": ["Low", "High"], "description": "How urgent"},
"items": {"type": "array", "items": {
"type": "object",
"properties": {"u": {"type": "string", "enum": ["Low", "High"]}}
}}
}
})
);
}
#[test]
fn generic_keeps_refs_of_recursive_types() {
let schema = json!({
"type": "object",
"properties": {"root": {"$ref": "#/$defs/Node"}, "kind": {"$ref": "#/$defs/Kind"}},
"$defs": {
"Node": {"type": "object", "properties": {
"children": {"type": "array", "items": {"$ref": "#/$defs/Node"}},
"kind": {"$ref": "#/$defs/Kind"}
}},
"Kind": {"type": "string", "enum": ["A"]}
}
});
let generic = schema_for_dialect(&schema, SchemaDialect::Generic);
let node = &generic["properties"]["root"];
assert_eq!(
node["properties"]["children"]["items"],
json!({"$ref": "#/$defs/Node"})
);
assert_eq!(
node["properties"]["kind"],
json!({"type": "string", "enum": ["A"]})
);
assert_eq!(
generic["properties"]["kind"],
json!({"type": "string", "enum": ["A"]})
);
let defs = generic["$defs"].as_object().unwrap();
assert_eq!(defs.keys().collect::<Vec<_>>(), ["Node"]);
assert_eq!(
defs["Node"]["properties"]["kind"],
json!({"type": "string", "enum": ["A"]})
);
}
#[test]
fn strict_requires_all_and_nulls_optionals() {
let schema = json!({
"$schema": "x",
"type": "object",
"properties": {
"a": {"type": "string"},
"b": {"type": ["string", "null"]},
"c": {"$ref": "#/$defs/C"},
"d": {"type": "integer", "format": "uint8", "minimum": 0}
},
"required": ["a", "d"],
"$defs": {
"C": {"oneOf": [{"type": "string", "const": "X"}, {"type": "object", "properties": {}}]}
}
});
let strict = schema_for_dialect(&schema, SchemaDialect::OpenAiStrict);
assert_eq!(strict["required"], json!(["a", "b", "c", "d"]));
assert_eq!(strict["additionalProperties"], json!(false));
assert_eq!(
strict["properties"]["b"],
json!({"type": ["string", "null"]})
);
assert_eq!(
strict["properties"]["c"],
json!({"anyOf": [{"$ref": "#/$defs/C"}, {"type": "null"}]})
);
assert_eq!(
strict["properties"]["d"],
json!({"type": "integer", "minimum": 0})
);
let c = &strict["$defs"]["C"];
assert_eq!(c["anyOf"][0], json!({"type": "string", "enum": ["X"]}));
assert_eq!(c["anyOf"][1]["additionalProperties"], json!(false));
assert!(strict.get("$schema").is_none());
}
#[test]
fn strict_keeps_supported_string_format() {
let schema = json!({"type": "string", "format": "date"});
assert_eq!(
schema_for_dialect(&schema, SchemaDialect::OpenAiStrict),
schema
);
}
}