codekingpro/portable-devtools
114k
1"""Utility functions for pydantic-settings sources."""2 3from __future__ import annotations as _annotations4 5from collections import deque6from collections.abc import Mapping, Sequence7from dataclasses import is_dataclass8from enum import Enum9from typing import Any, TypeVar, cast, get_args, get_origin10 11from pydantic import BaseModel, Json, RootModel, Secret12from pydantic._internal._utils import is_model_class13from pydantic.dataclasses import is_pydantic_dataclass14from pydantic.fields import FieldInfo15from pydantic.types import Strict16from typing_inspection import typing_objects17from typing_inspection.introspection import is_union_origin18 19from ..exceptions import SettingsError20from ..utils import _lenient_issubclass21from .types import EnvNoneType22 23 24def _get_env_var_key(key: str, case_sensitive: bool = False) -> str:25 return key if case_sensitive else key.lower()26 27 28def _parse_env_none_str(value: str | None, parse_none_str: str | None = None) -> str | None | EnvNoneType:29 return value if not (value == parse_none_str and parse_none_str is not None) else EnvNoneType(value)30 31 32def parse_env_vars(33 env_vars: Mapping[str, str | None],34 case_sensitive: bool = False,35 ignore_empty: bool = False,36 parse_none_str: str | None = None,37) -> Mapping[str, str | None]:38 return {39 _get_env_var_key(k, case_sensitive): _parse_env_none_str(v, parse_none_str)40 for k, v in env_vars.items()41 if not (ignore_empty and v == '')42 }43 44 45def _substitute_typevars(tp: Any, param_map: dict[Any, Any]) -> Any:46 """Substitute TypeVars in a type annotation with concrete types from param_map."""47 if isinstance(tp, TypeVar) and tp in param_map:48 return param_map[tp]49 args = get_args(tp)50 if not args:51 return tp52 new_args = tuple(_substitute_typevars(arg, param_map) for arg in args)53 if new_args == args:54 return tp55 origin = get_origin(tp)56 if origin is not None:57 try:58 return origin[new_args]59 except TypeError:60 # types.UnionType and similar are not directly subscriptable,61 # reconstruct using | operator62 import functools63 import operator64 65 return functools.reduce(operator.or_, new_args)66 return tp67 68 69def _resolve_type_alias(annotation: Any) -> Any:70 """Resolve a TypeAliasType to its underlying value, substituting type params if parameterized."""71 if typing_objects.is_typealiastype(annotation):72 return annotation.__value__73 origin = get_origin(annotation)74 if typing_objects.is_typealiastype(origin):75 type_params = getattr(origin, '__type_params__', ())76 type_args = get_args(annotation)77 value = origin.__value__78 if type_params and type_args:79 return _substitute_typevars(value, dict(zip(type_params, type_args)))80 return value81 return annotation82 83 84def _annotation_is_complex(annotation: Any, metadata: list[Any]) -> bool:85 # If the model is a root model, the root annotation should be used to86 # evaluate the complexity.87 annotation = _resolve_type_alias(annotation)88 if annotation is not None and _lenient_issubclass(annotation, RootModel) and annotation is not RootModel:89 annotation = cast('type[RootModel[Any]]', annotation)90 root_annotation = annotation.model_fields['root'].annotation91 if root_annotation is not None: # pragma: no branch92 annotation = root_annotation93 94 if any(isinstance(md, Json) for md in metadata): # type: ignore[misc]95 return False96 97 origin = get_origin(annotation)98 99 # Check if annotation is of the form Annotated[type, metadata].100 if typing_objects.is_annotated(origin):101 # Return result of recursive call on inner type.102 inner, *meta = get_args(annotation)103 return _annotation_is_complex(inner, meta)104 105 if origin is Secret:106 return False107 108 return (109 _annotation_is_complex_inner(annotation)110 or _annotation_is_complex_inner(origin)111 or hasattr(origin, '__pydantic_core_schema__')112 or hasattr(origin, '__get_pydantic_core_schema__')113 )114 115 116def _get_field_metadata(field: FieldInfo) -> list[Any]:117 annotation = _resolve_type_alias(field.annotation)118 metadata = field.metadata119 origin = get_origin(annotation)120 if typing_objects.is_annotated(origin):121 _, *meta = get_args(annotation)122 metadata += meta123 return metadata124 125 126def _annotation_is_complex_inner(annotation: type[Any] | None) -> bool:127 if _lenient_issubclass(annotation, (str, bytes)):128 return False129 130 return _lenient_issubclass(131 annotation, (BaseModel, Mapping, Sequence, tuple, set, frozenset, deque)132 ) or is_dataclass(annotation)133 134 135def _union_is_complex(annotation: type[Any] | None, metadata: list[Any]) -> bool:136 """Check if a union type contains any complex types."""137 for arg in get_args(annotation):138 if _annotation_is_complex(arg, metadata):139 return True140 # _annotation_is_complex doesn't handle bare Union types, so when an arg141 # is Annotated[Union[X, Y], ...], stripping Annotated yields a bare Union142 # that _annotation_is_complex can't evaluate. Recurse into it, but only143 # if the Annotated metadata doesn't suppress complexity (e.g. Json).144 inner = _strip_annotated(arg)145 if inner is not arg:146 _, *inner_meta = get_args(arg)147 if any(isinstance(md, Json) for md in inner_meta): # type: ignore[misc]148 continue149 if is_union_origin(get_origin(inner)):150 if _union_is_complex(inner, metadata):151 return True152 return False153 154 155def _union_has_strict_types(annotation: type[Any] | None) -> bool:156 """Check if a union type contains any strict-annotated types."""157 for arg in get_args(annotation):158 if typing_objects.is_annotated(get_origin(arg)):159 _, *meta = get_args(arg)160 if any(isinstance(m, Strict) for m in meta):161 return True162 return False163 164 165def _annotation_contains_types(166 annotation: type[Any] | None,167 types: tuple[Any, ...],168 is_include_origin: bool = True,169 is_strip_annotated: bool = False,170 is_instance: bool = False,171 collect: set[Any] | None = None,172) -> bool:173 """Check if a type annotation contains any of the specified types."""174 if is_strip_annotated:175 annotation = _strip_annotated(annotation)176 if is_include_origin is True:177 origin = get_origin(annotation)178 if origin in types:179 if collect is None:180 return True181 collect.add(annotation)182 if is_instance and any(isinstance(origin, type_) for type_ in types):183 if collect is None:184 return True185 collect.add(annotation)186 for type_ in get_args(annotation):187 if (188 _annotation_contains_types(189 type_,190 types,191 is_include_origin=True,192 is_strip_annotated=is_strip_annotated,193 is_instance=is_instance,194 collect=collect,195 )196 and collect is None197 ):198 return True199 if is_instance and any(isinstance(annotation, type_) for type_ in types):200 if collect is None:201 return True202 collect.add(annotation)203 if annotation in types:204 if collect is not None:205 collect.add(annotation)206 return True207 return False208 209 210def _strip_annotated(annotation: Any) -> Any:211 if typing_objects.is_annotated(get_origin(annotation)):212 return annotation.__origin__213 else:214 return annotation215 216 217def _annotation_enum_val_to_name(annotation: type[Any] | None, value: Any) -> str | None:218 for type_ in (annotation, get_origin(annotation), *get_args(annotation)):219 if _lenient_issubclass(type_, Enum):220 if value in type_.__members__.values():221 return type_(value).name222 return None223 224 225def _annotation_enum_name_to_val(annotation: type[Any] | None, name: Any) -> Any:226 for type_ in (annotation, get_origin(annotation), *get_args(annotation)):227 if _lenient_issubclass(type_, Enum):228 if name in type_.__members__.keys():229 return type_[name]230 return None231 232 233def _literal_has_numeric_enum(annotation: type[Any] | None) -> bool:234 """Check if annotation is a Literal type containing numeric Enum members (IntEnum, (int, Enum), (float, Enum))."""235 if typing_objects.is_literal(get_origin(annotation)):236 return any(isinstance(arg, (int, float)) and isinstance(arg, Enum) for arg in get_args(annotation))237 # Handle Annotated wrapping, e.g. Annotated[Literal[IntEnum.member], Field(...)]238 if typing_objects.is_annotated(get_origin(annotation)):239 inner = get_args(annotation)[0]240 return _literal_has_numeric_enum(inner)241 # Handle Union/Optional wrapping, e.g. Optional[Literal[IntEnum.member]]242 if is_union_origin(get_origin(annotation)):243 return any(_literal_has_numeric_enum(arg) for arg in get_args(annotation))244 return False245 246 247def _get_model_fields(model_cls: type[Any]) -> dict[str, Any]:248 """Get fields from a pydantic model or dataclass."""249 250 if is_pydantic_dataclass(model_cls) and hasattr(model_cls, '__pydantic_fields__'):251 return model_cls.__pydantic_fields__252 if is_model_class(model_cls):253 return model_cls.model_fields254 raise SettingsError(f'Error: {model_cls.__name__} is not subclass of BaseModel or pydantic.dataclasses.dataclass')255 256 257def _get_alias_names(258 field_name: str,259 field_info: Any,260 alias_path_args: dict[str, int | None] | None = None,261 case_sensitive: bool = True,262 populate_by_name: bool = False,263) -> tuple[tuple[str, ...], bool]:264 """Get alias names for a field, handling alias paths and case sensitivity."""265 from pydantic import AliasChoices, AliasPath266 267 alias_names: list[str] = []268 is_alias_path_only: bool = True269 if not any((field_info.alias, field_info.validation_alias)):270 alias_names += [field_name]271 is_alias_path_only = False272 else:273 new_alias_paths: list[AliasPath] = []274 for alias in (field_info.alias, field_info.validation_alias):275 if alias is None:276 continue277 elif isinstance(alias, str):278 alias_names.append(alias)279 is_alias_path_only = False280 elif isinstance(alias, AliasChoices):281 for name in alias.choices:282 if isinstance(name, str):283 alias_names.append(name)284 is_alias_path_only = False285 else:286 new_alias_paths.append(name)287 else:288 new_alias_paths.append(alias)289 for alias_path in new_alias_paths:290 name = cast(str, alias_path.path[0])291 name = name.lower() if not case_sensitive else name292 if alias_path_args is not None:293 alias_path_args[name] = (294 alias_path.path[1] if len(alias_path.path) > 1 and isinstance(alias_path.path[1], int) else None295 )296 if not alias_names and is_alias_path_only:297 alias_names.append(name)298 if populate_by_name and field_name not in alias_names:299 alias_names.append(field_name)300 is_alias_path_only = False301 if not case_sensitive:302 alias_names = [alias_name.lower() for alias_name in alias_names]303 return tuple(dict.fromkeys(alias_names)), is_alias_path_only304 305 306def _is_function(obj: Any) -> bool:307 """Check if an object is a function."""308 from types import BuiltinFunctionType, FunctionType309 310 return isinstance(obj, (FunctionType, BuiltinFunctionType))311 312 313__all__ = [314 '_annotation_contains_types',315 '_annotation_enum_name_to_val',316 '_annotation_enum_val_to_name',317 '_annotation_is_complex',318 '_annotation_is_complex_inner',319 '_get_alias_names',320 '_get_env_var_key',321 '_get_model_fields',322 '_is_function',323 '_literal_has_numeric_enum',324 '_parse_env_none_str',325 '_resolve_type_alias',326 '_strip_annotated',327 '_union_has_strict_types',328 '_union_is_complex',329 'parse_env_vars',330]331 