Skip to content

Commit 23341fe

Browse files
committed
handle forward ref and self-reference in JsonSchemaGenerator, add rule resolve_forward_refs for __origin__, add volatile param for TypeRegistry, fix NotImplementedError in Schema cls
1 parent 6df6b4e commit 23341fe

5 files changed

Lines changed: 184 additions & 56 deletions

File tree

tests/test_spec.py

Lines changed: 77 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,90 @@
1+
import utype
12
from utype.types import *
23
from utype.parser.rule import Rule
34

45

6+
class TargetSchema(utype.Schema):
7+
name: str
8+
ref: Optional['RefSchema'] = None
9+
ref_values: List['RefSchema'] = utype.Field(default_factory=list)
10+
11+
12+
class RefSchema(utype.Schema):
13+
value: int = None
14+
15+
16+
class InfiniteSchema(utype.Schema):
17+
name: str
18+
self: List['InfiniteSchema'] = utype.Field(default_factory=list)
19+
20+
521
class TestSpec:
622
def test_json_schema_parser(self):
723
from utype.specs.json_schema.parser import JsonSchemaParser
8-
from utype.specs.python.generator import PythonCodeGenerator
924
assert JsonSchemaParser({})() == Any
1025
assert JsonSchemaParser({'anyOf': [{}, {'type': 'null'}]})() in (Rule, Any)
1126
assert JsonSchemaParser({'type': 'object'})() == dict
1227
assert JsonSchemaParser({'type': 'array'})() == list
1328
assert JsonSchemaParser({'type': 'string'})() == str
1429
assert JsonSchemaParser({'type': 'string', 'format': 'date'})() == date
30+
31+
def test_schema_generator(self):
32+
class TestSchema(utype.Schema):
33+
int_val: int
34+
str_val: str
35+
bytes_val: bytes
36+
float_val: float
37+
bool_val: bool
38+
uuid_val: UUID
39+
# test nest types
40+
list_val: List[str] = utype.Field(default_factory=list) # test callable default
41+
union_val: Union[int, List[int]] = utype.Field(default_factory=list)
42+
43+
from utype.specs.json_schema.generator import JsonSchemaGenerator
44+
output = JsonSchemaGenerator(TestSchema)()
45+
assert output == {'type': 'object',
46+
'properties': {'int_val': {'type': 'integer'},
47+
'str_val': {'type': 'string'},
48+
'bytes_val': {'type': 'string', 'format': 'binary'},
49+
'float_val': {'type': 'number', 'format': 'float'},
50+
'bool_val': {'type': 'boolean'},
51+
'uuid_val': {'type': 'string', 'format': 'uuid'},
52+
'list_val': {'type': 'array', 'items': {'type': 'string'}},
53+
'union_val': {'anyOf': [{'type': 'integer'},
54+
{'type': 'array', 'items': {'type': 'integer'}}]}},
55+
'required': ['int_val',
56+
'str_val',
57+
'bytes_val',
58+
'float_val',
59+
'bool_val',
60+
'uuid_val']}
61+
62+
def test_forward_ref_generator(self):
63+
from utype.specs.json_schema.generator import JsonSchemaGenerator
64+
output = JsonSchemaGenerator(TargetSchema)()
65+
assert output == {'type': 'object',
66+
'properties': {'name': {'type': 'string'},
67+
'ref': {'anyOf': [{'type': 'object',
68+
'properties': {'value': {'type': 'integer'}}},
69+
{'type': 'null'}]},
70+
'ref_values': {'type': 'array',
71+
'items': {'type': 'object',
72+
'properties': {'value': {'type': 'integer'}}}}},
73+
'required': ['name']}
74+
75+
def test_recursive_generator(self):
76+
from utype.specs.json_schema.generator import JsonSchemaGenerator
77+
refs = {}
78+
ref_output = JsonSchemaGenerator(InfiniteSchema, defs=refs)()
79+
assert ref_output == {'$ref': '#/$defs/InfiniteSchema'}
80+
assert refs == {InfiniteSchema: {'type': 'object',
81+
'properties': {'name': {'type': 'string'},
82+
'self': {'type': 'array',
83+
'items': {'$ref': '#/$defs/InfiniteSchema'}}},
84+
'required': ['name']}}
85+
86+
output = JsonSchemaGenerator(InfiniteSchema)()
87+
assert output == {'type': 'object',
88+
'properties': {'name': {'type': 'string'},
89+
'self': {'type': 'array', 'items': {'$ref': '#/$defs/InfiniteSchema'}}},
90+
'required': ['name']}

utype/parser/rule.py

Lines changed: 58 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -177,19 +177,31 @@ def _parse_arg(mcs, arg):
177177
return Rule.annotate(__origin, *_new_args)
178178

179179
def resolve_forward_refs(cls):
180-
if not cls.combinator:
181-
return
180+
# if not cls.combinator:
181+
# return
182+
origin_resolved = False
183+
arg_resolved = False
184+
origin = getattr(cls, "__origin__", None)
185+
186+
if origin is not None:
187+
origin, origin_resolved = resolve_forward_type(origin)
188+
189+
if origin_resolved:
190+
setattr(cls, "__origin__", cls._parse_arg(origin))
191+
182192
args = []
183-
resolved = False
184193
for i, arg in enumerate(cls.args):
185194
arg, resolved = resolve_forward_type(arg)
186195
if resolved:
187196
arg = cls._parse_arg(arg)
197+
arg_resolved = True
188198
args.append(arg)
189-
if resolved:
199+
200+
if arg_resolved:
190201
# only adjust args if resolved
191202
setattr(cls, "__args__", tuple(args))
192-
return resolved
203+
204+
return origin_resolved or arg_resolved
193205

