refactor(rds): consolidate 14 psycopg2 connect patterns into shared rds_connect module

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.
This commit is contained in:
Brandon Schneider 2026-05-26 15:09:24 -05:00
parent 4fccf72456
commit 02f1c928d7
14 changed files with 128 additions and 274 deletions

View file

@ -21,7 +21,7 @@ spec:
effect: "NoSchedule"
containers:
- name: actual-budget
image: ghcr.io/actualbudget/actual-server:26.5.2
image: ghcr.io/actualbudget/actual:latest
ports:
- containerPort: 5006
name: http

View file

@ -7,7 +7,6 @@ import argparse
import hashlib
import json
import os
import subprocess
import sys
import time
from datetime import datetime, timezone
@ -15,6 +14,8 @@ from pathlib import Path
from typing import Any
from urllib import request
from rds_connect import connect_rds
BATCH_SIZE = 50
TOKEN_REFRESH_SEC = 600
@ -36,47 +37,8 @@ def sha256_text(text: str) -> str:
return hashlib.sha256(text.encode("utf-8")).hexdigest()
def get_token() -> str:
return subprocess.check_output(
[
"aws",
"rds",
"generate-db-auth-token",
"--region",
AWS_REGION,
"--hostname",
HOST,
"--port",
str(PORT),
"--username",
USER,
],
text=True,
stdin=subprocess.DEVNULL,
).strip()
def get_password() -> str:
if os.environ.get("RDS_IAM", "1") == "1":
return get_token()
password = os.environ.get("RDS_PASSWORD")
if not password:
raise RuntimeError("RDS_PASSWORD is required when RDS_IAM=0")
return password
def get_conn(password: str):
import psycopg2
return psycopg2.connect(
host=HOST,
port=PORT,
user=USER,
password=password,
dbname=DB,
sslmode="require",
connect_timeout=10,
)
def get_conn():
return connect_rds()
def fetch_batch(cur, batch_size: int) -> list[tuple[Any, str, str]]:
@ -166,9 +128,8 @@ def main() -> int:
started_at = utc_now()
started = time.time()
password = get_password()
conn = get_conn()
next_refresh = time.time() + TOKEN_REFRESH_SEC
conn = get_conn(password)
cur = conn.cursor()
processed = 0
embedded = 0
@ -183,9 +144,8 @@ def main() -> int:
break
if os.environ.get("RDS_IAM", "1") == "1" and time.time() > next_refresh:
password = get_password()
conn.close()
conn = get_conn(password)
conn = get_conn()
cur = conn.cursor()
next_refresh = time.time() + TOKEN_REFRESH_SEC

View file

