Skip to content

Commit ac4e066

Browse files
committed
fix: enforce dependent required fields
1 parent c0c4a3b commit ac4e066

4 files changed

Lines changed: 348 additions & 2 deletions

File tree

postprocess_models.py

Lines changed: 167 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414

1515
"""Post-generation fixes for constraints datamodel-code-generator ignores.
1616
17-
Nine constraint families are handled:
17+
Ten constraint families are handled:
1818
1919
* ``minProperties`` / ``maxProperties`` on an object schema WITH declared
2020
properties are dropped by the generator: every field is optional, so an
@@ -119,6 +119,15 @@
119119
itself ``allOf``-references, e.g. ``postal_address.json``, is not
120120
inspected).
121121
122+
* ``dependentRequired`` on an object is dropped entirely by the generator. For
123+
example, ``time_interval.json`` allows an empty fragment but requires ``opens``
124+
and ``closes`` to appear together; the generated ``TimeInterval`` instead
125+
accepts either field alone. The script scans root object schemas for valid
126+
dependent-field maps and injects a ``model_validator(mode="after")`` that uses
127+
property presence (``model_fields_set`` plus extra keys), not value truthiness.
128+
A rule naming a field absent from a projected generated class is skipped rather
129+
than approximated.
130+
122131
* ``additionalProperties: false`` on an object schema with named properties is
123132
normally overridden by the generator's ``--extra-fields=allow`` flag. The
124133
script detects schemas with ``additionalProperties: false`` and flips their
@@ -189,6 +198,26 @@ def {marker}(self):
189198
return self
190199
'''
191200

201+
_DEPENDENT_REQUIRED_MARKER = "_enforce_dependent_required"
202+
203+
_DEPENDENT_REQUIRED_TEMPLATE = '''
204+
@model_validator(mode="after")
205+
def {marker}(self):
206+
"""JSON Schema dependentRequired: enforce dependent fields."""
207+
rules = {rules!r}
208+
provided = self.model_fields_set | set(self.model_extra or {{}})
209+
for field, required_fields in rules.items():
210+
if field not in provided:
211+
continue
212+
for required in required_fields:
213+
if required not in provided:
214+
raise ValueError(
215+
f"Field {{required!r}} is required when {{field!r}} "
216+
"is provided (schema dependentRequired)"
217+
)
218+
return self
219+
'''
220+
192221
_UNIQUE_MARKER = "_enforce_unique_items"
193222

194223
_CONDITIONAL_REQUIRED_MARKER = "_enforce_conditional_required"
@@ -1392,6 +1421,98 @@ def inject_conditional_array_retyping(source, class_name, rules):
13921421
return _ensure_pydantic_import(out, "model_validator")
13931422

13941423

1424+
def find_root_dependent_required(schema_dir):
1425+
"""Map generated class names to root-level dependentRequired rules.
1426+
1427+
Only complete rules over declared properties are returned. Request variants
1428+
can project either side out; those rules are inapplicable to that generated
1429+
class and are skipped by ``inject_dependent_required``.
1430+
"""
1431+
found = {}
1432+
for path in sorted(Path(schema_dir).rglob("*.json")):
1433+
try:
1434+
schema = json.loads(path.read_text(encoding="utf-8"))
1435+
except (OSError, json.JSONDecodeError):
1436+
continue
1437+
if not isinstance(schema, dict):
1438+
continue
1439+
properties = schema.get("properties")
1440+
rules = schema.get("dependentRequired")
1441+
title = schema.get("title")
1442+
if (
1443+
not isinstance(properties, dict)
1444+
or not properties
1445+
or not isinstance(rules, dict)
1446+
or not rules
1447+
):
1448+
continue
1449+
normalized = {}
1450+
malformed = False
1451+
for field, required in rules.items():
1452+
if (
1453+
not isinstance(field, str)
1454+
or not isinstance(required, list)
1455+
or not required
1456+
or not all(isinstance(name, str) for name in required)
1457+
):
1458+
malformed = True
1459+
break
1460+
# Request variants may project either side out. Such a rule no
1461+
# longer applies to that generated class and is skipped silently.
1462+
if field in properties and all(
1463+
name in properties for name in required
1464+
):
1465+
normalized[field] = required
1466+
if malformed:
1467+
sys.stderr.write(
1468+
f" ! {path}: unsupported dependentRequired rule; skipped\n"
1469+
)
1470+
continue
1471+
if not normalized:
1472+
continue
1473+
if not isinstance(title, str) or not title:
1474+
sys.stderr.write(
1475+
f" ! {path}: root dependentRequired but no title; "
1476+
"cannot map to a class\n"
1477+
)
1478+
continue
1479+
found[_alias_name(title)] = normalized
1480+
return found
1481+
1482+
1483+
def inject_dependent_required(source, class_name, rules):
1484+
"""Inject root dependentRequired checks into one generated class."""
1485+
class_re = re.compile(rf"^class {re.escape(class_name)}\(", re.M)
1486+
match = class_re.search(source)
1487+
if not match:
1488+
return source
1489+
tail = re.compile(r"^\S", re.M)
1490+
end_match = tail.search(source, match.end())
1491+
end = end_match.start() if end_match else len(source)
1492+
class_body = source[match.start() : end]
1493+
if f"def {_DEPENDENT_REQUIRED_MARKER}(" in class_body:
1494+
return source
1495+
declared = {
1496+
field.group(1)
1497+
for field in re.finditer(r"^ (\w+): [^\n]+", class_body, re.M)
1498+
}
1499+
applicable = {
1500+
field: required
1501+
for field, required in rules.items()
1502+
if field in declared and all(name in declared for name in required)
1503+
}
1504+
if not applicable:
1505+
return source
1506+
method = _DEPENDENT_REQUIRED_TEMPLATE.format(
1507+
marker=_DEPENDENT_REQUIRED_MARKER,
1508+
rules=applicable,
1509+
)
1510+
body = source[:end].rstrip("\n")
1511+
rest = source[end:]
1512+
out = body + "\n" + method + ("\n" + rest if rest else "")
1513+
return _ensure_pydantic_import(out, "model_validator")
1514+
1515+
13951516
def find_unique_items_fields(schema_dir):
13961517
"""Map generated class names to fields carrying ``uniqueItems``.
13971518
@@ -1791,6 +1912,48 @@ def _patch_conditional_array_retyping():
17911912
return patched, 0
17921913