194206
def register_forward_refs(
195207
cls,
@@ -1322,6 +1334,7 @@ def annotate(
13221334
constraints: Dict[str, Any] = None,
13231335
global_vars: Dict[str, Any] = None,
13241336
forward_refs=None,
1337+
forward_key=None,
13251338
force_clear_refs=False,
13261339
options=None,
13271340
bound=None,
@@ -1378,6 +1391,7 @@ def annotate(
13781391
global_vars=global_vars,
13791392
forward_refs=forward_refs,
13801393
force_clear_refs=force_clear_refs,
1394+
forward_key=forward_key,
13811395
bound=bound
13821396
)
13831397
# this annotation can be a ForwardRef
@@ -1508,6 +1522,7 @@ def parse_annotation(
15081522
forward_refs=forward_refs,
15091523
global_vars=global_vars,
15101524
force_clear_refs=force_clear_refs,
1525+
forward_key=forward_key,
15111526
bound=bound
15121527
)
15131528
elif annotation:
@@ -1533,6 +1548,7 @@ def parse_annotation(
15331548
forward_refs=forward_refs,
15341549
global_vars=global_vars,
15351550
force_clear_refs=force_clear_refs,
1551+
forward_key=forward_key,
15361552
bound=bound
15371553
)
15381554
return None
@@ -1845,34 +1861,42 @@ def _parse_contains(cls, value, context: RuntimeContext):
18451861
@classmethod
18461862
def resolve_forward_refs(cls):
18471863
# an override version of LogicalType.resolve_forward_refs
1848-
if not cls.__args__:
1849-
return False
1850-
args = []
1851-
arg_transformers = []
1852-
resolved = False
1853-
for arg, trans in zip(cls.__args__, cls.__arg_transformers__):
1854-
if isinstance(arg, LogicalType):
1855-
# including the Rule class and LogicalType with combinator
1856-
if arg.resolve_forward_refs():
1857-
resolved = True
1858-
elif isinstance(arg, ForwardRef):
1859-
if arg.__forward_evaluated__:
1860-
arg = arg.__forward_value__
1861-
resolved = True
1862-
transformer = cls.transformer_cls.resolver_transformer(arg)
1863-
if not transformer:
1864-
warning_settings.warn(
1865-
f"{cls}: arg type: {arg} got no transformer resolved, "
1866-
f"will just pass {arg}(data) at runtime",
1867-
warning_settings.rule_no_arg_transformer
1868-
)
1869-
trans = transformer or trans
1870-
args.append(arg)
1871-
arg_transformers.append(trans)
1872-
if resolved:
1873-
cls.__args__ = tuple(args)
1874-
cls.__arg_transformers__ = tuple(arg_transformers)
1875-
return resolved
1864+
origin_resolved = False
1865+
args_resolved = False
1866+
1867+
if cls.__origin__:
1868+
origin, origin_resolved = resolve_forward_type(cls.__origin__)
1869+
if origin_resolved:
1870+
cls.__origin__ = origin
1871+
1872+
if cls.__args__:
1873+
args = []
1874+
arg_transformers = []
1875+
1876+
for arg, trans in zip(cls.__args__, cls.__arg_transformers__):
1877+
if isinstance(arg, LogicalType):
1878+
# including the Rule class and LogicalType with combinator
1879+
if arg.resolve_forward_refs():
1880+
args_resolved = True
1881+
elif isinstance(arg, ForwardRef):
1882+
if arg.__forward_evaluated__:
1883+
arg = arg.__forward_value__
1884+
args_resolved = True
1885+
transformer = cls.transformer_cls.resolver_transformer(arg)
1886+
if not transformer:
1887+
warning_settings.warn(
1888+
f"{cls}: arg type: {arg} got no transformer resolved, "
1889+
f"will just pass {arg}(data) at runtime",
1890+
warning_settings.rule_no_arg_transformer
1891+
)
1892+
trans = transformer or trans
1893+
args.append(arg)
1894+
arg_transformers.append(trans)
1895+
if args_resolved:
1896+
cls.__args__ = tuple(args)
1897+
cls.__arg_transformers__ = tuple(arg_transformers)
1898+
1899+
return origin_resolved or args_resolved
18761900

18771901
@classmethod
18781902
def resolve_args_parser(cls):
@@ -2058,6 +2082,6 @@ def transform_callable(transformer: TypeTransformer, value, t):
20582082
return value
20592083

20602084

2061-
@TypeTransformer.registry.register(metaclass=LogicalType)
2085+
@TypeTransformer.registry.register(metaclass=LogicalType, volatile=True)
20622086
def transform_rule(transformer: TypeTransformer, value, t: LogicalType):
20632087
return t(value, context=transformer.context)

utype/schema.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -95,7 +95,7 @@ def __name__(self):
9595
return self.__parser__.name
9696

9797
def __class_getitem__(cls, item):
98-
raise NotImplemented
98+
raise NotImplementedError
9999

100100
def __validate__(self):
101101
pass
@@ -184,7 +184,7 @@ def __class_getitem__(cls, item):
184184
# class _cls(cls):
185185
# __options__ = item
186186
# return _cls
187-
raise NotImplemented
187+
raise NotImplementedError
188188

189189
def __validate__(self):
190190
pass

0 commit comments

Comments
 (0)