From 9c65ed1a90d3d3b1dbb6ba6f5cc24bb7188b917d Mon Sep 17 00:00:00 2001 From: Steve Han Date: Mon, 10 Aug 2026 18:29:44 -0400 Subject: [PATCH 1/3] fix: normalize structured JSON numbers (#857) Signed-off-by: Steve Han --- .../processing/gsonschema/validators.py | 104 ++++++++++++++---- .../models/recipes/test_response_recipes.py | 23 ++++ .../processing/gsonschema/test_validators.py | 58 ++++++++++ 3 files changed, 161 insertions(+), 24 deletions(-) diff --git a/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py b/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py index 210d65cda..c21c3d0e6 100644 --- a/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py +++ b/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py @@ -143,35 +143,91 @@ def _get_decimal_info_from_anyof(schema: dict) -> tuple[bool, int | None]: return False, None -def normalize_decimal_fields(obj: DataObjectT, schema: JSONSchemaT) -> DataObjectT: - """Normalize Decimal-like anyOf fields to floats with proper precision.""" - if not isinstance(obj, dict): +def _resolve_local_ref(schema: JSONSchemaT, root_schema: JSONSchemaT) -> JSONSchemaT: + """Resolve local JSON Pointer references while preserving sibling keywords.""" + resolved_schema = schema + seen_refs: set[str] = set() + + while isinstance(resolved_schema, dict) and (ref := resolved_schema.get("$ref")): + if not isinstance(ref, str) or ref in seen_refs or not (ref == "#" or ref.startswith("#/")): + break + seen_refs.add(ref) + + target: Any = root_schema + if ref != "#": + for token in ref[2:].split("/"): + token = token.replace("~1", "/").replace("~0", "~") + if not isinstance(target, dict) or token not in target: + return resolved_schema + target = target[token] + if not isinstance(target, dict): + return resolved_schema + + siblings = {key: value for key, value in resolved_schema.items() if key != "$ref"} + resolved_schema = target | siblings + + return resolved_schema + + +def _normalize_numeric_fields( + obj: DataObjectT, + schema: JSONSchemaT, + root_schema: JSONSchemaT, + validator: Any, +) -> DataObjectT: + """Recursively canonicalize numeric values according to their JSON Schema types.""" + schema = _resolve_local_ref(schema, root_schema) + + is_decimal, decimal_places = _get_decimal_info_from_anyof(schema) + if is_decimal and isinstance(obj, (int, float, str)) and not isinstance(obj, bool): + value = Decimal(str(obj)) + if decimal_places is not None: + value = value.quantize(Decimal(f"0.{'0' * decimal_places}"), rounding=ROUND_HALF_UP) + return float(value) + + for keyword in ("oneOf", "anyOf"): + alternatives = schema.get(keyword) + if isinstance(alternatives, list): + for alternative in alternatives: + if validator.evolve(schema=alternative).is_valid(obj): + return _normalize_numeric_fields(obj, alternative, root_schema, validator) + + all_of = schema.get("allOf") + if isinstance(all_of, list): + for subschema in all_of: + obj = _normalize_numeric_fields(obj, subschema, root_schema, validator) + + schema_type = schema.get("type") + schema_types = {schema_type} if isinstance(schema_type, str) else set(schema_type or []) + if "integer" in schema_types and isinstance(obj, (int, float)) and not isinstance(obj, bool): + return int(obj) + if "number" in schema_types and isinstance(obj, (int, float)) and not isinstance(obj, bool): + return float(obj) + + if isinstance(obj, dict): + properties = schema.get("properties", {}) + additional_properties = schema.get("additionalProperties", {}) + for key, value in obj.items(): + field_schema = properties.get(key, additional_properties if isinstance(additional_properties, dict) else {}) + obj[key] = _normalize_numeric_fields(value, field_schema, root_schema, validator) return obj - defs = schema.get("$defs", {}) - obj_schema = defs.get(schema.get("$ref", "")[len("#/$defs/") :], schema) - props = obj_schema.get("properties", {}) - - for key, value in obj.items(): - field_schema = props.get(key, {}) - if "$ref" in field_schema: - field_schema = defs.get(field_schema["$ref"][len("#/$defs/") :], {}) - - if isinstance(value, dict): - obj[key] = normalize_decimal_fields(value, schema) - elif isinstance(value, list): - obj[key] = [normalize_decimal_fields(v, schema) if isinstance(v, dict) else v for v in value] - elif isinstance(value, (int, float, str)) and not isinstance(value, bool): - is_decimal, decimal_places = _get_decimal_info_from_anyof(field_schema) - if is_decimal: - d = Decimal(str(value)) - if decimal_places is not None: - d = d.quantize(Decimal(f"0.{'0' * decimal_places}"), rounding=ROUND_HALF_UP) - obj[key] = float(d) + if isinstance(obj, list): + prefix_items = schema.get("prefixItems", []) + item_schema = schema.get("items", {}) + for index, value in enumerate(obj): + field_schema = prefix_items[index] if index < len(prefix_items) else item_schema + obj[index] = _normalize_numeric_fields(value, field_schema, root_schema, validator) return obj +def normalize_numeric_fields(obj: DataObjectT, schema: JSONSchemaT) -> DataObjectT: + """Normalize JSON Schema numbers and integers to stable Python numeric types.""" + validator = _get_default_validator()(schema) + return _normalize_numeric_fields(obj, schema, schema, validator) + + ## We don't expect the outer data type (e.g. dict, list, or const) to be ## modified by the pruning action. @overload @@ -242,6 +298,6 @@ def validate( except lazy.jsonschema.ValidationError as exc: raise JSONSchemaValidationError(str(exc)) from exc - final_object = normalize_decimal_fields(final_object, schema) + final_object = normalize_numeric_fields(final_object, schema) return final_object diff --git a/packages/data-designer-engine/tests/engine/models/recipes/test_response_recipes.py b/packages/data-designer-engine/tests/engine/models/recipes/test_response_recipes.py index 89b7deaf7..f0fb1e5e9 100644 --- a/packages/data-designer-engine/tests/engine/models/recipes/test_response_recipes.py +++ b/packages/data-designer-engine/tests/engine/models/recipes/test_response_recipes.py @@ -45,6 +45,18 @@ class Foo(BaseModel): bar: Bar +class OverallScore(BaseModel): + score: float + + +class Evaluation(BaseModel): + overall: OverallScore + + +class Evaluations(BaseModel): + evaluations: list[Evaluation] + + def test_pydantic_response(): recipe = PydanticResponseRecipe(Foo) example = Foo(bar=Bar(baz=42)) @@ -108,6 +120,17 @@ def test_structured_response(): assert recipe.parse(response) == example +def test_structured_response_normalizes_nested_number_types(): + recipe = StructuredResponseRecipe(Evaluations.model_json_schema()) + response = recipe.generate_response_example({"evaluations": [{"overall": {"score": 9}}]}) + + result = recipe.parse(response) + + score = result["evaluations"][0]["overall"]["score"] + assert score == 9.0 + assert isinstance(score, float) + + def test_structured_response_extra_fields(): recipe = StructuredResponseRecipe(Foo.model_json_schema()) ## Make an example with an extra field in it -- should be pruned out. diff --git a/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py b/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py index ee9d18647..51097d156 100644 --- a/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py +++ b/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py @@ -3,6 +3,7 @@ import pytest +import data_designer.lazy_heavy_imports as lazy from data_designer.engine.processing.gsonschema.validators import JSONSchemaValidationError, validate @@ -302,3 +303,60 @@ def test_normalize_decimal_anyof_fields() -> None: result3 = validate({"name": "Gizmo", "price": "249.99"}, schema) assert result3["price"] == 249.99 assert isinstance(result3["price"], float) + + +NESTED_EVALUATION_SCHEMA = { + "$defs": { + "Criterion": { + "type": "object", + "properties": {"score": {"type": "integer"}}, + "required": ["score"], + }, + "Overall": { + "type": "object", + "properties": {"score": {"type": "number"}}, + "required": ["score"], + }, + "Evaluation": { + "type": "object", + "properties": { + "criterion": {"$ref": "#/$defs/Criterion"}, + "overall": {"$ref": "#/$defs/Overall"}, + }, + "required": ["criterion", "overall"], + }, + }, + "type": "object", + "properties": { + "evaluations": { + "type": "array", + "items": {"$ref": "#/$defs/Evaluation"}, + } + }, + "required": ["evaluations"], +} + + +def _evaluation_scores(overall_score: int | float) -> dict: + return {"evaluations": [{"criterion": {"score": 9.0}, "overall": {"score": overall_score}}]} + + +def test_normalize_nested_json_schema_numeric_types() -> None: + result = validate(_evaluation_scores(9), NESTED_EVALUATION_SCHEMA) + + criterion_score = result["evaluations"][0]["criterion"]["score"] + overall_score = result["evaluations"][0]["overall"]["score"] + assert criterion_score == 9 + assert isinstance(criterion_score, int) + assert overall_score == 9.0 + assert isinstance(overall_score, float) + + +def test_normalized_nested_numbers_have_compatible_parquet_schemas(tmp_path) -> None: + for batch_number, score in enumerate((9, 9.5)): + normalized = validate(_evaluation_scores(score), NESTED_EVALUATION_SCHEMA) + dataframe = lazy.pd.DataFrame({"qa_evaluations": [normalized]}) + dataframe.to_parquet(tmp_path / f"batch_{batch_number:05d}.parquet", index=False) + + combined = lazy.pd.read_parquet(tmp_path, dtype_backend="pyarrow") + assert len(combined) == 2 From ffc7f942708c17060800edb2c4c7cefd86c3643f Mon Sep 17 00:00:00 2001 From: Steve Han Date: Tue, 11 Aug 2026 10:00:53 -0400 Subject: [PATCH 2/3] fix: handle boolean JSON subschemas (#857) Signed-off-by: Steve Han --- .../processing/gsonschema/validators.py | 23 +++++--- .../processing/gsonschema/test_validators.py | 53 +++++++++++++++++++ 2 files changed, 69 insertions(+), 7 deletions(-) diff --git a/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py b/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py index c21c3d0e6..abccc1477 100644 --- a/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py +++ b/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py @@ -130,12 +130,12 @@ def _get_decimal_info_from_anyof(schema: dict) -> tuple[bool, int | None]: if not isinstance(any_of, list): return False, None - has_number = any(item.get("type") == "number" for item in any_of) + has_number = any(isinstance(item, dict) and item.get("type") == "number" for item in any_of) if not has_number: return False, None for item in any_of: - if item.get("type") == "string" and "pattern" in item: + if isinstance(item, dict) and item.get("type") == "string" and "pattern" in item: match = re.search(r"\\d\{0,(\d+)\}", item["pattern"]) if match: return True, int(match.group(1)) @@ -171,11 +171,14 @@ def _resolve_local_ref(schema: JSONSchemaT, root_schema: JSONSchemaT) -> JSONSch def _normalize_numeric_fields( obj: DataObjectT, - schema: JSONSchemaT, + schema: JSONSchemaT | bool, root_schema: JSONSchemaT, validator: Any, ) -> DataObjectT: """Recursively canonicalize numeric values according to their JSON Schema types.""" + if not isinstance(schema, dict): + return obj + schema = _resolve_local_ref(schema, root_schema) is_decimal, decimal_places = _get_decimal_info_from_anyof(schema) @@ -189,8 +192,10 @@ def _normalize_numeric_fields( alternatives = schema.get(keyword) if isinstance(alternatives, list): for alternative in alternatives: + # Full validation does not retain the matching union branch, so identify it again for normalization. if validator.evolve(schema=alternative).is_valid(obj): - return _normalize_numeric_fields(obj, alternative, root_schema, validator) + obj = _normalize_numeric_fields(obj, alternative, root_schema, validator) + break all_of = schema.get("allOf") if isinstance(all_of, list): @@ -206,17 +211,21 @@ def _normalize_numeric_fields( if isinstance(obj, dict): properties = schema.get("properties", {}) - additional_properties = schema.get("additionalProperties", {}) + additional_properties = schema.get("additionalProperties") for key, value in obj.items(): - field_schema = properties.get(key, additional_properties if isinstance(additional_properties, dict) else {}) + field_schema = properties.get(key, additional_properties) + if not isinstance(field_schema, (dict, bool)): + continue obj[key] = _normalize_numeric_fields(value, field_schema, root_schema, validator) return obj if isinstance(obj, list): prefix_items = schema.get("prefixItems", []) - item_schema = schema.get("items", {}) + item_schema = schema.get("items") for index, value in enumerate(obj): field_schema = prefix_items[index] if index < len(prefix_items) else item_schema + if not isinstance(field_schema, (dict, bool)): + continue obj[index] = _normalize_numeric_fields(value, field_schema, root_schema, validator) return obj diff --git a/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py b/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py index 51097d156..977ae0a92 100644 --- a/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py +++ b/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py @@ -305,6 +305,59 @@ def test_normalize_decimal_anyof_fields() -> None: assert isinstance(result3["price"], float) +@pytest.mark.parametrize( + "schema,data", + [ + ({"type": "array", "items": True}, [1, "two"]), + ( + {"type": "object", "properties": {"metadata": True}}, + {"metadata": {"score": 9}}, + ), + ], + ids=["boolean_items", "boolean_property"], +) +def test_boolean_subschema_does_not_crash_numeric_normalization(schema: dict, data: dict | list) -> None: + assert validate(data, schema) == data + + +def test_boolean_anyof_schema_does_not_crash_decimal_detection() -> None: + schema = {"anyOf": [True, {"type": "number"}]} + + assert validate(1, schema) == 1 + + +@pytest.mark.parametrize("keyword", ["oneOf", "anyOf"]) +def test_composition_normalization_includes_sibling_properties(keyword: str) -> None: + schema = { + "type": "object", + keyword: [{"required": ["a"]}, {"required": ["b"]}], + "properties": { + "a": {"type": "integer"}, + "score": {"type": "number"}, + }, + } + + result = validate({"a": 1, "score": 9}, schema) + + assert result["a"] == 1 + assert isinstance(result["a"], int) + assert result["score"] == 9.0 + assert isinstance(result["score"], float) + + +@pytest.mark.parametrize( + "value,expected_type", + [(9, float), (None, type(None))], + ids=["number", "null"], +) +def test_normalize_nullable_number(value, expected_type: type) -> None: + schema = {"anyOf": [{"type": "number"}, {"type": "null"}]} + + result = validate(value, schema) + + assert isinstance(result, expected_type) + + NESTED_EVALUATION_SCHEMA = { "$defs": { "Criterion": { From 137f71b2fba8ae5d3416511b9905a112bce2da89 Mon Sep 17 00:00:00 2001 From: Steve Han Date: Tue, 11 Aug 2026 13:17:19 -0400 Subject: [PATCH 3/3] fix: unify parquet schemas when reading (#857) Signed-off-by: Steve Han --- .../data_designer/config/utils/io_helpers.py | 22 ++-- .../tests/config/utils/test_io_helpers.py | 22 ++++ .../processing/gsonschema/validators.py | 117 ++++-------------- .../models/recipes/test_response_recipes.py | 23 ---- .../processing/gsonschema/test_validators.py | 111 ----------------- 5 files changed, 62 insertions(+), 233 deletions(-) diff --git a/packages/data-designer-config/src/data_designer/config/utils/io_helpers.py b/packages/data-designer-config/src/data_designer/config/utils/io_helpers.py index e71ec3a10..a8c5a61ee 100644 --- a/packages/data-designer-config/src/data_designer/config/utils/io_helpers.py +++ b/packages/data-designer-config/src/data_designer/config/utils/io_helpers.py @@ -112,23 +112,29 @@ def load_processor_dataset(processors_outputs_path: Path, processor_name: str) - def read_parquet_dataset(path: Path) -> pd.DataFrame: """Read a parquet dataset from a path. + Directory schemas are unified permissively before reading so compatible + physical type drift across files, such as nested integers and floats, is + promoted to a common representation. + Args: path: The path to the parquet dataset, can be either a file or a directory. Returns: The parquet dataset as a pandas DataFrame. """ - try: - return lazy.pd.read_parquet(path, dtype_backend="pyarrow") - except Exception as e: - if path.is_dir() and "Unsupported cast" in str(e): - logger.warning("Failed to read parquets as folder, falling back to individual files") + if path.is_dir() and (parquet_files := sorted(path.glob("*.parquet"))): + schemas = [lazy.pq.read_schema(file) for file in parquet_files] + try: + unified_schema = lazy.pa.unify_schemas(schemas, promote_options="permissive") + except (lazy.pa.ArrowInvalid, lazy.pa.ArrowTypeError): + logger.warning("Failed to unify parquet schemas, falling back to individual files") return lazy.pd.concat( - [lazy.pd.read_parquet(file, dtype_backend="pyarrow") for file in sorted(path.glob("*.parquet"))], + [lazy.pd.read_parquet(file, dtype_backend="pyarrow") for file in parquet_files], ignore_index=True, ) - else: - raise e + return lazy.pd.read_parquet(path, dtype_backend="pyarrow", schema=unified_schema) + + return lazy.pd.read_parquet(path, dtype_backend="pyarrow") def validate_dataset_file_path(file_path: str | Path, should_exist: bool = True) -> Path: diff --git a/packages/data-designer-config/tests/config/utils/test_io_helpers.py b/packages/data-designer-config/tests/config/utils/test_io_helpers.py index 1d65d945d..5f6ba1244 100644 --- a/packages/data-designer-config/tests/config/utils/test_io_helpers.py +++ b/packages/data-designer-config/tests/config/utils/test_io_helpers.py @@ -16,11 +16,33 @@ from data_designer.config.utils.io_helpers import ( _maybe_rewrite_url, is_http_url, + read_parquet_dataset, serialize_data, smart_load_yaml, ) +def test_read_parquet_dataset_unifies_nested_numeric_types(tmp_path) -> None: + evaluations = [ + {"evaluations": [{"overall": {"score": 9}}]}, + {"evaluations": [{"overall": {"score": 9.5}}]}, + ] + for batch_number, evaluation in enumerate(evaluations): + lazy.pd.DataFrame({"qa_evaluations": [evaluation]}).to_parquet( + tmp_path / f"batch_{batch_number:05d}.parquet", + index=False, + ) + + schemas = [lazy.pq.read_schema(file) for file in sorted(tmp_path.glob("*.parquet"))] + assert schemas[0] != schemas[1] + + result = read_parquet_dataset(tmp_path) + + scores = [row["evaluations"][0]["overall"]["score"] for row in result["qa_evaluations"]] + assert scores == [9.0, 9.5] + assert all(isinstance(score, float) for score in scores) + + def test_smart_load_yaml(): stub_dict = { "hello": "world", diff --git a/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py b/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py index abccc1477..210d65cda 100644 --- a/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py +++ b/packages/data-designer-engine/src/data_designer/engine/processing/gsonschema/validators.py @@ -130,12 +130,12 @@ def _get_decimal_info_from_anyof(schema: dict) -> tuple[bool, int | None]: if not isinstance(any_of, list): return False, None - has_number = any(isinstance(item, dict) and item.get("type") == "number" for item in any_of) + has_number = any(item.get("type") == "number" for item in any_of) if not has_number: return False, None for item in any_of: - if isinstance(item, dict) and item.get("type") == "string" and "pattern" in item: + if item.get("type") == "string" and "pattern" in item: match = re.search(r"\\d\{0,(\d+)\}", item["pattern"]) if match: return True, int(match.group(1)) @@ -143,100 +143,35 @@ def _get_decimal_info_from_anyof(schema: dict) -> tuple[bool, int | None]: return False, None -def _resolve_local_ref(schema: JSONSchemaT, root_schema: JSONSchemaT) -> JSONSchemaT: - """Resolve local JSON Pointer references while preserving sibling keywords.""" - resolved_schema = schema - seen_refs: set[str] = set() - - while isinstance(resolved_schema, dict) and (ref := resolved_schema.get("$ref")): - if not isinstance(ref, str) or ref in seen_refs or not (ref == "#" or ref.startswith("#/")): - break - seen_refs.add(ref) - - target: Any = root_schema - if ref != "#": - for token in ref[2:].split("/"): - token = token.replace("~1", "/").replace("~0", "~") - if not isinstance(target, dict) or token not in target: - return resolved_schema - target = target[token] - if not isinstance(target, dict): - return resolved_schema - - siblings = {key: value for key, value in resolved_schema.items() if key != "$ref"} - resolved_schema = target | siblings - - return resolved_schema - - -def _normalize_numeric_fields( - obj: DataObjectT, - schema: JSONSchemaT | bool, - root_schema: JSONSchemaT, - validator: Any, -) -> DataObjectT: - """Recursively canonicalize numeric values according to their JSON Schema types.""" - if not isinstance(schema, dict): - return obj - - schema = _resolve_local_ref(schema, root_schema) - - is_decimal, decimal_places = _get_decimal_info_from_anyof(schema) - if is_decimal and isinstance(obj, (int, float, str)) and not isinstance(obj, bool): - value = Decimal(str(obj)) - if decimal_places is not None: - value = value.quantize(Decimal(f"0.{'0' * decimal_places}"), rounding=ROUND_HALF_UP) - return float(value) - - for keyword in ("oneOf", "anyOf"): - alternatives = schema.get(keyword) - if isinstance(alternatives, list): - for alternative in alternatives: - # Full validation does not retain the matching union branch, so identify it again for normalization. - if validator.evolve(schema=alternative).is_valid(obj): - obj = _normalize_numeric_fields(obj, alternative, root_schema, validator) - break - - all_of = schema.get("allOf") - if isinstance(all_of, list): - for subschema in all_of: - obj = _normalize_numeric_fields(obj, subschema, root_schema, validator) - - schema_type = schema.get("type") - schema_types = {schema_type} if isinstance(schema_type, str) else set(schema_type or []) - if "integer" in schema_types and isinstance(obj, (int, float)) and not isinstance(obj, bool): - return int(obj) - if "number" in schema_types and isinstance(obj, (int, float)) and not isinstance(obj, bool): - return float(obj) - - if isinstance(obj, dict): - properties = schema.get("properties", {}) - additional_properties = schema.get("additionalProperties") - for key, value in obj.items(): - field_schema = properties.get(key, additional_properties) - if not isinstance(field_schema, (dict, bool)): - continue - obj[key] = _normalize_numeric_fields(value, field_schema, root_schema, validator) +def normalize_decimal_fields(obj: DataObjectT, schema: JSONSchemaT) -> DataObjectT: + """Normalize Decimal-like anyOf fields to floats with proper precision.""" + if not isinstance(obj, dict): return obj - if isinstance(obj, list): - prefix_items = schema.get("prefixItems", []) - item_schema = schema.get("items") - for index, value in enumerate(obj): - field_schema = prefix_items[index] if index < len(prefix_items) else item_schema - if not isinstance(field_schema, (dict, bool)): - continue - obj[index] = _normalize_numeric_fields(value, field_schema, root_schema, validator) + defs = schema.get("$defs", {}) + obj_schema = defs.get(schema.get("$ref", "")[len("#/$defs/") :], schema) + props = obj_schema.get("properties", {}) + + for key, value in obj.items(): + field_schema = props.get(key, {}) + if "$ref" in field_schema: + field_schema = defs.get(field_schema["$ref"][len("#/$defs/") :], {}) + + if isinstance(value, dict): + obj[key] = normalize_decimal_fields(value, schema) + elif isinstance(value, list): + obj[key] = [normalize_decimal_fields(v, schema) if isinstance(v, dict) else v for v in value] + elif isinstance(value, (int, float, str)) and not isinstance(value, bool): + is_decimal, decimal_places = _get_decimal_info_from_anyof(field_schema) + if is_decimal: + d = Decimal(str(value)) + if decimal_places is not None: + d = d.quantize(Decimal(f"0.{'0' * decimal_places}"), rounding=ROUND_HALF_UP) + obj[key] = float(d) return obj -def normalize_numeric_fields(obj: DataObjectT, schema: JSONSchemaT) -> DataObjectT: - """Normalize JSON Schema numbers and integers to stable Python numeric types.""" - validator = _get_default_validator()(schema) - return _normalize_numeric_fields(obj, schema, schema, validator) - - ## We don't expect the outer data type (e.g. dict, list, or const) to be ## modified by the pruning action. @overload @@ -307,6 +242,6 @@ def validate( except lazy.jsonschema.ValidationError as exc: raise JSONSchemaValidationError(str(exc)) from exc - final_object = normalize_numeric_fields(final_object, schema) + final_object = normalize_decimal_fields(final_object, schema) return final_object diff --git a/packages/data-designer-engine/tests/engine/models/recipes/test_response_recipes.py b/packages/data-designer-engine/tests/engine/models/recipes/test_response_recipes.py index f0fb1e5e9..89b7deaf7 100644 --- a/packages/data-designer-engine/tests/engine/models/recipes/test_response_recipes.py +++ b/packages/data-designer-engine/tests/engine/models/recipes/test_response_recipes.py @@ -45,18 +45,6 @@ class Foo(BaseModel): bar: Bar -class OverallScore(BaseModel): - score: float - - -class Evaluation(BaseModel): - overall: OverallScore - - -class Evaluations(BaseModel): - evaluations: list[Evaluation] - - def test_pydantic_response(): recipe = PydanticResponseRecipe(Foo) example = Foo(bar=Bar(baz=42)) @@ -120,17 +108,6 @@ def test_structured_response(): assert recipe.parse(response) == example -def test_structured_response_normalizes_nested_number_types(): - recipe = StructuredResponseRecipe(Evaluations.model_json_schema()) - response = recipe.generate_response_example({"evaluations": [{"overall": {"score": 9}}]}) - - result = recipe.parse(response) - - score = result["evaluations"][0]["overall"]["score"] - assert score == 9.0 - assert isinstance(score, float) - - def test_structured_response_extra_fields(): recipe = StructuredResponseRecipe(Foo.model_json_schema()) ## Make an example with an extra field in it -- should be pruned out. diff --git a/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py b/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py index 977ae0a92..ee9d18647 100644 --- a/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py +++ b/packages/data-designer-engine/tests/engine/processing/gsonschema/test_validators.py @@ -3,7 +3,6 @@ import pytest -import data_designer.lazy_heavy_imports as lazy from data_designer.engine.processing.gsonschema.validators import JSONSchemaValidationError, validate @@ -303,113 +302,3 @@ def test_normalize_decimal_anyof_fields() -> None: result3 = validate({"name": "Gizmo", "price": "249.99"}, schema) assert result3["price"] == 249.99 assert isinstance(result3["price"], float) - - -@pytest.mark.parametrize( - "schema,data", - [ - ({"type": "array", "items": True}, [1, "two"]), - ( - {"type": "object", "properties": {"metadata": True}}, - {"metadata": {"score": 9}}, - ), - ], - ids=["boolean_items", "boolean_property"], -) -def test_boolean_subschema_does_not_crash_numeric_normalization(schema: dict, data: dict | list) -> None: - assert validate(data, schema) == data - - -def test_boolean_anyof_schema_does_not_crash_decimal_detection() -> None: - schema = {"anyOf": [True, {"type": "number"}]} - - assert validate(1, schema) == 1 - - -@pytest.mark.parametrize("keyword", ["oneOf", "anyOf"]) -def test_composition_normalization_includes_sibling_properties(keyword: str) -> None: - schema = { - "type": "object", - keyword: [{"required": ["a"]}, {"required": ["b"]}], - "properties": { - "a": {"type": "integer"}, - "score": {"type": "number"}, - }, - } - - result = validate({"a": 1, "score": 9}, schema) - - assert result["a"] == 1 - assert isinstance(result["a"], int) - assert result["score"] == 9.0 - assert isinstance(result["score"], float) - - -@pytest.mark.parametrize( - "value,expected_type", - [(9, float), (None, type(None))], - ids=["number", "null"], -) -def test_normalize_nullable_number(value, expected_type: type) -> None: - schema = {"anyOf": [{"type": "number"}, {"type": "null"}]} - - result = validate(value, schema) - - assert isinstance(result, expected_type) - - -NESTED_EVALUATION_SCHEMA = { - "$defs": { - "Criterion": { - "type": "object", - "properties": {"score": {"type": "integer"}}, - "required": ["score"], - }, - "Overall": { - "type": "object", - "properties": {"score": {"type": "number"}}, - "required": ["score"], - }, - "Evaluation": { - "type": "object", - "properties": { - "criterion": {"$ref": "#/$defs/Criterion"}, - "overall": {"$ref": "#/$defs/Overall"}, - }, - "required": ["criterion", "overall"], - }, - }, - "type": "object", - "properties": { - "evaluations": { - "type": "array", - "items": {"$ref": "#/$defs/Evaluation"}, - } - }, - "required": ["evaluations"], -} - - -def _evaluation_scores(overall_score: int | float) -> dict: - return {"evaluations": [{"criterion": {"score": 9.0}, "overall": {"score": overall_score}}]} - - -def test_normalize_nested_json_schema_numeric_types() -> None: - result = validate(_evaluation_scores(9), NESTED_EVALUATION_SCHEMA) - - criterion_score = result["evaluations"][0]["criterion"]["score"] - overall_score = result["evaluations"][0]["overall"]["score"] - assert criterion_score == 9 - assert isinstance(criterion_score, int) - assert overall_score == 9.0 - assert isinstance(overall_score, float) - - -def test_normalized_nested_numbers_have_compatible_parquet_schemas(tmp_path) -> None: - for batch_number, score in enumerate((9, 9.5)): - normalized = validate(_evaluation_scores(score), NESTED_EVALUATION_SCHEMA) - dataframe = lazy.pd.DataFrame({"qa_evaluations": [normalized]}) - dataframe.to_parquet(tmp_path / f"batch_{batch_number:05d}.parquet", index=False) - - combined = lazy.pd.read_parquet(tmp_path, dtype_backend="pyarrow") - assert len(combined) == 2