17931914

1915+
def _patch_dependent_required():
1916+
"""Inject dependentRequired validators; return counts and status."""
1917+
rules_by_class = find_root_dependent_required(SCHEMA_DIR)
1918+
if not rules_by_class:
1919+
sys.stdout.write(
1920+
"postprocess: no dependentRequired constraints found\n"
1921+
)
1922+
return 0, 0
1923+
patched = 0
1924+
for class_name, rules in sorted(rules_by_class.items()):
1925+
hits = []
1926+
for path in sorted(OUTPUT_DIR.rglob("*.py")):
1927+
source = path.read_text(encoding="utf-8")
1928+
match = re.search(
1929+
rf"^class {re.escape(class_name)}\(", source, re.M
1930+
)
1931+
if match is None:
1932+
continue
1933+
tail = re.compile(r"^\S", re.M)
1934+
end_match = tail.search(source, match.end())
1935+
end = end_match.start() if end_match else len(source)
1936+
if (
1937+
f"def {_DEPENDENT_REQUIRED_MARKER}("
1938+
in source[match.start() : end]
1939+
):
1940+
hits.append(path)
1941+
continue
1942+
updated = inject_dependent_required(source, class_name, rules)
1943+
if updated == source:
1944+
continue
1945+
path.write_text(updated, encoding="utf-8")
1946+
patched += 1
1947+
hits.append(path)
1948+
label = (
1949+
", ".join(str(path) for path in hits) or "NO APPLICABLE CLASS FOUND"
1950+
)
1951+
sys.stdout.write(f" dependentRequired on '{class_name}' -> {label}\n")
1952+
if not hits:
1953+
return patched, 1
1954+
return patched, 0
1955+
1956+
17941957
def _patch_unique_items():
17951958
"""Inject uniqueItems validators; return (patched_count, exit_code)."""
17961959
unique_fields_by_class = find_unique_items_fields(SCHEMA_DIR)
@@ -1933,6 +2096,7 @@ def main():
19332096
patched_cr, rc_cr = _patch_conditional_required()
19342097
patched_cb, rc_cb = _patch_conditional_bounds()
19352098
patched_rt, rc_rt = _patch_conditional_array_retyping()
2099+
patched_dr, rc_dr = _patch_dependent_required()
19362100
patched_ui, rc_ui = _patch_unique_items()
19372101
patched_ef, rc_ef = _patch_extra_forbid()
19382102
total = (
@@ -1943,6 +2107,7 @@ def main():
19432107
+ patched_cr
19442108
+ patched_cb
19452109
+ patched_rt
2110+
+ patched_dr
19462111
+ patched_ui
19472112
+ patched_ef
19482113
)
@@ -1955,6 +2120,7 @@ def main():
19552120
or rc_cr
19562121
or rc_cb
19572122
or rc_rt
2123+
or rc_dr
19582124
or rc_ui
19592125
or rc_ef
19602126
)

src/ucp_sdk/models/schemas/common/types/time_interval.py

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818

1919
from __future__ import annotations
2020

21-
from pydantic import BaseModel, ConfigDict, Field
21+
from pydantic import BaseModel, ConfigDict, Field, model_validator
2222

2323

2424
class TimeInterval(BaseModel):
@@ -37,3 +37,19 @@ class TimeInterval(BaseModel):
3737
"""
3838
Closing time in 24-hour HH:MM format.
3939
"""
40+
41+
@model_validator(mode="after")
42+
def _enforce_dependent_required(self):
43+
"""JSON Schema dependentRequired: enforce dependent fields."""
44+
rules = {"opens": ["closes"], "closes": ["opens"]}
45+
provided = self.model_fields_set | set(self.model_extra or {})
46+
for field, required_fields in rules.items():
47+
if field not in provided:
48+
continue
49+
for required in required_fields:
50+
if required not in provided:
51+
raise ValueError(
52+
f"Field {required!r} is required when {field!r} "
53+
"is provided (schema dependentRequired)"
54+
)
55+
return self

src/ucp_sdk/models/schemas/shopping/types/fulfillment_method.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -109,3 +109,19 @@ def _enforce_conditional_item_retyping(self):
109109
f"when {rule['discriminator']} is {actual!r}"
110110
)
111111
return self
112+
113+
@model_validator(mode="after")
114+
def _enforce_dependent_required(self):
115+
"""JSON Schema dependentRequired: enforce dependent fields."""
116+
rules = {"destinations": ["type"]}
117+
provided = self.model_fields_set | set(self.model_extra or {})
118+
for field, required_fields in rules.items():
119+
if field not in provided:
120+
continue
121+
for required in required_fields:
122+
if required not in provided:
123+
raise ValueError(
124+
f"Field {required!r} is required when {field!r} "
125+
"is provided (schema dependentRequired)"
126+
)
127+
return self

0 commit comments

Comments
 (0)