mirror of
https://github.com/windmill-labs/windmill.git
synced 2026-08-18 16:02:10 +00:00
chore(aiagent): better handling of schemas (#7488)
* better handling of schemas * add tests * better * add tests * handle allof
This commit is contained in:
@@ -376,8 +376,8 @@ impl OpenAIQueryBuilder {
|
||||
.and_then(|schema| schema.properties.as_ref())
|
||||
.filter(|props| !props.is_empty())
|
||||
.map(|_| {
|
||||
let schema = args.output_schema.unwrap();
|
||||
let strict_schema = schema.clone().make_strict();
|
||||
let mut strict_schema = args.output_schema.unwrap().clone();
|
||||
strict_schema.make_strict();
|
||||
ResponsesApiTextFormat {
|
||||
format: ResponsesApiTextFormatConfig {
|
||||
r#type: "json_schema".to_string(),
|
||||
|
||||
@@ -72,8 +72,8 @@ impl OtherQueryBuilder {
|
||||
&& args.output_schema.is_some()
|
||||
&& !should_use_structured_output_tool
|
||||
{
|
||||
let schema = args.output_schema.unwrap();
|
||||
let strict_schema = schema.clone().make_strict();
|
||||
let mut strict_schema = args.output_schema.unwrap().clone();
|
||||
strict_schema.make_strict();
|
||||
Some(ResponseFormat {
|
||||
r#type: "json_schema".to_string(),
|
||||
json_schema: JsonSchemaFormat {
|
||||
|
||||
@@ -289,8 +289,16 @@ impl Default for SchemaType {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
#[serde(untagged)]
|
||||
pub enum AdditionalProperties {
|
||||
Bool(bool),
|
||||
Schema(Box<OpenAPISchema>),
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize, Default, Clone, Debug)]
|
||||
pub struct OpenAPISchema {
|
||||
// Core type fields
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub r#type: Option<SchemaType>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
@@ -309,7 +317,61 @@ pub struct OpenAPISchema {
|
||||
skip_serializing_if = "Option::is_none",
|
||||
rename = "additionalProperties"
|
||||
)]
|
||||
pub additional_properties: Option<bool>,
|
||||
pub additional_properties: Option<AdditionalProperties>,
|
||||
|
||||
// Schema metadata
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub title: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub description: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "$schema")]
|
||||
pub schema_url: Option<String>,
|
||||
|
||||
// References and definitions
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "$ref")]
|
||||
pub ref_path: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "$defs")]
|
||||
pub defs: Option<HashMap<String, Box<OpenAPISchema>>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub definitions: Option<HashMap<String, Box<OpenAPISchema>>>,
|
||||
|
||||
// Schema composition
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "allOf")]
|
||||
pub all_of: Option<Vec<Box<OpenAPISchema>>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "anyOf")]
|
||||
pub any_of: Option<Vec<Box<OpenAPISchema>>>,
|
||||
|
||||
// Value constraints
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub r#const: Option<serde_json::Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub default: Option<serde_json::Value>,
|
||||
|
||||
// String constraints
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub pattern: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "minLength")]
|
||||
pub min_length: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "maxLength")]
|
||||
pub max_length: Option<u64>,
|
||||
|
||||
// Number constraints
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub minimum: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub maximum: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "exclusiveMinimum")]
|
||||
pub exclusive_minimum: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "exclusiveMaximum")]
|
||||
pub exclusive_maximum: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "multipleOf")]
|
||||
pub multiple_of: Option<f64>,
|
||||
|
||||
// Array constraints
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "minItems")]
|
||||
pub min_items: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "maxItems")]
|
||||
pub max_items: Option<u64>,
|
||||
}
|
||||
|
||||
impl OpenAPISchema {
|
||||
@@ -415,32 +477,37 @@ impl OpenAPISchema {
|
||||
}
|
||||
|
||||
/// Makes this schema compatible with OpenAI's strict mode by:
|
||||
/// - Adding additionalProperties: false to all object types
|
||||
/// - Flattening allOf schemas (not supported by OpenAI strict mode)
|
||||
/// - Adding additionalProperties: false to all object types (if not already set)
|
||||
/// - Making non-required properties nullable
|
||||
/// - Ensuring all properties are in the required array
|
||||
pub fn make_strict(mut self) -> Self {
|
||||
pub fn make_strict(&mut self) {
|
||||
// First, flatten any allOf schemas since OpenAI strict mode doesn't support them
|
||||
self.flatten_all_of();
|
||||
|
||||
// Handle this schema if it's an object type
|
||||
if let Some(SchemaType::Single(ref type_str)) = self.r#type {
|
||||
if type_str == "object" {
|
||||
// Set additionalProperties to false
|
||||
self.additional_properties = Some(false);
|
||||
// Only set additionalProperties to false if not already set
|
||||
// If user provided a value (bool or schema), preserve it and let OpenAI handle it
|
||||
if self.additional_properties.is_none() {
|
||||
self.additional_properties = Some(AdditionalProperties::Bool(false));
|
||||
}
|
||||
|
||||
if let Some(properties) = self.properties.as_mut() {
|
||||
// Get original required fields
|
||||
let original_required = self.required.as_ref();
|
||||
let original_required = self.required.clone();
|
||||
|
||||
if let Some(required) = original_required {
|
||||
// Update properties to make non-required fields nullable
|
||||
for (key, prop) in properties.iter_mut() {
|
||||
let mut new_prop = (**prop).clone();
|
||||
// Make non-required fields nullable
|
||||
// Always iterate over properties to recursively process nested schemas
|
||||
for (key, prop) in properties.iter_mut() {
|
||||
// Make non-required fields nullable (only if there were required fields specified)
|
||||
if let Some(ref required) = original_required {
|
||||
if !required.contains(key) {
|
||||
new_prop = new_prop.make_nullable();
|
||||
prop.make_nullable();
|
||||
}
|
||||
// Recursively make nested schemas strict
|
||||
new_prop = new_prop.make_strict();
|
||||
*prop = Box::new(new_prop);
|
||||
}
|
||||
// Recursively make nested schemas strict
|
||||
prop.make_strict();
|
||||
}
|
||||
|
||||
// All properties must be in required array for strict mode
|
||||
@@ -451,21 +518,39 @@ impl OpenAPISchema {
|
||||
|
||||
// Recursively process nested schemas
|
||||
if let Some(ref mut items) = self.items {
|
||||
**items = items.as_ref().clone().make_strict();
|
||||
items.make_strict();
|
||||
}
|
||||
|
||||
if let Some(ref mut one_of) = self.one_of {
|
||||
*one_of = one_of
|
||||
.iter()
|
||||
.map(|schema| Box::new(schema.as_ref().clone().make_strict()))
|
||||
.collect();
|
||||
for schema in one_of.iter_mut() {
|
||||
schema.make_strict();
|
||||
}
|
||||
}
|
||||
|
||||
self
|
||||
// Process anyOf schemas (supported by OpenAI)
|
||||
if let Some(ref mut any_of) = self.any_of {
|
||||
for schema in any_of.iter_mut() {
|
||||
schema.make_strict();
|
||||
}
|
||||
}
|
||||
|
||||
// Process $defs schemas (definitions are used via $ref)
|
||||
if let Some(ref mut defs) = self.defs {
|
||||
for (_, schema) in defs.iter_mut() {
|
||||
schema.make_strict();
|
||||
}
|
||||
}
|
||||
|
||||
// Process definitions schemas
|
||||
if let Some(ref mut definitions) = self.definitions {
|
||||
for (_, schema) in definitions.iter_mut() {
|
||||
schema.make_strict();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Makes this property nullable by converting its type to a union with null
|
||||
pub fn make_nullable(mut self) -> Self {
|
||||
pub fn make_nullable(&mut self) {
|
||||
match self.r#type.take() {
|
||||
Some(SchemaType::Single(type_str)) => {
|
||||
if type_str != "null" {
|
||||
@@ -484,7 +569,147 @@ impl OpenAPISchema {
|
||||
self.r#type = Some(SchemaType::Single("null".into()));
|
||||
}
|
||||
}
|
||||
self
|
||||
}
|
||||
|
||||
/// Flattens all schemas in `allOf` into this schema, then removes `allOf`.
|
||||
/// This is needed because OpenAI strict mode doesn't support allOf.
|
||||
/// The allOf schemas are recursively flattened first, then merged.
|
||||
pub fn flatten_all_of(&mut self) {
|
||||
if let Some(all_of) = self.all_of.take() {
|
||||
for mut schema in all_of {
|
||||
// Recursively flatten any nested allOf in the schema being merged
|
||||
schema.flatten_all_of();
|
||||
self.merge_from(&schema);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Merges fields from another schema into this one.
|
||||
/// Properties from `other` are added if they don't already exist in `self`.
|
||||
/// Required fields from `other` are appended to `self`'s required array.
|
||||
fn merge_from(&mut self, other: &OpenAPISchema) {
|
||||
// Merge type (take other's if self has none)
|
||||
if self.r#type.is_none() {
|
||||
self.r#type = other.r#type.clone();
|
||||
}
|
||||
|
||||
// Merge properties (other's properties are added if key doesn't exist)
|
||||
if let Some(ref other_props) = other.properties {
|
||||
let props = self.properties.get_or_insert_with(HashMap::new);
|
||||
for (key, value) in other_props {
|
||||
props.entry(key.clone()).or_insert_with(|| value.clone());
|
||||
}
|
||||
}
|
||||
|
||||
// Merge required arrays (deduplicated)
|
||||
if let Some(ref other_required) = other.required {
|
||||
let required = self.required.get_or_insert_with(Vec::new);
|
||||
for item in other_required {
|
||||
if !required.contains(item) {
|
||||
required.push(item.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Merge items (for arrays) - take other's if self has none
|
||||
if self.items.is_none() && other.items.is_some() {
|
||||
self.items = other.items.clone();
|
||||
}
|
||||
|
||||
// Merge additionalProperties - take other's if self has none
|
||||
if self.additional_properties.is_none() && other.additional_properties.is_some() {
|
||||
self.additional_properties = other.additional_properties.clone();
|
||||
}
|
||||
|
||||
// Merge format - take other's if self has none
|
||||
if self.format.is_none() && other.format.is_some() {
|
||||
self.format = other.format.clone();
|
||||
}
|
||||
|
||||
// Merge enum values - combine them if both have enums
|
||||
if let Some(ref other_enum) = other.r#enum {
|
||||
let enums = self.r#enum.get_or_insert_with(Vec::new);
|
||||
for item in other_enum {
|
||||
if !enums.contains(item) {
|
||||
enums.push(item.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Merge title/description - take other's if self has none
|
||||
if self.title.is_none() && other.title.is_some() {
|
||||
self.title = other.title.clone();
|
||||
}
|
||||
if self.description.is_none() && other.description.is_some() {
|
||||
self.description = other.description.clone();
|
||||
}
|
||||
|
||||
// Merge $defs and definitions
|
||||
if let Some(ref other_defs) = other.defs {
|
||||
let defs = self.defs.get_or_insert_with(HashMap::new);
|
||||
for (key, value) in other_defs {
|
||||
defs.entry(key.clone()).or_insert_with(|| value.clone());
|
||||
}
|
||||
}
|
||||
if let Some(ref other_definitions) = other.definitions {
|
||||
let definitions = self.definitions.get_or_insert_with(HashMap::new);
|
||||
for (key, value) in other_definitions {
|
||||
definitions.entry(key.clone()).or_insert_with(|| value.clone());
|
||||
}
|
||||
}
|
||||
|
||||
// Merge numeric constraints - use the more restrictive values
|
||||
if other.minimum.is_some() {
|
||||
self.minimum = match (self.minimum, other.minimum) {
|
||||
(Some(a), Some(b)) => Some(a.max(b)),
|
||||
(None, Some(b)) => Some(b),
|
||||
(a, None) => a,
|
||||
};
|
||||
}
|
||||
if other.maximum.is_some() {
|
||||
self.maximum = match (self.maximum, other.maximum) {
|
||||
(Some(a), Some(b)) => Some(a.min(b)),
|
||||
(None, Some(b)) => Some(b),
|
||||
(a, None) => a,
|
||||
};
|
||||
}
|
||||
|
||||
// Merge string constraints
|
||||
if other.min_length.is_some() {
|
||||
self.min_length = match (self.min_length, other.min_length) {
|
||||
(Some(a), Some(b)) => Some(a.max(b)),
|
||||
(None, Some(b)) => Some(b),
|
||||
(a, None) => a,
|
||||
};
|
||||
}
|
||||
if other.max_length.is_some() {
|
||||
self.max_length = match (self.max_length, other.max_length) {
|
||||
(Some(a), Some(b)) => Some(a.min(b)),
|
||||
(None, Some(b)) => Some(b),
|
||||
(a, None) => a,
|
||||
};
|
||||
}
|
||||
|
||||
// Merge array constraints
|
||||
if other.min_items.is_some() {
|
||||
self.min_items = match (self.min_items, other.min_items) {
|
||||
(Some(a), Some(b)) => Some(a.max(b)),
|
||||
(None, Some(b)) => Some(b),
|
||||
(a, None) => a,
|
||||
};
|
||||
}
|
||||
if other.max_items.is_some() {
|
||||
self.max_items = match (self.max_items, other.max_items) {
|
||||
(Some(a), Some(b)) => Some(a.min(b)),
|
||||
(None, Some(b)) => Some(b),
|
||||
(a, None) => a,
|
||||
};
|
||||
}
|
||||
|
||||
// Merge pattern - take other's if self has none (can't really combine patterns)
|
||||
if self.pattern.is_none() && other.pattern.is_some() {
|
||||
self.pattern = other.pattern.clone();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -495,3 +720,490 @@ pub struct S3ObjectWithType {
|
||||
pub s3_object: S3Object,
|
||||
pub r#type: String,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::collections::HashMap;
|
||||
|
||||
/// Helper to create a simple string type schema
|
||||
fn string_schema() -> OpenAPISchema {
|
||||
OpenAPISchema {
|
||||
r#type: Some(SchemaType::Single("string".to_string())),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
/// Helper to create a simple integer type schema
|
||||
fn integer_schema() -> OpenAPISchema {
|
||||
OpenAPISchema {
|
||||
r#type: Some(SchemaType::Single("integer".to_string())),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
/// Helper to create an object schema with given properties
|
||||
fn object_schema(properties: Vec<(&str, OpenAPISchema)>) -> OpenAPISchema {
|
||||
OpenAPISchema {
|
||||
r#type: Some(SchemaType::Single("object".to_string())),
|
||||
properties: Some(
|
||||
properties
|
||||
.into_iter()
|
||||
.map(|(k, v)| (k.to_string(), Box::new(v)))
|
||||
.collect(),
|
||||
),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_adds_additional_properties_false() {
|
||||
let mut schema = object_schema(vec![("name", string_schema())]);
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
assert!(
|
||||
matches!(schema.additional_properties, Some(AdditionalProperties::Bool(false))),
|
||||
"Expected additionalProperties to be false"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_preserves_existing_additional_properties() {
|
||||
let mut schema = object_schema(vec![("name", string_schema())]);
|
||||
schema.additional_properties = Some(AdditionalProperties::Bool(true));
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
assert!(
|
||||
matches!(schema.additional_properties, Some(AdditionalProperties::Bool(true))),
|
||||
"Expected additionalProperties to remain true (user-specified)"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_all_properties_required() {
|
||||
let mut schema = object_schema(vec![
|
||||
("name", string_schema()),
|
||||
("age", integer_schema()),
|
||||
]);
|
||||
schema.required = Some(vec!["name".to_string()]); // Only name is required initially
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
let required = schema.required.as_ref().expect("required should be set");
|
||||
assert!(required.contains(&"name".to_string()), "name should be required");
|
||||
assert!(required.contains(&"age".to_string()), "age should be required");
|
||||
assert_eq!(required.len(), 2, "Should have exactly 2 required fields");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_non_required_becomes_nullable() {
|
||||
let mut schema = object_schema(vec![
|
||||
("name", string_schema()),
|
||||
("age", integer_schema()),
|
||||
]);
|
||||
schema.required = Some(vec!["name".to_string()]); // Only name is required
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
// age field should now be nullable (type becomes ["integer", "null"])
|
||||
let age_prop = schema
|
||||
.properties
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.get("age")
|
||||
.expect("age property should exist");
|
||||
|
||||
match &age_prop.r#type {
|
||||
Some(SchemaType::Multiple(types)) => {
|
||||
assert!(types.contains(&"integer".to_string()), "Should contain integer");
|
||||
assert!(types.contains(&"null".to_string()), "Should contain null");
|
||||
}
|
||||
_ => panic!("Expected age to have multiple types including null"),
|
||||
}
|
||||
|
||||
// name field should NOT be nullable (it was already required)
|
||||
let name_prop = schema
|
||||
.properties
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.get("name")
|
||||
.expect("name property should exist");
|
||||
|
||||
match &name_prop.r#type {
|
||||
Some(SchemaType::Single(t)) => {
|
||||
assert_eq!(t, "string", "name should remain a simple string type");
|
||||
}
|
||||
_ => panic!("Expected name to remain a single type"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_recursive_nested_objects() {
|
||||
let nested = object_schema(vec![("field", string_schema())]);
|
||||
let mut schema = object_schema(vec![("nested", nested)]);
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
// Nested object should also have additionalProperties: false
|
||||
let nested_prop = schema
|
||||
.properties
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.get("nested")
|
||||
.expect("nested property should exist");
|
||||
|
||||
assert!(
|
||||
matches!(nested_prop.additional_properties, Some(AdditionalProperties::Bool(false))),
|
||||
"Nested object should have additionalProperties: false"
|
||||
);
|
||||
|
||||
// Nested object should have all properties required
|
||||
let nested_required = nested_prop.required.as_ref().expect("nested required should be set");
|
||||
assert!(nested_required.contains(&"field".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_array_items() {
|
||||
let item_schema = object_schema(vec![("id", string_schema())]);
|
||||
let mut schema = OpenAPISchema {
|
||||
r#type: Some(SchemaType::Single("array".to_string())),
|
||||
items: Some(Box::new(item_schema)),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
let items = schema.items.as_ref().expect("items should exist");
|
||||
assert!(
|
||||
matches!(items.additional_properties, Some(AdditionalProperties::Bool(false))),
|
||||
"Array items should have additionalProperties: false"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_one_of() {
|
||||
let variant1 = object_schema(vec![("type", string_schema())]);
|
||||
let variant2 = object_schema(vec![("value", integer_schema())]);
|
||||
|
||||
let mut schema = OpenAPISchema {
|
||||
one_of: Some(vec![Box::new(variant1), Box::new(variant2)]),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
for (i, variant) in schema.one_of.as_ref().unwrap().iter().enumerate() {
|
||||
assert!(
|
||||
matches!(variant.additional_properties, Some(AdditionalProperties::Bool(false))),
|
||||
"oneOf variant {} should have additionalProperties: false",
|
||||
i
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_any_of() {
|
||||
let variant1 = object_schema(vec![("a", string_schema())]);
|
||||
let variant2 = object_schema(vec![("b", integer_schema())]);
|
||||
|
||||
let mut schema = OpenAPISchema {
|
||||
any_of: Some(vec![Box::new(variant1), Box::new(variant2)]),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
for (i, variant) in schema.any_of.as_ref().unwrap().iter().enumerate() {
|
||||
assert!(
|
||||
matches!(variant.additional_properties, Some(AdditionalProperties::Bool(false))),
|
||||
"anyOf variant {} should have additionalProperties: false",
|
||||
i
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_defs() {
|
||||
let def_schema = object_schema(vec![("prop", string_schema())]);
|
||||
let mut defs = HashMap::new();
|
||||
defs.insert("MyType".to_string(), Box::new(def_schema));
|
||||
|
||||
let mut schema = OpenAPISchema {
|
||||
defs: Some(defs),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
let my_type = schema
|
||||
.defs
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.get("MyType")
|
||||
.expect("MyType def should exist");
|
||||
|
||||
assert!(
|
||||
matches!(my_type.additional_properties, Some(AdditionalProperties::Bool(false))),
|
||||
"$defs schema should have additionalProperties: false"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_definitions() {
|
||||
let def_schema = object_schema(vec![("prop", string_schema())]);
|
||||
let mut definitions = HashMap::new();
|
||||
definitions.insert("MyType".to_string(), Box::new(def_schema));
|
||||
|
||||
let mut schema = OpenAPISchema {
|
||||
definitions: Some(definitions),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
let my_type = schema
|
||||
.definitions
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.get("MyType")
|
||||
.expect("MyType definition should exist");
|
||||
|
||||
assert!(
|
||||
matches!(my_type.additional_properties, Some(AdditionalProperties::Bool(false))),
|
||||
"definitions schema should have additionalProperties: false"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_nullable_single_type() {
|
||||
let mut schema = string_schema();
|
||||
|
||||
schema.make_nullable();
|
||||
|
||||
match &schema.r#type {
|
||||
Some(SchemaType::Multiple(types)) => {
|
||||
assert!(types.contains(&"string".to_string()));
|
||||
assert!(types.contains(&"null".to_string()));
|
||||
assert_eq!(types.len(), 2);
|
||||
}
|
||||
_ => panic!("Expected multiple types after make_nullable"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_nullable_already_null() {
|
||||
let mut schema = OpenAPISchema {
|
||||
r#type: Some(SchemaType::Single("null".to_string())),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.make_nullable();
|
||||
|
||||
// Should remain just "null"
|
||||
match &schema.r#type {
|
||||
Some(SchemaType::Single(t)) => {
|
||||
assert_eq!(t, "null");
|
||||
}
|
||||
_ => panic!("Expected single null type"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_nullable_multiple_types() {
|
||||
let mut schema = OpenAPISchema {
|
||||
r#type: Some(SchemaType::Multiple(vec![
|
||||
"string".to_string(),
|
||||
"integer".to_string(),
|
||||
])),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.make_nullable();
|
||||
|
||||
match &schema.r#type {
|
||||
Some(SchemaType::Multiple(types)) => {
|
||||
assert!(types.contains(&"string".to_string()));
|
||||
assert!(types.contains(&"integer".to_string()));
|
||||
assert!(types.contains(&"null".to_string()));
|
||||
assert_eq!(types.len(), 3);
|
||||
}
|
||||
_ => panic!("Expected multiple types after make_nullable"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_nullable_already_has_null() {
|
||||
let mut schema = OpenAPISchema {
|
||||
r#type: Some(SchemaType::Multiple(vec![
|
||||
"string".to_string(),
|
||||
"null".to_string(),
|
||||
])),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.make_nullable();
|
||||
|
||||
match &schema.r#type {
|
||||
Some(SchemaType::Multiple(types)) => {
|
||||
// Should not duplicate null
|
||||
assert_eq!(types.iter().filter(|t| *t == "null").count(), 1);
|
||||
assert_eq!(types.len(), 2);
|
||||
}
|
||||
_ => panic!("Expected multiple types"),
|
||||
}
|
||||
}
|
||||
|
||||
// ========== allOf flattening tests ==========
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_flattens_all_of_properties() {
|
||||
let schema1 = object_schema(vec![("name", string_schema())]);
|
||||
let schema2 = object_schema(vec![("age", integer_schema())]);
|
||||
|
||||
let mut schema = OpenAPISchema {
|
||||
all_of: Some(vec![Box::new(schema1), Box::new(schema2)]),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
// allOf should be removed
|
||||
assert!(schema.all_of.is_none(), "allOf should be removed after flattening");
|
||||
|
||||
// Properties should be merged
|
||||
let props = schema.properties.as_ref().expect("properties should exist");
|
||||
assert!(props.contains_key("name"), "Should have 'name' property");
|
||||
assert!(props.contains_key("age"), "Should have 'age' property");
|
||||
|
||||
// Type should be set to object
|
||||
match &schema.r#type {
|
||||
Some(SchemaType::Single(t)) => assert_eq!(t, "object"),
|
||||
_ => panic!("Expected single 'object' type"),
|
||||
}
|
||||
|
||||
// Should have additionalProperties: false (from make_strict)
|
||||
assert!(
|
||||
matches!(schema.additional_properties, Some(AdditionalProperties::Bool(false))),
|
||||
"Should have additionalProperties: false"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_merges_all_of_required() {
|
||||
let mut schema1 = object_schema(vec![("name", string_schema())]);
|
||||
schema1.required = Some(vec!["name".to_string()]);
|
||||
|
||||
let mut schema2 = object_schema(vec![("age", integer_schema())]);
|
||||
schema2.required = Some(vec!["age".to_string()]);
|
||||
|
||||
let mut schema = OpenAPISchema {
|
||||
all_of: Some(vec![Box::new(schema1), Box::new(schema2)]),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
// Both name and age should be in required (from merge + make_strict makes all required)
|
||||
let required = schema.required.as_ref().expect("required should be set");
|
||||
assert!(required.contains(&"name".to_string()), "name should be required");
|
||||
assert!(required.contains(&"age".to_string()), "age should be required");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_nested_all_of() {
|
||||
// Create a nested allOf structure:
|
||||
// allOf: [
|
||||
// { allOf: [{ properties: { a } }, { properties: { b } }] },
|
||||
// { properties: { c } }
|
||||
// ]
|
||||
let inner_schema1 = object_schema(vec![("a", string_schema())]);
|
||||
let inner_schema2 = object_schema(vec![("b", string_schema())]);
|
||||
let inner_all_of = OpenAPISchema {
|
||||
all_of: Some(vec![Box::new(inner_schema1), Box::new(inner_schema2)]),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let outer_schema = object_schema(vec![("c", string_schema())]);
|
||||
|
||||
let mut schema = OpenAPISchema {
|
||||
all_of: Some(vec![Box::new(inner_all_of), Box::new(outer_schema)]),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
// All three properties should be present after recursive flattening
|
||||
let props = schema.properties.as_ref().expect("properties should exist");
|
||||
assert!(props.contains_key("a"), "Should have 'a' property");
|
||||
assert!(props.contains_key("b"), "Should have 'b' property");
|
||||
assert!(props.contains_key("c"), "Should have 'c' property");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_all_of_with_base_properties() {
|
||||
// Schema has both its own properties AND allOf
|
||||
let all_of_schema = object_schema(vec![("extra", string_schema())]);
|
||||
|
||||
let mut schema = object_schema(vec![("base", string_schema())]);
|
||||
schema.all_of = Some(vec![Box::new(all_of_schema)]);
|
||||
|
||||
schema.make_strict();
|
||||
|
||||
// Both base and extra properties should exist
|
||||
let props = schema.properties.as_ref().expect("properties should exist");
|
||||
assert!(props.contains_key("base"), "Should have 'base' property");
|
||||
assert!(props.contains_key("extra"), "Should have 'extra' property");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_make_strict_all_of_merges_constraints() {
|
||||
let schema1 = OpenAPISchema {
|
||||
r#type: Some(SchemaType::Single("integer".to_string())),
|
||||
minimum: Some(0.0),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let schema2 = OpenAPISchema {
|
||||
r#type: Some(SchemaType::Single("integer".to_string())),
|
||||
minimum: Some(5.0), // More restrictive
|
||||
maximum: Some(100.0),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let mut schema = OpenAPISchema {
|
||||
all_of: Some(vec![Box::new(schema1), Box::new(schema2)]),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.flatten_all_of();
|
||||
|
||||
// Should take the more restrictive minimum (5.0)
|
||||
assert_eq!(schema.minimum, Some(5.0), "Should have more restrictive minimum");
|
||||
assert_eq!(schema.maximum, Some(100.0), "Should have maximum from schema2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_flatten_all_of_preserves_defs() {
|
||||
let def_schema = object_schema(vec![("field", string_schema())]);
|
||||
let mut defs = HashMap::new();
|
||||
defs.insert("MyType".to_string(), Box::new(def_schema));
|
||||
|
||||
let schema_with_defs = OpenAPISchema {
|
||||
defs: Some(defs),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let mut schema = OpenAPISchema {
|
||||
all_of: Some(vec![Box::new(schema_with_defs)]),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
schema.flatten_all_of();
|
||||
|
||||
// $defs should be merged
|
||||
let defs = schema.defs.as_ref().expect("defs should exist");
|
||||
assert!(defs.contains_key("MyType"), "Should have 'MyType' def");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -142,5 +142,311 @@ class TestOutputSchema:
|
||||
print(f"Output schema without tool result from {provider_config['name']}: {result}")
|
||||
|
||||
|
||||
class TestSchemaVariations:
|
||||
"""
|
||||
Test various schema features across all providers.
|
||||
|
||||
These tests verify that different JSON Schema features are correctly
|
||||
processed by make_strict() and accepted by providers.
|
||||
"""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider_config",
|
||||
ALL_PROVIDERS,
|
||||
ids=get_provider_ids(ALL_PROVIDERS),
|
||||
)
|
||||
def test_nested_objects_schema(
|
||||
self,
|
||||
client: AIAgentTestClient,
|
||||
setup_providers,
|
||||
provider_config,
|
||||
):
|
||||
"""Test deeply nested object structure."""
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"user": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"age": {"type": "integer"}
|
||||
}
|
||||
}
|
||||
},
|
||||
"required": ["user"]
|
||||
}
|
||||
|
||||
flow_value = create_ai_agent_flow(
|
||||
provider_input_transform=provider_config["input_transform"],
|
||||
system_prompt="Extract user info. Return structured data.",
|
||||
tools=[],
|
||||
output_schema=schema,
|
||||
)
|
||||
|
||||
result = client.run_preview_flow(
|
||||
flow_value=flow_value,
|
||||
args={"user_message": "User John is 25 years old"},
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert "user" in result or "John" in str(result)
|
||||
print(f"Nested objects result from {provider_config['name']}: {result}")
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider_config",
|
||||
ALL_PROVIDERS,
|
||||
ids=get_provider_ids(ALL_PROVIDERS),
|
||||
)
|
||||
def test_array_of_objects_schema(
|
||||
self,
|
||||
client: AIAgentTestClient,
|
||||
setup_providers,
|
||||
provider_config,
|
||||
):
|
||||
"""Test array with object items."""
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"items": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"id": {"type": "integer"},
|
||||
"label": {"type": "string"}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"required": ["items"]
|
||||
}
|
||||
|
||||
flow_value = create_ai_agent_flow(
|
||||
provider_input_transform=provider_config["input_transform"],
|
||||
system_prompt="Create a list of items. Return structured data.",
|
||||
tools=[],
|
||||
output_schema=schema,
|
||||
)
|
||||
|
||||
result = client.run_preview_flow(
|
||||
flow_value=flow_value,
|
||||
args={"user_message": "Create 2 items: Apple (id 1), Banana (id 2)"},
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert "items" in result or "Apple" in str(result)
|
||||
print(f"Array of objects result from {provider_config['name']}: {result}")
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider_config",
|
||||
ALL_PROVIDERS,
|
||||
ids=get_provider_ids(ALL_PROVIDERS),
|
||||
)
|
||||
def test_enum_schema(
|
||||
self,
|
||||
client: AIAgentTestClient,
|
||||
setup_providers,
|
||||
provider_config,
|
||||
):
|
||||
"""Test enum constraints."""
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"status": {
|
||||
"type": "string",
|
||||
"enum": ["pending", "approved", "rejected"]
|
||||
}
|
||||
},
|
||||
"required": ["status"]
|
||||
}
|
||||
|
||||
flow_value = create_ai_agent_flow(
|
||||
provider_input_transform=provider_config["input_transform"],
|
||||
system_prompt="Classify the request status. Return structured data.",
|
||||
tools=[],
|
||||
output_schema=schema,
|
||||
)
|
||||
|
||||
result = client.run_preview_flow(
|
||||
flow_value=flow_value,
|
||||
args={"user_message": "The request was accepted"},
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
status = result.get("status") if isinstance(result, dict) else None
|
||||
assert status in ["pending", "approved", "rejected"] or "approved" in str(result)
|
||||
print(f"Enum result from {provider_config['name']}: {result}")
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider_config",
|
||||
ALL_PROVIDERS,
|
||||
ids=get_provider_ids(ALL_PROVIDERS),
|
||||
)
|
||||
def test_optional_fields_schema(
|
||||
self,
|
||||
client: AIAgentTestClient,
|
||||
setup_providers,
|
||||
provider_config,
|
||||
):
|
||||
"""Test that optional fields are handled correctly (made nullable)."""
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"nickname": {"type": "string"} # Not in required - should be nullable
|
||||
},
|
||||
"required": ["name"]
|
||||
}
|
||||
|
||||
flow_value = create_ai_agent_flow(
|
||||
provider_input_transform=provider_config["input_transform"],
|
||||
system_prompt="Extract name info. Return structured data.",
|
||||
tools=[],
|
||||
output_schema=schema,
|
||||
)
|
||||
|
||||
result = client.run_preview_flow(
|
||||
flow_value=flow_value,
|
||||
args={"user_message": "The person is called Alice"},
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert "name" in result or "Alice" in str(result)
|
||||
print(f"Optional fields result from {provider_config['name']}: {result}")
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider_config",
|
||||
ALL_PROVIDERS,
|
||||
ids=get_provider_ids(ALL_PROVIDERS),
|
||||
)
|
||||
def test_number_constraints_schema(
|
||||
self,
|
||||
client: AIAgentTestClient,
|
||||
setup_providers,
|
||||
provider_config,
|
||||
):
|
||||
"""Test min/max constraints on numbers."""
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"rating": {
|
||||
"type": "number",
|
||||
"minimum": 1,
|
||||
"maximum": 5
|
||||
}
|
||||
},
|
||||
"required": ["rating"]
|
||||
}
|
||||
|
||||
flow_value = create_ai_agent_flow(
|
||||
provider_input_transform=provider_config["input_transform"],
|
||||
system_prompt="Provide a rating from 1 to 5. Return structured data.",
|
||||
tools=[],
|
||||
output_schema=schema,
|
||||
)
|
||||
|
||||
result = client.run_preview_flow(
|
||||
flow_value=flow_value,
|
||||
args={"user_message": "This is excellent, rate it highly"},
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
rating = result.get("rating") if isinstance(result, dict) else None
|
||||
if rating is not None:
|
||||
assert 1 <= rating <= 5, f"Rating {rating} out of bounds"
|
||||
print(f"Number constraints result from {provider_config['name']}: {result}")
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider_config",
|
||||
ALL_PROVIDERS,
|
||||
ids=get_provider_ids(ALL_PROVIDERS),
|
||||
)
|
||||
def test_definitions_ref_schema(
|
||||
self,
|
||||
client: AIAgentTestClient,
|
||||
setup_providers,
|
||||
provider_config,
|
||||
):
|
||||
"""Test $ref with definitions."""
|
||||
if provider_config["name"] == "google_ai":
|
||||
pytest.xfail("Google AI does not support $ref with definitions in output_schema")
|
||||
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"primary": {"$ref": "#/definitions/Color"},
|
||||
"secondary": {"$ref": "#/definitions/Color"}
|
||||
},
|
||||
"required": ["primary", "secondary"],
|
||||
"definitions": {
|
||||
"Color": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"hex": {"type": "string"}
|
||||
},
|
||||
"required": ["name", "hex"]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
flow_value = create_ai_agent_flow(
|
||||
provider_input_transform=provider_config["input_transform"],
|
||||
system_prompt="Provide color information. Return structured data with primary and secondary colors.",
|
||||
tools=[],
|
||||
output_schema=schema,
|
||||
)
|
||||
|
||||
result = client.run_preview_flow(
|
||||
flow_value=flow_value,
|
||||
args={"user_message": "Primary color is red (#FF0000), secondary is blue (#0000FF)"},
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert "primary" in result or "red" in str(result).lower()
|
||||
print(f"Definitions/ref result from {provider_config['name']}: {result}")
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider_config",
|
||||
ALL_PROVIDERS,
|
||||
ids=get_provider_ids(ALL_PROVIDERS),
|
||||
)
|
||||
def test_anyof_schema(
|
||||
self,
|
||||
client: AIAgentTestClient,
|
||||
setup_providers,
|
||||
provider_config,
|
||||
):
|
||||
"""Test anyOf for union types."""
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"value": {
|
||||
"anyOf": [
|
||||
{"type": "string"},
|
||||
{"type": "number"}
|
||||
]
|
||||
}
|
||||
},
|
||||
"required": ["value"]
|
||||
}
|
||||
|
||||
flow_value = create_ai_agent_flow(
|
||||
provider_input_transform=provider_config["input_transform"],
|
||||
system_prompt="Extract the value. Return structured data.",
|
||||
tools=[],
|
||||
output_schema=schema,
|
||||
)
|
||||
|
||||
result = client.run_preview_flow(
|
||||
flow_value=flow_value,
|
||||
args={"user_message": "The answer is 42"},
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert "value" in result or "42" in str(result)
|
||||
print(f"anyOf result from {provider_config['name']}: {result}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v", "-s"])
|
||||
|
||||
Reference in New Issue
Block a user