mirror of
https://github.com/allaunthefox/Research-Stack.git
synced 2026-07-31 03:05:21 +00:00
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:
parent
d18b414ea0
commit
7eda71868a
14 changed files with 128 additions and 274 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
91
4-Infrastructure/shim/rds_connect.py
Normal file
91
4-Infrastructure/shim/rds_connect.py
Normal 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)
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue