Research-Stack/4-Infrastructure/shim/batch_embed_artifacts.py
Brandon Schneider 02f1c928d7 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.
2026-05-26 15:09:34 -05:00

195 lines
6 KiB
Python

#!/usr/bin/env python3
"""Batch-embed ENE artifacts missing embeddings and emit a JSON receipt."""
from __future__ import annotations
import argparse
import hashlib
import json
import os
import sys
import time
from datetime import datetime, timezone
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
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")
OLLAMA_URL = os.environ.get("OLLAMA_URL", "http://100.85.244.73:11434").rstrip("/")
EMBED_MODEL = os.environ.get("EMBED_MODEL", "nomic-embed-text")
def utc_now() -> str:
return datetime.now(timezone.utc).isoformat()
def sha256_text(text: str) -> str:
return hashlib.sha256(text.encode("utf-8")).hexdigest()
def get_conn():
return connect_rds()
def fetch_batch(cur, batch_size: int) -> list[tuple[Any, str, str]]:
cur.execute(
"SELECT id, path, content FROM ene.artifacts "
"WHERE embedding IS NULL AND content IS NOT NULL AND content != '' "
"AND kind NOT IN ('verilog', 'shader', 'schema', 'binary') "
"ORDER BY path LIMIT %s",
(batch_size,),
)
return cur.fetchall()
def generate_embedding(text: str) -> list[float] | None:
payload = json.dumps({"model": EMBED_MODEL, "prompt": text[:2048]}).encode("utf-8")
req = request.Request(
f"{OLLAMA_URL}/api/embeddings",
data=payload,
headers={"Content-Type": "application/json"},
)
with request.urlopen(req, timeout=30) as resp:
data = json.loads(resp.read())
embedding = data.get("embedding")
if not isinstance(embedding, list):
return None
return embedding
def pgvector_literal(values: list[float]) -> str:
return "[" + ",".join(format(value, ".16g") for value in values) + "]"
def build_receipt(
*,
started_at: str,
processed: int,
embedded: int,
failed: list[dict[str, str]],
dry_run: bool,
limit: int | None,
duration_sec: float,
) -> dict[str, Any]:
receipt: dict[str, Any] = {
"schema": "ene_artifact_embedding_batch_receipt_v1",
"version": "1.0.0",
"generated_at_utc": utc_now(),
"started_at_utc": started_at,
"duration_ms": int(duration_sec * 1000),
"dry_run": dry_run,
"limit": limit,
"rds_host": HOST,
"rds_db": DB,
"ollama_url": OLLAMA_URL,
"embed_model": EMBED_MODEL,
"processed_count": processed,
"embedded_count": embedded,
"failed_count": len(failed),
"failures": failed[:20],
}
preimage = {k: v for k, v in receipt.items() if k != "receipt_hash"}
receipt["receipt_hash"] = sha256_text(json.dumps(preimage, sort_keys=True))
return receipt
def write_receipt(receipt: dict[str, Any], path: str | None) -> None:
encoded = json.dumps(receipt, sort_keys=True)
print(encoded)
if path:
out = Path(path)
out.parent.mkdir(parents=True, exist_ok=True)
out.write_text(encoded + "\n", encoding="utf-8")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--batch-size", type=int, default=BATCH_SIZE)
parser.add_argument("--limit", type=int, default=None)
parser.add_argument("--dry-run", action="store_true")
parser.add_argument("--receipt-out")
return parser.parse_args()
def main() -> int:
args = parse_args()
if args.batch_size <= 0:
raise SystemExit("--batch-size must be positive")
started_at = utc_now()
started = time.time()
conn = get_conn()
next_refresh = time.time() + TOKEN_REFRESH_SEC
cur = conn.cursor()
processed = 0
embedded = 0
failed: list[dict[str, str]] = []
attempted_ids: set[str] = set()
try:
rows = fetch_batch(cur, args.batch_size)
while rows and (args.limit is None or processed < args.limit):
rows = [row for row in rows if str(row[0]) not in attempted_ids]
if not rows:
break
if os.environ.get("RDS_IAM", "1") == "1" and time.time() > next_refresh:
conn.close()
conn = get_conn()
cur = conn.cursor()
next_refresh = time.time() + TOKEN_REFRESH_SEC
for aid, path, content in rows:
if args.limit is not None and processed >= args.limit:
break
attempted_ids.add(str(aid))
processed += 1
try:
if args.dry_run:
continue
embedding = generate_embedding(content)
if not embedding:
failed.append({"path": path, "error": "embedding response missing vector"})
continue
cur.execute(
"UPDATE ene.artifacts SET embedding = %s::vector WHERE id = %s",
(pgvector_literal(embedding), aid),
)
embedded += 1
except Exception as exc: # noqa: BLE001 - receipt records per-row failure.
failed.append({"path": path, "error": str(exc)})
if args.dry_run:
conn.rollback()
else:
conn.commit()
rows = fetch_batch(cur, args.batch_size)
finally:
cur.close()
conn.close()
receipt = build_receipt(
started_at=started_at,
processed=processed,
embedded=embedded,
failed=failed,
dry_run=args.dry_run,
limit=args.limit,
duration_sec=time.time() - started,
)
write_receipt(receipt, args.receipt_out)
return 0 if not failed else 2
if __name__ == "__main__":
sys.exit(main())