sources.py 92 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021102210231024102510261027102810291030103110321033103410351036103710381039104010411042104310441045104610471048104910501051105210531054105510561057105810591060106110621063106410651066106710681069107010711072107310741075107610771078107910801081108210831084108510861087108810891090109110921093109410951096109710981099110011011102110311041105110611071108110911101111111211131114111511161117111811191120112111221123112411251126112711281129113011311132113311341135113611371138113911401141114211431144114511461147114811491150115111521153115411551156115711581159116011611162116311641165116611671168116911701171117211731174117511761177117811791180118111821183118411851186118711881189119011911192119311941195119611971198119912001201120212031204120512061207120812091210121112121213121412151216121712181219122012211222122312241225122612271228122912301231123212331234123512361237123812391240124112421243124412451246124712481249125012511252125312541255125612571258125912601261126212631264126512661267126812691270127112721273127412751276127712781279128012811282128312841285128612871288128912901291129212931294129512961297129812991300130113021303130413051306130713081309131013111312131313141315131613171318131913201321132213231324132513261327132813291330133113321333133413351336133713381339134013411342134313441345134613471348134913501351135213531354135513561357135813591360136113621363136413651366136713681369137013711372137313741375137613771378137913801381138213831384138513861387138813891390139113921393139413951396139713981399140014011402140314041405140614071408140914101411141214131414141514161417141814191420142114221423142414251426142714281429143014311432143314341435143614371438143914401441144214431444144514461447144814491450145114521453145414551456145714581459146014611462146314641465146614671468146914701471147214731474147514761477147814791480148114821483148414851486148714881489149014911492149314941495149614971498149915001501150215031504150515061507150815091510151115121513151415151516151715181519152015211522152315241525152615271528152915301531153215331534153515361537153815391540154115421543154415451546154715481549155015511552155315541555155615571558155915601561156215631564156515661567156815691570157115721573157415751576157715781579158015811582158315841585158615871588158915901591159215931594159515961597159815991600160116021603160416051606160716081609161016111612161316141615161616171618161916201621162216231624162516261627162816291630163116321633163416351636163716381639164016411642164316441645164616471648164916501651165216531654165516561657165816591660166116621663166416651666166716681669167016711672167316741675167616771678167916801681168216831684168516861687168816891690169116921693169416951696169716981699170017011702170317041705170617071708170917101711171217131714171517161717171817191720172117221723172417251726172717281729173017311732173317341735173617371738173917401741174217431744174517461747174817491750175117521753175417551756175717581759176017611762176317641765176617671768176917701771177217731774177517761777177817791780178117821783178417851786178717881789179017911792179317941795179617971798179918001801180218031804180518061807180818091810181118121813181418151816181718181819182018211822182318241825182618271828182918301831183218331834183518361837183818391840184118421843184418451846184718481849185018511852185318541855185618571858185918601861186218631864186518661867186818691870187118721873187418751876187718781879188018811882188318841885188618871888188918901891189218931894189518961897189818991900190119021903190419051906190719081909191019111912191319141915191619171918191919201921192219231924192519261927192819291930193119321933193419351936193719381939194019411942194319441945194619471948194919501951195219531954195519561957195819591960196119621963196419651966196719681969197019711972197319741975197619771978197919801981198219831984198519861987198819891990199119921993199419951996199719981999200020012002200320042005200620072008200920102011201220132014201520162017201820192020202120222023202420252026202720282029203020312032203320342035203620372038203920402041204220432044204520462047204820492050205120522053205420552056205720582059206020612062206320642065206620672068206920702071207220732074207520762077207820792080208120822083208420852086208720882089209020912092209320942095209620972098209921002101210221032104210521062107210821092110211121122113211421152116211721182119212021212122212321242125212621272128212921302131213221332134213521362137213821392140214121422143214421452146214721482149215021512152215321542155215621572158215921602161216221632164216521662167216821692170217121722173217421752176217721782179218021812182218321842185218621872188218921902191219221932194219521962197219821992200220122022203220422052206220722082209221022112212221322142215221622172218221922202221222222232224222522262227222822292230223122322233223422352236223722382239224022412242224322442245
  1. from __future__ import annotations as _annotations
  2. import json
  3. import os
  4. import re
  5. import shlex
  6. import sys
  7. import typing
  8. import warnings
  9. from abc import ABC, abstractmethod
  10. if sys.version_info >= (3, 9):
  11. from argparse import BooleanOptionalAction
  12. from argparse import SUPPRESS, ArgumentParser, Namespace, RawDescriptionHelpFormatter, _SubParsersAction
  13. from collections import defaultdict, deque
  14. from dataclasses import asdict, is_dataclass
  15. from enum import Enum
  16. from pathlib import Path
  17. from textwrap import dedent
  18. from types import BuiltinFunctionType, FunctionType, SimpleNamespace
  19. from typing import (
  20. TYPE_CHECKING,
  21. Any,
  22. Callable,
  23. Dict,
  24. Generic,
  25. Iterator,
  26. Mapping,
  27. NoReturn,
  28. Optional,
  29. Sequence,
  30. TypeVar,
  31. Union,
  32. cast,
  33. overload,
  34. )
  35. import typing_extensions
  36. from dotenv import dotenv_values
  37. from pydantic import AliasChoices, AliasPath, BaseModel, Json, RootModel, TypeAdapter
  38. from pydantic._internal._repr import Representation
  39. from pydantic._internal._signature import _field_name_for_signature
  40. from pydantic._internal._typing_extra import WithArgsTypes, origin_is_union, typing_base
  41. from pydantic._internal._utils import deep_update, is_model_class, lenient_issubclass
  42. from pydantic.dataclasses import is_pydantic_dataclass
  43. from pydantic.fields import FieldInfo
  44. from pydantic_core import PydanticUndefined
  45. from typing_extensions import Annotated, _AnnotatedAlias, get_args, get_origin
  46. from pydantic_settings.utils import path_type_label
  47. if TYPE_CHECKING:
  48. if sys.version_info >= (3, 11):
  49. import tomllib
  50. else:
  51. tomllib = None
  52. import tomli
  53. import yaml
  54. from pydantic._internal._dataclasses import PydanticDataclass
  55. from pydantic_settings.main import BaseSettings
  56. PydanticModel = TypeVar('PydanticModel', bound=PydanticDataclass | BaseModel)
  57. else:
  58. yaml = None
  59. tomllib = None
  60. tomli = None
  61. PydanticModel = Any
  62. def import_yaml() -> None:
  63. global yaml
  64. if yaml is not None:
  65. return
  66. try:
  67. import yaml
  68. except ImportError as e:
  69. raise ImportError('PyYAML is not installed, run `pip install pydantic-settings[yaml]`') from e
  70. def import_toml() -> None:
  71. global tomli
  72. global tomllib
  73. if sys.version_info < (3, 11):
  74. if tomli is not None:
  75. return
  76. try:
  77. import tomli
  78. except ImportError as e:
  79. raise ImportError('tomli is not installed, run `pip install pydantic-settings[toml]`') from e
  80. else:
  81. if tomllib is not None:
  82. return
  83. import tomllib
  84. def import_azure_key_vault() -> None:
  85. global TokenCredential
  86. global SecretClient
  87. global ResourceNotFoundError
  88. try:
  89. from azure.core.credentials import TokenCredential
  90. from azure.core.exceptions import ResourceNotFoundError
  91. from azure.keyvault.secrets import SecretClient
  92. except ImportError as e:
  93. raise ImportError(
  94. 'Azure Key Vault dependencies are not installed, run `pip install pydantic-settings[azure-key-vault]`'
  95. ) from e
  96. DotenvType = Union[Path, str, Sequence[Union[Path, str]]]
  97. PathType = Union[Path, str, Sequence[Union[Path, str]]]
  98. DEFAULT_PATH: PathType = Path('')
  99. # This is used as default value for `_env_file` in the `BaseSettings` class and
  100. # `env_file` in `DotEnvSettingsSource` so the default can be distinguished from `None`.
  101. # See the docstring of `BaseSettings` for more details.
  102. ENV_FILE_SENTINEL: DotenvType = Path('')
  103. class SettingsError(ValueError):
  104. pass
  105. class _CliSubCommand:
  106. pass
  107. class _CliPositionalArg:
  108. pass
  109. class _CliImplicitFlag:
  110. pass
  111. class _CliExplicitFlag:
  112. pass
  113. class _CliInternalArgParser(ArgumentParser):
  114. def __init__(self, cli_exit_on_error: bool = True, **kwargs: Any) -> None:
  115. super().__init__(**kwargs)
  116. self._cli_exit_on_error = cli_exit_on_error
  117. def error(self, message: str) -> NoReturn:
  118. if not self._cli_exit_on_error:
  119. raise SettingsError(f'error parsing CLI: {message}')
  120. super().error(message)
  121. T = TypeVar('T')
  122. CliSubCommand = Annotated[Union[T, None], _CliSubCommand]
  123. CliPositionalArg = Annotated[T, _CliPositionalArg]
  124. _CliBoolFlag = TypeVar('_CliBoolFlag', bound=bool)
  125. CliImplicitFlag = Annotated[_CliBoolFlag, _CliImplicitFlag]
  126. CliExplicitFlag = Annotated[_CliBoolFlag, _CliExplicitFlag]
  127. CLI_SUPPRESS = SUPPRESS
  128. CliSuppress = Annotated[T, CLI_SUPPRESS]
  129. def get_subcommand(
  130. model: PydanticModel, is_required: bool = True, cli_exit_on_error: bool | None = None
  131. ) -> Optional[PydanticModel]:
  132. """
  133. Get the subcommand from a model.
  134. Args:
  135. model: The model to get the subcommand from.
  136. is_required: Determines whether a model must have subcommand set and raises error if not
  137. found. Defaults to `True`.
  138. cli_exit_on_error: Determines whether this function exits with error if no subcommand is found.
  139. Defaults to model_config `cli_exit_on_error` value if set. Otherwise, defaults to `True`.
  140. Returns:
  141. The subcommand model if found, otherwise `None`.
  142. Raises:
  143. SystemExit: When no subcommand is found and is_required=`True` and cli_exit_on_error=`True`
  144. (the default).
  145. SettingsError: When no subcommand is found and is_required=`True` and
  146. cli_exit_on_error=`False`.
  147. """
  148. model_cls = type(model)
  149. if cli_exit_on_error is None and is_model_class(model_cls):
  150. model_default = model_cls.model_config.get('cli_exit_on_error')
  151. if isinstance(model_default, bool):
  152. cli_exit_on_error = model_default
  153. if cli_exit_on_error is None:
  154. cli_exit_on_error = True
  155. subcommands: list[str] = []
  156. for field_name, field_info in _get_model_fields(model_cls).items():
  157. if _CliSubCommand in field_info.metadata:
  158. if getattr(model, field_name) is not None:
  159. return getattr(model, field_name)
  160. subcommands.append(field_name)
  161. if is_required:
  162. error_message = (
  163. f'Error: CLI subcommand is required {{{", ".join(subcommands)}}}'
  164. if subcommands
  165. else 'Error: CLI subcommand is required but no subcommands were found.'
  166. )
  167. raise SystemExit(error_message) if cli_exit_on_error else SettingsError(error_message)
  168. return None
  169. class EnvNoneType(str):
  170. pass
  171. class PydanticBaseSettingsSource(ABC):
  172. """
  173. Abstract base class for settings sources, every settings source classes should inherit from it.
  174. """
  175. def __init__(self, settings_cls: type[BaseSettings]):
  176. self.settings_cls = settings_cls
  177. self.config = settings_cls.model_config
  178. self._current_state: dict[str, Any] = {}
  179. self._settings_sources_data: dict[str, dict[str, Any]] = {}
  180. def _set_current_state(self, state: dict[str, Any]) -> None:
  181. """
  182. Record the state of settings from the previous settings sources. This should
  183. be called right before __call__.
  184. """
  185. self._current_state = state
  186. def _set_settings_sources_data(self, states: dict[str, dict[str, Any]]) -> None:
  187. """
  188. Record the state of settings from all previous settings sources. This should
  189. be called right before __call__.
  190. """
  191. self._settings_sources_data = states
  192. @property
  193. def current_state(self) -> dict[str, Any]:
  194. """
  195. The current state of the settings, populated by the previous settings sources.
  196. """
  197. return self._current_state
  198. @property
  199. def settings_sources_data(self) -> dict[str, dict[str, Any]]:
  200. """
  201. The state of all previous settings sources.
  202. """
  203. return self._settings_sources_data
  204. @abstractmethod
  205. def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]:
  206. """
  207. Gets the value, the key for model creation, and a flag to determine whether value is complex.
  208. This is an abstract method that should be overridden in every settings source classes.
  209. Args:
  210. field: The field.
  211. field_name: The field name.
  212. Returns:
  213. A tuple contains the key, value and a flag to determine whether value is complex.
  214. """
  215. pass
  216. def field_is_complex(self, field: FieldInfo) -> bool:
  217. """
  218. Checks whether a field is complex, in which case it will attempt to be parsed as JSON.
  219. Args:
  220. field: The field.
  221. Returns:
  222. Whether the field is complex.
  223. """
  224. return _annotation_is_complex(field.annotation, field.metadata)
  225. def prepare_field_value(self, field_name: str, field: FieldInfo, value: Any, value_is_complex: bool) -> Any:
  226. """
  227. Prepares the value of a field.
  228. Args:
  229. field_name: The field name.
  230. field: The field.
  231. value: The value of the field that has to be prepared.
  232. value_is_complex: A flag to determine whether value is complex.
  233. Returns:
  234. The prepared value.
  235. """
  236. if value is not None and (self.field_is_complex(field) or value_is_complex):
  237. return self.decode_complex_value(field_name, field, value)
  238. return value
  239. def decode_complex_value(self, field_name: str, field: FieldInfo, value: Any) -> Any:
  240. """
  241. Decode the value for a complex field
  242. Args:
  243. field_name: The field name.
  244. field: The field.
  245. value: The value of the field that has to be prepared.
  246. Returns:
  247. The decoded value for further preparation
  248. """
  249. return json.loads(value)
  250. @abstractmethod
  251. def __call__(self) -> dict[str, Any]:
  252. pass
  253. class DefaultSettingsSource(PydanticBaseSettingsSource):
  254. """
  255. Source class for loading default object values.
  256. Args:
  257. settings_cls: The Settings class.
  258. nested_model_default_partial_update: Whether to allow partial updates on nested model default object fields.
  259. Defaults to `False`.
  260. """
  261. def __init__(self, settings_cls: type[BaseSettings], nested_model_default_partial_update: bool | None = None):
  262. super().__init__(settings_cls)
  263. self.defaults: dict[str, Any] = {}
  264. self.nested_model_default_partial_update = (
  265. nested_model_default_partial_update
  266. if nested_model_default_partial_update is not None
  267. else self.config.get('nested_model_default_partial_update', False)
  268. )
  269. if self.nested_model_default_partial_update:
  270. for field_name, field_info in settings_cls.model_fields.items():
  271. if is_dataclass(type(field_info.default)):
  272. self.defaults[_field_name_for_signature(field_name, field_info)] = asdict(field_info.default)
  273. elif is_model_class(type(field_info.default)):
  274. self.defaults[_field_name_for_signature(field_name, field_info)] = field_info.default.model_dump()
  275. def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]:
  276. # Nothing to do here. Only implement the return statement to make mypy happy
  277. return None, '', False
  278. def __call__(self) -> dict[str, Any]:
  279. return self.defaults
  280. def __repr__(self) -> str:
  281. return f'DefaultSettingsSource(nested_model_default_partial_update={self.nested_model_default_partial_update})'
  282. class InitSettingsSource(PydanticBaseSettingsSource):
  283. """
  284. Source class for loading values provided during settings class initialization.
  285. """
  286. def __init__(
  287. self,
  288. settings_cls: type[BaseSettings],
  289. init_kwargs: dict[str, Any],
  290. nested_model_default_partial_update: bool | None = None,
  291. ):
  292. self.init_kwargs = init_kwargs
  293. super().__init__(settings_cls)
  294. self.nested_model_default_partial_update = (
  295. nested_model_default_partial_update
  296. if nested_model_default_partial_update is not None
  297. else self.config.get('nested_model_default_partial_update', False)
  298. )
  299. def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]:
  300. # Nothing to do here. Only implement the return statement to make mypy happy
  301. return None, '', False
  302. def __call__(self) -> dict[str, Any]:
  303. return (
  304. TypeAdapter(Dict[str, Any]).dump_python(self.init_kwargs)
  305. if self.nested_model_default_partial_update
  306. else self.init_kwargs
  307. )
  308. def __repr__(self) -> str:
  309. return f'InitSettingsSource(init_kwargs={self.init_kwargs!r})'
  310. class PydanticBaseEnvSettingsSource(PydanticBaseSettingsSource):
  311. def __init__(
  312. self,
  313. settings_cls: type[BaseSettings],
  314. case_sensitive: bool | None = None,
  315. env_prefix: str | None = None,
  316. env_ignore_empty: bool | None = None,
  317. env_parse_none_str: str | None = None,
  318. env_parse_enums: bool | None = None,
  319. ) -> None:
  320. super().__init__(settings_cls)
  321. self.case_sensitive = case_sensitive if case_sensitive is not None else self.config.get('case_sensitive', False)
  322. self.env_prefix = env_prefix if env_prefix is not None else self.config.get('env_prefix', '')
  323. self.env_ignore_empty = (
  324. env_ignore_empty if env_ignore_empty is not None else self.config.get('env_ignore_empty', False)
  325. )
  326. self.env_parse_none_str = (
  327. env_parse_none_str if env_parse_none_str is not None else self.config.get('env_parse_none_str')
  328. )
  329. self.env_parse_enums = env_parse_enums if env_parse_enums is not None else self.config.get('env_parse_enums')
  330. def _apply_case_sensitive(self, value: str) -> str:
  331. return value.lower() if not self.case_sensitive else value
  332. def _extract_field_info(self, field: FieldInfo, field_name: str) -> list[tuple[str, str, bool]]:
  333. """
  334. Extracts field info. This info is used to get the value of field from environment variables.
  335. It returns a list of tuples, each tuple contains:
  336. * field_key: The key of field that has to be used in model creation.
  337. * env_name: The environment variable name of the field.
  338. * value_is_complex: A flag to determine whether the value from environment variable
  339. is complex and has to be parsed.
  340. Args:
  341. field (FieldInfo): The field.
  342. field_name (str): The field name.
  343. Returns:
  344. list[tuple[str, str, bool]]: List of tuples, each tuple contains field_key, env_name, and value_is_complex.
  345. """
  346. field_info: list[tuple[str, str, bool]] = []
  347. if isinstance(field.validation_alias, (AliasChoices, AliasPath)):
  348. v_alias: str | list[str | int] | list[list[str | int]] | None = field.validation_alias.convert_to_aliases()
  349. else:
  350. v_alias = field.validation_alias
  351. if v_alias:
  352. if isinstance(v_alias, list): # AliasChoices, AliasPath
  353. for alias in v_alias:
  354. if isinstance(alias, str): # AliasPath
  355. field_info.append((alias, self._apply_case_sensitive(alias), True if len(alias) > 1 else False))
  356. elif isinstance(alias, list): # AliasChoices
  357. first_arg = cast(str, alias[0]) # first item of an AliasChoices must be a str
  358. field_info.append(
  359. (first_arg, self._apply_case_sensitive(first_arg), True if len(alias) > 1 else False)
  360. )
  361. else: # string validation alias
  362. field_info.append((v_alias, self._apply_case_sensitive(v_alias), False))
  363. if not v_alias or self.config.get('populate_by_name', False):
  364. if origin_is_union(get_origin(field.annotation)) and _union_is_complex(field.annotation, field.metadata):
  365. field_info.append((field_name, self._apply_case_sensitive(self.env_prefix + field_name), True))
  366. else:
  367. field_info.append((field_name, self._apply_case_sensitive(self.env_prefix + field_name), False))
  368. return field_info
  369. def _replace_field_names_case_insensitively(self, field: FieldInfo, field_values: dict[str, Any]) -> dict[str, Any]:
  370. """
  371. Replace field names in values dict by looking in models fields insensitively.
  372. By having the following models:
  373. ```py
  374. class SubSubSub(BaseModel):
  375. VaL3: str
  376. class SubSub(BaseModel):
  377. Val2: str
  378. SUB_sub_SuB: SubSubSub
  379. class Sub(BaseModel):
  380. VAL1: str
  381. SUB_sub: SubSub
  382. class Settings(BaseSettings):
  383. nested: Sub
  384. model_config = SettingsConfigDict(env_nested_delimiter='__')
  385. ```
  386. Then:
  387. _replace_field_names_case_insensitively(
  388. field,
  389. {"val1": "v1", "sub_SUB": {"VAL2": "v2", "sub_SUB_sUb": {"vAl3": "v3"}}}
  390. )
  391. Returns {'VAL1': 'v1', 'SUB_sub': {'Val2': 'v2', 'SUB_sub_SuB': {'VaL3': 'v3'}}}
  392. """
  393. values: dict[str, Any] = {}
  394. for name, value in field_values.items():
  395. sub_model_field: FieldInfo | None = None
  396. annotation = field.annotation
  397. # If field is Optional, we need to find the actual type
  398. args = get_args(annotation)
  399. if origin_is_union(get_origin(field.annotation)) and len(args) == 2 and type(None) in args:
  400. for arg in args:
  401. if arg is not None:
  402. annotation = arg
  403. break
  404. # This is here to make mypy happy
  405. # Item "None" of "Optional[Type[Any]]" has no attribute "model_fields"
  406. if not annotation or not hasattr(annotation, 'model_fields'):
  407. values[name] = value
  408. continue
  409. # Find field in sub model by looking in fields case insensitively
  410. for sub_model_field_name, f in annotation.model_fields.items():
  411. if not f.validation_alias and sub_model_field_name.lower() == name.lower():
  412. sub_model_field = f
  413. break
  414. if not sub_model_field:
  415. values[name] = value
  416. continue
  417. if lenient_issubclass(sub_model_field.annotation, BaseModel) and isinstance(value, dict):
  418. values[sub_model_field_name] = self._replace_field_names_case_insensitively(sub_model_field, value)
  419. else:
  420. values[sub_model_field_name] = value
  421. return values
  422. def _replace_env_none_type_values(self, field_value: dict[str, Any]) -> dict[str, Any]:
  423. """
  424. Recursively parse values that are of "None" type(EnvNoneType) to `None` type(None).
  425. """
  426. values: dict[str, Any] = {}
  427. for key, value in field_value.items():
  428. if not isinstance(value, EnvNoneType):
  429. values[key] = value if not isinstance(value, dict) else self._replace_env_none_type_values(value)
  430. else:
  431. values[key] = None
  432. return values
  433. def __call__(self) -> dict[str, Any]:
  434. data: dict[str, Any] = {}
  435. for field_name, field in self.settings_cls.model_fields.items():
  436. try:
  437. field_value, field_key, value_is_complex = self.get_field_value(field, field_name)
  438. except Exception as e:
  439. raise SettingsError(
  440. f'error getting value for field "{field_name}" from source "{self.__class__.__name__}"'
  441. ) from e
  442. try:
  443. field_value = self.prepare_field_value(field_name, field, field_value, value_is_complex)
  444. except ValueError as e:
  445. raise SettingsError(
  446. f'error parsing value for field "{field_name}" from source "{self.__class__.__name__}"'
  447. ) from e
  448. if field_value is not None:
  449. if self.env_parse_none_str is not None:
  450. if isinstance(field_value, dict):
  451. field_value = self._replace_env_none_type_values(field_value)
  452. elif isinstance(field_value, EnvNoneType):
  453. field_value = None
  454. if (
  455. not self.case_sensitive
  456. # and lenient_issubclass(field.annotation, BaseModel)
  457. and isinstance(field_value, dict)
  458. ):
  459. data[field_key] = self._replace_field_names_case_insensitively(field, field_value)
  460. else:
  461. data[field_key] = field_value
  462. return data
  463. class SecretsSettingsSource(PydanticBaseEnvSettingsSource):
  464. """
  465. Source class for loading settings values from secret files.
  466. """
  467. def __init__(
  468. self,
  469. settings_cls: type[BaseSettings],
  470. secrets_dir: PathType | None = None,
  471. case_sensitive: bool | None = None,
  472. env_prefix: str | None = None,
  473. env_ignore_empty: bool | None = None,
  474. env_parse_none_str: str | None = None,
  475. env_parse_enums: bool | None = None,
  476. ) -> None:
  477. super().__init__(
  478. settings_cls, case_sensitive, env_prefix, env_ignore_empty, env_parse_none_str, env_parse_enums
  479. )
  480. self.secrets_dir = secrets_dir if secrets_dir is not None else self.config.get('secrets_dir')
  481. def __call__(self) -> dict[str, Any]:
  482. """
  483. Build fields from "secrets" files.
  484. """
  485. secrets: dict[str, str | None] = {}
  486. if self.secrets_dir is None:
  487. return secrets
  488. secrets_dirs = [self.secrets_dir] if isinstance(self.secrets_dir, (str, os.PathLike)) else self.secrets_dir
  489. secrets_paths = [Path(p).expanduser() for p in secrets_dirs]
  490. self.secrets_paths = []
  491. for path in secrets_paths:
  492. if not path.exists():
  493. warnings.warn(f'directory "{path}" does not exist')
  494. else:
  495. self.secrets_paths.append(path)
  496. if not len(self.secrets_paths):
  497. return secrets
  498. for path in self.secrets_paths:
  499. if not path.is_dir():
  500. raise SettingsError(f'secrets_dir must reference a directory, not a {path_type_label(path)}')
  501. return super().__call__()
  502. @classmethod
  503. def find_case_path(cls, dir_path: Path, file_name: str, case_sensitive: bool) -> Path | None:
  504. """
  505. Find a file within path's directory matching filename, optionally ignoring case.
  506. Args:
  507. dir_path: Directory path.
  508. file_name: File name.
  509. case_sensitive: Whether to search for file name case sensitively.
  510. Returns:
  511. Whether file path or `None` if file does not exist in directory.
  512. """
  513. for f in dir_path.iterdir():
  514. if f.name == file_name:
  515. return f
  516. elif not case_sensitive and f.name.lower() == file_name.lower():
  517. return f
  518. return None
  519. def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]:
  520. """
  521. Gets the value for field from secret file and a flag to determine whether value is complex.
  522. Args:
  523. field: The field.
  524. field_name: The field name.
  525. Returns:
  526. A tuple contains the key, value if the file exists otherwise `None`, and
  527. a flag to determine whether value is complex.
  528. """
  529. for field_key, env_name, value_is_complex in self._extract_field_info(field, field_name):
  530. # paths reversed to match the last-wins behaviour of `env_file`
  531. for secrets_path in reversed(self.secrets_paths):
  532. path = self.find_case_path(secrets_path, env_name, self.case_sensitive)
  533. if not path:
  534. # path does not exist, we currently don't return a warning for this
  535. continue
  536. if path.is_file():
  537. return path.read_text().strip(), field_key, value_is_complex
  538. else:
  539. warnings.warn(
  540. f'attempted to load secret file "{path}" but found a {path_type_label(path)} instead.',
  541. stacklevel=4,
  542. )
  543. return None, field_key, value_is_complex
  544. def __repr__(self) -> str:
  545. return f'SecretsSettingsSource(secrets_dir={self.secrets_dir!r})'
  546. class EnvSettingsSource(PydanticBaseEnvSettingsSource):
  547. """
  548. Source class for loading settings values from environment variables.
  549. """
  550. def __init__(
  551. self,
  552. settings_cls: type[BaseSettings],
  553. case_sensitive: bool | None = None,
  554. env_prefix: str | None = None,
  555. env_nested_delimiter: str | None = None,
  556. env_ignore_empty: bool | None = None,
  557. env_parse_none_str: str | None = None,
  558. env_parse_enums: bool | None = None,
  559. ) -> None:
  560. super().__init__(
  561. settings_cls, case_sensitive, env_prefix, env_ignore_empty, env_parse_none_str, env_parse_enums
  562. )
  563. self.env_nested_delimiter = (
  564. env_nested_delimiter if env_nested_delimiter is not None else self.config.get('env_nested_delimiter')
  565. )
  566. self.env_prefix_len = len(self.env_prefix)
  567. self.env_vars = self._load_env_vars()
  568. def _load_env_vars(self) -> Mapping[str, str | None]:
  569. return parse_env_vars(os.environ, self.case_sensitive, self.env_ignore_empty, self.env_parse_none_str)
  570. def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]:
  571. """
  572. Gets the value for field from environment variables and a flag to determine whether value is complex.
  573. Args:
  574. field: The field.
  575. field_name: The field name.
  576. Returns:
  577. A tuple contains the key, value if the file exists otherwise `None`, and
  578. a flag to determine whether value is complex.
  579. """
  580. env_val: str | None = None
  581. for field_key, env_name, value_is_complex in self._extract_field_info(field, field_name):
  582. env_val = self.env_vars.get(env_name)
  583. if env_val is not None:
  584. break
  585. return env_val, field_key, value_is_complex
  586. def prepare_field_value(self, field_name: str, field: FieldInfo, value: Any, value_is_complex: bool) -> Any:
  587. """
  588. Prepare value for the field.
  589. * Extract value for nested field.
  590. * Deserialize value to python object for complex field.
  591. Args:
  592. field: The field.
  593. field_name: The field name.
  594. Returns:
  595. A tuple contains prepared value for the field.
  596. Raises:
  597. ValuesError: When There is an error in deserializing value for complex field.
  598. """
  599. is_complex, allow_parse_failure = self._field_is_complex(field)
  600. if self.env_parse_enums:
  601. enum_val = _annotation_enum_name_to_val(field.annotation, value)
  602. value = value if enum_val is None else enum_val
  603. if is_complex or value_is_complex:
  604. if isinstance(value, EnvNoneType):
  605. return value
  606. elif value is None:
  607. # field is complex but no value found so far, try explode_env_vars
  608. env_val_built = self.explode_env_vars(field_name, field, self.env_vars)
  609. if env_val_built:
  610. return env_val_built
  611. else:
  612. # field is complex and there's a value, decode that as JSON, then add explode_env_vars
  613. try:
  614. value = self.decode_complex_value(field_name, field, value)
  615. except ValueError as e:
  616. if not allow_parse_failure:
  617. raise e
  618. if isinstance(value, dict):
  619. return deep_update(value, self.explode_env_vars(field_name, field, self.env_vars))
  620. else:
  621. return value
  622. elif value is not None:
  623. # simplest case, field is not complex, we only need to add the value if it was found
  624. return value
  625. def _field_is_complex(self, field: FieldInfo) -> tuple[bool, bool]:
  626. """
  627. Find out if a field is complex, and if so whether JSON errors should be ignored
  628. """
  629. if self.field_is_complex(field):
  630. allow_parse_failure = False
  631. elif origin_is_union(get_origin(field.annotation)) and _union_is_complex(field.annotation, field.metadata):
  632. allow_parse_failure = True
  633. else:
  634. return False, False
  635. return True, allow_parse_failure
  636. # Default value of `case_sensitive` is `None`, because we don't want to break existing behavior.
  637. # We have to change the method to a non-static method and use
  638. # `self.case_sensitive` instead in V3.
  639. def next_field(
  640. self, field: FieldInfo | Any | None, key: str, case_sensitive: bool | None = None
  641. ) -> FieldInfo | None:
  642. """
  643. Find the field in a sub model by key(env name)
  644. By having the following models:
  645. ```py
  646. class SubSubModel(BaseSettings):
  647. dvals: Dict
  648. class SubModel(BaseSettings):
  649. vals: list[str]
  650. sub_sub_model: SubSubModel
  651. class Cfg(BaseSettings):
  652. sub_model: SubModel
  653. ```
  654. Then:
  655. next_field(sub_model, 'vals') Returns the `vals` field of `SubModel` class
  656. next_field(sub_model, 'sub_sub_model') Returns `sub_sub_model` field of `SubModel` class
  657. Args:
  658. field: The field.
  659. key: The key (env name).
  660. case_sensitive: Whether to search for key case sensitively.
  661. Returns:
  662. Field if it finds the next field otherwise `None`.
  663. """
  664. if not field:
  665. return None
  666. annotation = field.annotation if isinstance(field, FieldInfo) else field
  667. if origin_is_union(get_origin(annotation)) or isinstance(annotation, WithArgsTypes):
  668. for type_ in get_args(annotation):
  669. type_has_key = self.next_field(type_, key, case_sensitive)
  670. if type_has_key:
  671. return type_has_key
  672. elif is_model_class(annotation) or is_pydantic_dataclass(annotation):
  673. fields = _get_model_fields(annotation)
  674. # `case_sensitive is None` is here to be compatible with the old behavior.
  675. # Has to be removed in V3.
  676. for field_name, f in fields.items():
  677. for _, env_name, _ in self._extract_field_info(f, field_name):
  678. if case_sensitive is None or case_sensitive:
  679. if field_name == key or env_name == key:
  680. return f
  681. elif field_name.lower() == key.lower() or env_name.lower() == key.lower():
  682. return f
  683. return None
  684. def explode_env_vars(self, field_name: str, field: FieldInfo, env_vars: Mapping[str, str | None]) -> dict[str, Any]:
  685. """
  686. Process env_vars and extract the values of keys containing env_nested_delimiter into nested dictionaries.
  687. This is applied to a single field, hence filtering by env_var prefix.
  688. Args:
  689. field_name: The field name.
  690. field: The field.
  691. env_vars: Environment variables.
  692. Returns:
  693. A dictionary contains extracted values from nested env values.
  694. """
  695. is_dict = lenient_issubclass(get_origin(field.annotation), dict)
  696. prefixes = [
  697. f'{env_name}{self.env_nested_delimiter}' for _, env_name, _ in self._extract_field_info(field, field_name)
  698. ]
  699. result: dict[str, Any] = {}
  700. for env_name, env_val in env_vars.items():
  701. if not any(env_name.startswith(prefix) for prefix in prefixes):
  702. continue
  703. # we remove the prefix before splitting in case the prefix has characters in common with the delimiter
  704. env_name_without_prefix = env_name[self.env_prefix_len :]
  705. _, *keys, last_key = env_name_without_prefix.split(self.env_nested_delimiter)
  706. env_var = result
  707. target_field: FieldInfo | None = field
  708. for key in keys:
  709. target_field = self.next_field(target_field, key, self.case_sensitive)
  710. if isinstance(env_var, dict):
  711. env_var = env_var.setdefault(key, {})
  712. # get proper field with last_key
  713. target_field = self.next_field(target_field, last_key, self.case_sensitive)
  714. # check if env_val maps to a complex field and if so, parse the env_val
  715. if (target_field or is_dict) and env_val:
  716. if target_field:
  717. is_complex, allow_json_failure = self._field_is_complex(target_field)
  718. else:
  719. # nested field type is dict
  720. is_complex, allow_json_failure = True, True
  721. if is_complex:
  722. try:
  723. env_val = self.decode_complex_value(last_key, target_field, env_val) # type: ignore
  724. except ValueError as e:
  725. if not allow_json_failure:
  726. raise e
  727. if isinstance(env_var, dict):
  728. if last_key not in env_var or not isinstance(env_val, EnvNoneType) or env_var[last_key] == {}:
  729. env_var[last_key] = env_val
  730. return result
  731. def __repr__(self) -> str:
  732. return (
  733. f'EnvSettingsSource(env_nested_delimiter={self.env_nested_delimiter!r}, '
  734. f'env_prefix_len={self.env_prefix_len!r})'
  735. )
  736. class DotEnvSettingsSource(EnvSettingsSource):
  737. """
  738. Source class for loading settings values from env files.
  739. """
  740. def __init__(
  741. self,
  742. settings_cls: type[BaseSettings],
  743. env_file: DotenvType | None = ENV_FILE_SENTINEL,
  744. env_file_encoding: str | None = None,
  745. case_sensitive: bool | None = None,
  746. env_prefix: str | None = None,
  747. env_nested_delimiter: str | None = None,
  748. env_ignore_empty: bool | None = None,
  749. env_parse_none_str: str | None = None,
  750. env_parse_enums: bool | None = None,
  751. ) -> None:
  752. self.env_file = env_file if env_file != ENV_FILE_SENTINEL else settings_cls.model_config.get('env_file')
  753. self.env_file_encoding = (
  754. env_file_encoding if env_file_encoding is not None else settings_cls.model_config.get('env_file_encoding')
  755. )
  756. super().__init__(
  757. settings_cls,
  758. case_sensitive,
  759. env_prefix,
  760. env_nested_delimiter,
  761. env_ignore_empty,
  762. env_parse_none_str,
  763. env_parse_enums,
  764. )
  765. def _load_env_vars(self) -> Mapping[str, str | None]:
  766. return self._read_env_files()
  767. @staticmethod
  768. def _static_read_env_file(
  769. file_path: Path,
  770. *,
  771. encoding: str | None = None,
  772. case_sensitive: bool = False,
  773. ignore_empty: bool = False,
  774. parse_none_str: str | None = None,
  775. ) -> Mapping[str, str | None]:
  776. file_vars: dict[str, str | None] = dotenv_values(file_path, encoding=encoding or 'utf8')
  777. return parse_env_vars(file_vars, case_sensitive, ignore_empty, parse_none_str)
  778. def _read_env_file(
  779. self,
  780. file_path: Path,
  781. ) -> Mapping[str, str | None]:
  782. return self._static_read_env_file(
  783. file_path,
  784. encoding=self.env_file_encoding,
  785. case_sensitive=self.case_sensitive,
  786. ignore_empty=self.env_ignore_empty,
  787. parse_none_str=self.env_parse_none_str,
  788. )
  789. def _read_env_files(self) -> Mapping[str, str | None]:
  790. env_files = self.env_file
  791. if env_files is None:
  792. return {}
  793. if isinstance(env_files, (str, os.PathLike)):
  794. env_files = [env_files]
  795. dotenv_vars: dict[str, str | None] = {}
  796. for env_file in env_files:
  797. env_path = Path(env_file).expanduser()
  798. if env_path.is_file():
  799. dotenv_vars.update(self._read_env_file(env_path))
  800. return dotenv_vars
  801. def __call__(self) -> dict[str, Any]:
  802. data: dict[str, Any] = super().__call__()
  803. is_extra_allowed = self.config.get('extra') != 'forbid'
  804. # As `extra` config is allowed in dotenv settings source, We have to
  805. # update data with extra env variables from dotenv file.
  806. for env_name, env_value in self.env_vars.items():
  807. if not env_value or env_name in data:
  808. continue
  809. env_used = False
  810. for field_name, field in self.settings_cls.model_fields.items():
  811. for _, field_env_name, _ in self._extract_field_info(field, field_name):
  812. if env_name == field_env_name or (
  813. (
  814. _annotation_is_complex(field.annotation, field.metadata)
  815. or (
  816. origin_is_union(get_origin(field.annotation))
  817. and _union_is_complex(field.annotation, field.metadata)
  818. )
  819. )
  820. and env_name.startswith(field_env_name)
  821. ):
  822. env_used = True
  823. break
  824. if env_used:
  825. break
  826. if not env_used:
  827. if is_extra_allowed and env_name.startswith(self.env_prefix):
  828. # env_prefix should be respected and removed from the env_name
  829. normalized_env_name = env_name[len(self.env_prefix) :]
  830. data[normalized_env_name] = env_value
  831. else:
  832. data[env_name] = env_value
  833. return data
  834. def __repr__(self) -> str:
  835. return (
  836. f'DotEnvSettingsSource(env_file={self.env_file!r}, env_file_encoding={self.env_file_encoding!r}, '
  837. f'env_nested_delimiter={self.env_nested_delimiter!r}, env_prefix_len={self.env_prefix_len!r})'
  838. )
  839. class CliSettingsSource(EnvSettingsSource, Generic[T]):
  840. """
  841. Source class for loading settings values from CLI.
  842. Note:
  843. A `CliSettingsSource` connects with a `root_parser` object by using the parser methods to add
  844. `settings_cls` fields as command line arguments. The `CliSettingsSource` internal parser representation
  845. is based upon the `argparse` parsing library, and therefore, requires the parser methods to support
  846. the same attributes as their `argparse` library counterparts.
  847. Args:
  848. cli_prog_name: The CLI program name to display in help text. Defaults to `None` if cli_parse_args is `None`.
  849. Otherwse, defaults to sys.argv[0].
  850. cli_parse_args: The list of CLI arguments to parse. Defaults to None.
  851. If set to `True`, defaults to sys.argv[1:].
  852. cli_parse_none_str: The CLI string value that should be parsed (e.g. "null", "void", "None", etc.) into `None`
  853. type(None). Defaults to "null" if cli_avoid_json is `False`, and "None" if cli_avoid_json is `True`.
  854. cli_hide_none_type: Hide `None` values in CLI help text. Defaults to `False`.
  855. cli_avoid_json: Avoid complex JSON objects in CLI help text. Defaults to `False`.
  856. cli_enforce_required: Enforce required fields at the CLI. Defaults to `False`.
  857. cli_use_class_docs_for_groups: Use class docstrings in CLI group help text instead of field descriptions.
  858. Defaults to `False`.
  859. cli_exit_on_error: Determines whether or not the internal parser exits with error info when an error occurs.
  860. Defaults to `True`.
  861. cli_prefix: Prefix for command line arguments added under the root parser. Defaults to "".
  862. cli_flag_prefix_char: The flag prefix character to use for CLI optional arguments. Defaults to '-'.
  863. cli_implicit_flags: Whether `bool` fields should be implicitly converted into CLI boolean flags.
  864. (e.g. --flag, --no-flag). Defaults to `False`.
  865. cli_ignore_unknown_args: Whether to ignore unknown CLI args and parse only known ones. Defaults to `False`.
  866. case_sensitive: Whether CLI "--arg" names should be read with case-sensitivity. Defaults to `True`.
  867. Note: Case-insensitive matching is only supported on the internal root parser and does not apply to CLI
  868. subcommands.
  869. root_parser: The root parser object.
  870. parse_args_method: The root parser parse args method. Defaults to `argparse.ArgumentParser.parse_args`.
  871. add_argument_method: The root parser add argument method. Defaults to `argparse.ArgumentParser.add_argument`.
  872. add_argument_group_method: The root parser add argument group method.
  873. Defaults to `argparse.ArgumentParser.add_argument_group`.
  874. add_parser_method: The root parser add new parser (sub-command) method.
  875. Defaults to `argparse._SubParsersAction.add_parser`.
  876. add_subparsers_method: The root parser add subparsers (sub-commands) method.
  877. Defaults to `argparse.ArgumentParser.add_subparsers`.
  878. formatter_class: A class for customizing the root parser help text. Defaults to `argparse.RawDescriptionHelpFormatter`.
  879. """
  880. def __init__(
  881. self,
  882. settings_cls: type[BaseSettings],
  883. cli_prog_name: str | None = None,
  884. cli_parse_args: bool | list[str] | tuple[str, ...] | None = None,
  885. cli_parse_none_str: str | None = None,
  886. cli_hide_none_type: bool | None = None,
  887. cli_avoid_json: bool | None = None,
  888. cli_enforce_required: bool | None = None,
  889. cli_use_class_docs_for_groups: bool | None = None,
  890. cli_exit_on_error: bool | None = None,
  891. cli_prefix: str | None = None,
  892. cli_flag_prefix_char: str | None = None,
  893. cli_implicit_flags: bool | None = None,
  894. cli_ignore_unknown_args: bool | None = None,
  895. case_sensitive: bool | None = True,
  896. root_parser: Any = None,
  897. parse_args_method: Callable[..., Any] | None = None,
  898. add_argument_method: Callable[..., Any] | None = ArgumentParser.add_argument,
  899. add_argument_group_method: Callable[..., Any] | None = ArgumentParser.add_argument_group,
  900. add_parser_method: Callable[..., Any] | None = _SubParsersAction.add_parser,
  901. add_subparsers_method: Callable[..., Any] | None = ArgumentParser.add_subparsers,
  902. formatter_class: Any = RawDescriptionHelpFormatter,
  903. ) -> None:
  904. self.cli_prog_name = (
  905. cli_prog_name if cli_prog_name is not None else settings_cls.model_config.get('cli_prog_name', sys.argv[0])
  906. )
  907. self.cli_hide_none_type = (
  908. cli_hide_none_type
  909. if cli_hide_none_type is not None
  910. else settings_cls.model_config.get('cli_hide_none_type', False)
  911. )
  912. self.cli_avoid_json = (
  913. cli_avoid_json if cli_avoid_json is not None else settings_cls.model_config.get('cli_avoid_json', False)
  914. )
  915. if not cli_parse_none_str:
  916. cli_parse_none_str = 'None' if self.cli_avoid_json is True else 'null'
  917. self.cli_parse_none_str = cli_parse_none_str
  918. self.cli_enforce_required = (
  919. cli_enforce_required
  920. if cli_enforce_required is not None
  921. else settings_cls.model_config.get('cli_enforce_required', False)
  922. )
  923. self.cli_use_class_docs_for_groups = (
  924. cli_use_class_docs_for_groups
  925. if cli_use_class_docs_for_groups is not None
  926. else settings_cls.model_config.get('cli_use_class_docs_for_groups', False)
  927. )
  928. self.cli_exit_on_error = (
  929. cli_exit_on_error
  930. if cli_exit_on_error is not None
  931. else settings_cls.model_config.get('cli_exit_on_error', True)
  932. )
  933. self.cli_prefix = cli_prefix if cli_prefix is not None else settings_cls.model_config.get('cli_prefix', '')
  934. self.cli_flag_prefix_char = (
  935. cli_flag_prefix_char
  936. if cli_flag_prefix_char is not None
  937. else settings_cls.model_config.get('cli_flag_prefix_char', '-')
  938. )
  939. self._cli_flag_prefix = self.cli_flag_prefix_char * 2
  940. if self.cli_prefix:
  941. if cli_prefix.startswith('.') or cli_prefix.endswith('.') or not cli_prefix.replace('.', '').isidentifier(): # type: ignore
  942. raise SettingsError(f'CLI settings source prefix is invalid: {cli_prefix}')
  943. self.cli_prefix += '.'
  944. self.cli_implicit_flags = (
  945. cli_implicit_flags
  946. if cli_implicit_flags is not None
  947. else settings_cls.model_config.get('cli_implicit_flags', False)
  948. )
  949. self.cli_ignore_unknown_args = (
  950. cli_ignore_unknown_args
  951. if cli_ignore_unknown_args is not None
  952. else settings_cls.model_config.get('cli_ignore_unknown_args', False)
  953. )
  954. case_sensitive = case_sensitive if case_sensitive is not None else True
  955. if not case_sensitive and root_parser is not None:
  956. raise SettingsError('Case-insensitive matching is only supported on the internal root parser')
  957. super().__init__(
  958. settings_cls,
  959. env_nested_delimiter='.',
  960. env_parse_none_str=self.cli_parse_none_str,
  961. env_parse_enums=True,
  962. env_prefix=self.cli_prefix,
  963. case_sensitive=case_sensitive,
  964. )
  965. root_parser = (
  966. _CliInternalArgParser(
  967. cli_exit_on_error=self.cli_exit_on_error,
  968. prog=self.cli_prog_name,
  969. description=None if settings_cls.__doc__ is None else dedent(settings_cls.__doc__),
  970. formatter_class=formatter_class,
  971. prefix_chars=self.cli_flag_prefix_char,
  972. )
  973. if root_parser is None
  974. else root_parser
  975. )
  976. self._connect_root_parser(
  977. root_parser=root_parser,
  978. parse_args_method=parse_args_method,
  979. add_argument_method=add_argument_method,
  980. add_argument_group_method=add_argument_group_method,
  981. add_parser_method=add_parser_method,
  982. add_subparsers_method=add_subparsers_method,
  983. formatter_class=formatter_class,
  984. )
  985. if cli_parse_args not in (None, False):
  986. if cli_parse_args is True:
  987. cli_parse_args = sys.argv[1:]
  988. elif not isinstance(cli_parse_args, (list, tuple)):
  989. raise SettingsError(
  990. f'cli_parse_args must be List[str] or Tuple[str, ...], recieved {type(cli_parse_args)}'
  991. )
  992. self._load_env_vars(parsed_args=self._parse_args(self.root_parser, cli_parse_args))
  993. @overload
  994. def __call__(self) -> dict[str, Any]: ...
  995. @overload
  996. def __call__(self, *, args: list[str] | tuple[str, ...] | bool) -> CliSettingsSource[T]:
  997. """
  998. Parse and load the command line arguments list into the CLI settings source.
  999. Args:
  1000. args:
  1001. The command line arguments to parse and load. Defaults to `None`, which means do not parse
  1002. command line arguments. If set to `True`, defaults to sys.argv[1:]. If set to `False`, does
  1003. not parse command line arguments.
  1004. Returns:
  1005. CliSettingsSource: The object instance itself.
  1006. """
  1007. ...
  1008. @overload
  1009. def __call__(self, *, parsed_args: Namespace | SimpleNamespace | dict[str, Any]) -> CliSettingsSource[T]:
  1010. """
  1011. Loads parsed command line arguments into the CLI settings source.
  1012. Note:
  1013. The parsed args must be in `argparse.Namespace`, `SimpleNamespace`, or vars dictionary
  1014. (e.g., vars(argparse.Namespace)) format.
  1015. Args:
  1016. parsed_args: The parsed args to load.
  1017. Returns:
  1018. CliSettingsSource: The object instance itself.
  1019. """
  1020. ...
  1021. def __call__(
  1022. self,
  1023. *,
  1024. args: list[str] | tuple[str, ...] | bool | None = None,
  1025. parsed_args: Namespace | SimpleNamespace | dict[str, list[str] | str] | None = None,
  1026. ) -> dict[str, Any] | CliSettingsSource[T]:
  1027. if args is not None and parsed_args is not None:
  1028. raise SettingsError('`args` and `parsed_args` are mutually exclusive')
  1029. elif args is not None:
  1030. if args is False:
  1031. return self._load_env_vars(parsed_args={})
  1032. if args is True:
  1033. args = sys.argv[1:]
  1034. return self._load_env_vars(parsed_args=self._parse_args(self.root_parser, args))
  1035. elif parsed_args is not None:
  1036. return self._load_env_vars(parsed_args=parsed_args)
  1037. else:
  1038. return super().__call__()
  1039. @overload
  1040. def _load_env_vars(self) -> Mapping[str, str | None]: ...
  1041. @overload
  1042. def _load_env_vars(self, *, parsed_args: Namespace | SimpleNamespace | dict[str, Any]) -> CliSettingsSource[T]:
  1043. """
  1044. Loads the parsed command line arguments into the CLI environment settings variables.
  1045. Note:
  1046. The parsed args must be in `argparse.Namespace`, `SimpleNamespace`, or vars dictionary
  1047. (e.g., vars(argparse.Namespace)) format.
  1048. Args:
  1049. parsed_args: The parsed args to load.
  1050. Returns:
  1051. CliSettingsSource: The object instance itself.
  1052. """
  1053. ...
  1054. def _load_env_vars(
  1055. self, *, parsed_args: Namespace | SimpleNamespace | dict[str, list[str] | str] | None = None
  1056. ) -> Mapping[str, str | None] | CliSettingsSource[T]:
  1057. if parsed_args is None:
  1058. return {}
  1059. if isinstance(parsed_args, (Namespace, SimpleNamespace)):
  1060. parsed_args = vars(parsed_args)
  1061. selected_subcommands: list[str] = []
  1062. for field_name, val in parsed_args.items():
  1063. if isinstance(val, list):
  1064. parsed_args[field_name] = self._merge_parsed_list(val, field_name)
  1065. elif field_name.endswith(':subcommand') and val is not None:
  1066. subcommand_name = field_name.split(':')[0] + val
  1067. subcommand_dest = self._cli_subcommands[field_name][subcommand_name]
  1068. selected_subcommands.append(subcommand_dest)
  1069. for subcommands in self._cli_subcommands.values():
  1070. for subcommand_dest in subcommands.values():
  1071. if subcommand_dest not in selected_subcommands:
  1072. parsed_args[subcommand_dest] = self.cli_parse_none_str
  1073. parsed_args = {key: val for key, val in parsed_args.items() if not key.endswith(':subcommand')}
  1074. if selected_subcommands:
  1075. last_selected_subcommand = max(selected_subcommands, key=len)
  1076. if not any(field_name for field_name in parsed_args.keys() if f'{last_selected_subcommand}.' in field_name):
  1077. parsed_args[last_selected_subcommand] = '{}'
  1078. self.env_vars = parse_env_vars(
  1079. cast(Mapping[str, str], parsed_args),
  1080. self.case_sensitive,
  1081. self.env_ignore_empty,
  1082. self.cli_parse_none_str,
  1083. )
  1084. return self
  1085. def _get_merge_parsed_list_types(
  1086. self, parsed_list: list[str], field_name: str
  1087. ) -> tuple[Optional[type], Optional[type]]:
  1088. merge_type = self._cli_dict_args.get(field_name, list)
  1089. if (
  1090. merge_type is list
  1091. or not origin_is_union(get_origin(merge_type))
  1092. or not any(
  1093. type_
  1094. for type_ in get_args(merge_type)
  1095. if type_ is not type(None) and get_origin(type_) not in (dict, Mapping)
  1096. )
  1097. ):
  1098. inferred_type = merge_type
  1099. else:
  1100. inferred_type = list if parsed_list and (len(parsed_list) > 1 or parsed_list[0].startswith('[')) else str
  1101. return merge_type, inferred_type
  1102. def _merge_parsed_list(self, parsed_list: list[str], field_name: str) -> str:
  1103. try:
  1104. merged_list: list[str] = []
  1105. is_last_consumed_a_value = False
  1106. merge_type, inferred_type = self._get_merge_parsed_list_types(parsed_list, field_name)
  1107. for val in parsed_list:
  1108. if not isinstance(val, str):
  1109. # If val is not a string, it's from an external parser and we can ignore parsing the rest of the
  1110. # list.
  1111. break
  1112. val = val.strip()
  1113. if val.startswith('[') and val.endswith(']'):
  1114. val = val[1:-1].strip()
  1115. while val:
  1116. val = val.strip()
  1117. if val.startswith(','):
  1118. val = self._consume_comma(val, merged_list, is_last_consumed_a_value)
  1119. is_last_consumed_a_value = False
  1120. else:
  1121. if val.startswith('{') or val.startswith('['):
  1122. val = self._consume_object_or_array(val, merged_list)
  1123. else:
  1124. try:
  1125. val = self._consume_string_or_number(val, merged_list, merge_type)
  1126. except ValueError as e:
  1127. if merge_type is inferred_type:
  1128. raise e
  1129. merge_type = inferred_type
  1130. val = self._consume_string_or_number(val, merged_list, merge_type)
  1131. is_last_consumed_a_value = True
  1132. if not is_last_consumed_a_value:
  1133. val = self._consume_comma(val, merged_list, is_last_consumed_a_value)
  1134. if merge_type is str:
  1135. return merged_list[0]
  1136. elif merge_type is list:
  1137. return f'[{",".join(merged_list)}]'
  1138. else:
  1139. merged_dict: dict[str, str] = {}
  1140. for item in merged_list:
  1141. merged_dict.update(json.loads(item))
  1142. return json.dumps(merged_dict)
  1143. except Exception as e:
  1144. raise SettingsError(f'Parsing error encountered for {field_name}: {e}')
  1145. def _consume_comma(self, item: str, merged_list: list[str], is_last_consumed_a_value: bool) -> str:
  1146. if not is_last_consumed_a_value:
  1147. merged_list.append('""')
  1148. return item[1:]
  1149. def _consume_object_or_array(self, item: str, merged_list: list[str]) -> str:
  1150. count = 1
  1151. close_delim = '}' if item.startswith('{') else ']'
  1152. for consumed in range(1, len(item)):
  1153. if item[consumed] in ('{', '['):
  1154. count += 1
  1155. elif item[consumed] in ('}', ']'):
  1156. count -= 1
  1157. if item[consumed] == close_delim and count == 0:
  1158. merged_list.append(item[: consumed + 1])
  1159. return item[consumed + 1 :]
  1160. raise SettingsError(f'Missing end delimiter "{close_delim}"')
  1161. def _consume_string_or_number(self, item: str, merged_list: list[str], merge_type: type[Any] | None) -> str:
  1162. consumed = 0 if merge_type is not str else len(item)
  1163. is_find_end_quote = False
  1164. while consumed < len(item):
  1165. if item[consumed] == '"' and (consumed == 0 or item[consumed - 1] != '\\'):
  1166. is_find_end_quote = not is_find_end_quote
  1167. if not is_find_end_quote and item[consumed] == ',':
  1168. break
  1169. consumed += 1
  1170. if is_find_end_quote:
  1171. raise SettingsError('Mismatched quotes')
  1172. val_string = item[:consumed].strip()
  1173. if merge_type in (list, str):
  1174. try:
  1175. float(val_string)
  1176. except ValueError:
  1177. if val_string == self.cli_parse_none_str:
  1178. val_string = 'null'
  1179. if val_string not in ('true', 'false', 'null') and not val_string.startswith('"'):
  1180. val_string = f'"{val_string}"'
  1181. merged_list.append(val_string)
  1182. else:
  1183. key, val = (kv for kv in val_string.split('=', 1))
  1184. if key.startswith('"') and not key.endswith('"') and not val.startswith('"') and val.endswith('"'):
  1185. raise ValueError(f'Dictionary key=val parameter is a quoted string: {val_string}')
  1186. key, val = key.strip('"'), val.strip('"')
  1187. merged_list.append(json.dumps({key: val}))
  1188. return item[consumed:]
  1189. def _get_sub_models(self, model: type[BaseModel], field_name: str, field_info: FieldInfo) -> list[type[BaseModel]]:
  1190. field_types: tuple[Any, ...] = (
  1191. (field_info.annotation,) if not get_args(field_info.annotation) else get_args(field_info.annotation)
  1192. )
  1193. if self.cli_hide_none_type:
  1194. field_types = tuple([type_ for type_ in field_types if type_ is not type(None)])
  1195. sub_models: list[type[BaseModel]] = []
  1196. for type_ in field_types:
  1197. if _annotation_contains_types(type_, (_CliSubCommand,), is_include_origin=False):
  1198. raise SettingsError(f'CliSubCommand is not outermost annotation for {model.__name__}.{field_name}')
  1199. elif _annotation_contains_types(type_, (_CliPositionalArg,), is_include_origin=False):
  1200. raise SettingsError(f'CliPositionalArg is not outermost annotation for {model.__name__}.{field_name}')
  1201. if is_model_class(type_) or is_pydantic_dataclass(type_):
  1202. sub_models.append(type_) # type: ignore
  1203. return sub_models
  1204. def _get_alias_names(
  1205. self, field_name: str, field_info: FieldInfo, alias_path_args: dict[str, str]
  1206. ) -> tuple[tuple[str, ...], bool]:
  1207. alias_names: list[str] = []
  1208. is_alias_path_only: bool = True
  1209. if not any((field_info.alias, field_info.validation_alias)):
  1210. alias_names += [field_name]
  1211. is_alias_path_only = False
  1212. else:
  1213. new_alias_paths: list[AliasPath] = []
  1214. for alias in (field_info.alias, field_info.validation_alias):
  1215. if alias is None:
  1216. continue
  1217. elif isinstance(alias, str):
  1218. alias_names.append(alias)
  1219. is_alias_path_only = False
  1220. elif isinstance(alias, AliasChoices):
  1221. for name in alias.choices:
  1222. if isinstance(name, str):
  1223. alias_names.append(name)
  1224. is_alias_path_only = False
  1225. else:
  1226. new_alias_paths.append(name)
  1227. else:
  1228. new_alias_paths.append(alias)
  1229. for alias_path in new_alias_paths:
  1230. name = cast(str, alias_path.path[0])
  1231. name = name.lower() if not self.case_sensitive else name
  1232. alias_path_args[name] = 'dict' if len(alias_path.path) > 2 else 'list'
  1233. if not alias_names and is_alias_path_only:
  1234. alias_names.append(name)
  1235. if not self.case_sensitive:
  1236. alias_names = [alias_name.lower() for alias_name in alias_names]
  1237. return tuple(dict.fromkeys(alias_names)), is_alias_path_only
  1238. def _verify_cli_flag_annotations(self, model: type[BaseModel], field_name: str, field_info: FieldInfo) -> None:
  1239. if _CliImplicitFlag in field_info.metadata:
  1240. cli_flag_name = 'CliImplicitFlag'
  1241. elif _CliExplicitFlag in field_info.metadata:
  1242. cli_flag_name = 'CliExplicitFlag'
  1243. else:
  1244. return
  1245. if field_info.annotation is not bool:
  1246. raise SettingsError(f'{cli_flag_name} argument {model.__name__}.{field_name} is not of type bool')
  1247. elif sys.version_info < (3, 9) and (
  1248. field_info.default is PydanticUndefined and field_info.default_factory is None
  1249. ):
  1250. raise SettingsError(
  1251. f'{cli_flag_name} argument {model.__name__}.{field_name} must have default for python versions < 3.9'
  1252. )
  1253. def _sort_arg_fields(self, model: type[BaseModel]) -> list[tuple[str, FieldInfo]]:
  1254. positional_args, subcommand_args, optional_args = [], [], []
  1255. for field_name, field_info in _get_model_fields(model).items():
  1256. if _CliSubCommand in field_info.metadata:
  1257. if not field_info.is_required():
  1258. raise SettingsError(f'subcommand argument {model.__name__}.{field_name} has a default value')
  1259. else:
  1260. alias_names, *_ = self._get_alias_names(field_name, field_info, {})
  1261. if len(alias_names) > 1:
  1262. raise SettingsError(f'subcommand argument {model.__name__}.{field_name} has multiple aliases')
  1263. field_types = [type_ for type_ in get_args(field_info.annotation) if type_ is not type(None)]
  1264. for field_type in field_types:
  1265. if not (is_model_class(field_type) or is_pydantic_dataclass(field_type)):
  1266. raise SettingsError(
  1267. f'subcommand argument {model.__name__}.{field_name} has type not derived from BaseModel'
  1268. )
  1269. subcommand_args.append((field_name, field_info))
  1270. elif _CliPositionalArg in field_info.metadata:
  1271. if not field_info.is_required():
  1272. raise SettingsError(f'positional argument {model.__name__}.{field_name} has a default value')
  1273. else:
  1274. alias_names, *_ = self._get_alias_names(field_name, field_info, {})
  1275. if len(alias_names) > 1:
  1276. raise SettingsError(f'positional argument {model.__name__}.{field_name} has multiple aliases')
  1277. positional_args.append((field_name, field_info))
  1278. else:
  1279. self._verify_cli_flag_annotations(model, field_name, field_info)
  1280. optional_args.append((field_name, field_info))
  1281. return positional_args + subcommand_args + optional_args
  1282. @property
  1283. def root_parser(self) -> T:
  1284. """The connected root parser instance."""
  1285. return self._root_parser
  1286. def _connect_parser_method(
  1287. self, parser_method: Callable[..., Any] | None, method_name: str, *args: Any, **kwargs: Any
  1288. ) -> Callable[..., Any]:
  1289. if (
  1290. parser_method is not None
  1291. and self.case_sensitive is False
  1292. and method_name == 'parsed_args_method'
  1293. and isinstance(self._root_parser, _CliInternalArgParser)
  1294. ):
  1295. def parse_args_insensitive_method(
  1296. root_parser: _CliInternalArgParser,
  1297. args: list[str] | tuple[str, ...] | None = None,
  1298. namespace: Namespace | None = None,
  1299. ) -> Any:
  1300. insensitive_args = []
  1301. for arg in shlex.split(shlex.join(args)) if args else []:
  1302. flag_prefix = rf'\{self.cli_flag_prefix_char}{{1,2}}'
  1303. matched = re.match(rf'^({flag_prefix}[^\s=]+)(.*)', arg)
  1304. if matched:
  1305. arg = matched.group(1).lower() + matched.group(2)
  1306. insensitive_args.append(arg)
  1307. return parser_method(root_parser, insensitive_args, namespace) # type: ignore
  1308. return parse_args_insensitive_method
  1309. elif parser_method is None:
  1310. def none_parser_method(*args: Any, **kwargs: Any) -> Any:
  1311. raise SettingsError(
  1312. f'cannot connect CLI settings source root parser: {method_name} is set to `None` but is needed for connecting'
  1313. )
  1314. return none_parser_method
  1315. else:
  1316. return parser_method
  1317. def _connect_root_parser(
  1318. self,
  1319. root_parser: T,
  1320. parse_args_method: Callable[..., Any] | None,
  1321. add_argument_method: Callable[..., Any] | None = ArgumentParser.add_argument,
  1322. add_argument_group_method: Callable[..., Any] | None = ArgumentParser.add_argument_group,
  1323. add_parser_method: Callable[..., Any] | None = _SubParsersAction.add_parser,
  1324. add_subparsers_method: Callable[..., Any] | None = ArgumentParser.add_subparsers,
  1325. formatter_class: Any = RawDescriptionHelpFormatter,
  1326. ) -> None:
  1327. def _parse_known_args(*args: Any, **kwargs: Any) -> Namespace:
  1328. return ArgumentParser.parse_known_args(*args, **kwargs)[0]
  1329. self._root_parser = root_parser
  1330. if parse_args_method is None:
  1331. parse_args_method = _parse_known_args if self.cli_ignore_unknown_args else ArgumentParser.parse_args
  1332. self._parse_args = self._connect_parser_method(parse_args_method, 'parsed_args_method')
  1333. self._add_argument = self._connect_parser_method(add_argument_method, 'add_argument_method')
  1334. self._add_argument_group = self._connect_parser_method(add_argument_group_method, 'add_argument_group_method')
  1335. self._add_parser = self._connect_parser_method(add_parser_method, 'add_parser_method')
  1336. self._add_subparsers = self._connect_parser_method(add_subparsers_method, 'add_subparsers_method')
  1337. self._formatter_class = formatter_class
  1338. self._cli_dict_args: dict[str, type[Any] | None] = {}
  1339. self._cli_subcommands: defaultdict[str, dict[str, str]] = defaultdict(dict)
  1340. self._add_parser_args(
  1341. parser=self.root_parser,
  1342. model=self.settings_cls,
  1343. added_args=[],
  1344. arg_prefix=self.env_prefix,
  1345. subcommand_prefix=self.env_prefix,
  1346. group=None,
  1347. alias_prefixes=[],
  1348. model_default=PydanticUndefined,
  1349. )
  1350. def _add_parser_args(
  1351. self,
  1352. parser: Any,
  1353. model: type[BaseModel],
  1354. added_args: list[str],
  1355. arg_prefix: str,
  1356. subcommand_prefix: str,
  1357. group: Any,
  1358. alias_prefixes: list[str],
  1359. model_default: Any,
  1360. ) -> ArgumentParser:
  1361. subparsers: Any = None
  1362. alias_path_args: dict[str, str] = {}
  1363. for field_name, field_info in self._sort_arg_fields(model):
  1364. sub_models: list[type[BaseModel]] = self._get_sub_models(model, field_name, field_info)
  1365. alias_names, is_alias_path_only = self._get_alias_names(field_name, field_info, alias_path_args)
  1366. preferred_alias = alias_names[0]
  1367. if _CliSubCommand in field_info.metadata:
  1368. for model in sub_models:
  1369. subcommand_alias = model.__name__ if len(sub_models) > 1 else preferred_alias
  1370. subcommand_name = f'{arg_prefix}{subcommand_alias}'
  1371. subcommand_dest = f'{arg_prefix}{preferred_alias}'
  1372. self._cli_subcommands[f'{arg_prefix}:subcommand'][subcommand_name] = subcommand_dest
  1373. subcommand_help = None if len(sub_models) > 1 else field_info.description
  1374. if self.cli_use_class_docs_for_groups:
  1375. subcommand_help = None if model.__doc__ is None else dedent(model.__doc__)
  1376. subparsers = (
  1377. self._add_subparsers(
  1378. parser,
  1379. title='subcommands',
  1380. dest=f'{arg_prefix}:subcommand',
  1381. description=field_info.description if len(sub_models) > 1 else None,
  1382. )
  1383. if subparsers is None
  1384. else subparsers
  1385. )
  1386. if hasattr(subparsers, 'metavar'):
  1387. subparsers.metavar = (
  1388. f'{subparsers.metavar[:-1]},{subcommand_alias}}}'
  1389. if subparsers.metavar
  1390. else f'{{{subcommand_alias}}}'
  1391. )
  1392. self._add_parser_args(
  1393. parser=self._add_parser(
  1394. subparsers,
  1395. subcommand_alias,
  1396. help=subcommand_help,
  1397. formatter_class=self._formatter_class,
  1398. description=None if model.__doc__ is None else dedent(model.__doc__),
  1399. ),
  1400. model=model,
  1401. added_args=[],
  1402. arg_prefix=f'{arg_prefix}{preferred_alias}.',
  1403. subcommand_prefix=f'{subcommand_prefix}{preferred_alias}.',
  1404. group=None,
  1405. alias_prefixes=[],
  1406. model_default=PydanticUndefined,
  1407. )
  1408. else:
  1409. flag_prefix: str = self._cli_flag_prefix
  1410. is_append_action = _annotation_contains_types(
  1411. field_info.annotation, (list, set, dict, Sequence, Mapping), is_strip_annotated=True
  1412. )
  1413. is_parser_submodel = sub_models and not is_append_action
  1414. kwargs: dict[str, Any] = {}
  1415. kwargs['default'] = CLI_SUPPRESS
  1416. kwargs['help'] = self._help_format(field_name, field_info, model_default)
  1417. kwargs['metavar'] = self._metavar_format(field_info.annotation)
  1418. kwargs['required'] = (
  1419. self.cli_enforce_required and field_info.is_required() and model_default is PydanticUndefined
  1420. )
  1421. kwargs['dest'] = (
  1422. # Strip prefix if validation alias is set and value is not complex.
  1423. # Related https://github.com/pydantic/pydantic-settings/pull/25
  1424. f'{arg_prefix}{preferred_alias}'[self.env_prefix_len :]
  1425. if arg_prefix and field_info.validation_alias is not None and not is_parser_submodel
  1426. else f'{arg_prefix}{preferred_alias}'
  1427. )
  1428. if kwargs['dest'] in added_args:
  1429. continue
  1430. if is_append_action:
  1431. kwargs['action'] = 'append'
  1432. if _annotation_contains_types(field_info.annotation, (dict, Mapping), is_strip_annotated=True):
  1433. self._cli_dict_args[kwargs['dest']] = field_info.annotation
  1434. arg_names = self._get_arg_names(arg_prefix, subcommand_prefix, alias_prefixes, alias_names)
  1435. if _CliPositionalArg in field_info.metadata:
  1436. kwargs['metavar'] = preferred_alias.upper()
  1437. arg_names = [kwargs['dest']]
  1438. del kwargs['dest']
  1439. del kwargs['required']
  1440. flag_prefix = ''
  1441. self._convert_bool_flag(kwargs, field_info, model_default)
  1442. if is_parser_submodel:
  1443. self._add_parser_submodels(
  1444. parser,
  1445. sub_models,
  1446. added_args,
  1447. arg_prefix,
  1448. subcommand_prefix,
  1449. flag_prefix,
  1450. arg_names,
  1451. kwargs,
  1452. field_name,
  1453. field_info,
  1454. alias_names,
  1455. model_default=model_default,
  1456. )
  1457. elif not is_alias_path_only:
  1458. if group is not None:
  1459. if isinstance(group, dict):
  1460. group = self._add_argument_group(parser, **group)
  1461. added_args += list(arg_names)
  1462. self._add_argument(group, *(f'{flag_prefix[:len(name)]}{name}' for name in arg_names), **kwargs)
  1463. else:
  1464. added_args += list(arg_names)
  1465. self._add_argument(
  1466. parser, *(f'{flag_prefix[:len(name)]}{name}' for name in arg_names), **kwargs
  1467. )
  1468. self._add_parser_alias_paths(parser, alias_path_args, added_args, arg_prefix, subcommand_prefix, group)
  1469. return parser
  1470. def _convert_bool_flag(self, kwargs: dict[str, Any], field_info: FieldInfo, model_default: Any) -> None:
  1471. if kwargs['metavar'] == 'bool':
  1472. default = None
  1473. if field_info.default is not PydanticUndefined:
  1474. default = field_info.default
  1475. if model_default is not PydanticUndefined:
  1476. default = model_default
  1477. if sys.version_info >= (3, 9) or isinstance(default, bool):
  1478. if (self.cli_implicit_flags or _CliImplicitFlag in field_info.metadata) and (
  1479. _CliExplicitFlag not in field_info.metadata
  1480. ):
  1481. del kwargs['metavar']
  1482. kwargs['action'] = (
  1483. BooleanOptionalAction if sys.version_info >= (3, 9) else f'store_{str(not default).lower()}'
  1484. )
  1485. def _get_arg_names(
  1486. self, arg_prefix: str, subcommand_prefix: str, alias_prefixes: list[str], alias_names: tuple[str, ...]
  1487. ) -> list[str]:
  1488. arg_names: list[str] = []
  1489. for prefix in [arg_prefix] + alias_prefixes:
  1490. for name in alias_names:
  1491. arg_names.append(
  1492. f'{prefix}{name}'
  1493. if subcommand_prefix == self.env_prefix
  1494. else f'{prefix.replace(subcommand_prefix, "", 1)}{name}'
  1495. )
  1496. return arg_names
  1497. def _add_parser_submodels(
  1498. self,
  1499. parser: Any,
  1500. sub_models: list[type[BaseModel]],
  1501. added_args: list[str],
  1502. arg_prefix: str,
  1503. subcommand_prefix: str,
  1504. flag_prefix: str,
  1505. arg_names: list[str],
  1506. kwargs: dict[str, Any],
  1507. field_name: str,
  1508. field_info: FieldInfo,
  1509. alias_names: tuple[str, ...],
  1510. model_default: Any,
  1511. ) -> None:
  1512. model_group: Any = None
  1513. model_group_kwargs: dict[str, Any] = {}
  1514. model_group_kwargs['title'] = f'{arg_names[0]} options'
  1515. model_group_kwargs['description'] = field_info.description
  1516. if self.cli_use_class_docs_for_groups and len(sub_models) == 1:
  1517. model_group_kwargs['description'] = None if sub_models[0].__doc__ is None else dedent(sub_models[0].__doc__)
  1518. if model_default is not PydanticUndefined:
  1519. if is_model_class(type(model_default)) or is_pydantic_dataclass(type(model_default)):
  1520. model_default = getattr(model_default, field_name)
  1521. else:
  1522. if field_info.default is not PydanticUndefined:
  1523. model_default = field_info.default
  1524. elif field_info.default_factory is not None:
  1525. model_default = field_info.default_factory
  1526. if model_default is None:
  1527. desc_header = f'default: {self.cli_parse_none_str} (undefined)'
  1528. if model_group_kwargs['description'] is not None:
  1529. model_group_kwargs['description'] = dedent(f'{desc_header}\n{model_group_kwargs["description"]}')
  1530. else:
  1531. model_group_kwargs['description'] = desc_header
  1532. preferred_alias = alias_names[0]
  1533. if not self.cli_avoid_json:
  1534. added_args.append(arg_names[0])
  1535. kwargs['help'] = f'set {arg_names[0]} from JSON string'
  1536. model_group = self._add_argument_group(parser, **model_group_kwargs)
  1537. self._add_argument(model_group, *(f'{flag_prefix}{name}' for name in arg_names), **kwargs)
  1538. for model in sub_models:
  1539. self._add_parser_args(
  1540. parser=parser,
  1541. model=model,
  1542. added_args=added_args,
  1543. arg_prefix=f'{arg_prefix}{preferred_alias}.',
  1544. subcommand_prefix=subcommand_prefix,
  1545. group=model_group if model_group else model_group_kwargs,
  1546. alias_prefixes=[f'{arg_prefix}{name}.' for name in alias_names[1:]],
  1547. model_default=model_default,
  1548. )
  1549. def _add_parser_alias_paths(
  1550. self,
  1551. parser: Any,
  1552. alias_path_args: dict[str, str],
  1553. added_args: list[str],
  1554. arg_prefix: str,
  1555. subcommand_prefix: str,
  1556. group: Any,
  1557. ) -> None:
  1558. if alias_path_args:
  1559. context = parser
  1560. if group is not None:
  1561. context = self._add_argument_group(parser, **group) if isinstance(group, dict) else group
  1562. is_nested_alias_path = arg_prefix.endswith('.')
  1563. arg_prefix = arg_prefix[:-1] if is_nested_alias_path else arg_prefix
  1564. for name, metavar in alias_path_args.items():
  1565. name = '' if is_nested_alias_path else name
  1566. arg_name = (
  1567. f'{arg_prefix}{name}'
  1568. if subcommand_prefix == self.env_prefix
  1569. else f'{arg_prefix.replace(subcommand_prefix, "", 1)}{name}'
  1570. )
  1571. kwargs: dict[str, Any] = {}
  1572. kwargs['default'] = CLI_SUPPRESS
  1573. kwargs['help'] = 'pydantic alias path'
  1574. kwargs['dest'] = f'{arg_prefix}{name}'
  1575. if metavar == 'dict' or is_nested_alias_path:
  1576. kwargs['metavar'] = 'dict'
  1577. else:
  1578. kwargs['action'] = 'append'
  1579. kwargs['metavar'] = 'list'
  1580. if arg_name not in added_args:
  1581. added_args.append(arg_name)
  1582. self._add_argument(context, f'{self._cli_flag_prefix}{arg_name}', **kwargs)
  1583. def _get_modified_args(self, obj: Any) -> tuple[str, ...]:
  1584. if not self.cli_hide_none_type:
  1585. return get_args(obj)
  1586. else:
  1587. return tuple([type_ for type_ in get_args(obj) if type_ is not type(None)])
  1588. def _metavar_format_choices(self, args: list[str], obj_qualname: str | None = None) -> str:
  1589. if 'JSON' in args:
  1590. args = args[: args.index('JSON') + 1] + [arg for arg in args[args.index('JSON') + 1 :] if arg != 'JSON']
  1591. metavar = ','.join(args)
  1592. if obj_qualname:
  1593. return f'{obj_qualname}[{metavar}]'
  1594. else:
  1595. return metavar if len(args) == 1 else f'{{{metavar}}}'
  1596. def _metavar_format_recurse(self, obj: Any) -> str:
  1597. """Pretty metavar representation of a type. Adapts logic from `pydantic._repr.display_as_type`."""
  1598. obj = _strip_annotated(obj)
  1599. if _is_function(obj):
  1600. # If function is locally defined use __name__ instead of __qualname__
  1601. return obj.__name__ if '<locals>' in obj.__qualname__ else obj.__qualname__
  1602. elif obj is ...:
  1603. return '...'
  1604. elif isinstance(obj, Representation):
  1605. return repr(obj)
  1606. elif isinstance(obj, typing_extensions.TypeAliasType):
  1607. return str(obj)
  1608. if not isinstance(obj, (typing_base, WithArgsTypes, type)):
  1609. obj = obj.__class__
  1610. if origin_is_union(get_origin(obj)):
  1611. return self._metavar_format_choices(list(map(self._metavar_format_recurse, self._get_modified_args(obj))))
  1612. elif get_origin(obj) in (typing_extensions.Literal, typing.Literal):
  1613. return self._metavar_format_choices(list(map(str, self._get_modified_args(obj))))
  1614. elif lenient_issubclass(obj, Enum):
  1615. return self._metavar_format_choices([val.name for val in obj])
  1616. elif isinstance(obj, WithArgsTypes):
  1617. return self._metavar_format_choices(
  1618. list(map(self._metavar_format_recurse, self._get_modified_args(obj))), obj_qualname=obj.__qualname__
  1619. )
  1620. elif obj is type(None):
  1621. return self.cli_parse_none_str
  1622. elif is_model_class(obj):
  1623. return 'JSON'
  1624. elif isinstance(obj, type):
  1625. return obj.__qualname__
  1626. else:
  1627. return repr(obj).replace('typing.', '').replace('typing_extensions.', '')
  1628. def _metavar_format(self, obj: Any) -> str:
  1629. return self._metavar_format_recurse(obj).replace(', ', ',')
  1630. def _help_format(self, field_name: str, field_info: FieldInfo, model_default: Any) -> str:
  1631. _help = field_info.description if field_info.description else ''
  1632. if _help == CLI_SUPPRESS or CLI_SUPPRESS in field_info.metadata:
  1633. return CLI_SUPPRESS
  1634. if field_info.is_required() and model_default in (PydanticUndefined, None):
  1635. if _CliPositionalArg not in field_info.metadata:
  1636. ifdef = 'ifdef: ' if model_default is None else ''
  1637. _help += f' ({ifdef}required)' if _help else f'({ifdef}required)'
  1638. else:
  1639. default = f'(default: {self.cli_parse_none_str})'
  1640. if is_model_class(type(model_default)) or is_pydantic_dataclass(type(model_default)):
  1641. default = f'(default: {getattr(model_default, field_name)})'
  1642. elif model_default not in (PydanticUndefined, None) and _is_function(model_default):
  1643. default = f'(default factory: {self._metavar_format(model_default)})'
  1644. elif field_info.default not in (PydanticUndefined, None):
  1645. enum_name = _annotation_enum_val_to_name(field_info.annotation, field_info.default)
  1646. default = f'(default: {field_info.default if enum_name is None else enum_name})'
  1647. elif field_info.default_factory is not None:
  1648. default = f'(default factory: {self._metavar_format(field_info.default_factory)})'
  1649. _help += f' {default}' if _help else default
  1650. return _help.replace('%', '%%') if issubclass(type(self._root_parser), ArgumentParser) else _help
  1651. class ConfigFileSourceMixin(ABC):
  1652. def _read_files(self, files: PathType | None) -> dict[str, Any]:
  1653. if files is None:
  1654. return {}
  1655. if isinstance(files, (str, os.PathLike)):
  1656. files = [files]
  1657. vars: dict[str, Any] = {}
  1658. for file in files:
  1659. file_path = Path(file).expanduser()
  1660. if file_path.is_file():
  1661. vars.update(self._read_file(file_path))
  1662. return vars
  1663. @abstractmethod
  1664. def _read_file(self, path: Path) -> dict[str, Any]:
  1665. pass
  1666. class JsonConfigSettingsSource(InitSettingsSource, ConfigFileSourceMixin):
  1667. """
  1668. A source class that loads variables from a JSON file
  1669. """
  1670. def __init__(
  1671. self,
  1672. settings_cls: type[BaseSettings],
  1673. json_file: PathType | None = DEFAULT_PATH,
  1674. json_file_encoding: str | None = None,
  1675. ):
  1676. self.json_file_path = json_file if json_file != DEFAULT_PATH else settings_cls.model_config.get('json_file')
  1677. self.json_file_encoding = (
  1678. json_file_encoding
  1679. if json_file_encoding is not None
  1680. else settings_cls.model_config.get('json_file_encoding')
  1681. )
  1682. self.json_data = self._read_files(self.json_file_path)
  1683. super().__init__(settings_cls, self.json_data)
  1684. def _read_file(self, file_path: Path) -> dict[str, Any]:
  1685. with open(file_path, encoding=self.json_file_encoding) as json_file:
  1686. return json.load(json_file)
  1687. class TomlConfigSettingsSource(InitSettingsSource, ConfigFileSourceMixin):
  1688. """
  1689. A source class that loads variables from a TOML file
  1690. """
  1691. def __init__(
  1692. self,
  1693. settings_cls: type[BaseSettings],
  1694. toml_file: PathType | None = DEFAULT_PATH,
  1695. ):
  1696. self.toml_file_path = toml_file if toml_file != DEFAULT_PATH else settings_cls.model_config.get('toml_file')
  1697. self.toml_data = self._read_files(self.toml_file_path)
  1698. super().__init__(settings_cls, self.toml_data)
  1699. def _read_file(self, file_path: Path) -> dict[str, Any]:
  1700. import_toml()
  1701. with open(file_path, mode='rb') as toml_file:
  1702. if sys.version_info < (3, 11):
  1703. return tomli.load(toml_file)
  1704. return tomllib.load(toml_file)
  1705. class PyprojectTomlConfigSettingsSource(TomlConfigSettingsSource):
  1706. """
  1707. A source class that loads variables from a `pyproject.toml` file.
  1708. """
  1709. def __init__(
  1710. self,
  1711. settings_cls: type[BaseSettings],
  1712. toml_file: Path | None = None,
  1713. ) -> None:
  1714. self.toml_file_path = self._pick_pyproject_toml_file(
  1715. toml_file, settings_cls.model_config.get('pyproject_toml_depth', 0)
  1716. )
  1717. self.toml_table_header: tuple[str, ...] = settings_cls.model_config.get(
  1718. 'pyproject_toml_table_header', ('tool', 'pydantic-settings')
  1719. )
  1720. self.toml_data = self._read_files(self.toml_file_path)
  1721. for key in self.toml_table_header:
  1722. self.toml_data = self.toml_data.get(key, {})
  1723. super(TomlConfigSettingsSource, self).__init__(settings_cls, self.toml_data)
  1724. @staticmethod
  1725. def _pick_pyproject_toml_file(provided: Path | None, depth: int) -> Path:
  1726. """Pick a `pyproject.toml` file path to use.
  1727. Args:
  1728. provided: Explicit path provided when instantiating this class.
  1729. depth: Number of directories up the tree to check of a pyproject.toml.
  1730. """
  1731. if provided:
  1732. return provided.resolve()
  1733. rv = Path.cwd() / 'pyproject.toml'
  1734. count = 0
  1735. if not rv.is_file():
  1736. child = rv.parent.parent / 'pyproject.toml'
  1737. while count < depth:
  1738. if child.is_file():
  1739. return child
  1740. if str(child.parent) == rv.root:
  1741. break # end discovery after checking system root once
  1742. child = child.parent.parent / 'pyproject.toml'
  1743. count += 1
  1744. return rv
  1745. class YamlConfigSettingsSource(InitSettingsSource, ConfigFileSourceMixin):
  1746. """
  1747. A source class that loads variables from a yaml file
  1748. """
  1749. def __init__(
  1750. self,
  1751. settings_cls: type[BaseSettings],
  1752. yaml_file: PathType | None = DEFAULT_PATH,
  1753. yaml_file_encoding: str | None = None,
  1754. ):
  1755. self.yaml_file_path = yaml_file if yaml_file != DEFAULT_PATH else settings_cls.model_config.get('yaml_file')
  1756. self.yaml_file_encoding = (
  1757. yaml_file_encoding
  1758. if yaml_file_encoding is not None
  1759. else settings_cls.model_config.get('yaml_file_encoding')
  1760. )
  1761. self.yaml_data = self._read_files(self.yaml_file_path)
  1762. super().__init__(settings_cls, self.yaml_data)
  1763. def _read_file(self, file_path: Path) -> dict[str, Any]:
  1764. import_yaml()
  1765. with open(file_path, encoding=self.yaml_file_encoding) as yaml_file:
  1766. return yaml.safe_load(yaml_file) or {}
  1767. class AzureKeyVaultMapping(Mapping[str, Optional[str]]):
  1768. _loaded_secrets: dict[str, str | None]
  1769. _secret_client: SecretClient # type: ignore
  1770. _secret_names: list[str]
  1771. def __init__(
  1772. self,
  1773. secret_client: SecretClient, # type: ignore
  1774. ) -> None:
  1775. self._loaded_secrets = {}
  1776. self._secret_client = secret_client
  1777. self._secret_names: list[str] = [secret.name for secret in self._secret_client.list_properties_of_secrets()]
  1778. def __getitem__(self, key: str) -> str | None:
  1779. if key not in self._loaded_secrets:
  1780. try:
  1781. self._loaded_secrets[key] = self._secret_client.get_secret(key).value
  1782. except Exception:
  1783. raise KeyError(key)
  1784. return self._loaded_secrets[key]
  1785. def __len__(self) -> int:
  1786. return len(self._secret_names)
  1787. def __iter__(self) -> Iterator[str]:
  1788. return iter(self._secret_names)
  1789. class AzureKeyVaultSettingsSource(EnvSettingsSource):
  1790. _url: str
  1791. _credential: TokenCredential # type: ignore
  1792. _secret_client: SecretClient # type: ignore
  1793. def __init__(
  1794. self,
  1795. settings_cls: type[BaseSettings],
  1796. url: str,
  1797. credential: TokenCredential, # type: ignore
  1798. env_prefix: str | None = None,
  1799. env_parse_none_str: str | None = None,
  1800. env_parse_enums: bool | None = None,
  1801. ) -> None:
  1802. import_azure_key_vault()
  1803. self._url = url
  1804. self._credential = credential
  1805. super().__init__(
  1806. settings_cls,
  1807. case_sensitive=True,
  1808. env_prefix=env_prefix,
  1809. env_nested_delimiter='--',
  1810. env_ignore_empty=False,
  1811. env_parse_none_str=env_parse_none_str,
  1812. env_parse_enums=env_parse_enums,
  1813. )
  1814. def _load_env_vars(self) -> Mapping[str, Optional[str]]:
  1815. secret_client = SecretClient(vault_url=self._url, credential=self._credential) # type: ignore
  1816. return AzureKeyVaultMapping(secret_client)
  1817. def __repr__(self) -> str:
  1818. return f'AzureKeyVaultSettingsSource(url={self._url!r}, ' f'env_nested_delimiter={self.env_nested_delimiter!r})'
  1819. def _get_env_var_key(key: str, case_sensitive: bool = False) -> str:
  1820. return key if case_sensitive else key.lower()
  1821. def _parse_env_none_str(value: str | None, parse_none_str: str | None = None) -> str | None | EnvNoneType:
  1822. return value if not (value == parse_none_str and parse_none_str is not None) else EnvNoneType(value)
  1823. def parse_env_vars(
  1824. env_vars: Mapping[str, str | None],
  1825. case_sensitive: bool = False,
  1826. ignore_empty: bool = False,
  1827. parse_none_str: str | None = None,
  1828. ) -> Mapping[str, str | None]:
  1829. return {
  1830. _get_env_var_key(k, case_sensitive): _parse_env_none_str(v, parse_none_str)
  1831. for k, v in env_vars.items()
  1832. if not (ignore_empty and v == '')
  1833. }
  1834. def read_env_file(
  1835. file_path: Path,
  1836. *,
  1837. encoding: str | None = None,
  1838. case_sensitive: bool = False,
  1839. ignore_empty: bool = False,
  1840. parse_none_str: str | None = None,
  1841. ) -> Mapping[str, str | None]:
  1842. warnings.warn(
  1843. 'read_env_file will be removed in the next version, use DotEnvSettingsSource._static_read_env_file if you must',
  1844. DeprecationWarning,
  1845. )
  1846. return DotEnvSettingsSource._static_read_env_file(
  1847. file_path,
  1848. encoding=encoding,
  1849. case_sensitive=case_sensitive,
  1850. ignore_empty=ignore_empty,
  1851. parse_none_str=parse_none_str,
  1852. )
  1853. def _annotation_is_complex(annotation: type[Any] | None, metadata: list[Any]) -> bool:
  1854. # If the model is a root model, the root annotation should be used to
  1855. # evaluate the complexity.
  1856. try:
  1857. if annotation is not None and issubclass(annotation, RootModel):
  1858. # In some rare cases (see test_root_model_as_field),
  1859. # the root attribute is not available. For these cases, python 3.8 and 3.9
  1860. # return 'RootModelRootType'.
  1861. root_annotation = annotation.__annotations__.get('root', None)
  1862. if root_annotation is not None and root_annotation != 'RootModelRootType':
  1863. annotation = root_annotation
  1864. except TypeError:
  1865. pass
  1866. if any(isinstance(md, Json) for md in metadata): # type: ignore[misc]
  1867. return False
  1868. # Check if annotation is of the form Annotated[type, metadata].
  1869. if isinstance(annotation, _AnnotatedAlias):
  1870. # Return result of recursive call on inner type.
  1871. inner, *meta = get_args(annotation)
  1872. return _annotation_is_complex(inner, meta)
  1873. origin = get_origin(annotation)
  1874. return (
  1875. _annotation_is_complex_inner(annotation)
  1876. or _annotation_is_complex_inner(origin)
  1877. or hasattr(origin, '__pydantic_core_schema__')
  1878. or hasattr(origin, '__get_pydantic_core_schema__')
  1879. )
  1880. def _annotation_is_complex_inner(annotation: type[Any] | None) -> bool:
  1881. if lenient_issubclass(annotation, (str, bytes)):
  1882. return False
  1883. return lenient_issubclass(annotation, (BaseModel, Mapping, Sequence, tuple, set, frozenset, deque)) or is_dataclass(
  1884. annotation
  1885. )
  1886. def _union_is_complex(annotation: type[Any] | None, metadata: list[Any]) -> bool:
  1887. return any(_annotation_is_complex(arg, metadata) for arg in get_args(annotation))
  1888. def _annotation_contains_types(
  1889. annotation: type[Any] | None,
  1890. types: tuple[Any, ...],
  1891. is_include_origin: bool = True,
  1892. is_strip_annotated: bool = False,
  1893. ) -> bool:
  1894. if is_strip_annotated:
  1895. annotation = _strip_annotated(annotation)
  1896. if is_include_origin is True and get_origin(annotation) in types:
  1897. return True
  1898. for type_ in get_args(annotation):
  1899. if _annotation_contains_types(type_, types, is_include_origin=True, is_strip_annotated=is_strip_annotated):
  1900. return True
  1901. return annotation in types
  1902. def _strip_annotated(annotation: Any) -> Any:
  1903. while get_origin(annotation) == Annotated:
  1904. annotation = get_args(annotation)[0]
  1905. return annotation
  1906. def _annotation_enum_val_to_name(annotation: type[Any] | None, value: Any) -> Optional[str]:
  1907. for type_ in (annotation, get_origin(annotation), *get_args(annotation)):
  1908. if lenient_issubclass(type_, Enum):
  1909. if value in tuple(val.value for val in type_):
  1910. return type_(value).name
  1911. return None
  1912. def _annotation_enum_name_to_val(annotation: type[Any] | None, name: Any) -> Any:
  1913. for type_ in (annotation, get_origin(annotation), *get_args(annotation)):
  1914. if lenient_issubclass(type_, Enum):
  1915. if name in tuple(val.name for val in type_):
  1916. return type_[name]
  1917. return None
  1918. def _get_model_fields(model_cls: type[Any]) -> dict[str, FieldInfo]:
  1919. if is_pydantic_dataclass(model_cls) and hasattr(model_cls, '__pydantic_fields__'):
  1920. return model_cls.__pydantic_fields__
  1921. if is_model_class(model_cls):
  1922. return model_cls.model_fields
  1923. raise SettingsError(f'Error: {model_cls.__name__} is not subclass of BaseModel or pydantic.dataclasses.dataclass')
  1924. def _is_function(obj: Any) -> bool:
  1925. return isinstance(obj, (FunctionType, BuiltinFunctionType))