Загрузить файлы в «venv/Lib/site-packages/pydantic_settings/sources/providers»
This commit is contained in:
@@ -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']
|
||||||
@@ -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']
|
||||||
@@ -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']
|
||||||
@@ -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)
|
||||||
@@ -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']
|
||||||
Reference in New Issue
Block a user