Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
base.py583 linesDownload Raw Back to sources
1"""Base classes and core functionality for pydantic-settings sources."""2 3from __future__ import annotations as _annotations4 5import json6from abc import ABC, abstractmethod7from collections.abc import Sequence8from dataclasses import asdict, is_dataclass9from pathlib import Path10from typing import TYPE_CHECKING, Any, cast, get_args11 12from pydantic import AliasChoices, AliasPath, BaseModel, TypeAdapter13from pydantic._internal._typing_extra import (  # type: ignore[attr-defined]14    get_origin,15)16from pydantic._internal._utils import deep_update, is_model_class17from pydantic.fields import FieldInfo18from typing_inspection.introspection import is_union_origin19 20from ..exceptions import SettingsError21from ..utils import _lenient_issubclass22from .types import EnvNoneType, EnvPrefixTarget, ForceDecode, NoDecode, PathType, PydanticModel, _CliSubCommand23from .utils import (24    _annotation_is_complex,25    _get_alias_names,26    _get_field_metadata,27    _get_model_fields,28    _resolve_type_alias,29    _strip_annotated,30    _union_is_complex,31)32 33if TYPE_CHECKING:34    from pydantic_settings.main import BaseSettings35 36 37def get_subcommand(38    model: PydanticModel,39    is_required: bool = True,40    cli_exit_on_error: bool | None = None,41    _suppress_errors: list[SettingsError | SystemExit] | None = None,42) -> PydanticModel | None:43    """44    Get the subcommand from a model.45 46    Args:47        model: The model to get the subcommand from.48        is_required: Determines whether a model must have subcommand set and raises error if not49            found. Defaults to `True`.50        cli_exit_on_error: Determines whether this function exits with error if no subcommand is found.51            Defaults to model_config `cli_exit_on_error` value if set. Otherwise, defaults to `True`.52 53    Returns:54        The subcommand model if found, otherwise `None`.55 56    Raises:57        SystemExit: When no subcommand is found and is_required=`True` and cli_exit_on_error=`True`58            (the default).59        SettingsError: When no subcommand is found and is_required=`True` and60            cli_exit_on_error=`False`.61    """62 63    model_cls = type(model)64    if cli_exit_on_error is None and is_model_class(model_cls):65        model_default = model_cls.model_config.get('cli_exit_on_error')66        if isinstance(model_default, bool):67            cli_exit_on_error = model_default68    if cli_exit_on_error is None:69        cli_exit_on_error = True70 71    subcommands: list[str] = []72    for field_name, field_info in _get_model_fields(model_cls).items():73        if _CliSubCommand in field_info.metadata:74            if getattr(model, field_name) is not None:75                return getattr(model, field_name)76            subcommands.append(field_name)77 78    if is_required:79        error_message = (80            f'Error: CLI subcommand is required {{{", ".join(subcommands)}}}'81            if subcommands82            else 'Error: CLI subcommand is required but no subcommands were found.'83        )84        err = SystemExit(error_message) if cli_exit_on_error else SettingsError(error_message)85        if _suppress_errors is None:86            raise err87        _suppress_errors.append(err)88 89    return None90 91 92class PydanticBaseSettingsSource(ABC):93    """94    Abstract base class for settings sources, every settings source classes should inherit from it.95    """96 97    def __init__(self, settings_cls: type[BaseSettings]):98        self.settings_cls = settings_cls99        self.config = settings_cls.model_config100        self._current_state: dict[str, Any] = {}101        self._settings_sources_data: dict[str, dict[str, Any]] = {}102 103    def _set_current_state(self, state: dict[str, Any]) -> None:104        """105        Record the state of settings from the previous settings sources. This should106        be called right before __call__.107        """108        self._current_state = state109 110    def _set_settings_sources_data(self, states: dict[str, dict[str, Any]]) -> None:111        """112        Record the state of settings from all previous settings sources. This should113        be called right before __call__.114        """115        self._settings_sources_data = states116 117    @property118    def current_state(self) -> dict[str, Any]:119        """120        The current state of the settings, populated by the previous settings sources.121        """122        return self._current_state123 124    @property125    def settings_sources_data(self) -> dict[str, dict[str, Any]]:126        """127        The state of all previous settings sources.128        """129        return self._settings_sources_data130 131    @abstractmethod132    def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]:133        """134        Gets the value, the key for model creation, and a flag to determine whether value is complex.135 136        This is an abstract method that should be overridden in every settings source classes.137 138        Args:139            field: The field.140            field_name: The field name.141 142        Returns:143            A tuple that contains the value, key and a flag to determine whether value is complex.144        """145        pass146 147    def field_is_complex(self, field: FieldInfo) -> bool:148        """149        Checks whether a field is complex, in which case it will attempt to be parsed as JSON.150 151        Args:152            field: The field.153 154        Returns:155            Whether the field is complex.156        """157        return _annotation_is_complex(field.annotation, field.metadata)158 159    def prepare_field_value(self, field_name: str, field: FieldInfo, value: Any, value_is_complex: bool) -> Any:160        """161        Prepares the value of a field.162 163        Args:164            field_name: The field name.165            field: The field.166            value: The value of the field that has to be prepared.167            value_is_complex: A flag to determine whether value is complex.168 169        Returns:170            The prepared value.171        """172        if value is not None and (self.field_is_complex(field) or value_is_complex):173            return self.decode_complex_value(field_name, field, value)174        return value175 176    def decode_complex_value(self, field_name: str, field: FieldInfo, value: Any) -> Any:177        """178        Decode the value for a complex field179 180        Args:181            field_name: The field name.182            field: The field.183            value: The value of the field that has to be prepared.184 185        Returns:186            The decoded value for further preparation187        """188        if field and (189            NoDecode in _get_field_metadata(field)190            or (self.config.get('enable_decoding') is False and ForceDecode not in field.metadata)191        ):192            return value193 194        return json.loads(value)195 196    @abstractmethod197    def __call__(self) -> dict[str, Any]:198        pass199 200 201class ConfigFileSourceMixin(ABC):202    def _read_files(self, files: PathType | None, deep_merge: bool = False) -> dict[str, Any]:203        if files is None:204            return {}205        if not isinstance(files, Sequence) or isinstance(files, str):206            files = [files]207        vars: dict[str, Any] = {}208        for file in files:209            if isinstance(file, str):210                file_path = Path(file)211            else:212                file_path = file213            if isinstance(file_path, Path):214                file_path = file_path.expanduser()215 216            if not file_path.is_file():217                continue218 219            updating_vars = self._read_file(file_path)220            if deep_merge:221                vars = deep_update(vars, updating_vars)222            else:223                vars.update(updating_vars)224        return vars225 226    @abstractmethod227    def _read_file(self, path: Path) -> dict[str, Any]:228        pass229 230 231class DefaultSettingsSource(PydanticBaseSettingsSource):232    """233    Source class for loading default object values.234 235    Args:236        settings_cls: The Settings class.237        nested_model_default_partial_update: Whether to allow partial updates on nested model default object fields.238            Defaults to `False`.239    """240 241    def __init__(self, settings_cls: type[BaseSettings], nested_model_default_partial_update: bool | None = None):242        super().__init__(settings_cls)243        self.defaults: dict[str, Any] = {}244        self.nested_model_default_partial_update = (245            nested_model_default_partial_update246            if nested_model_default_partial_update is not None247            else self.config.get('nested_model_default_partial_update', False)248        )249        if self.nested_model_default_partial_update:250            for field_name, field_info in settings_cls.model_fields.items():251                alias_names, *_ = _get_alias_names(field_name, field_info)252                preferred_alias = alias_names[0]253                if is_dataclass(type(field_info.default)):254                    self.defaults[preferred_alias] = asdict(field_info.default)255                elif is_model_class(type(field_info.default)):256                    self.defaults[preferred_alias] = field_info.default.model_dump()257 258    def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]:259        # Nothing to do here. Only implement the return statement to make mypy happy260        return None, '', False261 262    def __call__(self) -> dict[str, Any]:263        return self.defaults264 265    def __repr__(self) -> str:266        return (267            f'{self.__class__.__name__}(nested_model_default_partial_update={self.nested_model_default_partial_update})'268        )269 270 271class InitSettingsSource(PydanticBaseSettingsSource):272    """273    Source class for loading values provided during settings class initialization.274    """275 276    def __init__(277        self,278        settings_cls: type[BaseSettings],279        init_kwargs: dict[str, Any],280        nested_model_default_partial_update: bool | None = None,281    ):282        self.init_kwargs = {}283        init_kwarg_names = set(init_kwargs.keys())284        for field_name, field_info in settings_cls.model_fields.items():285            alias_names, *_ = _get_alias_names(field_name, field_info)286            # When populate_by_name is True, allow using the field name as an input key,287            # but normalize to the preferred alias to keep keys consistent across sources.288            matchable_names = set(alias_names)289            include_name = settings_cls.model_config.get('populate_by_name', False) or settings_cls.model_config.get(290                'validate_by_name', False291            )292            if include_name:293                matchable_names.add(field_name)294            init_kwarg_name = init_kwarg_names & matchable_names295            if init_kwarg_name:296                preferred_alias = alias_names[0] if alias_names else field_name297                # Choose provided key deterministically: prefer the first alias in alias_names order;298                # fall back to field_name if allowed and provided.299                provided_key = next((alias for alias in alias_names if alias in init_kwarg_names), None)300                if provided_key is None and include_name and field_name in init_kwarg_names:301                    provided_key = field_name302                # provided_key should not be None here because init_kwarg_name is non-empty303                assert provided_key is not None304                init_kwarg_names -= init_kwarg_name305                self.init_kwargs[preferred_alias] = init_kwargs[provided_key]306        # Include any remaining init kwargs (e.g., extras) unchanged307        # Note: If populate_by_name is True and the provided key is the field name, but308        # no alias exists, we keep it as-is so it can be processed as extra if allowed.309        self.init_kwargs.update({key: val for key, val in init_kwargs.items() if key in init_kwarg_names})310 311        super().__init__(settings_cls)312        self.nested_model_default_partial_update = (313            nested_model_default_partial_update314            if nested_model_default_partial_update is not None315            else self.config.get('nested_model_default_partial_update', False)316        )317 318    def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]:319        # Nothing to do here. Only implement the return statement to make mypy happy320        return None, '', False321 322    def __call__(self) -> dict[str, Any]:323        return (324            TypeAdapter(dict[str, Any]).dump_python(self.init_kwargs)325            if self.nested_model_default_partial_update326            else self.init_kwargs327        )328 329    def __repr__(self) -> str:330        return f'{self.__class__.__name__}(init_kwargs={self.init_kwargs!r})'331 332 333class PydanticBaseEnvSettingsSource(PydanticBaseSettingsSource):334    def __init__(335        self,336        settings_cls: type[BaseSettings],337        case_sensitive: bool | None = None,338        env_prefix: str | None = None,339        env_prefix_target: EnvPrefixTarget | None = None,340        env_ignore_empty: bool | None = None,341        env_parse_none_str: str | None = None,342        env_parse_enums: bool | None = None,343    ) -> None:344        super().__init__(settings_cls)345        self.case_sensitive = case_sensitive if case_sensitive is not None else self.config.get('case_sensitive', False)346        self.env_prefix = env_prefix if env_prefix is not None else self.config.get('env_prefix', '')347        self.env_prefix_target = (348            env_prefix_target if env_prefix_target is not None else self.config.get('env_prefix_target', 'variable')349        )350        self.env_ignore_empty = (351            env_ignore_empty if env_ignore_empty is not None else self.config.get('env_ignore_empty', False)352        )353        self.env_parse_none_str = (354            env_parse_none_str if env_parse_none_str is not None else self.config.get('env_parse_none_str')355        )356        self.env_parse_enums = env_parse_enums if env_parse_enums is not None else self.config.get('env_parse_enums')357 358    def _apply_case_sensitive(self, value: str) -> str:359        return value.lower() if not self.case_sensitive else value360 361    def _extract_field_info(self, field: FieldInfo, field_name: str) -> list[tuple[str, str, bool]]:362        """363        Extracts field info. This info is used to get the value of field from environment variables.364 365        It returns a list of tuples, each tuple contains:366            * field_key: The key of field that has to be used in model creation.367            * env_name: The environment variable name of the field.368            * value_is_complex: A flag to determine whether the value from environment variable369              is complex and has to be parsed.370 371        Args:372            field (FieldInfo): The field.373            field_name (str): The field name.374 375        Returns:376            list[tuple[str, str, bool]]: List of tuples, each tuple contains field_key, env_name, and value_is_complex.377        """378        field_info: list[tuple[str, str, bool]] = []379        if isinstance(field.validation_alias, (AliasChoices, AliasPath)):380            v_alias: str | list[str | int] | list[list[str | int]] | None = field.validation_alias.convert_to_aliases()381        else:382            v_alias = field.validation_alias383 384        if v_alias:385            env_prefix = self.env_prefix if self.env_prefix_target in ('alias', 'all') else ''386            if isinstance(v_alias, list):  # AliasChoices, AliasPath387                for alias in v_alias:388                    if isinstance(alias, str):  # AliasPath389                        field_info.append(390                            (alias, self._apply_case_sensitive(env_prefix + alias), True if len(alias) > 1 else False)391                        )392                    elif isinstance(alias, list):  # AliasChoices393                        first_arg = cast(str, alias[0])  # first item of an AliasChoices must be a str394                        field_info.append(395                            (396                                first_arg,397                                self._apply_case_sensitive(env_prefix + first_arg),398                                True if len(alias) > 1 else False,399                            )400                        )401            else:  # string validation alias402                field_info.append((v_alias, self._apply_case_sensitive(env_prefix + v_alias), False))403 404        if not v_alias or self.config.get('populate_by_name', False) or self.config.get('validate_by_name', False):405            annotation = _strip_annotated(_resolve_type_alias(field.annotation))406            env_prefix = self.env_prefix if self.env_prefix_target in ('variable', 'all') else ''407            if is_union_origin(get_origin(annotation)) and _union_is_complex(annotation, field.metadata):408                field_info.append((field_name, self._apply_case_sensitive(env_prefix + field_name), True))409            else:410                field_info.append((field_name, self._apply_case_sensitive(env_prefix + field_name), False))411 412        return field_info413 414    def _replace_field_names_case_insensitively(self, field: FieldInfo, field_values: dict[str, Any]) -> dict[str, Any]:415        """416        Replace field names in values dict by looking in models fields insensitively.417 418        By having the following models:419 420            ```py421            class SubSubSub(BaseModel):422                VaL3: str423 424            class SubSub(BaseModel):425                Val2: str426                SUB_sub_SuB: SubSubSub427 428            class Sub(BaseModel):429                VAL1: str430                SUB_sub: SubSub431 432            class Settings(BaseSettings):433                nested: Sub434 435                model_config = SettingsConfigDict(env_nested_delimiter='__')436            ```437 438        Then:439            _replace_field_names_case_insensitively(440                field,441                {"val1": "v1", "sub_SUB": {"VAL2": "v2", "sub_SUB_sUb": {"vAl3": "v3"}}}442            )443            Returns {'VAL1': 'v1', 'SUB_sub': {'Val2': 'v2', 'SUB_sub_SuB': {'VaL3': 'v3'}}}444        """445        values: dict[str, Any] = {}446 447        for name, value in field_values.items():448            sub_model_field: FieldInfo | None = None449 450            annotation = field.annotation451 452            # If field is Optional, we need to find the actual type453            if is_union_origin(get_origin(field.annotation)):454                args = get_args(annotation)455                if len(args) == 2 and type(None) in args:456                    for arg in args:457                        if arg is not None:458                            annotation = arg459                            break460 461            # This is here to make mypy happy462            # Item "None" of "Optional[Type[Any]]" has no attribute "model_fields"463            if not annotation or not hasattr(annotation, 'model_fields'):464                values[name] = value465                continue466            else:467                model_fields: dict[str, FieldInfo] = annotation.model_fields468 469            # Find field in sub model by looking in fields case insensitively470            field_key: str | None = None471            for sub_model_field_name, sub_model_field in model_fields.items():472                aliases, _ = _get_alias_names(sub_model_field_name, sub_model_field)473                _search = (alias for alias in aliases if alias.lower() == name.lower())474                if field_key := next(_search, None):475                    break476 477            if not field_key:478                values[name] = value479                continue480 481            if (482                sub_model_field is not None483                and _lenient_issubclass(sub_model_field.annotation, BaseModel)484                and isinstance(value, dict)485            ):486                values[field_key] = self._replace_field_names_case_insensitively(sub_model_field, value)487            else:488                values[field_key] = value489 490        return values491 492    def _replace_env_none_type_values(self, field_value: dict[str, Any]) -> dict[str, Any]:493        """494        Recursively parse values that are of "None" type(EnvNoneType) to `None` type(None).495        """496        values: dict[str, Any] = {}497 498        for key, value in field_value.items():499            if not isinstance(value, EnvNoneType):500                values[key] = value if not isinstance(value, dict) else self._replace_env_none_type_values(value)501            else:502                values[key] = None503 504        return values505 506    def _get_resolved_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]:507        """508        Gets the value, the preferred alias key for model creation, and a flag to determine whether value509        is complex.510 511        Note:512            In V3, this method should either be made public, or, this method should be removed and the513            abstract method get_field_value should be updated to include a "use_preferred_alias" flag.514 515        Args:516            field: The field.517            field_name: The field name.518 519        Returns:520            A tuple that contains the value, preferred key and a flag to determine whether value is complex.521        """522        field_value, field_key, value_is_complex = self.get_field_value(field, field_name)523        if not (524            value_is_complex525            or (526                (self.config.get('populate_by_name', False) or self.config.get('validate_by_name', False))527                and (field_key == field_name)528            )529        ):530            field_infos = self._extract_field_info(field, field_name)531            preferred_key, _, preferred_is_complex = field_infos[0]532            # Only normalize to preferred_key when it's a simple string alias.533            # When the preferred key comes from an AliasPath (complex entry), skip normalization534            # to avoid using the AliasPath's first element as the key (see #766).535            if not preferred_is_complex:536                return field_value, preferred_key, value_is_complex537        return field_value, field_key, value_is_complex538 539    def __call__(self) -> dict[str, Any]:540        data: dict[str, Any] = {}541 542        for field_name, field in self.settings_cls.model_fields.items():543            try:544                field_value, field_key, value_is_complex = self._get_resolved_field_value(field, field_name)545            except Exception as e:546                raise SettingsError(547                    f'error getting value for field "{field_name}" from source "{self.__class__.__name__}"'548                ) from e549 550            try:551                field_value = self.prepare_field_value(field_name, field, field_value, value_is_complex)552            except ValueError as e:553                raise SettingsError(554                    f'error parsing value for field "{field_name}" from source "{self.__class__.__name__}"'555                ) from e556 557            if field_value is not None:558                if self.env_parse_none_str is not None:559                    if isinstance(field_value, dict):560                        field_value = self._replace_env_none_type_values(field_value)561                    elif isinstance(field_value, EnvNoneType):562                        field_value = None563                if (564                    not self.case_sensitive565                    # and _lenient_issubclass(field.annotation, BaseModel)566                    and isinstance(field_value, dict)567                ):568                    data[field_key] = self._replace_field_names_case_insensitively(field, field_value)569                else:570                    data[field_key] = field_value571 572        return data573 574 575__all__ = [576    'ConfigFileSourceMixin',577    'DefaultSettingsSource',578    'InitSettingsSource',579    'PydanticBaseEnvSettingsSource',580    'PydanticBaseSettingsSource',581    'SettingsError',582]583 
codekingpro/portable-devtools · Team Ai