ifixkart-backend/venv/lib/python3.12/site-packages/pydantic_settings/sources/utils.py

376 lines
14 KiB
Python

"""Utility functions for pydantic-settings sources."""
from __future__ import annotations as _annotations
import warnings
from collections import deque
from collections.abc import Mapping, Sequence
from dataclasses import is_dataclass
from enum import Enum
from typing import Any, TypedDict, TypeVar, cast, get_args, get_origin
from pydantic import BaseModel, Json, RootModel, Secret
from pydantic._internal._utils import is_model_class
from pydantic.dataclasses import is_pydantic_dataclass
from pydantic.fields import FieldInfo
from pydantic.types import Strict
from typing_inspection import typing_objects
from typing_inspection.introspection import is_union_origin
from ..exceptions import IncompleteFieldDefinitionWarning, SettingsError
from ..utils import _lenient_issubclass
from .types import EnvNoneType
class InitState(TypedDict, total=False):
"""State shared between settings sources during a single settings resolution."""
field_info_ids: set[int]
"""The `id()`s of the incomplete `FieldInfo` instances that were already warned about."""
def _warn_if_field_info_incomplete(field_info: FieldInfo, field_name: str, init_state: InitState) -> None:
"""Warn if the field is incomplete, i.e. its annotation contains unresolved forward references.
An incomplete `FieldInfo` instance is unsafe to inspect — any of its attributes (annotation,
aliases, metadata, default) may rely on the unresolved annotation, so settings sources may
silently fail to resolve the field's value. Each instance is only warned about once per
`init_state`, so that a field accessed by multiple sources during a single settings
resolution doesn't emit duplicate warnings.
"""
if getattr(field_info, '_complete', True):
return
warned_ids = init_state.setdefault('field_info_ids', set())
if id(field_info) in warned_ids:
return
warned_ids.add(id(field_info))
warnings.warn(
f'Field {field_name!r} has an incomplete definition: its annotation contains an unresolved '
'forward reference, so settings sources may fail to correctly resolve its value. '
'Call `model_rebuild()` on the model where the field is defined, once all the referenced '
'types are defined.',
IncompleteFieldDefinitionWarning,
)
def _get_env_var_key(key: str, case_sensitive: bool = False) -> str:
return key if case_sensitive else key.lower()
def _parse_env_none_str(value: str | None, parse_none_str: str | None = None) -> str | None | EnvNoneType:
return value if not (value == parse_none_str and parse_none_str is not None) else EnvNoneType(value)
def parse_env_vars(
env_vars: Mapping[str, str | None],
case_sensitive: bool = False,
ignore_empty: bool = False,
parse_none_str: str | None = None,
) -> Mapping[str, str | None]:
return {
_get_env_var_key(k, case_sensitive): _parse_env_none_str(v, parse_none_str)
for k, v in env_vars.items()
if not (ignore_empty and v == '')
}
def _substitute_typevars(tp: Any, param_map: dict[Any, Any]) -> Any:
"""Substitute TypeVars in a type annotation with concrete types from param_map."""
if isinstance(tp, TypeVar) and tp in param_map:
return param_map[tp]
args = get_args(tp)
if not args:
return tp
new_args = tuple(_substitute_typevars(arg, param_map) for arg in args)
if new_args == args:
return tp
origin = get_origin(tp)
if origin is not None:
try:
return origin[new_args]
except TypeError:
# types.UnionType and similar are not directly subscriptable,
# reconstruct using | operator
import functools
import operator
return functools.reduce(operator.or_, new_args)
return tp
def _resolve_type_alias(annotation: Any) -> Any:
"""Resolve a TypeAliasType to its underlying value, substituting type params if parameterized."""
if typing_objects.is_typealiastype(annotation):
return annotation.__value__
origin = get_origin(annotation)
if typing_objects.is_typealiastype(origin):
type_params = getattr(origin, '__type_params__', ())
type_args = get_args(annotation)
value = origin.__value__
if type_params and type_args:
return _substitute_typevars(value, dict(zip(type_params, type_args)))
return value
return annotation
def _annotation_is_complex(annotation: Any, metadata: list[Any], init_state: InitState | None = None) -> bool:
# If the model is a root model, the root annotation should be used to
# evaluate the complexity.
annotation = _resolve_type_alias(annotation)
if annotation is not None and _lenient_issubclass(annotation, RootModel) and annotation is not RootModel:
annotation = cast('type[RootModel[Any]]', annotation)
root_field = annotation.model_fields['root']
if init_state is not None:
_warn_if_field_info_incomplete(root_field, f'{annotation.__name__}.root', init_state)
root_annotation = root_field.annotation
if root_annotation is not None: # pragma: no branch
annotation = root_annotation
if any(isinstance(md, Json) for md in metadata): # type: ignore[misc]
return False
origin = get_origin(annotation)
# Check if annotation is of the form Annotated[type, metadata].
if typing_objects.is_annotated(origin):
# Return result of recursive call on inner type.
inner, *meta = get_args(annotation)
return _annotation_is_complex(inner, meta, init_state)
if _lenient_issubclass(origin, Secret):
return False
return (
_annotation_is_complex_inner(annotation)
or _annotation_is_complex_inner(origin)
or hasattr(origin, '__pydantic_core_schema__')
or hasattr(origin, '__get_pydantic_core_schema__')
)
def _get_field_metadata(field: FieldInfo) -> list[Any]:
annotation = _resolve_type_alias(field.annotation)
metadata = field.metadata
origin = get_origin(annotation)
if typing_objects.is_annotated(origin):
_, *meta = get_args(annotation)
metadata += meta
return metadata
def _annotation_is_complex_inner(annotation: type[Any] | None) -> bool:
if _lenient_issubclass(annotation, (str, bytes)):
return False
return _lenient_issubclass(
annotation, (BaseModel, Mapping, Sequence, tuple, set, frozenset, deque)
) or is_dataclass(annotation)
def _union_is_complex(annotation: type[Any] | None, metadata: list[Any], init_state: InitState | None = None) -> bool:
"""Check if a union type contains any complex types."""
for arg in get_args(annotation):
if _annotation_is_complex(arg, metadata, init_state):
return True
# _annotation_is_complex doesn't handle bare Union types, so when an arg
# is Annotated[Union[X, Y], ...], stripping Annotated yields a bare Union
# that _annotation_is_complex can't evaluate. Recurse into it, but only
# if the Annotated metadata doesn't suppress complexity (e.g. Json).
inner = _strip_annotated(arg)
if inner is not arg:
_, *inner_meta = get_args(arg)
if any(isinstance(md, Json) for md in inner_meta): # type: ignore[misc]
continue
if is_union_origin(get_origin(inner)) and _union_is_complex(inner, metadata, init_state):
return True
return False
def _union_has_strict_types(annotation: type[Any] | None) -> bool:
"""Check if a union type contains any strict-annotated types."""
for arg in get_args(annotation):
if typing_objects.is_annotated(get_origin(arg)):
_, *meta = get_args(arg)
if any(isinstance(m, Strict) for m in meta):
return True
return False
def _annotation_contains_types(
annotation: type[Any] | None,
types: tuple[Any, ...],
is_include_origin: bool = True,
is_strip_annotated: bool = False,
is_instance: bool = False,
collect: set[Any] | None = None,
) -> bool:
"""Check if a type annotation contains any of the specified types."""
if is_strip_annotated:
annotation = _strip_annotated(annotation)
if is_include_origin is True:
origin = get_origin(annotation)
if origin in types:
if collect is None:
return True
collect.add(annotation)
if is_instance and any(isinstance(origin, type_) for type_ in types):
if collect is None:
return True
collect.add(annotation)
for type_ in get_args(annotation):
if (
_annotation_contains_types(
type_,
types,
is_include_origin=True,
is_strip_annotated=is_strip_annotated,
is_instance=is_instance,
collect=collect,
)
and collect is None
):
return True
if is_instance and any(isinstance(annotation, type_) for type_ in types):
if collect is None:
return True
collect.add(annotation)
if annotation in types:
if collect is not None:
collect.add(annotation)
return True
return False
def _strip_annotated(annotation: Any) -> Any:
if typing_objects.is_annotated(get_origin(annotation)):
return annotation.__origin__
else:
return annotation
def _annotation_enum_val_to_name(annotation: type[Any] | None, value: Any) -> str | None:
for type_ in (annotation, get_origin(annotation)):
if _lenient_issubclass(type_, Enum):
enum_type = cast('type[Enum]', type_)
if value in enum_type.__members__.values():
return enum_type(value).name
for arg in get_args(annotation):
enum_name = _annotation_enum_val_to_name(arg, value)
if enum_name is not None:
return enum_name
return None
def _annotation_enum_name_to_val(annotation: type[Any] | None, name: Any) -> Any:
for type_ in (annotation, get_origin(annotation)):
if _lenient_issubclass(type_, Enum):
enum_type = cast('type[Enum]', type_)
if name in enum_type.__members__:
return enum_type[name]
for arg in get_args(annotation):
enum_val = _annotation_enum_name_to_val(arg, name)
if enum_val is not None:
return enum_val
return None
def _literal_has_numeric_enum(annotation: type[Any] | None) -> bool:
"""Check if annotation is a Literal type containing numeric Enum members (IntEnum, (int, Enum), (float, Enum))."""
if typing_objects.is_literal(get_origin(annotation)):
return any(isinstance(arg, (int, float)) and isinstance(arg, Enum) for arg in get_args(annotation))
# Handle Annotated wrapping, e.g. Annotated[Literal[IntEnum.member], Field(...)]
if typing_objects.is_annotated(get_origin(annotation)):
inner = get_args(annotation)[0]
return _literal_has_numeric_enum(inner)
# Handle Union/Optional wrapping, e.g. Optional[Literal[IntEnum.member]]
if is_union_origin(get_origin(annotation)):
return any(_literal_has_numeric_enum(arg) for arg in get_args(annotation))
return False
def _get_model_fields(model_cls: type[Any]) -> dict[str, FieldInfo]:
"""Get fields from a pydantic model or dataclass."""
if is_pydantic_dataclass(model_cls) and hasattr(model_cls, '__pydantic_fields__'):
return model_cls.__pydantic_fields__
if is_model_class(model_cls):
return model_cls.model_fields
raise SettingsError(f'Error: {model_cls.__name__} is not subclass of BaseModel or pydantic.dataclasses.dataclass')
def _get_alias_names(
field_name: str,
field_info: FieldInfo,
alias_path_args: dict[str, int | None] | None = None,
case_sensitive: bool = True,
populate_by_name: bool = False,
) -> tuple[tuple[str, ...], bool]:
"""Get alias names for a field, handling alias paths and case sensitivity."""
from pydantic import AliasChoices, AliasPath
alias_names: list[str] = []
is_alias_path_only: bool = True
if not any((field_info.alias, field_info.validation_alias)):
alias_names += [field_name]
is_alias_path_only = False
else:
new_alias_paths: list[AliasPath] = []
for alias in (field_info.alias, field_info.validation_alias):
if alias is None:
continue
elif isinstance(alias, str):
alias_names.append(alias)
is_alias_path_only = False
elif isinstance(alias, AliasChoices):
for name in alias.choices:
if isinstance(name, str):
alias_names.append(name)
is_alias_path_only = False
else:
new_alias_paths.append(name)
else:
new_alias_paths.append(alias)
for alias_path in new_alias_paths:
name = cast(str, alias_path.path[0])
name = name.lower() if not case_sensitive else name
if alias_path_args is not None:
alias_path_args[name] = (
alias_path.path[1] if len(alias_path.path) > 1 and isinstance(alias_path.path[1], int) else None
)
if not alias_names and is_alias_path_only:
alias_names.append(name)
if populate_by_name and field_name not in alias_names:
alias_names.append(field_name)
is_alias_path_only = False
if not case_sensitive:
alias_names = [alias_name.lower() for alias_name in alias_names]
return tuple(dict.fromkeys(alias_names)), is_alias_path_only
def _is_function(obj: Any) -> bool:
"""Check if an object is a function."""
from types import BuiltinFunctionType, FunctionType
return isinstance(obj, (FunctionType, BuiltinFunctionType))
__all__ = [
'InitState',
'_annotation_contains_types',
'_annotation_enum_name_to_val',
'_annotation_enum_val_to_name',
'_annotation_is_complex',
'_annotation_is_complex_inner',
'_get_alias_names',
'_get_env_var_key',
'_get_model_fields',
'_is_function',
'_literal_has_numeric_enum',
'_parse_env_none_str',
'_resolve_type_alias',
'_strip_annotated',
'_union_has_strict_types',
'_union_is_complex',
'_warn_if_field_info_incomplete',
'parse_env_vars',
]