"""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', ]