config.py 4.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138
  1. from __future__ import annotations
  2. import os
  3. import typing
  4. import warnings
  5. from pathlib import Path
  6. class undefined:
  7. pass
  8. class EnvironError(Exception):
  9. pass
  10. class Environ(typing.MutableMapping[str, str]):
  11. def __init__(self, environ: typing.MutableMapping[str, str] = os.environ):
  12. self._environ = environ
  13. self._has_been_read: set[str] = set()
  14. def __getitem__(self, key: str) -> str:
  15. self._has_been_read.add(key)
  16. return self._environ.__getitem__(key)
  17. def __setitem__(self, key: str, value: str) -> None:
  18. if key in self._has_been_read:
  19. raise EnvironError(f"Attempting to set environ['{key}'], but the value has already been read.")
  20. self._environ.__setitem__(key, value)
  21. def __delitem__(self, key: str) -> None:
  22. if key in self._has_been_read:
  23. raise EnvironError(f"Attempting to delete environ['{key}'], but the value has already been read.")
  24. self._environ.__delitem__(key)
  25. def __iter__(self) -> typing.Iterator[str]:
  26. return iter(self._environ)
  27. def __len__(self) -> int:
  28. return len(self._environ)
  29. environ = Environ()
  30. T = typing.TypeVar("T")
  31. class Config:
  32. def __init__(
  33. self,
  34. env_file: str | Path | None = None,
  35. environ: typing.Mapping[str, str] = environ,
  36. env_prefix: str = "",
  37. ) -> None:
  38. self.environ = environ
  39. self.env_prefix = env_prefix
  40. self.file_values: dict[str, str] = {}
  41. if env_file is not None:
  42. if not os.path.isfile(env_file):
  43. warnings.warn(f"Config file '{env_file}' not found.")
  44. else:
  45. self.file_values = self._read_file(env_file)
  46. @typing.overload
  47. def __call__(self, key: str, *, default: None) -> str | None: ...
  48. @typing.overload
  49. def __call__(self, key: str, cast: type[T], default: T = ...) -> T: ...
  50. @typing.overload
  51. def __call__(self, key: str, cast: type[str] = ..., default: str = ...) -> str: ...
  52. @typing.overload
  53. def __call__(
  54. self,
  55. key: str,
  56. cast: typing.Callable[[typing.Any], T] = ...,
  57. default: typing.Any = ...,
  58. ) -> T: ...
  59. @typing.overload
  60. def __call__(self, key: str, cast: type[str] = ..., default: T = ...) -> T | str: ...
  61. def __call__(
  62. self,
  63. key: str,
  64. cast: typing.Callable[[typing.Any], typing.Any] | None = None,
  65. default: typing.Any = undefined,
  66. ) -> typing.Any:
  67. return self.get(key, cast, default)
  68. def get(
  69. self,
  70. key: str,
  71. cast: typing.Callable[[typing.Any], typing.Any] | None = None,
  72. default: typing.Any = undefined,
  73. ) -> typing.Any:
  74. key = self.env_prefix + key
  75. if key in self.environ:
  76. value = self.environ[key]
  77. return self._perform_cast(key, value, cast)
  78. if key in self.file_values:
  79. value = self.file_values[key]
  80. return self._perform_cast(key, value, cast)
  81. if default is not undefined:
  82. return self._perform_cast(key, default, cast)
  83. raise KeyError(f"Config '{key}' is missing, and has no default.")
  84. def _read_file(self, file_name: str | Path) -> dict[str, str]:
  85. file_values: dict[str, str] = {}
  86. with open(file_name) as input_file:
  87. for line in input_file.readlines():
  88. line = line.strip()
  89. if "=" in line and not line.startswith("#"):
  90. key, value = line.split("=", 1)
  91. key = key.strip()
  92. value = value.strip().strip("\"'")
  93. file_values[key] = value
  94. return file_values
  95. def _perform_cast(
  96. self,
  97. key: str,
  98. value: typing.Any,
  99. cast: typing.Callable[[typing.Any], typing.Any] | None = None,
  100. ) -> typing.Any:
  101. if cast is None or value is None:
  102. return value
  103. elif cast is bool and isinstance(value, str):
  104. mapping = {"true": True, "1": True, "false": False, "0": False}
  105. value = value.lower()
  106. if value not in mapping:
  107. raise ValueError(f"Config '{key}' has value '{value}'. Not a valid bool.")
  108. return mapping[value]
  109. try:
  110. return cast(value)
  111. except (TypeError, ValueError):
  112. raise ValueError(f"Config '{key}' has value '{value}'. Not a valid {cast.__name__}.")