Merge pull request #897 from owtaylor/openapi-recursive-components

This commit is contained in:
Jeremiah Lowin 2025-06-20 18:43:27 -04:00 committed by GitHub
commit 337c2f8db5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 111 additions and 31 deletions

View file

@ -1,6 +1,7 @@
from __future__ import annotations
import copy
from collections import defaultdict
def _prune_param(schema: dict, param: str) -> dict:
@ -24,25 +25,77 @@ def _prune_param(schema: dict, param: str) -> dict:
return schema
def _prune_unused_defs(schema: dict) -> dict:
"""Walk the schema and prune unused defs."""
root_defs: set[str] = set()
referenced_by: defaultdict[str, list] = defaultdict(list)
defs = schema.get("$defs")
if defs is None:
return schema
def walk(
node: object, current_def: str | None = None, skip_defs: bool = False
) -> None:
if isinstance(node, dict):
# Process $ref for definition tracking
ref = node.get("$ref")
if isinstance(ref, str) and ref.startswith("#/$defs/"):
def_name = ref.split("/")[-1]
if current_def:
referenced_by[def_name].append(current_def)
else:
root_defs.add(def_name)
# Walk children
for k, v in node.items():
if skip_defs and k == "$defs":
continue
walk(v, current_def=current_def)
elif isinstance(node, list):
for v in node:
walk(v)
# Traverse the schema once, skipping the $defs
walk(schema, skip_defs=True)
# Now figure out what defs reference other defs
for def_name, value in defs.items():
walk(value, current_def=def_name)
# Figure out what defs were referenced directly or recursively
def def_is_referenced(def_name):
if def_name in root_defs:
return True
references = referenced_by.get(def_name)
if references:
for reference in references:
if def_is_referenced(reference):
return True
return False
# Remove orphaned definitions if requested
for def_name in list(defs):
if not def_is_referenced(def_name):
defs.pop(def_name)
if not defs:
schema.pop("$defs", None)
return schema
def _walk_and_prune(
schema: dict,
prune_defs: bool = False,
prune_titles: bool = False,
prune_additional_properties: bool = False,
) -> dict:
"""Walk the schema and optionally prune titles, unused definitions, and additionalProperties: false."""
# Will only be used if prune_defs is True
used_defs: set[str] = set()
"""Walk the schema and optionally prune titles and additionalProperties: false."""
def walk(node: object) -> None:
if isinstance(node, dict):
# Process $ref for definition tracking
if prune_defs:
ref = node.get("$ref")
if isinstance(ref, str) and ref.startswith("#/$defs/"):
used_defs.add(ref.split("/")[-1])
# Remove title if requested
if prune_titles and "title" in node:
node.pop("title")
@ -62,18 +115,8 @@ def _walk_and_prune(
for v in node:
walk(v)
# Traverse the schema once
walk(schema)
# Remove orphaned definitions if requested
if prune_defs:
defs = schema.get("$defs", {})
for def_name in list(defs):
if def_name not in used_defs:
defs.pop(def_name)
if not defs:
schema.pop("$defs", None)
return schema
@ -109,12 +152,13 @@ def compress_schema(
schema = _prune_param(schema, param=param)
# Do a single walk to handle pruning operations
if prune_defs or prune_titles or prune_additional_properties:
if prune_titles or prune_additional_properties:
schema = _walk_and_prune(
schema,
prune_defs=prune_defs,
prune_titles=prune_titles,
prune_additional_properties=prune_additional_properties,
)
if prune_defs:
schema = _prune_unused_defs(schema)
return schema

View file

@ -262,16 +262,18 @@ class OpenAPIParser(
if isinstance(resolved_schema, (self.schema_cls)):
# Convert schema to dictionary
return resolved_schema.model_dump(
result = resolved_schema.model_dump(
mode="json", by_alias=True, exclude_none=True
)
elif isinstance(resolved_schema, dict):
return resolved_schema
result = resolved_schema
else:
logger.warning(
f"Expected Schema after resolving, got {type(resolved_schema)}. Returning empty dict."
)
return {}
result = {}
return _replace_ref_with_defs(result)
except Exception as e:
logger.error(f"Failed to extract schema as dict: {e}", exc_info=False)
return {}