mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-22 21:44:18 +02:00
Dedupe discriminator-required helper across schema converters (#4362)
This commit is contained in:
parent
0ca2c0d115
commit
5fa4f32cca
2 changed files with 34 additions and 51 deletions
|
|
@ -90,24 +90,7 @@ def _strip_discriminator(obj: Any) -> Any:
|
|||
if isinstance(obj, dict):
|
||||
skip = "discriminator" in obj and ("anyOf" in obj or "oneOf" in obj)
|
||||
if skip:
|
||||
discriminator = obj.get("discriminator")
|
||||
property_name = (
|
||||
discriminator.get("propertyName")
|
||||
if isinstance(discriminator, dict)
|
||||
else None
|
||||
)
|
||||
if isinstance(property_name, str):
|
||||
obj = obj.copy()
|
||||
for key in ("anyOf", "oneOf"):
|
||||
variants = obj.get(key)
|
||||
if not isinstance(variants, list):
|
||||
continue
|
||||
obj[key] = [
|
||||
_require_property(variant, property_name)
|
||||
if isinstance(variant, dict)
|
||||
else variant
|
||||
for variant in variants
|
||||
]
|
||||
obj = require_discriminator_property(obj)
|
||||
# Keys that hold instance data, not sub-schemas — don't recurse.
|
||||
_DATA_KEYS = {"default", "const", "examples", "enum"}
|
||||
return {
|
||||
|
|
@ -130,6 +113,37 @@ def _require_property(schema: dict[str, Any], property_name: str) -> dict[str, A
|
|||
return schema
|
||||
|
||||
|
||||
def require_discriminator_property(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Keep an OpenAPI discriminator's tag mandatory after the keyword is dropped.
|
||||
|
||||
Returns a copy of *schema* with ``discriminator.propertyName`` added to each
|
||||
``anyOf``/``oneOf`` variant's ``required`` list. A Pydantic discriminated
|
||||
union whose tag has a default omits that tag from ``required``; without this,
|
||||
an untagged payload passes the generated schema but fails later in the source
|
||||
model with ``union_tag_not_found``. No-op if there is no string
|
||||
``propertyName``.
|
||||
"""
|
||||
discriminator = schema.get("discriminator")
|
||||
if not isinstance(discriminator, dict):
|
||||
return schema
|
||||
property_name = discriminator.get("propertyName")
|
||||
if not isinstance(property_name, str):
|
||||
return schema
|
||||
|
||||
result = schema.copy()
|
||||
for key in ("anyOf", "oneOf"):
|
||||
variants = result.get(key)
|
||||
if not isinstance(variants, list):
|
||||
continue
|
||||
result[key] = [
|
||||
_require_property(variant, property_name)
|
||||
if isinstance(variant, dict)
|
||||
else variant
|
||||
for variant in variants
|
||||
]
|
||||
return result
|
||||
|
||||
|
||||
def dereference_refs(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Resolve all $ref references in a JSON schema by inlining definitions.
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ for our specific use case.
|
|||
|
||||
from typing import Any
|
||||
|
||||
from fastmcp.utilities.json_schema import require_discriminator_property
|
||||
from fastmcp.utilities.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
|
@ -97,7 +98,7 @@ def convert_openapi_schema_to_json_schema(
|
|||
result["anyOf"] = result.pop("oneOf")
|
||||
|
||||
# Step 3: Preserve discriminator tag presence before removing the keyword
|
||||
result = _require_discriminator_property(result)
|
||||
result = require_discriminator_property(result)
|
||||
|
||||
# Step 4: Remove OpenAPI-specific fields
|
||||
for field in OPENAPI_SPECIFIC_FIELDS:
|
||||
|
|
@ -152,38 +153,6 @@ def convert_openapi_schema_to_json_schema(
|
|||
return result
|
||||
|
||||
|
||||
def _require_discriminator_property(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Keep discriminator tags mandatory after dropping OpenAPI metadata."""
|
||||
discriminator = schema.get("discriminator")
|
||||
if not isinstance(discriminator, dict):
|
||||
return schema
|
||||
property_name = discriminator.get("propertyName")
|
||||
if not isinstance(property_name, str):
|
||||
return schema
|
||||
|
||||
result = schema.copy()
|
||||
for key in ("anyOf", "oneOf"):
|
||||
variants = result.get(key)
|
||||
if not isinstance(variants, list):
|
||||
continue
|
||||
result[key] = [
|
||||
_require_property(variant, property_name)
|
||||
if isinstance(variant, dict)
|
||||
else variant
|
||||
for variant in variants
|
||||
]
|
||||
return result
|
||||
|
||||
|
||||
def _require_property(schema: dict[str, Any], property_name: str) -> dict[str, Any]:
|
||||
required = schema.get("required")
|
||||
if required is None:
|
||||
return {**schema, "required": [property_name]}
|
||||
if isinstance(required, list) and property_name not in required:
|
||||
return {**schema, "required": [*required, property_name]}
|
||||
return schema
|
||||
|
||||
|
||||
def _convert_nullable_field(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Convert OpenAPI nullable field to JSON Schema type array."""
|
||||
if "nullable" not in schema:
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue