typedlm 0.1.0

Typed, testable LLM programs for Rust
Documentation
//! Rewriting generated schemas into the flavour a provider accepts.

use serde_json::{Map, Value};

use crate::SchemaDialect;

/// String formats accepted by OpenAI-compatible strict structured outputs.
const STRICT_STRING_FORMATS: &[&str] = &[
    "date-time",
    "time",
    "date",
    "duration",
    "email",
    "hostname",
    "ipv4",
    "ipv6",
    "uuid",
];

/// Converts a schema generated by `schemars` into `dialect`.
///
/// Validation always runs against the original schema; this only shapes what is
/// sent to the provider.
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
}

/// Replaces local `$ref`s with the referenced definition. Servers that render the
/// schema into the prompt (common for tool parameters) leave references unresolved,
/// and models then guess the shape. Recursive types keep their `$ref` and the
/// definitions they need.
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();
                        // Keywords next to `$ref` (e.g. a field description) win.
                        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());
                            // A kept definition is itself expanded, keeping its own cycles.
                            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(),
    }
}

/// Expands each definition with itself on the stack, so its own recursion stays a `$ref`.
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)
}

/// OpenAI strict mode: every property required (optional ones made nullable),
/// `additionalProperties: false`, `anyOf` instead of `oneOf`, `enum` instead of
/// `const`, only supported string formats.
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));
}

/// Lets an optional property be `null`, since strict mode cannot omit it.
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
        );
    }
}