diff --git a/venv/Lib/site-packages/pydantic_settings/sources/providers/env.py b/venv/Lib/site-packages/pydantic_settings/sources/providers/env.py new file mode 100644 index 0000000..3146366 --- /dev/null +++ b/venv/Lib/site-packages/pydantic_settings/sources/providers/env.py @@ -0,0 +1,329 @@ +from __future__ import annotations as _annotations + +import json +import os +from collections.abc import Mapping +from typing import ( + TYPE_CHECKING, + Any, + get_args, + get_origin, +) + +from pydantic import Json, TypeAdapter, ValidationError +from pydantic._internal._utils import deep_update, is_model_class +from pydantic.dataclasses import is_pydantic_dataclass +from pydantic.fields import FieldInfo +from typing_inspection.introspection import is_union_origin + +from ...utils import _lenient_issubclass +from ..base import PydanticBaseEnvSettingsSource +from ..types import EnvNoneType, EnvPrefixTarget +from ..utils import ( + _annotation_contains_types, + _annotation_enum_name_to_val, + _annotation_is_complex, + _get_model_fields, + _literal_has_numeric_enum, + _union_has_strict_types, + _union_is_complex, + parse_env_vars, +) + +if TYPE_CHECKING: + from pydantic_settings.main import BaseSettings + + +class EnvSettingsSource(PydanticBaseEnvSettingsSource): + """ + Source class for loading settings values from environment variables. + """ + + def __init__( + self, + settings_cls: type[BaseSettings], + case_sensitive: bool | None = None, + env_prefix: str | None = None, + env_prefix_target: EnvPrefixTarget | None = None, + env_nested_delimiter: str | None = None, + env_nested_max_split: int | None = None, + env_ignore_empty: bool | None = None, + env_parse_none_str: str | None = None, + env_parse_enums: bool | None = None, + ) -> None: + super().__init__( + settings_cls, + case_sensitive, + env_prefix, + env_prefix_target, + env_ignore_empty, + env_parse_none_str, + env_parse_enums, + ) + self.env_nested_delimiter = ( + env_nested_delimiter if env_nested_delimiter is not None else self.config.get('env_nested_delimiter') + ) + self.env_nested_max_split = ( + env_nested_max_split if env_nested_max_split is not None else self.config.get('env_nested_max_split') + ) + self.maxsplit = (self.env_nested_max_split or 0) - 1 + self.env_prefix_len = len(self.env_prefix) + + self.env_vars = self._load_env_vars() + + def _load_env_vars(self) -> Mapping[str, str | None]: + return parse_env_vars(os.environ, self.case_sensitive, self.env_ignore_empty, self.env_parse_none_str) + + def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]: + """ + Gets the value for field from environment variables and a flag to determine whether value is complex. + + Args: + field: The field. + field_name: The field name. + + Returns: + A tuple that contains the value (`None` if not found), key, and + a flag to determine whether value is complex. + """ + + env_val: str | None = None + for field_key, env_name, value_is_complex in self._extract_field_info(field, field_name): + env_val = self.env_vars.get(env_name) + if env_val is not None: + break + + return env_val, field_key, value_is_complex + + def prepare_field_value(self, field_name: str, field: FieldInfo, value: Any, value_is_complex: bool) -> Any: + """ + Prepare value for the field. + + * Extract value for nested field. + * Deserialize value to python object for complex field. + + Args: + field: The field. + field_name: The field name. + + Returns: + A tuple contains prepared value for the field. + + Raises: + ValuesError: When There is an error in deserializing value for complex field. + """ + is_complex, allow_parse_failure = self._field_is_complex(field) + if self.env_parse_enums: + enum_val = _annotation_enum_name_to_val(field.annotation, value) + value = value if enum_val is None else enum_val + + if is_complex or value_is_complex: + if isinstance(value, EnvNoneType): + return value + elif value is None: + # field is complex but no value found so far, try explode_env_vars + env_val_built = self.explode_env_vars(field_name, field, self.env_vars) + if env_val_built: + return env_val_built + else: + # field is complex and there's a value, decode that as JSON, then add explode_env_vars + try: + value = self.decode_complex_value(field_name, field, value) + except ValueError as e: + if not allow_parse_failure: + raise e + + if isinstance(value, dict): + return deep_update(value, self.explode_env_vars(field_name, field, self.env_vars)) + else: + return value + elif value is not None: + # simplest case, field is not complex, we only need to add the value if it was found + return self._coerce_env_val_strict(field, value) + + def _field_is_complex(self, field: FieldInfo) -> tuple[bool, bool]: + """ + Find out if a field is complex, and if so whether JSON errors should be ignored + """ + if self.field_is_complex(field): + allow_parse_failure = False + elif is_union_origin(get_origin(field.annotation)) and _union_is_complex(field.annotation, field.metadata): + allow_parse_failure = True + else: + return False, False + + return True, allow_parse_failure + + # Default value of `case_sensitive` is `None`, because we don't want to break existing behavior. + # We have to change the method to a non-static method and use + # `self.case_sensitive` instead in V3. + def next_field( + self, field: FieldInfo | Any | None, key: str, case_sensitive: bool | None = None + ) -> FieldInfo | None: + """ + Find the field in a sub model by key(env name) + + By having the following models: + + ```py + class SubSubModel(BaseSettings): + dvals: Dict + + class SubModel(BaseSettings): + vals: list[str] + sub_sub_model: SubSubModel + + class Cfg(BaseSettings): + sub_model: SubModel + ``` + + Then: + next_field(sub_model, 'vals') Returns the `vals` field of `SubModel` class + next_field(sub_model, 'sub_sub_model') Returns `sub_sub_model` field of `SubModel` class + + Args: + field: The field. + key: The key (env name). + case_sensitive: Whether to search for key case sensitively. + + Returns: + Field if it finds the next field otherwise `None`. + """ + if not field: + return None + + annotation = field.annotation if isinstance(field, FieldInfo) else field + for type_ in get_args(annotation): + type_has_key = self.next_field(type_, key, case_sensitive) + if type_has_key: + return type_has_key + if _lenient_issubclass(get_origin(annotation), dict): + # get value type if it's a dict + return get_args(annotation)[-1] + elif is_model_class(annotation) or is_pydantic_dataclass(annotation): # type: ignore[arg-type] + fields = _get_model_fields(annotation) + # `case_sensitive is None` is here to be compatible with the old behavior. + # Has to be removed in V3. + for field_name, f in fields.items(): + for _, env_name, _ in self._extract_field_info(f, field_name): + if case_sensitive is None or case_sensitive: + if field_name == key or env_name == key: + return f + elif field_name.lower() == key.lower() or env_name.lower() == key.lower(): + return f + return None + + def explode_env_vars(self, field_name: str, field: FieldInfo, env_vars: Mapping[str, str | None]) -> dict[str, Any]: # noqa: C901 + """ + Process env_vars and extract the values of keys containing env_nested_delimiter into nested dictionaries. + + This is applied to a single field, hence filtering by env_var prefix. + + Args: + field_name: The field name. + field: The field. + env_vars: Environment variables. + + Returns: + A dictionary contains extracted values from nested env values. + """ + if not self.env_nested_delimiter: + return {} + + ann = field.annotation + is_dict = ann is dict or _lenient_issubclass(get_origin(ann), dict) + + prefixes = [ + f'{env_name}{self.env_nested_delimiter}' for _, env_name, _ in self._extract_field_info(field, field_name) + ] + result: dict[str, Any] = {} + for env_name, env_val in env_vars.items(): + try: + prefix = next(prefix for prefix in prefixes if env_name.startswith(prefix)) + except StopIteration: + continue + # we remove the prefix before splitting in case the prefix has characters in common with the delimiter + env_name_without_prefix = env_name[len(prefix) :] + *keys, last_key = env_name_without_prefix.split(self.env_nested_delimiter, self.maxsplit) + env_var = result + target_field: FieldInfo | None = field + for key in keys: + target_field = self.next_field(target_field, key, self.case_sensitive) + if isinstance(env_var, dict): + env_var = env_var.setdefault(key, {}) + + # get proper field with last_key + target_field = self.next_field(target_field, last_key, self.case_sensitive) + + # check if env_val maps to a complex field and if so, parse the env_val + if (target_field or is_dict) and env_val: + if isinstance(target_field, FieldInfo): + is_complex, allow_json_failure = self._field_is_complex(target_field) + if self.env_parse_enums: + enum_val = _annotation_enum_name_to_val(target_field.annotation, env_val) + env_val = env_val if enum_val is None else enum_val + elif target_field: + # target_field is a raw type (e.g. from dict value type annotation) + is_complex = _annotation_is_complex(target_field, []) + allow_json_failure = True + else: + # nested field type is dict + is_complex, allow_json_failure = True, True + if is_complex: + try: + field_info = target_field if isinstance(target_field, FieldInfo) else None + env_val = self.decode_complex_value(last_key, field_info, env_val) # type: ignore + except ValueError as e: + if not allow_json_failure: + raise e + if isinstance(env_var, dict): + if last_key not in env_var or not isinstance(env_val, EnvNoneType) or env_var[last_key] == {}: + env_var[last_key] = self._coerce_env_val_strict(target_field, env_val) + return result + + def _coerce_env_val_strict(self, field: FieldInfo | None, value: Any) -> Any: + """ + Coerce environment string values based on field annotation if model config is `strict=True` + or if the field annotation contains strict-annotated types (e.g. Optional[StrictBool]). + + Args: + field: The field. + value: The value to coerce. + + Returns: + The coerced value if successful, otherwise the original value. + """ + try: + should_coerce = self.config.get('strict') + if not should_coerce and isinstance(field, FieldInfo): + should_coerce = ( + is_union_origin(get_origin(field.annotation)) and _union_has_strict_types(field.annotation) + ) or _literal_has_numeric_enum(field.annotation) + if should_coerce and isinstance(value, str) and isinstance(field, FieldInfo): + if value == self.env_parse_none_str: + return value + if not _annotation_contains_types(field.annotation, (Json,), is_instance=True): + try: + return TypeAdapter(field.annotation).validate_python(value) + except ValidationError: + # Try JSON decoding as fallback (e.g. 'true' -> True for StrictBool) + try: + decoded = json.loads(value) + except (ValueError, json.JSONDecodeError): + raise + if not isinstance(decoded, str): + return TypeAdapter(field.annotation).validate_python(decoded) + raise + except ValidationError: + # Allow validation error to be raised at time of instantiation + pass + return value + + def __repr__(self) -> str: + return ( + f'{self.__class__.__name__}(env_nested_delimiter={self.env_nested_delimiter!r}, ' + f'env_prefix_len={self.env_prefix_len!r})' + ) + + +__all__ = ['EnvSettingsSource'] diff --git a/venv/Lib/site-packages/pydantic_settings/sources/providers/gcp.py b/venv/Lib/site-packages/pydantic_settings/sources/providers/gcp.py new file mode 100644 index 0000000..0885eb8 --- /dev/null +++ b/venv/Lib/site-packages/pydantic_settings/sources/providers/gcp.py @@ -0,0 +1,241 @@ +from __future__ import annotations as _annotations + +import warnings +from collections.abc import Iterator, Mapping +from functools import cached_property +from typing import TYPE_CHECKING, Any + +from pydantic.fields import FieldInfo + +from ..types import SecretVersion +from .env import EnvSettingsSource + +if TYPE_CHECKING: + from google.auth import default as google_auth_default + from google.auth.credentials import Credentials + from google.cloud.secretmanager import SecretManagerServiceClient + + from pydantic_settings.main import BaseSettings +else: + Credentials = None + SecretManagerServiceClient = None + google_auth_default = None + + +def import_gcp_secret_manager() -> None: + global Credentials + global SecretManagerServiceClient + global google_auth_default + + try: + from google.auth import default as google_auth_default + from google.auth.credentials import Credentials + + with warnings.catch_warnings(): + warnings.filterwarnings('ignore', category=FutureWarning) + from google.cloud.secretmanager import SecretManagerServiceClient + except ImportError as e: # pragma: no cover + raise ImportError( + 'GCP Secret Manager dependencies are not installed, run `pip install pydantic-settings[gcp-secret-manager]`' + ) from e + + +class GoogleSecretManagerMapping(Mapping[str, str | None]): + _loaded_secrets: dict[str, str | None] + _secret_client: SecretManagerServiceClient + + def __init__(self, secret_client: SecretManagerServiceClient, project_id: str, case_sensitive: bool) -> None: + self._loaded_secrets = {} + self._secret_client = secret_client + self._project_id = project_id + self._case_sensitive = case_sensitive + + @property + def _gcp_project_path(self) -> str: + return self._secret_client.common_project_path(self._project_id) + + def _select_case_insensitive_secret(self, lower_name: str, candidates: list[str]) -> str: + if len(candidates) == 1: + return candidates[0] + + # Sort to ensure deterministic selection (prefer lowercase / ASCII last) + candidates.sort() + winner = candidates[-1] + warnings.warn( + f"Secret collision: Found multiple secrets {candidates} normalizing to '{lower_name}'. " + f"Using '{winner}' for case-insensitive lookup.", + UserWarning, + stacklevel=2, + ) + return winner + + @cached_property + def _secret_name_map(self) -> dict[str, str]: + mapping: dict[str, str] = {} + # Group secrets by normalized name to detect collisions + normalized_groups: dict[str, list[str]] = {} + + secrets = self._secret_client.list_secrets(parent=self._gcp_project_path) + for secret in secrets: + name = self._secret_client.parse_secret_path(secret.name).get('secret', '') + mapping[name] = name + + if not self._case_sensitive: + lower_name = name.lower() + if lower_name not in normalized_groups: + normalized_groups[lower_name] = [] + normalized_groups[lower_name].append(name) + + if not self._case_sensitive: + for lower_name, candidates in normalized_groups.items(): + mapping[lower_name] = self._select_case_insensitive_secret(lower_name, candidates) + + return mapping + + @property + def _secret_names(self) -> list[str]: + return list(self._secret_name_map.keys()) + + def _secret_version_path(self, key: str, version: str = 'latest') -> str: + return self._secret_client.secret_version_path(self._project_id, key, version) + + def _get_secret_value(self, gcp_secret_name: str, version: str = 'latest') -> str | None: + try: + return self._secret_client.access_secret_version( + name=self._secret_version_path(gcp_secret_name, version) + ).payload.data.decode('UTF-8') + except Exception: + return None + + def __getitem__(self, key: str) -> str | None: + if key in self._loaded_secrets: + return self._loaded_secrets[key] + + gcp_secret_name = self._secret_name_map.get(key) + if gcp_secret_name is None and not self._case_sensitive: + gcp_secret_name = self._secret_name_map.get(key.lower()) + + if gcp_secret_name: + self._loaded_secrets[key] = self._get_secret_value(gcp_secret_name) + else: + raise KeyError(key) + + return self._loaded_secrets[key] + + def __len__(self) -> int: + return len(self._secret_names) + + def __iter__(self) -> Iterator[str]: + return iter(self._secret_names) + + +class GoogleSecretManagerSettingsSource(EnvSettingsSource): + _credentials: Credentials + _secret_client: SecretManagerServiceClient + _project_id: str + + def __init__( + self, + settings_cls: type[BaseSettings], + credentials: Credentials | None = None, + project_id: str | None = None, + env_prefix: str | None = None, + env_parse_none_str: str | None = None, + env_parse_enums: bool | None = None, + secret_client: SecretManagerServiceClient | None = None, + case_sensitive: bool | None = True, + ) -> None: + # Import Google Packages if they haven't already been imported + if SecretManagerServiceClient is None or Credentials is None or google_auth_default is None: + import_gcp_secret_manager() + + # If credentials or project_id are not passed, then + # try to get them from the default function + if not credentials or not project_id: + _creds, _project_id = google_auth_default() + + # Set the credentials and/or project id if they weren't specified + if credentials is None: + credentials = _creds + + if project_id is None: + if isinstance(_project_id, str): + project_id = _project_id + else: + raise AttributeError( + 'project_id is required to be specified either as an argument or from the google.auth.default. See https://google-auth.readthedocs.io/en/master/reference/google.auth.html#google.auth.default' + ) + + self._credentials: Credentials = credentials + self._project_id: str = project_id + + if secret_client: + self._secret_client = secret_client + else: + self._secret_client = SecretManagerServiceClient(credentials=self._credentials) + + super().__init__( + settings_cls, + case_sensitive=case_sensitive, + env_prefix=env_prefix, + env_ignore_empty=False, + env_parse_none_str=env_parse_none_str, + env_parse_enums=env_parse_enums, + ) + + def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]: + """Override get_field_value to get the secret value from GCP Secret Manager. + Look for a SecretVersion metadata field to specify a particular SecretVersion. + + Args: + field: The field to get the value for + field_name: The declared name of the field + + Returns: + A tuple of (value, key, value_is_complex), where `key` is the identifier used + to populate the model (either the field name or an alias, depending on + configuration). + """ + + secret_version = next((m.version for m in field.metadata if isinstance(m, SecretVersion)), None) + + # If a secret version is specified, try to get that specific version of the secret from + # GCP Secret Manager via the GoogleSecretManagerMapping. This allows different versions + # of the same secret name to be retrieved independently and cached in the GoogleSecretManagerMapping + if secret_version and isinstance(self.env_vars, GoogleSecretManagerMapping): + for field_key, env_name, value_is_complex in self._extract_field_info(field, field_name): + gcp_secret_name = self.env_vars._secret_name_map.get(env_name) + if gcp_secret_name is None and not self.case_sensitive: + gcp_secret_name = self.env_vars._secret_name_map.get(env_name.lower()) + + if gcp_secret_name: + env_val = self.env_vars._get_secret_value(gcp_secret_name, secret_version) + if env_val is not None: + # If populate_by_name is enabled, return field_name to allow multiple fields + # with the same alias but different versions to be distinguished + if self.settings_cls.model_config.get('populate_by_name'): + return env_val, field_name, value_is_complex + return env_val, field_key, value_is_complex + + # If a secret version is specified but not found, we should not fall back to "latest" (default behavior) + # as that would be incorrect. We return None to indicate the value was not found. + return None, field_name, False + + val, key, is_complex = super().get_field_value(field, field_name) + + # If populate_by_name is enabled, we need to return the field_name as the key + # without this being enabled, you cannot load two secrets with the same name but different versions + if self.settings_cls.model_config.get('populate_by_name') and val is not None: + return val, field_name, is_complex + return val, key, is_complex + + def _load_env_vars(self) -> Mapping[str, str | None]: + return GoogleSecretManagerMapping( + self._secret_client, project_id=self._project_id, case_sensitive=self.case_sensitive + ) + + def __repr__(self) -> str: + return f'{self.__class__.__name__}(project_id={self._project_id!r}, env_nested_delimiter={self.env_nested_delimiter!r})' + + +__all__ = ['GoogleSecretManagerSettingsSource', 'GoogleSecretManagerMapping'] diff --git a/venv/Lib/site-packages/pydantic_settings/sources/providers/json.py b/venv/Lib/site-packages/pydantic_settings/sources/providers/json.py new file mode 100644 index 0000000..a3b73a9 --- /dev/null +++ b/venv/Lib/site-packages/pydantic_settings/sources/providers/json.py @@ -0,0 +1,48 @@ +"""JSON file settings source.""" + +from __future__ import annotations as _annotations + +import json +from pathlib import Path +from typing import ( + TYPE_CHECKING, + Any, +) + +from ..base import ConfigFileSourceMixin, InitSettingsSource +from ..types import DEFAULT_PATH, PathType + +if TYPE_CHECKING: + from pydantic_settings.main import BaseSettings + + +class JsonConfigSettingsSource(InitSettingsSource, ConfigFileSourceMixin): + """ + A source class that loads variables from a JSON file + """ + + def __init__( + self, + settings_cls: type[BaseSettings], + json_file: PathType | None = DEFAULT_PATH, + json_file_encoding: str | None = None, + deep_merge: bool = False, + ): + self.json_file_path = json_file if json_file != DEFAULT_PATH else settings_cls.model_config.get('json_file') + self.json_file_encoding = ( + json_file_encoding + if json_file_encoding is not None + else settings_cls.model_config.get('json_file_encoding') + ) + self.json_data = self._read_files(self.json_file_path, deep_merge=deep_merge) + super().__init__(settings_cls, self.json_data) + + def _read_file(self, file_path: Path) -> dict[str, Any]: + with file_path.open(encoding=self.json_file_encoding) as json_file: + return json.load(json_file) + + def __repr__(self) -> str: + return f'{self.__class__.__name__}(json_file={self.json_file_path})' + + +__all__ = ['JsonConfigSettingsSource'] diff --git a/venv/Lib/site-packages/pydantic_settings/sources/providers/nested_secrets.py b/venv/Lib/site-packages/pydantic_settings/sources/providers/nested_secrets.py new file mode 100644 index 0000000..97d5afc --- /dev/null +++ b/venv/Lib/site-packages/pydantic_settings/sources/providers/nested_secrets.py @@ -0,0 +1,210 @@ +import os +import warnings +from collections.abc import Iterator +from functools import reduce +from pathlib import Path +from typing import TYPE_CHECKING, Any, Literal, Optional + +from ...exceptions import SettingsError +from ...utils import path_type_label +from ..base import PydanticBaseSettingsSource +from ..utils import parse_env_vars +from .env import EnvSettingsSource +from .secrets import SecretsSettingsSource + +if TYPE_CHECKING: + from ...main import BaseSettings + from ...sources import PathType + + +SECRETS_DIR_MAX_SIZE = 16 * 2**20 # 16 MiB seems to be a reasonable default + + +class NestedSecretsSettingsSource(EnvSettingsSource): + def __init__( + self, + file_secret_settings: PydanticBaseSettingsSource | SecretsSettingsSource, + secrets_dir: Optional['PathType'] = None, + secrets_dir_missing: Literal['ok', 'warn', 'error'] | None = None, + secrets_dir_max_size: int | None = None, + secrets_case_sensitive: bool | None = None, + secrets_prefix: str | None = None, + secrets_nested_delimiter: str | None = None, + secrets_nested_subdir: bool | None = None, + # args for compatibility with SecretsSettingsSource, don't use directly + case_sensitive: bool | None = None, + env_prefix: str | None = None, + ) -> None: + # We allow the first argument to be settings_cls like original + # SecretsSettingsSource. However, it is recommended to pass + # SecretsSettingsSource instance instead (as it is shown in usage examples), + # otherwise `_secrets_dir` arg passed to Settings() constructor will be ignored. + settings_cls: type[BaseSettings] = getattr( + file_secret_settings, + 'settings_cls', + file_secret_settings, # type: ignore[arg-type] + ) + # config options + conf = settings_cls.model_config + self.secrets_dir: PathType | None = first_not_none( + getattr(file_secret_settings, 'secrets_dir', None), + secrets_dir, + conf.get('secrets_dir'), + ) + self.secrets_dir_missing: Literal['ok', 'warn', 'error'] = first_not_none( + secrets_dir_missing, + conf.get('secrets_dir_missing'), + 'warn', + ) + if self.secrets_dir_missing not in ('ok', 'warn', 'error'): + raise SettingsError(f'invalid secrets_dir_missing value: {self.secrets_dir_missing}') + self.secrets_dir_max_size: int = first_not_none( + secrets_dir_max_size, + conf.get('secrets_dir_max_size'), + SECRETS_DIR_MAX_SIZE, + ) + self.case_sensitive: bool = first_not_none( + secrets_case_sensitive, + conf.get('secrets_case_sensitive'), + case_sensitive, + conf.get('case_sensitive'), + False, + ) + self.secrets_prefix: str = first_not_none( + secrets_prefix, + conf.get('secrets_prefix'), + env_prefix, + conf.get('env_prefix'), + '', + ) + + # nested options + self.secrets_nested_delimiter: str | None = first_not_none( + secrets_nested_delimiter, + conf.get('secrets_nested_delimiter'), + conf.get('env_nested_delimiter'), + ) + self.secrets_nested_subdir: bool = first_not_none( + secrets_nested_subdir, + conf.get('secrets_nested_subdir'), + False, + ) + if self.secrets_nested_subdir: + if secrets_nested_delimiter or conf.get('secrets_nested_delimiter'): + raise SettingsError('Options secrets_nested_delimiter and secrets_nested_subdir are mutually exclusive') + else: + self.secrets_nested_delimiter = os.sep + + # ensure valid secrets_path + if self.secrets_dir is None: + paths = [] + elif isinstance(self.secrets_dir, (Path, str)): + paths = [self.secrets_dir] + else: + paths = list(self.secrets_dir) + self.secrets_paths: list[Path] = [Path(p).expanduser().resolve() for p in paths] + for path in self.secrets_paths: + self.validate_secrets_path(path) + + # construct parent + super().__init__( + settings_cls, + case_sensitive=self.case_sensitive, + env_prefix=self.secrets_prefix, + env_nested_delimiter=self.secrets_nested_delimiter, + env_ignore_empty=False, # match SecretsSettingsSource behaviour + env_parse_enums=True, # we can pass everything here, it will still behave as "True" + env_parse_none_str=None, # match SecretsSettingsSource behaviour + ) + self.env_parse_none_str = None # update manually because of None + + # update parent members + if not len(self.secrets_paths): + self.env_vars = {} + else: + secrets = reduce( + lambda d1, d2: dict((*d1.items(), *d2.items())), + (self.load_secrets(p) for p in self.secrets_paths), + ) + self.env_vars = parse_env_vars( + secrets, + self.case_sensitive, + self.env_ignore_empty, + self.env_parse_none_str, + ) + + def validate_secrets_path(self, path: Path) -> None: + if not path.exists(): + if self.secrets_dir_missing == 'ok': + pass + elif self.secrets_dir_missing == 'warn': + warnings.warn(f'directory "{path}" does not exist', stacklevel=2) + elif self.secrets_dir_missing == 'error': + raise SettingsError(f'directory "{path}" does not exist') + else: + raise ValueError # unreachable, checked before + else: + if not path.is_dir(): + raise SettingsError(f'secrets_dir must reference a directory, not a {path_type_label(path)}') + secrets_dir_size = sum(f.stat().st_size for f in self._iter_secret_files(path)) + if secrets_dir_size > self.secrets_dir_max_size: + raise SettingsError(f'secrets_dir size is above {self.secrets_dir_max_size} bytes') + + @staticmethod + def _iter_secret_files(path: Path) -> Iterator[Path]: + """Yield the secret files contained in ``path``. + + ``path`` is expected to already be resolved. The directory tree is walked + explicitly so that symbolic links are handled safely: + + * a file is only yielded if its real location stays within ``path``; entries + that resolve outside of it (e.g. through a symbolic link) are skipped, so + they neither contribute to the ``secrets_dir_max_size`` accounting nor get + loaded; + * each real directory is visited at most once, so cyclic or repeated + symlinks cannot make the walk loop and inflate the size accounting or the + number of loaded secrets. + + Because the size check and the loader share this iterator, they always see + the same set of files. + """ + seen_dirs: set[Path] = set() + + def walk(directory: Path) -> Iterator[Path]: + # Guard against symlink loops / a directory reachable through multiple + # links being traversed more than once. + resolved_dir = directory.resolve() + if resolved_dir in seen_dirs: + return + seen_dirs.add(resolved_dir) + try: + entries = sorted(directory.iterdir()) + except OSError: + return + for entry in entries: + resolved = entry.resolve() + if resolved.is_dir(): + # Only descend into directories that stay within secrets_dir. + # A symlinked directory pointing outside of ``path`` is not + # followed at all, so we never walk (potentially large) external + # trees and never read files from outside secrets_dir. + if resolved == path or path in resolved.parents: + yield from walk(entry) + elif resolved.is_file() and path in resolved.parents: + # Defense in depth: a file whose real location escapes + # secrets_dir (e.g. a symlink pointing outside of ``path``) is + # skipped from both the size accounting and the load. + yield entry + + yield from walk(path) + + @classmethod + def load_secrets(cls, path: Path) -> dict[str, str]: + return {str(p.relative_to(path)): p.read_text().strip() for p in cls._iter_secret_files(path)} + + def __repr__(self) -> str: + return f'NestedSecretsSettingsSource(secrets_dir={self.secrets_dir!r})' + + +def first_not_none(*objs: Any) -> Any: + return next(filter(lambda o: o is not None, objs), None) diff --git a/venv/Lib/site-packages/pydantic_settings/sources/providers/pyproject.py b/venv/Lib/site-packages/pydantic_settings/sources/providers/pyproject.py new file mode 100644 index 0000000..bb02cbb --- /dev/null +++ b/venv/Lib/site-packages/pydantic_settings/sources/providers/pyproject.py @@ -0,0 +1,62 @@ +"""Pyproject TOML file settings source.""" + +from __future__ import annotations as _annotations + +from pathlib import Path +from typing import ( + TYPE_CHECKING, +) + +from .toml import TomlConfigSettingsSource + +if TYPE_CHECKING: + from pydantic_settings.main import BaseSettings + + +class PyprojectTomlConfigSettingsSource(TomlConfigSettingsSource): + """ + A source class that loads variables from a `pyproject.toml` file. + """ + + def __init__( + self, + settings_cls: type[BaseSettings], + toml_file: Path | None = None, + ) -> None: + self.toml_file_path = self._pick_pyproject_toml_file( + toml_file, settings_cls.model_config.get('pyproject_toml_depth', 0) + ) + self.toml_table_header: tuple[str, ...] = settings_cls.model_config.get( + 'pyproject_toml_table_header', ('tool', 'pydantic-settings') + ) + self.toml_data = self._read_files(self.toml_file_path) + for key in self.toml_table_header: + self.toml_data = self.toml_data.get(key, {}) + super(TomlConfigSettingsSource, self).__init__(settings_cls, self.toml_data) + + @staticmethod + def _pick_pyproject_toml_file(provided: Path | None, depth: int) -> Path: + """Pick a `pyproject.toml` file path to use. + + Args: + provided: Explicit path provided when instantiating this class. + depth: Number of directories up the tree to check of a pyproject.toml. + + """ + if provided: + return provided.resolve() + rv = Path.cwd() / 'pyproject.toml' + count = 0 + if not rv.is_file(): + child = rv.parent.parent / 'pyproject.toml' + while count < depth: + if child.is_file(): + return child + if str(child.parent) == rv.root: + break # end discovery after checking system root once + child = child.parent.parent / 'pyproject.toml' + count += 1 + return rv + + +__all__ = ['PyprojectTomlConfigSettingsSource']