diff --git a/backend/windmill-worker/src/ai/providers/openai.rs b/backend/windmill-worker/src/ai/providers/openai.rs index ae31cc6c83..97967b9c50 100644 --- a/backend/windmill-worker/src/ai/providers/openai.rs +++ b/backend/windmill-worker/src/ai/providers/openai.rs @@ -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(), diff --git a/backend/windmill-worker/src/ai/providers/other.rs b/backend/windmill-worker/src/ai/providers/other.rs index c8c2fef1f0..38e1d924aa 100644 --- a/backend/windmill-worker/src/ai/providers/other.rs +++ b/backend/windmill-worker/src/ai/providers/other.rs @@ -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 { diff --git a/backend/windmill-worker/src/ai/types.rs b/backend/windmill-worker/src/ai/types.rs index f28331cdb8..b48b4a73d1 100644 --- a/backend/windmill-worker/src/ai/types.rs +++ b/backend/windmill-worker/src/ai/types.rs @@ -289,8 +289,16 @@ impl Default for SchemaType { } } +#[derive(Serialize, Deserialize, Clone, Debug)] +#[serde(untagged)] +pub enum AdditionalProperties { + Bool(bool), + Schema(Box), +} + #[derive(Serialize, Deserialize, Default, Clone, Debug)] pub struct OpenAPISchema { + // Core type fields #[serde(skip_serializing_if = "Option::is_none")] pub r#type: Option, #[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, + pub additional_properties: Option, + + // Schema metadata + #[serde(skip_serializing_if = "Option::is_none")] + pub title: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub description: Option, + #[serde(skip_serializing_if = "Option::is_none", rename = "$schema")] + pub schema_url: Option, + + // References and definitions + #[serde(skip_serializing_if = "Option::is_none", rename = "$ref")] + pub ref_path: Option, + #[serde(skip_serializing_if = "Option::is_none", rename = "$defs")] + pub defs: Option>>, + #[serde(skip_serializing_if = "Option::is_none")] + pub definitions: Option>>, + + // Schema composition + #[serde(skip_serializing_if = "Option::is_none", rename = "allOf")] + pub all_of: Option>>, + #[serde(skip_serializing_if = "Option::is_none", rename = "anyOf")] + pub any_of: Option>>, + + // Value constraints + #[serde(skip_serializing_if = "Option::is_none")] + pub r#const: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub default: Option, + + // String constraints + #[serde(skip_serializing_if = "Option::is_none")] + pub pattern: Option, + #[serde(skip_serializing_if = "Option::is_none", rename = "minLength")] + pub min_length: Option, + #[serde(skip_serializing_if = "Option::is_none", rename = "maxLength")] + pub max_length: Option, + + // Number constraints + #[serde(skip_serializing_if = "Option::is_none")] + pub minimum: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub maximum: Option, + #[serde(skip_serializing_if = "Option::is_none", rename = "exclusiveMinimum")] + pub exclusive_minimum: Option, + #[serde(skip_serializing_if = "Option::is_none", rename = "exclusiveMaximum")] + pub exclusive_maximum: Option, + #[serde(skip_serializing_if = "Option::is_none", rename = "multipleOf")] + pub multiple_of: Option, + + // Array constraints + #[serde(skip_serializing_if = "Option::is_none", rename = "minItems")] + pub min_items: Option, + #[serde(skip_serializing_if = "Option::is_none", rename = "maxItems")] + pub max_items: Option, } 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"); + } +} diff --git a/integration_tests/ai_agent_tests/test_output_schema.py b/integration_tests/ai_agent_tests/test_output_schema.py index d0ef3ef510..a44804e725 100644 --- a/integration_tests/ai_agent_tests/test_output_schema.py +++ b/integration_tests/ai_agent_tests/test_output_schema.py @@ -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"])