database_urls.py 3.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121
  1. """Fail-closed validation for production PostgreSQL identities and templates."""
  2. from __future__ import annotations
  3. import argparse
  4. import re
  5. import sys
  6. from dataclasses import dataclass
  7. from pathlib import Path
  8. from urllib.parse import unquote, urlsplit
  9. _BAD_PERCENT_ESCAPE = re.compile(r"%(?![0-9a-fA-F]{2})")
  10. _PLACEHOLDER_VALUES = {
  11. "change-me",
  12. "changeme",
  13. "database",
  14. "database_name",
  15. "dbname",
  16. "host",
  17. "hostname",
  18. "pass",
  19. "password",
  20. "replace",
  21. "user",
  22. "username",
  23. }
  24. @dataclass(frozen=True)
  25. class DatabaseUrlParts:
  26. username: str
  27. password: str
  28. host: str
  29. database: str
  30. def is_template_secret(value: str) -> bool:
  31. """Return true only for empty or explicit template values, after decoding."""
  32. normalized = unquote(str(value or "")).strip().lower()
  33. return (
  34. not normalized
  35. or normalized in _PLACEHOLDER_VALUES
  36. or normalized.startswith(("replace-", "your-"))
  37. or "${" in normalized
  38. or normalized.startswith("<")
  39. or normalized.endswith(">")
  40. )
  41. def validate_postgresql_url(value: str, field_name: str) -> DatabaseUrlParts:
  42. """Parse a PostgreSQL URL and reject empty/template identity components."""
  43. raw = str(value or "").strip()
  44. if not raw or _BAD_PERCENT_ESCAPE.search(raw):
  45. raise RuntimeError(f"{field_name} is missing or malformed")
  46. try:
  47. parsed = urlsplit(raw)
  48. port = parsed.port
  49. except (TypeError, ValueError) as exc:
  50. raise RuntimeError(f"{field_name} is malformed") from exc
  51. if parsed.scheme.split("+", 1)[0].lower() not in {"postgres", "postgresql"}:
  52. raise RuntimeError(f"{field_name} must use PostgreSQL")
  53. values = {
  54. "username": unquote(parsed.username or ""),
  55. "password": unquote(parsed.password or ""),
  56. "host": unquote(parsed.hostname or ""),
  57. "database": unquote(parsed.path.lstrip("/")),
  58. }
  59. if port is not None and not 1 <= port <= 65535:
  60. raise RuntimeError(f"{field_name} has an invalid port")
  61. for component, component_value in values.items():
  62. if is_template_secret(component_value):
  63. raise RuntimeError(f"{field_name} has an invalid {component}")
  64. return DatabaseUrlParts(**values)
  65. def _read_env_file(path: str | Path) -> dict[str, str]:
  66. values: dict[str, str] = {}
  67. for raw_line in Path(path).read_text(encoding="utf-8-sig").splitlines():
  68. line = raw_line.strip()
  69. if not line or line.startswith("#") or "=" not in line:
  70. continue
  71. name, value = line.split("=", 1)
  72. values[name.strip()] = value.strip().strip('"').strip("'")
  73. return values
  74. def validate_database_environment(values: dict[str, str]) -> None:
  75. runtime_parts = None
  76. for field_name in (
  77. "DB_ROLE_INIT_DATABASE_URL",
  78. "MIGRATION_DATABASE_URL",
  79. "DATABASE_URL",
  80. ):
  81. parts = validate_postgresql_url(values.get(field_name, ""), field_name)
  82. if field_name == "DATABASE_URL":
  83. runtime_parts = parts
  84. runtime_password = values.get("DATAOPS_RUNTIME_PASSWORD", "")
  85. if is_template_secret(runtime_password):
  86. raise RuntimeError("DATAOPS_RUNTIME_PASSWORD is missing or a template value")
  87. if runtime_parts is None or runtime_parts.password != runtime_password:
  88. raise RuntimeError("DATAOPS_RUNTIME_PASSWORD must match DATABASE_URL")
  89. def load_and_validate_database_env(path: str | Path) -> None:
  90. validate_database_environment(_read_env_file(path))
  91. def main(argv: list[str] | None = None) -> int:
  92. parser = argparse.ArgumentParser()
  93. parser.add_argument("--env-file", required=True)
  94. args = parser.parse_args(argv)
  95. try:
  96. load_and_validate_database_env(args.env_file)
  97. except (OSError, RuntimeError) as exc:
  98. print(f"database configuration rejected: {exc}", file=sys.stderr)
  99. return 2
  100. return 0
  101. if __name__ == "__main__":
  102. raise SystemExit(main())