Files
call-center/scripts/postgres_dev_preflight.py

317 lines
12 KiB
Python

from __future__ import annotations
import argparse
from dataclasses import asdict, dataclass
import json
import os
from pathlib import Path
import socket
import sys
from urllib.parse import urlparse
import httpx
from sqlalchemy import create_engine, text
ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
from services.shared.schema_migrations import applied_migration_versions, list_migration_files, validate_schema_migrations_applied
from services.shared.security import issue_app_token
REQUIRED_PROXY_HEALTH_SERVICES = [
"auth",
"customer",
"interaction",
"routing",
"voice",
"telegram",
"whatsapp",
"recording",
"ivr",
"ai",
"ai-voice-runtime",
"kb",
"reporting",
"supervisor",
]
@dataclass
class CheckResult:
name: str
ok: bool
details: str
def _load_env_file(path: Path) -> None:
if not path.exists() or not path.is_file():
return
for line in path.read_text(encoding="utf-8").splitlines():
stripped = line.strip()
if not stripped or stripped.startswith("#") or "=" not in stripped:
continue
key, value = stripped.split("=", 1)
key = key.strip()
if key and key not in os.environ:
os.environ[key] = value.strip()
def _bool_env(name: str, default: bool) -> bool:
raw = os.getenv(name)
if raw is None:
return default
return raw.strip().lower() in {"1", "true", "yes", "on"}
def _normalize_database_url(url: str) -> str:
if url.startswith("postgres://"):
return "postgresql+psycopg://" + url[len("postgres://") :]
if url.startswith("postgresql://") and "+psycopg" not in url:
return "postgresql+psycopg://" + url[len("postgresql://") :]
return url
def _ops_headers() -> dict[str, str]:
if _bool_env("ALLOW_LEGACY_HEADER_AUTH", True):
return {"X-User": "admin", "X-Role": "admin"}
token = issue_app_token(
subject="ops:postgres-dev-preflight",
username="postgres-dev-preflight",
role="admin",
auth_source="service",
provider="postgres-preflight",
ttl_seconds=300,
)
return {"Authorization": f"Bearer {token}"}
def _check_tcp_endpoint(name: str, url: str, default_port: int) -> CheckResult:
parsed = urlparse(url)
host = parsed.hostname
port = parsed.port or default_port
if not host:
return CheckResult(name, False, "host is missing")
try:
with socket.create_connection((host, port), timeout=3.0):
pass
except OSError as exc:
return CheckResult(name, False, f"{host}:{port} unreachable: {exc}")
return CheckResult(name, True, f"{host}:{port} reachable")
def _check_database_connection(database_url: str) -> CheckResult:
engine = create_engine(
_normalize_database_url(database_url),
future=True,
pool_pre_ping=True,
connect_args={"connect_timeout": 5},
)
try:
with engine.connect() as conn:
conn.execute(text("SELECT 1"))
except Exception as exc: # noqa: BLE001
return CheckResult("database_connection", False, str(exc))
finally:
engine.dispose()
return CheckResult("database_connection", True, "SELECT 1 succeeded")
def _check_schema_migrations(database_url: str) -> CheckResult:
engine = create_engine(
_normalize_database_url(database_url),
future=True,
pool_pre_ping=True,
connect_args={"connect_timeout": 5},
)
try:
validate_schema_migrations_applied(db_engine=engine)
expected = list_migration_files(db_engine=engine)
applied = applied_migration_versions(db_engine=engine)
except Exception as exc: # noqa: BLE001
return CheckResult("schema_migrations", False, str(exc))
finally:
engine.dispose()
return CheckResult("schema_migrations", True, f"applied={len(applied)}/{len(expected)}")
def _http_preflight_checks(base_url: str) -> list[CheckResult]:
results: list[CheckResult] = []
normalized = base_url.rstrip("/")
ops_headers = _ops_headers()
read_headers = dict(ops_headers)
with httpx.Client(base_url=normalized, timeout=10, trust_env=False) as client:
try:
response = client.get("/health")
ok = response.status_code == 200 and response.json().get("status") == "ok"
results.append(CheckResult("gateway_health", ok, f"status={response.status_code}"))
except Exception as exc: # noqa: BLE001
return [CheckResult("gateway_health", False, str(exc))]
try:
response = client.get("/registry")
payload = response.json() if response.status_code == 200 else {}
services = payload.get("services", {}) if isinstance(payload, dict) else {}
missing = [name for name in REQUIRED_PROXY_HEALTH_SERVICES if name not in services]
ok = response.status_code == 200 and not missing
details = f"status={response.status_code}"
if missing:
details += f", missing={missing}"
results.append(CheckResult("gateway_registry", ok, details))
except Exception as exc: # noqa: BLE001
results.append(CheckResult("gateway_registry", False, str(exc)))
return results
services_to_check = list(REQUIRED_PROXY_HEALTH_SERVICES)
if _bool_env("EVENT_BUS_ENABLED", False):
services_to_check.append("event-bus")
if _bool_env("ASTERISK_BRIDGE_ENABLED", False):
services_to_check.append("asterisk-bridge")
for service in services_to_check:
try:
response = client.get(f"/proxy/{service}/health", headers=ops_headers)
payload = response.json() if response.status_code == 200 else {}
status_value = str(payload.get("status") or "").strip().lower() if isinstance(payload, dict) else ""
ok = response.status_code == 200 and status_value == "ok"
results.append(CheckResult(f"{service}_health", ok, f"status={response.status_code}"))
except Exception as exc: # noqa: BLE001
results.append(CheckResult(f"{service}_health", False, str(exc)))
try:
response = client.post(
"/proxy/auth/auth/login",
json={"username": "admin", "password": "admin123"},
)
payload = response.json() if response.status_code == 200 else {}
access_token = str(payload.get("access_token") or "").strip() if isinstance(payload, dict) else ""
if access_token:
read_headers = {"Authorization": f"Bearer {access_token}"}
ok = response.status_code == 200 and bool(access_token)
results.append(CheckResult("auth_login", ok, f"status={response.status_code}"))
except Exception as exc: # noqa: BLE001
results.append(CheckResult("auth_login", False, str(exc)))
try:
response = client.get("/proxy/routing/queues", headers=read_headers)
ok = response.status_code == 200 and isinstance(response.json(), list)
results.append(CheckResult("routing_read", ok, f"status={response.status_code}"))
except Exception as exc: # noqa: BLE001
results.append(CheckResult("routing_read", False, str(exc)))
try:
response = client.get("/proxy/voice/integrations/voice/events?limit=1", headers=read_headers)
ok = response.status_code == 200 and isinstance(response.json(), list)
results.append(CheckResult("voice_read", ok, f"status={response.status_code}"))
except Exception as exc: # noqa: BLE001
results.append(CheckResult("voice_read", False, str(exc)))
try:
response = client.get("/proxy/whatsapp/integrations/whatsapp/threads", headers=read_headers)
ok = response.status_code == 200 and isinstance(response.json(), list)
results.append(CheckResult("whatsapp_read", ok, f"status={response.status_code}"))
except Exception as exc: # noqa: BLE001
results.append(CheckResult("whatsapp_read", False, str(exc)))
return results
def run_preflight(
*,
database_url: str,
base_url: str | None,
require_rabbitmq: bool,
) -> list[CheckResult]:
results: list[CheckResult] = []
raw_database_url = str(database_url or "").strip()
if not raw_database_url:
return [CheckResult("database_url", False, "DATABASE_URL is required")]
if not raw_database_url.startswith(("postgres://", "postgresql://", "postgresql+psycopg://")):
return [CheckResult("database_url", False, "DATABASE_URL must point to PostgreSQL")]
results.append(CheckResult("database_url", True, "PostgreSQL URL detected"))
mode = str(os.getenv("SCHEMA_MANAGEMENT_MODE", "")).strip().lower()
if mode != "migrations":
results.append(CheckResult("schema_management_mode", False, "SCHEMA_MANAGEMENT_MODE must be set to migrations"))
return results
results.append(CheckResult("schema_management_mode", True, "migrations"))
results.append(_check_tcp_endpoint("postgres_tcp", raw_database_url, 5432))
if not results[-1].ok:
return results
results.append(_check_database_connection(raw_database_url))
if not results[-1].ok:
return results
results.append(_check_schema_migrations(raw_database_url))
if not results[-1].ok:
return results
if require_rabbitmq or _bool_env("EVENT_BUS_ENABLED", False):
event_bus_url = str(os.getenv("EVENT_BUS_URL", "")).strip()
if not event_bus_url:
results.append(CheckResult("rabbitmq_tcp", False, "EVENT_BUS_URL is required when RabbitMQ check is enabled"))
else:
results.append(_check_tcp_endpoint("rabbitmq_tcp", event_bus_url, 5672))
if base_url:
results.extend(_http_preflight_checks(base_url))
return results
def main() -> int:
parser = argparse.ArgumentParser(description="PostgreSQL dev preflight checks")
parser.add_argument("--env-file", action="append", default=[".env.production"])
parser.add_argument("--database-url", default=os.getenv("DATABASE_URL", ""))
parser.add_argument("--base-url", default="")
parser.add_argument("--require-rabbitmq", action="store_true")
parser.add_argument("--json", action="store_true")
args = parser.parse_args()
for raw_path in args.env_file:
path = Path(raw_path)
if not path.is_absolute():
path = (ROOT / path).resolve()
_load_env_file(path)
database_url = str(args.database_url or os.getenv("DATABASE_URL", "")).strip()
base_url = str(args.base_url or "").strip() or None
results = run_preflight(
database_url=database_url,
base_url=base_url,
require_rabbitmq=bool(args.require_rabbitmq),
)
failures = [item for item in results if not item.ok]
if args.json:
print(
json.dumps(
{
"ok": not failures,
"checks": [asdict(item) for item in results],
},
ensure_ascii=False,
indent=2,
)
)
return 0 if not failures else 1
for item in results:
label = "[ok]" if item.ok else "[fail]"
print(f"{label} {item.name}: {item.details}")
if failures:
print("[FAIL] PostgreSQL dev preflight failed")
return 1
print("[PASS] PostgreSQL dev preflight passed")
return 0
if __name__ == "__main__":
raise SystemExit(main())