openai/openai-python

Public

mirrored from https://github.com/openai/openai-pythonAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
49e84301eccc8cb1028bea4c9e69456e2ebf5c2e

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

tests/test_models.py

1017lines · modecode

1import json
2from typing import TYPE_CHECKING, Any, Dict, List, Union, Iterable, Optional, cast
3from datetime import datetime, timezone
4from collections import deque
5from typing_extensions import Literal, Annotated, TypedDict, TypeAliasType
6
7import pytest
8import pydantic
9from pydantic import Field
10
11from openai._utils import PropertyInfo
12from openai._compat import PYDANTIC_V1, parse_obj, model_dump, model_json
13from openai._models import DISCRIMINATOR_CACHE, BaseModel, EagerIterable, construct_type
14
15
16class BasicModel(BaseModel):
17 foo: str
18
19
20@pytest.mark.parametrize("value", ["hello", 1], ids=["correct type", "mismatched"])
21def test_basic(value: object) -> None:
22 m = BasicModel.construct(foo=value)
23 assert m.foo == value
24
25
26def test_directly_nested_model() -> None:
27 class NestedModel(BaseModel):
28 nested: BasicModel
29
30 m = NestedModel.construct(nested={"foo": "Foo!"})
31 assert m.nested.foo == "Foo!"
32
33 # mismatched types
34 m = NestedModel.construct(nested="hello!")
35 assert cast(Any, m.nested) == "hello!"
36
37
38def test_optional_nested_model() -> None:
39 class NestedModel(BaseModel):
40 nested: Optional[BasicModel]
41
42 m1 = NestedModel.construct(nested=None)
43 assert m1.nested is None
44
45 m2 = NestedModel.construct(nested={"foo": "bar"})
46 assert m2.nested is not None
47 assert m2.nested.foo == "bar"
48
49 # mismatched types
50 m3 = NestedModel.construct(nested={"foo"})
51 assert isinstance(cast(Any, m3.nested), set)
52 assert cast(Any, m3.nested) == {"foo"}
53
54
55def test_list_nested_model() -> None:
56 class NestedModel(BaseModel):
57 nested: List[BasicModel]
58
59 m = NestedModel.construct(nested=[{"foo": "bar"}, {"foo": "2"}])
60 assert m.nested is not None
61 assert isinstance(m.nested, list)
62 assert len(m.nested) == 2
63 assert m.nested[0].foo == "bar"
64 assert m.nested[1].foo == "2"
65
66 # mismatched types
67 m = NestedModel.construct(nested=True)
68 assert cast(Any, m.nested) is True
69
70 m = NestedModel.construct(nested=[False])
71 assert cast(Any, m.nested) == [False]
72
73
74def test_optional_list_nested_model() -> None:
75 class NestedModel(BaseModel):
76 nested: Optional[List[BasicModel]]
77
78 m1 = NestedModel.construct(nested=[{"foo": "bar"}, {"foo": "2"}])
79 assert m1.nested is not None
80 assert isinstance(m1.nested, list)
81 assert len(m1.nested) == 2
82 assert m1.nested[0].foo == "bar"
83 assert m1.nested[1].foo == "2"
84
85 m2 = NestedModel.construct(nested=None)
86 assert m2.nested is None
87
88 # mismatched types
89 m3 = NestedModel.construct(nested={1})
90 assert cast(Any, m3.nested) == {1}
91
92 m4 = NestedModel.construct(nested=[False])
93 assert cast(Any, m4.nested) == [False]
94
95
96def test_list_optional_items_nested_model() -> None:
97 class NestedModel(BaseModel):
98 nested: List[Optional[BasicModel]]
99
100 m = NestedModel.construct(nested=[None, {"foo": "bar"}])
101 assert m.nested is not None
102 assert isinstance(m.nested, list)
103 assert len(m.nested) == 2
104 assert m.nested[0] is None
105 assert m.nested[1] is not None
106 assert m.nested[1].foo == "bar"
107
108 # mismatched types
109 m3 = NestedModel.construct(nested="foo")
110 assert cast(Any, m3.nested) == "foo"
111
112 m4 = NestedModel.construct(nested=[False])
113 assert cast(Any, m4.nested) == [False]
114
115
116def test_list_mismatched_type() -> None:
117 class NestedModel(BaseModel):
118 nested: List[str]
119
120 m = NestedModel.construct(nested=False)
121 assert cast(Any, m.nested) is False
122
123
124def test_raw_dictionary() -> None:
125 class NestedModel(BaseModel):
126 nested: Dict[str, str]
127
128 m = NestedModel.construct(nested={"hello": "world"})
129 assert m.nested == {"hello": "world"}
130
131 # mismatched types
132 m = NestedModel.construct(nested=False)
133 assert cast(Any, m.nested) is False
134
135
136def test_nested_dictionary_model() -> None:
137 class NestedModel(BaseModel):
138 nested: Dict[str, BasicModel]
139
140 m = NestedModel.construct(nested={"hello": {"foo": "bar"}})
141 assert isinstance(m.nested, dict)
142 assert m.nested["hello"].foo == "bar"
143
144 # mismatched types
145 m = NestedModel.construct(nested={"hello": False})
146 assert cast(Any, m.nested["hello"]) is False
147
148
149def test_unknown_fields() -> None:
150 m1 = BasicModel.construct(foo="foo", unknown=1)
151 assert m1.foo == "foo"
152 assert cast(Any, m1).unknown == 1
153
154 m2 = BasicModel.construct(foo="foo", unknown={"foo_bar": True})
155 assert m2.foo == "foo"
156 assert cast(Any, m2).unknown == {"foo_bar": True}
157
158 assert model_dump(m2) == {"foo": "foo", "unknown": {"foo_bar": True}}
159
160
161def test_strict_validation_unknown_fields() -> None:
162 class Model(BaseModel):
163 foo: str
164
165 model = parse_obj(Model, dict(foo="hello!", user="Robert"))
166 assert model.foo == "hello!"
167 assert cast(Any, model).user == "Robert"
168
169 assert model_dump(model) == {"foo": "hello!", "user": "Robert"}
170
171
172def test_aliases() -> None:
173 class Model(BaseModel):
174 my_field: int = Field(alias="myField")
175
176 m = Model.construct(myField=1)
177 assert m.my_field == 1
178
179 # mismatched types
180 m = Model.construct(myField={"hello": False})
181 assert cast(Any, m.my_field) == {"hello": False}
182
183
184def test_repr() -> None:
185 model = BasicModel(foo="bar")
186 assert str(model) == "BasicModel(foo='bar')"
187 assert repr(model) == "BasicModel(foo='bar')"
188
189
190def test_repr_nested_model() -> None:
191 class Child(BaseModel):
192 name: str
193 age: int
194
195 class Parent(BaseModel):
196 name: str
197 child: Child
198
199 model = Parent(name="Robert", child=Child(name="Foo", age=5))
200 assert str(model) == "Parent(name='Robert', child=Child(name='Foo', age=5))"
201 assert repr(model) == "Parent(name='Robert', child=Child(name='Foo', age=5))"
202
203
204def test_optional_list() -> None:
205 class Submodel(BaseModel):
206 name: str
207
208 class Model(BaseModel):
209 items: Optional[List[Submodel]]
210
211 m = Model.construct(items=None)
212 assert m.items is None
213
214 m = Model.construct(items=[])
215 assert m.items == []
216
217 m = Model.construct(items=[{"name": "Robert"}])
218 assert m.items is not None
219 assert len(m.items) == 1
220 assert m.items[0].name == "Robert"
221
222
223def test_nested_union_of_models() -> None:
224 class Submodel1(BaseModel):
225 bar: bool
226
227 class Submodel2(BaseModel):
228 thing: str
229
230 class Model(BaseModel):
231 foo: Union[Submodel1, Submodel2]
232
233 m = Model.construct(foo={"thing": "hello"})
234 assert isinstance(m.foo, Submodel2)
235 assert m.foo.thing == "hello"
236
237
238def test_nested_union_of_mixed_types() -> None:
239 class Submodel1(BaseModel):
240 bar: bool
241
242 class Model(BaseModel):
243 foo: Union[Submodel1, Literal[True], Literal["CARD_HOLDER"]]
244
245 m = Model.construct(foo=True)
246 assert m.foo is True
247
248 m = Model.construct(foo="CARD_HOLDER")
249 assert m.foo == "CARD_HOLDER"
250
251 m = Model.construct(foo={"bar": False})
252 assert isinstance(m.foo, Submodel1)
253 assert m.foo.bar is False
254
255
256def test_nested_union_multiple_variants() -> None:
257 class Submodel1(BaseModel):
258 bar: bool
259
260 class Submodel2(BaseModel):
261 thing: str
262
263 class Submodel3(BaseModel):
264 foo: int
265
266 class Model(BaseModel):
267 foo: Union[Submodel1, Submodel2, None, Submodel3]
268
269 m = Model.construct(foo={"thing": "hello"})
270 assert isinstance(m.foo, Submodel2)
271 assert m.foo.thing == "hello"
272
273 m = Model.construct(foo=None)
274 assert m.foo is None
275
276 m = Model.construct()
277 assert m.foo is None
278
279 m = Model.construct(foo={"foo": "1"})
280 assert isinstance(m.foo, Submodel3)
281 assert m.foo.foo == 1
282
283
284def test_nested_union_invalid_data() -> None:
285 class Submodel1(BaseModel):
286 level: int
287
288 class Submodel2(BaseModel):
289 name: str
290
291 class Model(BaseModel):
292 foo: Union[Submodel1, Submodel2]
293
294 m = Model.construct(foo=True)
295 assert cast(bool, m.foo) is True
296
297 m = Model.construct(foo={"name": 3})
298 if PYDANTIC_V1:
299 assert isinstance(m.foo, Submodel2)
300 assert m.foo.name == "3"
301 else:
302 assert isinstance(m.foo, Submodel1)
303 assert m.foo.name == 3 # type: ignore
304
305
306def test_list_of_unions() -> None:
307 class Submodel1(BaseModel):
308 level: int
309
310 class Submodel2(BaseModel):
311 name: str
312
313 class Model(BaseModel):
314 items: List[Union[Submodel1, Submodel2]]
315
316 m = Model.construct(items=[{"level": 1}, {"name": "Robert"}])
317 assert len(m.items) == 2
318 assert isinstance(m.items[0], Submodel1)
319 assert m.items[0].level == 1
320 assert isinstance(m.items[1], Submodel2)
321 assert m.items[1].name == "Robert"
322
323 m = Model.construct(items=[{"level": -1}, 156])
324 assert len(m.items) == 2
325 assert isinstance(m.items[0], Submodel1)
326 assert m.items[0].level == -1
327 assert cast(Any, m.items[1]) == 156
328
329
330def test_union_of_lists() -> None:
331 class SubModel1(BaseModel):
332 level: int
333
334 class SubModel2(BaseModel):
335 name: str
336
337 class Model(BaseModel):
338 items: Union[List[SubModel1], List[SubModel2]]
339
340 # with one valid entry
341 m = Model.construct(items=[{"name": "Robert"}])
342 assert len(m.items) == 1
343 assert isinstance(m.items[0], SubModel2)
344 assert m.items[0].name == "Robert"
345
346 # with two entries pointing to different types
347 m = Model.construct(items=[{"level": 1}, {"name": "Robert"}])
348 assert len(m.items) == 2
349 assert isinstance(m.items[0], SubModel1)
350 assert m.items[0].level == 1
351 assert isinstance(m.items[1], SubModel1)
352 assert cast(Any, m.items[1]).name == "Robert"
353
354 # with two entries pointing to *completely* different types
355 m = Model.construct(items=[{"level": -1}, 156])
356 assert len(m.items) == 2
357 assert isinstance(m.items[0], SubModel1)
358 assert m.items[0].level == -1
359 assert cast(Any, m.items[1]) == 156
360
361
362def test_dict_of_union() -> None:
363 class SubModel1(BaseModel):
364 name: str
365
366 class SubModel2(BaseModel):
367 foo: str
368
369 class Model(BaseModel):
370 data: Dict[str, Union[SubModel1, SubModel2]]
371
372 m = Model.construct(data={"hello": {"name": "there"}, "foo": {"foo": "bar"}})
373 assert len(list(m.data.keys())) == 2
374 assert isinstance(m.data["hello"], SubModel1)
375 assert m.data["hello"].name == "there"
376 assert isinstance(m.data["foo"], SubModel2)
377 assert m.data["foo"].foo == "bar"
378
379 # TODO: test mismatched type
380
381
382def test_double_nested_union() -> None:
383 class SubModel1(BaseModel):
384 name: str
385
386 class SubModel2(BaseModel):
387 bar: str
388
389 class Model(BaseModel):
390 data: Dict[str, List[Union[SubModel1, SubModel2]]]
391
392 m = Model.construct(data={"foo": [{"bar": "baz"}, {"name": "Robert"}]})
393 assert len(m.data["foo"]) == 2
394
395 entry1 = m.data["foo"][0]
396 assert isinstance(entry1, SubModel2)
397 assert entry1.bar == "baz"
398
399 entry2 = m.data["foo"][1]
400 assert isinstance(entry2, SubModel1)
401 assert entry2.name == "Robert"
402
403 # TODO: test mismatched type
404
405
406def test_union_of_dict() -> None:
407 class SubModel1(BaseModel):
408 name: str
409
410 class SubModel2(BaseModel):
411 foo: str
412
413 class Model(BaseModel):
414 data: Union[Dict[str, SubModel1], Dict[str, SubModel2]]
415
416 m = Model.construct(data={"hello": {"name": "there"}, "foo": {"foo": "bar"}})
417 assert len(list(m.data.keys())) == 2
418 assert isinstance(m.data["hello"], SubModel1)
419 assert m.data["hello"].name == "there"
420 assert isinstance(m.data["foo"], SubModel1)
421 assert cast(Any, m.data["foo"]).foo == "bar"
422
423
424def test_iso8601_datetime() -> None:
425 class Model(BaseModel):
426 created_at: datetime
427
428 expected = datetime(2019, 12, 27, 18, 11, 19, 117000, tzinfo=timezone.utc)
429
430 if PYDANTIC_V1:
431 expected_json = '{"created_at": "2019-12-27T18:11:19.117000+00:00"}'
432 else:
433 expected_json = '{"created_at":"2019-12-27T18:11:19.117000Z"}'
434
435 model = Model.construct(created_at="2019-12-27T18:11:19.117Z")
436 assert model.created_at == expected
437 assert model_json(model) == expected_json
438
439 model = parse_obj(Model, dict(created_at="2019-12-27T18:11:19.117Z"))
440 assert model.created_at == expected
441 assert model_json(model) == expected_json
442
443
444def test_does_not_coerce_int() -> None:
445 class Model(BaseModel):
446 bar: int
447
448 assert Model.construct(bar=1).bar == 1
449 assert Model.construct(bar=10.9).bar == 10.9
450 assert Model.construct(bar="19").bar == "19" # type: ignore[comparison-overlap]
451 assert Model.construct(bar=False).bar is False
452
453
454def test_int_to_float_safe_conversion() -> None:
455 class Model(BaseModel):
456 float_field: float
457
458 m = Model.construct(float_field=10)
459 assert m.float_field == 10.0
460 assert isinstance(m.float_field, float)
461
462 m = Model.construct(float_field=10.12)
463 assert m.float_field == 10.12
464 assert isinstance(m.float_field, float)
465
466 # number too big
467 m = Model.construct(float_field=2**53 + 1)
468 assert m.float_field == 2**53 + 1
469 assert isinstance(m.float_field, int)
470
471
472def test_deprecated_alias() -> None:
473 class Model(BaseModel):
474 resource_id: str = Field(alias="model_id")
475
476 @property
477 def model_id(self) -> str:
478 return self.resource_id
479
480 m = Model.construct(model_id="id")
481 assert m.model_id == "id"
482 assert m.resource_id == "id"
483 assert m.resource_id is m.model_id
484
485 m = parse_obj(Model, {"model_id": "id"})
486 assert m.model_id == "id"
487 assert m.resource_id == "id"
488 assert m.resource_id is m.model_id
489
490
491def test_omitted_fields() -> None:
492 class Model(BaseModel):
493 resource_id: Optional[str] = None
494
495 m = Model.construct()
496 assert m.resource_id is None
497 assert "resource_id" not in m.model_fields_set
498
499 m = Model.construct(resource_id=None)
500 assert m.resource_id is None
501 assert "resource_id" in m.model_fields_set
502
503 m = Model.construct(resource_id="foo")
504 assert m.resource_id == "foo"
505 assert "resource_id" in m.model_fields_set
506
507
508def test_to_dict() -> None:
509 class Model(BaseModel):
510 foo: Optional[str] = Field(alias="FOO", default=None)
511
512 m = Model(FOO="hello")
513 assert m.to_dict() == {"FOO": "hello"}
514 assert m.to_dict(use_api_names=False) == {"foo": "hello"}
515
516 m2 = Model()
517 assert m2.to_dict() == {}
518 assert m2.to_dict(exclude_unset=False) == {"FOO": None}
519 assert m2.to_dict(exclude_unset=False, exclude_none=True) == {}
520 assert m2.to_dict(exclude_unset=False, exclude_defaults=True) == {}
521
522 m3 = Model(FOO=None)
523 assert m3.to_dict() == {"FOO": None}
524 assert m3.to_dict(exclude_none=True) == {}
525 assert m3.to_dict(exclude_defaults=True) == {}
526
527 class Model2(BaseModel):
528 created_at: datetime
529
530 time_str = "2024-03-21T11:39:01.275859"
531 m4 = Model2.construct(created_at=time_str)
532 assert m4.to_dict(mode="python") == {"created_at": datetime.fromisoformat(time_str)}
533 assert m4.to_dict(mode="json") == {"created_at": time_str}
534
535 if PYDANTIC_V1:
536 with pytest.raises(ValueError, match="warnings is only supported in Pydantic v2"):
537 m.to_dict(warnings=False)
538
539
540def test_forwards_compat_model_dump_method() -> None:
541 class Model(BaseModel):
542 foo: Optional[str] = Field(alias="FOO", default=None)
543
544 m = Model(FOO="hello")
545 assert m.model_dump() == {"foo": "hello"}
546 assert m.model_dump(include={"bar"}) == {}
547 assert m.model_dump(exclude={"foo"}) == {}
548 assert m.model_dump(by_alias=True) == {"FOO": "hello"}
549
550 m2 = Model()
551 assert m2.model_dump() == {"foo": None}
552 assert m2.model_dump(exclude_unset=True) == {}
553 assert m2.model_dump(exclude_none=True) == {}
554 assert m2.model_dump(exclude_defaults=True) == {}
555
556 m3 = Model(FOO=None)
557 assert m3.model_dump() == {"foo": None}
558 assert m3.model_dump(exclude_none=True) == {}
559
560 if PYDANTIC_V1:
561 with pytest.raises(ValueError, match="round_trip is only supported in Pydantic v2"):
562 m.model_dump(round_trip=True)
563
564 with pytest.raises(ValueError, match="warnings is only supported in Pydantic v2"):
565 m.model_dump(warnings=False)
566
567
568def test_compat_method_no_error_for_warnings() -> None:
569 class Model(BaseModel):
570 foo: Optional[str]
571
572 m = Model(foo="hello")
573 assert isinstance(model_dump(m, warnings=False), dict)
574
575
576def test_to_json() -> None:
577 class Model(BaseModel):
578 foo: Optional[str] = Field(alias="FOO", default=None)
579
580 m = Model(FOO="hello")
581 assert json.loads(m.to_json()) == {"FOO": "hello"}
582 assert json.loads(m.to_json(use_api_names=False)) == {"foo": "hello"}
583
584 if PYDANTIC_V1:
585 assert m.to_json(indent=None) == '{"FOO": "hello"}'
586 else:
587 assert m.to_json(indent=None) == '{"FOO":"hello"}'
588
589 m2 = Model()
590 assert json.loads(m2.to_json()) == {}
591 assert json.loads(m2.to_json(exclude_unset=False)) == {"FOO": None}
592 assert json.loads(m2.to_json(exclude_unset=False, exclude_none=True)) == {}
593 assert json.loads(m2.to_json(exclude_unset=False, exclude_defaults=True)) == {}
594
595 m3 = Model(FOO=None)
596 assert json.loads(m3.to_json()) == {"FOO": None}
597 assert json.loads(m3.to_json(exclude_none=True)) == {}
598
599 if PYDANTIC_V1:
600 with pytest.raises(ValueError, match="warnings is only supported in Pydantic v2"):
601 m.to_json(warnings=False)
602
603
604def test_forwards_compat_model_dump_json_method() -> None:
605 class Model(BaseModel):
606 foo: Optional[str] = Field(alias="FOO", default=None)
607
608 m = Model(FOO="hello")
609 assert json.loads(m.model_dump_json()) == {"foo": "hello"}
610 assert json.loads(m.model_dump_json(include={"bar"})) == {}
611 assert json.loads(m.model_dump_json(include={"foo"})) == {"foo": "hello"}
612 assert json.loads(m.model_dump_json(by_alias=True)) == {"FOO": "hello"}
613
614 assert m.model_dump_json(indent=2) == '{\n "foo": "hello"\n}'
615
616 m2 = Model()
617 assert json.loads(m2.model_dump_json()) == {"foo": None}
618 assert json.loads(m2.model_dump_json(exclude_unset=True)) == {}
619 assert json.loads(m2.model_dump_json(exclude_none=True)) == {}
620 assert json.loads(m2.model_dump_json(exclude_defaults=True)) == {}
621
622 m3 = Model(FOO=None)
623 assert json.loads(m3.model_dump_json()) == {"foo": None}
624 assert json.loads(m3.model_dump_json(exclude_none=True)) == {}
625
626 if PYDANTIC_V1:
627 with pytest.raises(ValueError, match="round_trip is only supported in Pydantic v2"):
628 m.model_dump_json(round_trip=True)
629
630 with pytest.raises(ValueError, match="warnings is only supported in Pydantic v2"):
631 m.model_dump_json(warnings=False)
632
633
634def test_type_compat() -> None:
635 # our model type can be assigned to Pydantic's model type
636
637 def takes_pydantic(model: pydantic.BaseModel) -> None: # noqa: ARG001
638 ...
639
640 class OurModel(BaseModel):
641 foo: Optional[str] = None
642
643 takes_pydantic(OurModel())
644
645
646def test_annotated_types() -> None:
647 class Model(BaseModel):
648 value: str
649
650 m = construct_type(
651 value={"value": "foo"},
652 type_=cast(Any, Annotated[Model, "random metadata"]),
653 )
654 assert isinstance(m, Model)
655 assert m.value == "foo"
656
657
658def test_discriminated_unions_invalid_data() -> None:
659 class A(BaseModel):
660 type: Literal["a"]
661
662 data: str
663
664 class B(BaseModel):
665 type: Literal["b"]
666
667 data: int
668
669 m = construct_type(
670 value={"type": "b", "data": "foo"},
671 type_=cast(Any, Annotated[Union[A, B], PropertyInfo(discriminator="type")]),
672 )
673 assert isinstance(m, B)
674 assert m.type == "b"
675 assert m.data == "foo" # type: ignore[comparison-overlap]
676
677 m = construct_type(
678 value={"type": "a", "data": 100},
679 type_=cast(Any, Annotated[Union[A, B], PropertyInfo(discriminator="type")]),
680 )
681 assert isinstance(m, A)
682 assert m.type == "a"
683 if PYDANTIC_V1:
684 # pydantic v1 automatically converts inputs to strings
685 # if the expected type is a str
686 assert m.data == "100"
687 else:
688 assert m.data == 100 # type: ignore[comparison-overlap]
689
690
691def test_discriminated_unions_unknown_variant() -> None:
692 class A(BaseModel):
693 type: Literal["a"]
694
695 data: str
696
697 class B(BaseModel):
698 type: Literal["b"]
699
700 data: int
701
702 m = construct_type(
703 value={"type": "c", "data": None, "new_thing": "bar"},
704 type_=cast(Any, Annotated[Union[A, B], PropertyInfo(discriminator="type")]),
705 )
706
707 # just chooses the first variant
708 assert isinstance(m, A)
709 assert m.type == "c" # type: ignore[comparison-overlap]
710 assert m.data == None # type: ignore[unreachable]
711 assert m.new_thing == "bar"
712
713
714def test_discriminated_unions_invalid_data_nested_unions() -> None:
715 class A(BaseModel):
716 type: Literal["a"]
717
718 data: str
719
720 class B(BaseModel):
721 type: Literal["b"]
722
723 data: int
724
725 class C(BaseModel):
726 type: Literal["c"]
727
728 data: bool
729
730 m = construct_type(
731 value={"type": "b", "data": "foo"},
732 type_=cast(Any, Annotated[Union[Union[A, B], C], PropertyInfo(discriminator="type")]),
733 )
734 assert isinstance(m, B)
735 assert m.type == "b"
736 assert m.data == "foo" # type: ignore[comparison-overlap]
737
738 m = construct_type(
739 value={"type": "c", "data": "foo"},
740 type_=cast(Any, Annotated[Union[Union[A, B], C], PropertyInfo(discriminator="type")]),
741 )
742 assert isinstance(m, C)
743 assert m.type == "c"
744 assert m.data == "foo" # type: ignore[comparison-overlap]
745
746
747def test_discriminated_unions_with_aliases_invalid_data() -> None:
748 class A(BaseModel):
749 foo_type: Literal["a"] = Field(alias="type")
750
751 data: str
752
753 class B(BaseModel):
754 foo_type: Literal["b"] = Field(alias="type")
755
756 data: int
757
758 m = construct_type(
759 value={"type": "b", "data": "foo"},
760 type_=cast(Any, Annotated[Union[A, B], PropertyInfo(discriminator="foo_type")]),
761 )
762 assert isinstance(m, B)
763 assert m.foo_type == "b"
764 assert m.data == "foo" # type: ignore[comparison-overlap]
765
766 m = construct_type(
767 value={"type": "a", "data": 100},
768 type_=cast(Any, Annotated[Union[A, B], PropertyInfo(discriminator="foo_type")]),
769 )
770 assert isinstance(m, A)
771 assert m.foo_type == "a"
772 if PYDANTIC_V1:
773 # pydantic v1 automatically converts inputs to strings
774 # if the expected type is a str
775 assert m.data == "100"
776 else:
777 assert m.data == 100 # type: ignore[comparison-overlap]
778
779
780def test_discriminated_unions_overlapping_discriminators_invalid_data() -> None:
781 class A(BaseModel):
782 type: Literal["a"]
783
784 data: bool
785
786 class B(BaseModel):
787 type: Literal["a"]
788
789 data: int
790
791 m = construct_type(
792 value={"type": "a", "data": "foo"},
793 type_=cast(Any, Annotated[Union[A, B], PropertyInfo(discriminator="type")]),
794 )
795 assert isinstance(m, B)
796 assert m.type == "a"
797 assert m.data == "foo" # type: ignore[comparison-overlap]
798
799
800def test_discriminated_unions_invalid_data_uses_cache() -> None:
801 class A(BaseModel):
802 type: Literal["a"]
803
804 data: str
805
806 class B(BaseModel):
807 type: Literal["b"]
808
809 data: int
810
811 UnionType = cast(Any, Union[A, B])
812
813 assert not DISCRIMINATOR_CACHE.get(UnionType)
814
815 m = construct_type(
816 value={"type": "b", "data": "foo"}, type_=cast(Any, Annotated[UnionType, PropertyInfo(discriminator="type")])
817 )
818 assert isinstance(m, B)
819 assert m.type == "b"
820 assert m.data == "foo" # type: ignore[comparison-overlap]
821
822 discriminator = DISCRIMINATOR_CACHE.get(UnionType)
823 assert discriminator is not None
824
825 m = construct_type(
826 value={"type": "b", "data": "foo"}, type_=cast(Any, Annotated[UnionType, PropertyInfo(discriminator="type")])
827 )
828 assert isinstance(m, B)
829 assert m.type == "b"
830 assert m.data == "foo" # type: ignore[comparison-overlap]
831
832 # if the discriminator details object stays the same between invocations then
833 # we hit the cache
834 assert DISCRIMINATOR_CACHE.get(UnionType) is discriminator
835
836
837@pytest.mark.skipif(PYDANTIC_V1, reason="TypeAliasType is not supported in Pydantic v1")
838def test_type_alias_type() -> None:
839 Alias = TypeAliasType("Alias", str) # pyright: ignore
840
841 class Model(BaseModel):
842 alias: Alias
843 union: Union[int, Alias]
844
845 m = construct_type(value={"alias": "foo", "union": "bar"}, type_=Model)
846 assert isinstance(m, Model)
847 assert isinstance(m.alias, str)
848 assert m.alias == "foo"
849 assert isinstance(m.union, str)
850 assert m.union == "bar"
851
852
853@pytest.mark.skipif(PYDANTIC_V1, reason="TypeAliasType is not supported in Pydantic v1")
854def test_field_named_cls() -> None:
855 class Model(BaseModel):
856 cls: str
857
858 m = construct_type(value={"cls": "foo"}, type_=Model)
859 assert isinstance(m, Model)
860 assert isinstance(m.cls, str)
861
862
863def test_discriminated_union_case() -> None:
864 class A(BaseModel):
865 type: Literal["a"]
866
867 data: bool
868
869 class B(BaseModel):
870 type: Literal["b"]
871
872 data: List[Union[A, object]]
873
874 class ModelA(BaseModel):
875 type: Literal["modelA"]
876
877 data: int
878
879 class ModelB(BaseModel):
880 type: Literal["modelB"]
881
882 required: str
883
884 data: Union[A, B]
885
886 # when constructing ModelA | ModelB, value data doesn't match ModelB exactly - missing `required`
887 m = construct_type(
888 value={"type": "modelB", "data": {"type": "a", "data": True}},
889 type_=cast(Any, Annotated[Union[ModelA, ModelB], PropertyInfo(discriminator="type")]),
890 )
891
892 assert isinstance(m, ModelB)
893
894
895def test_nested_discriminated_union() -> None:
896 class InnerType1(BaseModel):
897 type: Literal["type_1"]
898
899 class InnerModel(BaseModel):
900 inner_value: str
901
902 class InnerType2(BaseModel):
903 type: Literal["type_2"]
904 some_inner_model: InnerModel
905
906 class Type1(BaseModel):
907 base_type: Literal["base_type_1"]
908 value: Annotated[
909 Union[
910 InnerType1,
911 InnerType2,
912 ],
913 PropertyInfo(discriminator="type"),
914 ]
915
916 class Type2(BaseModel):
917 base_type: Literal["base_type_2"]
918
919 T = Annotated[
920 Union[
921 Type1,
922 Type2,
923 ],
924 PropertyInfo(discriminator="base_type"),
925 ]
926
927 model = construct_type(
928 type_=T,
929 value={
930 "base_type": "base_type_1",
931 "value": {
932 "type": "type_2",
933 },
934 },
935 )
936 assert isinstance(model, Type1)
937 assert isinstance(model.value, InnerType2)
938
939
940@pytest.mark.skipif(PYDANTIC_V1, reason="this is only supported in pydantic v2 for now")
941def test_extra_properties() -> None:
942 class Item(BaseModel):
943 prop: int
944
945 class Model(BaseModel):
946 __pydantic_extra__: Dict[str, Item] = Field(init=False) # pyright: ignore[reportIncompatibleVariableOverride]
947
948 other: str
949
950 if TYPE_CHECKING:
951
952 def __getattr__(self, attr: str) -> Item: ...
953
954 model = construct_type(
955 type_=Model,
956 value={
957 "a": {"prop": 1},
958 "other": "foo",
959 },
960 )
961 assert isinstance(model, Model)
962 assert model.a.prop == 1
963 assert isinstance(model.a, Item)
964 assert model.other == "foo"
965
966
967# NOTE: Workaround for Pydantic Iterable behavior.
968# Iterable fields are replaced with a ValidatorIterator and may be consumed
969# during serialization, which can cause subsequent dumps to return empty data.
970# See: https://github.com/pydantic/pydantic/issues/9541
971@pytest.mark.parametrize(
972 "data, expected_validated",
973 [
974 ([1, 2, 3], [1, 2, 3]),
975 ((1, 2, 3), (1, 2, 3)),
976 (set([1, 2, 3]), set([1, 2, 3])),
977 (iter([1, 2, 3]), [1, 2, 3]),
978 ([], []),
979 ((x for x in [1, 2, 3]), [1, 2, 3]),
980 (map(lambda x: x, [1, 2, 3]), [1, 2, 3]),
981 (frozenset([1, 2, 3]), frozenset([1, 2, 3])),
982 (deque([1, 2, 3]), deque([1, 2, 3])),
983 ],
984 ids=["list", "tuple", "set", "iterator", "empty", "generator", "map", "frozenset", "deque"],
985)
986@pytest.mark.skipif(PYDANTIC_V1, reason="this is only supported in pydantic v2")
987def test_iterable_construction(data: Iterable[int], expected_validated: Iterable[int]) -> None:
988 class TypeWithIterable(TypedDict):
989 items: EagerIterable[int]
990
991 class Model(BaseModel):
992 data: TypeWithIterable
993
994 m = Model.model_validate({"data": {"items": data}})
995 assert m.data["items"] == expected_validated
996
997 # Verify repeated dumps don't lose data (the original bug)
998 assert m.model_dump()["data"]["items"] == list(expected_validated)
999 assert m.model_dump()["data"]["items"] == list(expected_validated)
1000
1001
1002@pytest.mark.skipif(PYDANTIC_V1, reason="this is only supported in pydantic v2")
1003def test_iterable_construction_str_falls_back_to_list() -> None:
1004 # str is iterable (over chars), but str(list_of_chars) produces the list's repr
1005 # rather than reconstructing a string from items. We special-case str to fall
1006 # back to list instead of attempting reconstruction.
1007 class TypeWithIterable(TypedDict):
1008 items: EagerIterable[str]
1009
1010 class Model(BaseModel):
1011 data: TypeWithIterable
1012
1013 m = Model.model_validate({"data": {"items": "hello"}})
1014
1015 # falls back to list of chars rather than calling str(["h", "e", "l", "l", "o"])
1016 assert m.data["items"] == ["h", "e", "l", "l", "o"]
1017 assert m.model_dump()["data"]["items"] == ["h", "e", "l", "l", "o"]