codekingpro/portable-devtools
114k
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 