mirror of
https://github.com/allaunthefox/Research-Stack.git
synced 2026-07-31 03:05:21 +00:00
Creates 4-Infrastructure/shim/rds_connect.py with a single connect_rds()
function that resolves connection parameters in priority order:
1. explicit kwargs
2. DATABASE_URL env var (postgres://user:pass@host:port/dbname?sslmode=...)
3. individual RDS_* env vars (RDS_HOST, RDS_PORT, RDS_USER, etc.)
4. built-in defaults
Auth resolution (when password is empty or RDS_IAM=1):
1. RDS_IAM_TOKEN env var (pre-computed)
2. boto3 SDK generate_db_auth_token (preferred)
3. subprocess aws rds generate-db-auth-token (fallback)
4. RDS_PASSWORD env var (non-IAM)
Replaces 8 connection pattern variants across 14 active shims:
- subprocess + RDS_IAM_TOKEN fallback: pist_trace_classify_mcp, joint_classifier,
pist_prove_and_classify, ingest_57_flexures
- boto3 SDK: ene_wiki_body_reingest, ene_migrate_and_tag, dataset_ingest_rds
- subprocess + RDS_PASSWORD: batch_embed_artifacts, sync_wiki_to_rds, seed_flexure_dataset
- RDS_IAM_AUTH: pist_classify
- bashrc parsed: credential_loader
v1.4a benchmark confirmed at 100% after refactor.
91 lines
3.4 KiB
Python
91 lines
3.4 KiB
Python
#!/usr/bin/env python3
|
|
"""Shared RDS connection helper — resolves env vars, IAM auth, DATABASE_URL."""
|
|
|
|
import os, subprocess
|
|
from urllib.parse import urlparse
|
|
|
|
|
|
def _resolve_params() -> dict:
|
|
"""Resolve connection parameters from env, preferring DATABASE_URL."""
|
|
du = os.environ.get("DATABASE_URL", "").strip()
|
|
if du:
|
|
p = urlparse(du)
|
|
params = {
|
|
"host": p.hostname or "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com",
|
|
"port": p.port or 5432,
|
|
"user": p.username or "postgres",
|
|
"password": p.password or "",
|
|
"dbname": p.path.lstrip("/") if p.path else "postgres",
|
|
"sslmode": "require",
|
|
}
|
|
# Extract sslmode from query string
|
|
if p.query:
|
|
for q in p.query.split("&"):
|
|
if "=" in q:
|
|
k, v = q.split("=", 1)
|
|
if k == "sslmode":
|
|
params["sslmode"] = v
|
|
return params
|
|
|
|
host = os.environ.get("RDS_HOST", "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com")
|
|
port = int(os.environ.get("RDS_PORT", "5432"))
|
|
user = os.environ.get("RDS_USER", "postgres")
|
|
dbname = os.environ.get("RDS_DB", os.environ.get("RDS_DBNAME", "postgres"))
|
|
sslmode = os.environ.get("RDS_SSLMODE", "require")
|
|
password = os.environ.get("RDS_PASSWORD", "")
|
|
return {"host": host, "port": port, "user": user,
|
|
"password": password, "dbname": dbname, "sslmode": sslmode}
|
|
|
|
|
|
def _get_iam_token(host: str, port: int, user: str, region: str) -> str:
|
|
"""Generate IAM token via boto3 (preferred) or AWS CLI (fallback)."""
|
|
try:
|
|
import boto3
|
|
return boto3.client("rds", region_name=region).generate_db_auth_token(
|
|
DBHostname=host, Port=port, DBUsername=user, Region=region)
|
|
except ImportError:
|
|
return subprocess.check_output([
|
|
"aws", "rds", "generate-db-auth-token",
|
|
"--region", region, "--hostname", host,
|
|
"--port", str(port), "--username", user,
|
|
], text=True).strip()
|
|
|
|
|
|
def connect_rds(**overrides):
|
|
"""Connect to RDS. Override any resolved param via kwargs.
|
|
|
|
Resolution order per field:
|
|
1. explicit **override
|
|
2. DATABASE_URL env var
|
|
3. individual RDS_* env vars
|
|
4. built-in defaults
|
|
|
|
Auth resolution (when password is empty):
|
|
1. RDS_IAM_TOKEN env var (pre-computed)
|
|
2. boto3 SDK generate_db_auth_token (RDS_IAM=1 or RDS_IAM_AUTH=1)
|
|
3. subprocess aws rds generate-db-auth-token
|
|
4. RDS_PASSWORD env var
|
|
"""
|
|
p = _resolve_params()
|
|
p.update(overrides)
|
|
|
|
# Determine if IAM auth is requested
|
|
iam_requested = (
|
|
os.environ.get("RDS_IAM", "").strip() == "1" or
|
|
os.environ.get("RDS_IAM_AUTH", "").strip() == "1"
|
|
)
|
|
|
|
# Resolve password
|
|
if not p.get("password") or iam_requested:
|
|
token = os.environ.get("RDS_IAM_TOKEN", "").strip()
|
|
if token:
|
|
p["password"] = token
|
|
elif iam_requested or not p.get("password"):
|
|
region = os.environ.get("AWS_REGION", "us-east-1")
|
|
p["password"] = _get_iam_token(p["host"], p["port"], p["user"], region)
|
|
|
|
import psycopg2
|
|
kw = {k: p[k] for k in ("host", "port", "user", "password", "dbname", "sslmode")}
|
|
if "connect_timeout" in p:
|
|
kw["connect_timeout"] = p["connect_timeout"]
|
|
return psycopg2.connect(**kw)
|