v2.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493
  1. import re
  2. import warnings
  3. from collections.abc import Sequence
  4. from copy import copy
  5. from dataclasses import dataclass, is_dataclass
  6. from enum import Enum
  7. from functools import lru_cache
  8. from typing import (
  9. Annotated,
  10. Any,
  11. Literal,
  12. Union,
  13. cast,
  14. get_args,
  15. get_origin,
  16. )
  17. from fastapi._compat import lenient_issubclass, shared
  18. from fastapi.openapi.constants import REF_TEMPLATE
  19. from fastapi.types import IncEx, ModelNameMap, UnionType
  20. from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, create_model
  21. from pydantic import PydanticSchemaGenerationError as PydanticSchemaGenerationError
  22. from pydantic import PydanticUndefinedAnnotation as PydanticUndefinedAnnotation
  23. from pydantic import ValidationError as ValidationError
  24. from pydantic._internal import _typing_extra as _pydantic_typing_extra
  25. from pydantic._internal._schema_generation_shared import ( # type: ignore[attr-defined]
  26. GetJsonSchemaHandler as GetJsonSchemaHandler,
  27. )
  28. from pydantic.fields import FieldInfo as FieldInfo
  29. from pydantic.json_schema import GenerateJsonSchema as _GenerateJsonSchema
  30. from pydantic.json_schema import JsonSchemaValue as JsonSchemaValue
  31. from pydantic_core import CoreSchema as CoreSchema
  32. from pydantic_core import PydanticUndefined
  33. from pydantic_core import Url as Url
  34. from pydantic_core.core_schema import (
  35. with_info_plain_validator_function as with_info_plain_validator_function,
  36. )
  37. RequiredParam = PydanticUndefined
  38. Undefined = PydanticUndefined
  39. def evaluate_forwardref(
  40. value: Any,
  41. globalns: dict[str, Any] | None = None,
  42. localns: dict[str, Any] | None = None,
  43. ) -> Any:
  44. # eval_type_lenient has been deprecated since Pydantic v2.10.0b1 (PR #10530)
  45. try_eval_type = getattr(_pydantic_typing_extra, "try_eval_type", None)
  46. if try_eval_type is not None:
  47. return try_eval_type(value, globalns, localns)[0]
  48. return _pydantic_typing_extra.eval_type_lenient( # ty: ignore[deprecated]
  49. value, globalns, localns
  50. )
  51. class GenerateJsonSchema(_GenerateJsonSchema):
  52. # TODO: remove when this is merged (or equivalent): https://github.com/pydantic/pydantic/pull/12841
  53. # and dropping support for any version of Pydantic before that one (so, in a very long time)
  54. def bytes_schema(self, schema: CoreSchema) -> JsonSchemaValue:
  55. json_schema = {"type": "string", "contentMediaType": "application/octet-stream"}
  56. bytes_mode = (
  57. self._config.ser_json_bytes
  58. if self.mode == "serialization"
  59. else self._config.val_json_bytes
  60. )
  61. if bytes_mode == "base64":
  62. json_schema["contentEncoding"] = "base64"
  63. self.update_with_validations(json_schema, schema, self.ValidationsMapping.bytes)
  64. return json_schema
  65. # TODO: remove when dropping support for Pydantic < v2.12.3
  66. _Attrs = {
  67. "default": ...,
  68. "default_factory": None,
  69. "alias": None,
  70. "alias_priority": None,
  71. "validation_alias": None,
  72. "serialization_alias": None,
  73. "title": None,
  74. "field_title_generator": None,
  75. "description": None,
  76. "examples": None,
  77. "exclude": None,
  78. "exclude_if": None,
  79. "discriminator": None,
  80. "deprecated": None,
  81. "json_schema_extra": None,
  82. "frozen": None,
  83. "validate_default": None,
  84. "repr": True,
  85. "init": None,
  86. "init_var": None,
  87. "kw_only": None,
  88. }
  89. # TODO: remove when dropping support for Pydantic < v2.12.3
  90. def asdict(field_info: FieldInfo) -> dict[str, Any]:
  91. attributes = {}
  92. for attr in _Attrs:
  93. value = getattr(field_info, attr, Undefined)
  94. if value is not Undefined:
  95. attributes[attr] = value
  96. return {
  97. "annotation": field_info.annotation,
  98. "metadata": field_info.metadata,
  99. "attributes": attributes,
  100. }
  101. @dataclass
  102. class ModelField:
  103. field_info: FieldInfo
  104. name: str
  105. mode: Literal["validation", "serialization"] = "validation"
  106. config: ConfigDict | None = None
  107. @property
  108. def alias(self) -> str:
  109. a = self.field_info.alias
  110. return a if a is not None else self.name
  111. @property
  112. def validation_alias(self) -> str | None:
  113. va = self.field_info.validation_alias
  114. if isinstance(va, str) and va:
  115. return va
  116. return None
  117. @property
  118. def serialization_alias(self) -> str | None:
  119. sa = self.field_info.serialization_alias
  120. return sa or None
  121. @property
  122. def default(self) -> Any:
  123. return self.get_default()
  124. def __post_init__(self) -> None:
  125. with warnings.catch_warnings():
  126. # Pydantic >= 2.12.0 warns about field specific metadata that is unused
  127. # (e.g. `TypeAdapter(Annotated[int, Field(alias='b')])`). In some cases, we
  128. # end up building the type adapter from a model field annotation so we
  129. # need to ignore the warning:
  130. if shared.PYDANTIC_VERSION_MINOR_TUPLE >= (2, 12):
  131. from pydantic.warnings import UnsupportedFieldAttributeWarning
  132. warnings.simplefilter(
  133. "ignore", category=UnsupportedFieldAttributeWarning
  134. )
  135. # TODO: remove after setting the min Pydantic to v2.12.3
  136. # that adds asdict(), and use self.field_info.asdict() instead
  137. field_dict = asdict(self.field_info)
  138. annotated_args = (
  139. field_dict["annotation"],
  140. *field_dict["metadata"],
  141. # this FieldInfo needs to be created again so that it doesn't include
  142. # the old field info metadata and only the rest of the attributes
  143. Field(**field_dict["attributes"]),
  144. )
  145. self._type_adapter: TypeAdapter[Any] = TypeAdapter(
  146. Annotated[annotated_args], # ty: ignore[invalid-type-form]
  147. config=self.config,
  148. )
  149. def get_default(self) -> Any:
  150. if self.field_info.is_required():
  151. return Undefined
  152. return self.field_info.get_default(call_default_factory=True)
  153. def validate(
  154. self,
  155. value: Any,
  156. values: dict[str, Any] = {}, # noqa: B006
  157. *,
  158. loc: tuple[int | str, ...] = (),
  159. ) -> tuple[Any, list[dict[str, Any]]]:
  160. try:
  161. return (
  162. self._type_adapter.validate_python(value, from_attributes=True),
  163. [],
  164. )
  165. except ValidationError as exc:
  166. return None, _regenerate_error_with_loc(
  167. errors=exc.errors(include_url=False), loc_prefix=loc
  168. )
  169. def serialize(
  170. self,
  171. value: Any,
  172. *,
  173. mode: Literal["json", "python"] = "json",
  174. include: IncEx | None = None,
  175. exclude: IncEx | None = None,
  176. by_alias: bool = True,
  177. exclude_unset: bool = False,
  178. exclude_defaults: bool = False,
  179. exclude_none: bool = False,
  180. ) -> Any:
  181. # What calls this code passes a value that already called
  182. # self._type_adapter.validate_python(value)
  183. return self._type_adapter.dump_python(
  184. value,
  185. mode=mode,
  186. include=include,
  187. exclude=exclude,
  188. by_alias=by_alias,
  189. exclude_unset=exclude_unset,
  190. exclude_defaults=exclude_defaults,
  191. exclude_none=exclude_none,
  192. )
  193. def serialize_json(
  194. self,
  195. value: Any,
  196. *,
  197. include: IncEx | None = None,
  198. exclude: IncEx | None = None,
  199. by_alias: bool = True,
  200. exclude_unset: bool = False,
  201. exclude_defaults: bool = False,
  202. exclude_none: bool = False,
  203. ) -> bytes:
  204. # What calls this code passes a value that already called
  205. # self._type_adapter.validate_python(value)
  206. # This uses Pydantic's dump_json() which serializes directly to JSON
  207. # bytes in one pass (via Rust), avoiding the intermediate Python dict
  208. # step of dump_python(mode="json") + json.dumps().
  209. return self._type_adapter.dump_json(
  210. value,
  211. include=include,
  212. exclude=exclude,
  213. by_alias=by_alias,
  214. exclude_unset=exclude_unset,
  215. exclude_defaults=exclude_defaults,
  216. exclude_none=exclude_none,
  217. )
  218. def __hash__(self) -> int:
  219. # Each ModelField is unique for our purposes, to allow making a dict from
  220. # ModelField to its JSON Schema.
  221. return id(self)
  222. def _has_computed_fields(field: ModelField) -> bool:
  223. computed_fields = field._type_adapter.core_schema.get("schema", {}).get(
  224. "computed_fields", []
  225. )
  226. return len(computed_fields) > 0
  227. def get_schema_from_model_field(
  228. *,
  229. field: ModelField,
  230. model_name_map: ModelNameMap,
  231. field_mapping: dict[
  232. tuple[ModelField, Literal["validation", "serialization"]], JsonSchemaValue
  233. ],
  234. separate_input_output_schemas: bool = True,
  235. ) -> dict[str, Any]:
  236. override_mode: Literal["validation"] | None = (
  237. None
  238. if (separate_input_output_schemas or _has_computed_fields(field))
  239. else "validation"
  240. )
  241. field_alias = (
  242. (field.validation_alias or field.alias)
  243. if field.mode == "validation"
  244. else (field.serialization_alias or field.alias)
  245. )
  246. # This expects that GenerateJsonSchema was already used to generate the definitions
  247. json_schema = field_mapping[(field, override_mode or field.mode)]
  248. if "$ref" not in json_schema:
  249. # TODO remove when deprecating Pydantic v1
  250. # Ref: https://github.com/pydantic/pydantic/blob/d61792cc42c80b13b23e3ffa74bc37ec7c77f7d1/pydantic/schema.py#L207
  251. json_schema["title"] = field.field_info.title or field_alias.title().replace(
  252. "_", " "
  253. )
  254. return json_schema
  255. def get_definitions(
  256. *,
  257. fields: Sequence[ModelField],
  258. model_name_map: ModelNameMap,
  259. separate_input_output_schemas: bool = True,
  260. ) -> tuple[
  261. dict[tuple[ModelField, Literal["validation", "serialization"]], JsonSchemaValue],
  262. dict[str, dict[str, Any]],
  263. ]:
  264. schema_generator = GenerateJsonSchema(ref_template=REF_TEMPLATE)
  265. validation_fields = [field for field in fields if field.mode == "validation"]
  266. serialization_fields = [field for field in fields if field.mode == "serialization"]
  267. flat_validation_models = get_flat_models_from_fields(
  268. validation_fields, known_models=set()
  269. )
  270. flat_serialization_models = get_flat_models_from_fields(
  271. serialization_fields, known_models=set()
  272. )
  273. flat_validation_model_fields = [
  274. ModelField(
  275. field_info=FieldInfo(annotation=model),
  276. name=model.__name__,
  277. mode="validation",
  278. )
  279. for model in flat_validation_models
  280. ]
  281. flat_serialization_model_fields = [
  282. ModelField(
  283. field_info=FieldInfo(annotation=model),
  284. name=model.__name__,
  285. mode="serialization",
  286. )
  287. for model in flat_serialization_models
  288. ]
  289. flat_model_fields = flat_validation_model_fields + flat_serialization_model_fields
  290. input_types = {f.field_info.annotation for f in fields}
  291. unique_flat_model_fields = {
  292. f for f in flat_model_fields if f.field_info.annotation not in input_types
  293. }
  294. inputs = [
  295. (
  296. field,
  297. (
  298. field.mode
  299. if (separate_input_output_schemas or _has_computed_fields(field))
  300. else "validation"
  301. ),
  302. field._type_adapter.core_schema,
  303. )
  304. for field in list(fields) + list(unique_flat_model_fields)
  305. ]
  306. field_mapping, definitions = schema_generator.generate_definitions(inputs=inputs)
  307. for item_def in cast(dict[str, dict[str, Any]], definitions).values():
  308. if "description" in item_def:
  309. item_description = cast(str, item_def["description"]).split("\f")[0]
  310. item_def["description"] = item_description
  311. # definitions: dict[DefsRef, dict[str, Any]]
  312. # but mypy complains about general str in other places that are not declared as
  313. # DefsRef, although DefsRef is just str:
  314. # DefsRef = NewType('DefsRef', str)
  315. # So, a cast to simplify the types here
  316. return field_mapping, cast(dict[str, dict[str, Any]], definitions)
  317. def is_scalar_field(field: ModelField) -> bool:
  318. from fastapi import params
  319. return shared.field_annotation_is_scalar(
  320. field.field_info.annotation
  321. ) and not isinstance(field.field_info, params.Body)
  322. def copy_field_info(*, field_info: FieldInfo, annotation: Any) -> FieldInfo:
  323. cls = type(field_info)
  324. merged_field_info = cls.from_annotation(annotation)
  325. new_field_info = copy(field_info)
  326. new_field_info.metadata = merged_field_info.metadata
  327. new_field_info.annotation = merged_field_info.annotation
  328. return new_field_info
  329. def serialize_sequence_value(*, field: ModelField, value: Any) -> Sequence[Any]:
  330. origin_type = get_origin(field.field_info.annotation) or field.field_info.annotation
  331. if origin_type is Union or origin_type is UnionType: # Handle optional sequences
  332. union_args = get_args(field.field_info.annotation)
  333. for union_arg in union_args:
  334. if union_arg is type(None):
  335. continue
  336. origin_type = get_origin(union_arg) or union_arg
  337. break
  338. assert issubclass(origin_type, shared.sequence_types) # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
  339. return shared.sequence_annotation_to_type[origin_type](value) # type: ignore[no-any-return,index] # ty: ignore[invalid-return-type]
  340. def get_missing_field_error(loc: tuple[int | str, ...]) -> dict[str, Any]:
  341. error = ValidationError.from_exception_data(
  342. "Field required", [{"type": "missing", "loc": loc, "input": {}}]
  343. ).errors(include_url=False)[0]
  344. error["input"] = None
  345. return error # type: ignore[return-value] # ty: ignore[invalid-return-type]
  346. def create_body_model(
  347. *, fields: Sequence[ModelField], model_name: str
  348. ) -> type[BaseModel]:
  349. field_params = {f.name: (f.field_info.annotation, f.field_info) for f in fields}
  350. BodyModel: type[BaseModel] = create_model(model_name, **field_params) # type: ignore[call-overload] # ty: ignore[no-matching-overload]
  351. return BodyModel
  352. def get_model_fields(model: type[BaseModel]) -> list[ModelField]:
  353. model_fields: list[ModelField] = []
  354. for name, field_info in model.model_fields.items():
  355. type_ = field_info.annotation
  356. if lenient_issubclass(type_, (BaseModel, dict)) or is_dataclass(type_):
  357. model_config = None
  358. else:
  359. model_config = model.model_config
  360. model_fields.append(
  361. ModelField(
  362. field_info=field_info,
  363. name=name,
  364. config=model_config,
  365. )
  366. )
  367. return model_fields
  368. @lru_cache
  369. def get_cached_model_fields(model: type[BaseModel]) -> list[ModelField]:
  370. return get_model_fields(model)
  371. # Duplicate of several schema functions from Pydantic v1 to make them compatible with
  372. # Pydantic v2 and allow mixing the models
  373. TypeModelOrEnum = type["BaseModel"] | type[Enum]
  374. TypeModelSet = set[TypeModelOrEnum]
  375. def normalize_name(name: str) -> str:
  376. return re.sub(r"[^a-zA-Z0-9.\-_]", "_", name)
  377. def get_model_name_map(unique_models: TypeModelSet) -> dict[TypeModelOrEnum, str]:
  378. name_model_map = {}
  379. for model in unique_models:
  380. model_name = normalize_name(model.__name__)
  381. name_model_map[model_name] = model
  382. return {v: k for k, v in name_model_map.items()}
  383. def get_flat_models_from_model(
  384. model: type["BaseModel"], known_models: TypeModelSet | None = None
  385. ) -> TypeModelSet:
  386. known_models = known_models or set()
  387. fields = get_model_fields(model)
  388. get_flat_models_from_fields(fields, known_models=known_models)
  389. return known_models
  390. def get_flat_models_from_annotation(
  391. annotation: Any, known_models: TypeModelSet
  392. ) -> TypeModelSet:
  393. origin = get_origin(annotation)
  394. if origin is not None:
  395. for arg in get_args(annotation):
  396. if lenient_issubclass(arg, (BaseModel, Enum)):
  397. if arg not in known_models:
  398. known_models.add(arg) # type: ignore[arg-type]
  399. if lenient_issubclass(arg, BaseModel):
  400. get_flat_models_from_model(arg, known_models=known_models)
  401. else:
  402. get_flat_models_from_annotation(arg, known_models=known_models)
  403. return known_models
  404. def get_flat_models_from_field(
  405. field: ModelField, known_models: TypeModelSet
  406. ) -> TypeModelSet:
  407. field_type = field.field_info.annotation
  408. if lenient_issubclass(field_type, BaseModel):
  409. if field_type in known_models:
  410. return known_models
  411. known_models.add(field_type)
  412. get_flat_models_from_model(field_type, known_models=known_models)
  413. elif lenient_issubclass(field_type, Enum):
  414. known_models.add(field_type)
  415. else:
  416. get_flat_models_from_annotation(field_type, known_models=known_models)
  417. return known_models
  418. def get_flat_models_from_fields(
  419. fields: Sequence[ModelField], known_models: TypeModelSet
  420. ) -> TypeModelSet:
  421. for field in fields:
  422. get_flat_models_from_field(field, known_models=known_models)
  423. return known_models
  424. def _regenerate_error_with_loc(
  425. *, errors: Sequence[Any], loc_prefix: tuple[str | int, ...]
  426. ) -> list[dict[str, Any]]:
  427. updated_loc_errors: list[Any] = [
  428. {**err, "loc": loc_prefix + err.get("loc", ())} for err in errors
  429. ]
  430. return updated_loc_errors