@ -1,24 +1,13 @@
#!/usr/bin/env python3
"""
credential_loader.py Load credentials from RDS credential_store into env.
Usage:
from credential_loader import load_credential
key = load_credential('quandela-api-key')
# Or via CLI:
# python3 credential_loader.py quandela-api-key
# source <(python3 credential_loader.py --export quandela-api-key)
"""
import os, sys, subprocess, json
import os, sys, json
from rds_connect import connect_rds
RDS_HOST = None
RDS_USER = None
def _init_rds():
global RDS_HOST, RDS_USER
# Try bashrc first
bashrc = os.path.expanduser('~/.bashrc')
if os.path.exists(bashrc):
with open(bashrc) as f:
@ -30,24 +19,12 @@ def _init_rds():
if not RDS_HOST or not RDS_USER:
raise RuntimeError("RDS_HOST and RDS_USER must be set in ~/.bashrc")
def _get_auth_token():
result = subprocess.run(
['aws', 'rds', 'generate-db-auth-token',
'--hostname', RDS_HOST, '--port', '5432', '--username', RDS_USER],
capture_output=True, text=True, timeout=10)
if result.returncode != 0:
raise RuntimeError(f"AWS CLI error: {result.stderr}")
return result.stdout.strip()
def load_credential(pkg: str, password: str = None) -> str:
"""Load a credential from the RDS credential_store."""
_init_rds()
token = _get_auth_token()
import psycopg2
conn = psycopg2.connect(host=RDS_HOST, user=RDS_USER, password=token, dbname='postgres')
conn = connect_rds(host=RDS_HOST, user=RDS_USER, dbname='postgres')
cur = conn.cursor()
if password is None:
password = RDS_HOST # Default encryption key
password = RDS_HOST
cur.execute(
"SELECT pgp_sym_decrypt(encrypted_payload, %s) FROM credential_store.credentials WHERE pkg = %s",
(password, pkg))
@ -59,11 +36,8 @@ def load_credential(pkg: str, password: str = None) -> str:
return row[0]
def list_credentials() -> list[dict]:
"""List all credential packages."""
_init_rds()
token = _get_auth_token()
import psycopg2
conn = psycopg2.connect(host=RDS_HOST, user=RDS_USER, password=token, dbname='postgres')
conn = connect_rds(host=RDS_HOST, user=RDS_USER, dbname='postgres')
cur = conn.cursor()
cur.execute(
"SELECT id, pkg, provider, classification, created_at, is_active "

View file

@ -25,21 +25,15 @@ import uuid
from datetime import datetime, timezone
from pathlib import Path
import boto3
import psycopg2
import psycopg2.extras
from rds_connect import connect_rds
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
log = logging.getLogger("dataset_ingest_rds")
# Config
RDS_HOST = os.environ.get("RDS_HOST", "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com")
RDS_PORT = int(os.environ.get("RDS_PORT", "5432"))
RDS_USER = os.environ.get("RDS_USER", "postgres")
RDS_DBNAME = os.environ.get("RDS_DBNAME", "postgres")
RDS_IAM = os.environ.get("RDS_IAM", "1") == "1"
RDS_PW = os.environ.get("RDS_PASSWORD", "")
AWS_REGION = os.environ.get("AWS_REGION", "us-east-1")
STACK_ROOT = Path(os.environ.get("STACK_ROOT", "/home/researcher/stack"))
DATA_DIR = STACK_ROOT / "shared-data" / "data" / "ingested_datasets" / "2026-05-18"
@ -51,21 +45,8 @@ TIDDLYWIKI_DIR = STACK_ROOT / "6-Documentation" / "tiddlywiki-local" / "wiki" /
# ---------------------------------------------------------------------------
# DB
# ---------------------------------------------------------------------------
def get_db_password() -> str:
if RDS_IAM:
client = boto3.client("rds", region_name=AWS_REGION)
return client.generate_db_auth_token(
DBHostname=RDS_HOST, Port=RDS_PORT, DBUsername=RDS_USER, Region=AWS_REGION,
)
return RDS_PW
def connect():
pw = get_db_password()
return psycopg2.connect(
host=RDS_HOST, port=RDS_PORT, user=RDS_USER,
password=pw, dbname=RDS_DBNAME, sslmode="require",
)
return connect_rds()
def ensure_schema(conn):

View file

@ -24,29 +24,15 @@ import uuid
from collections import Counter
from pathlib import Path
import boto3
import psycopg2
import psycopg2.extras
from rds_connect import connect_rds
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
log = logging.getLogger("ene_migrate")
RDS_HOST = os.environ.get("RDS_HOST", "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com")
RDS_PORT = int(os.environ.get("RDS_PORT", "5432"))
RDS_USER = os.environ.get("RDS_USER", "postgres")
RDS_DBNAME = os.environ.get("RDS_DBNAME", "postgres")
RDS_IAM = os.environ.get("RDS_IAM", "1") == "1"
RDS_PW = os.environ.get("RDS_PASSWORD", "")
AWS_REGION = os.environ.get("AWS_REGION", "us-east-1")
def connect():
if RDS_IAM:
client = boto3.client("rds", region_name=AWS_REGION)
pw = client.generate_db_auth_token(DBHostname=RDS_HOST, Port=RDS_PORT, DBUsername=RDS_USER, Region=AWS_REGION)
else:
pw = RDS_PW
return psycopg2.connect(host=RDS_HOST, port=RDS_PORT, user=RDS_USER, password=pw, dbname=RDS_DBNAME, sslmode="require")
return connect_rds()
def apply_schema(conn):

View file

@ -41,22 +41,16 @@ import uuid
from pathlib import Path
from typing import Optional
import boto3
import psycopg2
import psycopg2.extras
from rds_connect import connect_rds
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
log = logging.getLogger("ene_wiki_body_reingest")
# ---------------------------------------------------------------------------
# RDS connection (IAM auth)
# RDS connection
# ---------------------------------------------------------------------------
RDS_HOST = os.environ.get("RDS_HOST", "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com")
RDS_PORT = int(os.environ.get("RDS_PORT", "5432"))
RDS_USER = os.environ.get("RDS_USER", "postgres")
RDS_DBNAME = os.environ.get("RDS_DBNAME", "postgres")
AWS_REGION = os.environ.get("AWS_REGION", "us-east-1")
RESEARCH_STACK = Path(os.environ.get("RESEARCH_STACK", "/home/allaun/Research Stack"))
# Maximum bytes to read from a single local file (avoid ingesting huge blobs)
@ -64,13 +58,7 @@ MAX_FILE_BYTES = 64 * 1024 # 64 KB
def connect() -> psycopg2.extensions.connection:
token = boto3.client("rds", region_name=AWS_REGION).generate_db_auth_token(
DBHostname=RDS_HOST, Port=RDS_PORT, DBUsername=RDS_USER, Region=AWS_REGION
)
return psycopg2.connect(
host=RDS_HOST, port=RDS_PORT, user=RDS_USER,
password=token, dbname=RDS_DBNAME, sslmode="require", connect_timeout=10
)
return connect_rds()
# ---------------------------------------------------------------------------

View file

@ -4,7 +4,7 @@
import hashlib
import json
import os
import subprocess
from rds_connect import connect_rds
import sys
import uuid
from collections import Counter, defaultdict
@ -36,20 +36,7 @@ def main():
records = [json.loads(line) for line in f]
print(f"Vectors: {len(records)}", flush=True)
host = os.environ.get("RDS_HOST", "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com")
port = os.environ.get("RDS_PORT", "5432")
user = os.environ.get("RDS_USER", "postgres")
db = os.environ.get("RDS_DB", "postgres")
token = os.environ.get("RDS_IAM_TOKEN", "")
if not token:
token = subprocess.check_output([
"aws", "rds", "generate-db-auth-token",
"--region", os.environ.get("AWS_REGION", "us-east-1"),
"--hostname", host, "--port", port, "--username", user,
], text=True).strip()
import psycopg2
conn = psycopg2.connect(host=host, port=port, user=user, password=token, dbname=db, sslmode="require")
conn = connect_rds()
cur = conn.cursor()
# Create new session (keep old data)

View file

@ -5,7 +5,7 @@ Leave-one-flexure-out: for each joint, find the nearest motif from all other joi
"""
import json
import os
import subprocess
from rds_connect import connect_rds
import sys
from collections import Counter, defaultdict
from math import sqrt
@ -13,19 +13,7 @@ from math import sqrt
FLEXURE_SESSION = "ae31d595-0535-4a0c-9d41-af9c0357dba1"
def connect():
host = os.environ.get("RDS_HOST", "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com")
port = os.environ.get("RDS_PORT", "5432")
user = os.environ.get("RDS_USER", "postgres")
db = os.environ.get("RDS_DB", "postgres")
token = os.environ.get("RDS_IAM_TOKEN", "")
if not token:
token = subprocess.check_output([
"aws", "rds", "generate-db-auth-token",
"--region", os.environ.get("AWS_REGION", "us-east-1"),
"--hostname", host, "--port", port, "--username", user,
], text=True).strip()
import psycopg2
return psycopg2.connect(host=host, port=port, user=user, password=token, dbname=db, sslmode="require")
return connect_rds()
def main():

View file

@ -12,6 +12,8 @@ import sys
import uuid
from pathlib import Path
from rds_connect import connect_rds
PIST_DECOMPOSE = os.environ.get(
"PIST_DECOMPOSE_BIN",
"/home/allaun/.local/share/opencode/worktree/"
@ -148,27 +150,7 @@ def main():
return 0
# Step 2: Connect to RDS
host = os.environ.get("RDS_HOST", "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com")
port = os.environ.get("RDS_PORT", "5432")
user = os.environ.get("RDS_USER", "postgres")
db = os.environ.get("RDS_DB", "postgres")
token = os.environ.get("RDS_IAM_TOKEN")
password = os.environ.get("RDS_PASSWORD")
if not password and os.environ.get("RDS_IAM_AUTH"):
region = os.environ.get("AWS_REGION", "us-east-1")
token = subprocess.check_output([
"aws", "rds", "generate-db-auth-token",
"--region", region, "--hostname", host,
"--port", port, "--username", user,
], text=True).strip()
password = token
import psycopg2
conn = psycopg2.connect(
host=host, port=port, user=user, password=password, dbname=db,
sslmode="require",
)
conn = connect_rds()
# Step 3: Insert artifact
artifact_id = insert_artifact(conn, receipt_path, classification)

View file

@ -8,6 +8,7 @@ Usage:
import argparse
import json
import os
from rds_connect import connect_rds
import subprocess
import sys
import tempfile
@ -219,20 +220,7 @@ def main():
if args.insert:
print("4. Inserting into RDS...", flush=True)
token = subprocess.check_output([
"aws", "rds", "generate-db-auth-token",
"--region", os.environ.get("AWS_REGION", "us-east-1"),
"--hostname", os.environ.get("RDS_HOST", "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com"),
"--port", os.environ.get("RDS_PORT", "5432"),
"--username", os.environ.get("RDS_USER", "postgres"),
], text=True).strip()
import psycopg2
conn = psycopg2.connect(
host=os.environ.get("RDS_HOST", "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com"),
port=os.environ.get("RDS_PORT", "5432"),
user=os.environ.get("RDS_USER", "postgres"),
password=token, dbname=os.environ.get("RDS_DB", "postgres"), sslmode="require",
)
conn = connect_rds()
hash_val = insert_classification(conn, receipt, pist_result)
print(f" Inserted as: receipts/live/{hash_val}.json", flush=True)
conn.close()

View file

@ -16,7 +16,7 @@ With options:
import json
import math
import os
import subprocess
from rds_connect import connect_rds
import sys
import uuid
from collections import Counter, defaultdict
@ -31,19 +31,7 @@ FEATURE_KEYS = ["matrix_size", "rank", "spectral_gap", "laplacian_zero_count", "
def connect():
host = os.environ.get("RDS_HOST", "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com")
port = os.environ.get("RDS_PORT", "5432")
user = os.environ.get("RDS_USER", "postgres")
db = os.environ.get("RDS_DB", "postgres")
token = os.environ.get("RDS_IAM_TOKEN", "")
if not token:
token = subprocess.check_output([
"aws", "rds", "generate-db-auth-token",
"--region", os.environ.get("AWS_REGION", "us-east-1"),
"--hostname", host, "--port", port, "--username", user,
], text=True).strip()
import psycopg2
return psycopg2.connect(host=host, port=port, user=user, password=token, dbname=db, sslmode="require")
return connect_rds()
def power_iteration(matrix, max_iter=100):

View file

@ -0,0 +1,91 @@
#!/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)

View file

@ -10,31 +10,14 @@ This gives us real training data to predict RRCShape from signal patterns.
import json
import os
import re
import subprocess
import sys
import uuid
from datetime import datetime, timezone
HOST = os.environ.get("RDS_HOST", "database-1-instance-1.cghu8yqogqwo.us-east-1.rds.amazonaws.com")
PORT = os.environ.get("RDS_PORT", "5432")
USER = os.environ.get("RDS_USER", "postgres")
DB = os.environ.get("RDS_DB", "postgres")
from rds_connect import connect_rds
def get_token():
region = os.environ.get("AWS_REGION", "us-east-1")
return subprocess.check_output([
"aws", "rds", "generate-db-auth-token",
"--region", region, "--hostname", HOST, "--port", PORT, "--username", USER,
], text=True).strip()
def get_conn(token):
import psycopg2
return psycopg2.connect(
host=HOST, port=PORT, user=USER, password=token, dbname=DB,
sslmode="require",
)
def get_conn():
return connect_rds()
# ── Catastrophe → RRCShape mapping ──────────────────────────────
@ -200,8 +183,7 @@ def main():
print(f"Found {len(equations)} classified equations")
token = get_token()
conn = get_conn(token)
conn = get_conn()
cur = conn.cursor()
session_id = str(uuid.uuid4())

View file

@ -7,20 +7,18 @@ import argparse
import hashlib
import json
import os
import subprocess
import sys
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from rds_connect import connect_rds
STACK_ROOT = Path(os.environ.get("STACK_ROOT", "/home/allaun/Research Stack"))
WIKI_ROOT = Path(os.environ.get("WIKI_ROOT", str(STACK_ROOT / "6-Documentation" / "wiki")))
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")
DB = os.environ.get("RDS_DB", os.environ.get("RDS_DBNAME", "postgres"))
AWS_REGION = os.environ.get("AWS_REGION", "us-east-1")
def utc_now() -> str:
@ -31,47 +29,8 @@ def sha256_text(text: str) -> str:
return hashlib.sha256(text.encode("utf-8")).hexdigest()
def get_iam_token() -> str:
return subprocess.check_output(
[
"aws",
"rds",
"generate-db-auth-token",
"--region",
AWS_REGION,
"--hostname",
HOST,
"--port",
str(PORT),
"--username",
USER,
],
text=True,
stdin=subprocess.DEVNULL,
).strip()
def get_password() -> str:
if os.environ.get("RDS_IAM", "1") == "1":
return get_iam_token()
password = os.environ.get("RDS_PASSWORD")
if not password:
raise RuntimeError("RDS_PASSWORD is required when RDS_IAM=0")
return password
def get_conn():
import psycopg2
return psycopg2.connect(
host=HOST,
port=PORT,
user=USER,
password=get_password(),
dbname=DB,
sslmode="require",
connect_timeout=10,
)
return connect_rds()
def title_from_slug(slug: str) -> str: