From 7eda71868a3ad0b8a1844e34b76e3b947bfdbe1f Mon Sep 17 00:00:00 2001 From: Brandon Schneider Date: Tue, 26 May 2026 15:09:24 -0500 Subject: [PATCH] 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. --- .../manifests/actual-budget/deployment.yaml | 2 +- .../shim/batch_embed_artifacts.py | 52 ++--------- 4-Infrastructure/shim/credential_loader.py | 36 +------- 4-Infrastructure/shim/dataset_ingest_rds.py | 23 +---- 4-Infrastructure/shim/ene_migrate_and_tag.py | 18 +--- .../shim/ene_wiki_body_reingest.py | 18 +--- 4-Infrastructure/shim/ingest_57_flexures.py | 17 +--- 4-Infrastructure/shim/joint_classifier.py | 16 +--- 4-Infrastructure/shim/pist_classify.py | 24 +---- .../shim/pist_prove_and_classify.py | 16 +--- .../shim/pist_trace_classify_mcp.py | 16 +--- 4-Infrastructure/shim/rds_connect.py | 91 +++++++++++++++++++ 4-Infrastructure/shim/seed_flexure_dataset.py | 26 +----- 4-Infrastructure/shim/sync_wiki_to_rds.py | 47 +--------- 14 files changed, 128 insertions(+), 274 deletions(-) create mode 100644 4-Infrastructure/shim/rds_connect.py diff --git a/4-Infrastructure/k3s-flake/manifests/actual-budget/deployment.yaml b/4-Infrastructure/k3s-flake/manifests/actual-budget/deployment.yaml index 4d93019c..3d653ed3 100644 --- a/4-Infrastructure/k3s-flake/manifests/actual-budget/deployment.yaml +++ b/4-Infrastructure/k3s-flake/manifests/actual-budget/deployment.yaml @@ -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 diff --git a/4-Infrastructure/shim/batch_embed_artifacts.py b/4-Infrastructure/shim/batch_embed_artifacts.py index bfde8a54..6f8e7137 100644 --- a/4-Infrastructure/shim/batch_embed_artifacts.py +++ b/4-Infrastructure/shim/batch_embed_artifacts.py @@ -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 diff --git a/4-Infrastructure/shim/credential_loader.py b/4-Infrastructure/shim/credential_loader.py index 83c88dc2..c194bd02 100644 --- a/4-Infrastructure/shim/credential_loader.py +++ b/4-Infrastructure/shim/credential_loader.py @@ -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 " diff --git a/4-Infrastructure/shim/dataset_ingest_rds.py b/4-Infrastructure/shim/dataset_ingest_rds.py index 07e48cd3..8ac2d35f 100644 --- a/4-Infrastructure/shim/dataset_ingest_rds.py +++ b/4-Infrastructure/shim/dataset_ingest_rds.py @@ -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): diff --git a/4-Infrastructure/shim/ene_migrate_and_tag.py b/4-Infrastructure/shim/ene_migrate_and_tag.py index 57237f29..76de5a86 100644 --- a/4-Infrastructure/shim/ene_migrate_and_tag.py +++ b/4-Infrastructure/shim/ene_migrate_and_tag.py @@ -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): diff --git a/4-Infrastructure/shim/ene_wiki_body_reingest.py b/4-Infrastructure/shim/ene_wiki_body_reingest.py index 68b4998e..b08cf03b 100644 --- a/4-Infrastructure/shim/ene_wiki_body_reingest.py +++ b/4-Infrastructure/shim/ene_wiki_body_reingest.py @@ -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() # --------------------------------------------------------------------------- diff --git a/4-Infrastructure/shim/ingest_57_flexures.py b/4-Infrastructure/shim/ingest_57_flexures.py index c393bc0a..73e8a589 100644 --- a/4-Infrastructure/shim/ingest_57_flexures.py +++ b/4-Infrastructure/shim/ingest_57_flexures.py @@ -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) diff --git a/4-Infrastructure/shim/joint_classifier.py b/4-Infrastructure/shim/joint_classifier.py index 546d1156..2048b6c8 100644 --- a/4-Infrastructure/shim/joint_classifier.py +++ b/4-Infrastructure/shim/joint_classifier.py @@ -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(): diff --git a/4-Infrastructure/shim/pist_classify.py b/4-Infrastructure/shim/pist_classify.py index 7e0029de..107c45c9 100644 --- a/4-Infrastructure/shim/pist_classify.py +++ b/4-Infrastructure/shim/pist_classify.py @@ -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) diff --git a/4-Infrastructure/shim/pist_prove_and_classify.py b/4-Infrastructure/shim/pist_prove_and_classify.py index 09b56bf9..8403c06c 100644 --- a/4-Infrastructure/shim/pist_prove_and_classify.py +++ b/4-Infrastructure/shim/pist_prove_and_classify.py @@ -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() diff --git a/4-Infrastructure/shim/pist_trace_classify_mcp.py b/4-Infrastructure/shim/pist_trace_classify_mcp.py index d7cfb869..d4b3595c 100644 --- a/4-Infrastructure/shim/pist_trace_classify_mcp.py +++ b/4-Infrastructure/shim/pist_trace_classify_mcp.py @@ -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): diff --git a/4-Infrastructure/shim/rds_connect.py b/4-Infrastructure/shim/rds_connect.py new file mode 100644 index 00000000..88cab4c9 --- /dev/null +++ b/4-Infrastructure/shim/rds_connect.py @@ -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) diff --git a/4-Infrastructure/shim/seed_flexure_dataset.py b/4-Infrastructure/shim/seed_flexure_dataset.py index c18d6ceb..b7de428e 100644 --- a/4-Infrastructure/shim/seed_flexure_dataset.py +++ b/4-Infrastructure/shim/seed_flexure_dataset.py @@ -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()) diff --git a/4-Infrastructure/shim/sync_wiki_to_rds.py b/4-Infrastructure/shim/sync_wiki_to_rds.py index b5574ab6..4d7259fc 100644 --- a/4-Infrastructure/shim/sync_wiki_to_rds.py +++ b/4-Infrastructure/shim/sync_wiki_to_rds.py @@ -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: