"""Database operations for the AI Research Oracle. Handles connections, schema initialization, and common queries. """ import os import sqlite3 from typing import Optional from oracle.config import DB_PATH, SCHEMA_PATH, SOURCE_TIERS def get_connection(db_path: Optional[str] = None) -> sqlite3.Connection: """Open a database connection.""" path = db_path or str(DB_PATH) conn = sqlite3.connect(path) conn.row_factory = sqlite3.Row return conn def get_ro_connection(db_path: Optional[str] = None) -> sqlite3.Connection: """Open a read-only database connection.""" path = db_path or str(DB_PATH) return sqlite3.connect(f"file:{path}?mode=ro", uri=True) def init_db(db_path: Optional[str] = None, schema_path: Optional[str] = None) -> sqlite3.Connection: """Initialize or open the database, applying schema if it exists. Schema uses CREATE IF NOT EXISTS so repeated calls are idempotent. Also runs World Monitor migration columns (content_hash, verdict, freshness). """ conn = sqlite3.connect(db_path or str(DB_PATH)) sp = schema_path or str(SCHEMA_PATH) if os.path.exists(sp): with open(sp) as f: conn.executescript(f.read()) conn.commit() # World Monitor migration columns (idempotent) for col in [ "content_hash TEXT DEFAULT ''", "verdict TEXT DEFAULT ''", "source_tier INTEGER DEFAULT 2", ]: try: conn.execute(f"ALTER TABLE entries ADD COLUMN {col}") except sqlite3.OperationalError: pass # already exists conn.commit() return conn def migrate_world_monitor(conn: Optional[sqlite3.Connection] = None) -> dict: """Apply World Monitor schema migrations + backfill. Returns: {columns_added: int, hashes_backfilled: int, verdicts_set: int} """ from oracle.dedup import backfill_hashes, apply_verdicts, get_source_tier c = conn or get_connection() cur = c.cursor() # Check which columns already exist cur.execute("PRAGMA table_info(entries)") existing = {row[1] for row in cur.fetchall()} columns_to_add = [] if "content_hash" not in existing: columns_to_add.append("content_hash TEXT DEFAULT ''") if "verdict" not in existing: columns_to_add.append("verdict TEXT DEFAULT ''") if "source_tier" not in existing: columns_to_add.append("source_tier INTEGER DEFAULT 2") added = 0 for col_def in columns_to_add: try: cur.execute(f"ALTER TABLE entries ADD COLUMN {col_def}") added += 1 except sqlite3.OperationalError: pass # race condition or already exists c.commit() # Backfill content hashes hashes = backfill_hashes(c) # Backfill source tiers — reset first so the WHERE clause catches everything cur.execute("UPDATE entries SET source_tier = 0") c.commit() for source_name, tier_info in SOURCE_TIERS.items(): cur.execute( "UPDATE entries SET source_tier = ? WHERE source = ?", (tier_info["tier"], source_name), ) c.commit() # Set verdicts verdicts = apply_verdicts(c) return {"columns_added": added, "hashes_backfilled": hashes, "verdicts_set": verdicts} def get_stats(conn: sqlite3.Connection) -> dict: """Return database statistics.""" cur = conn.cursor() cur.execute("SELECT COUNT(*) FROM entries") total = cur.fetchone()[0] cur.execute(""" SELECT source, COUNT(*) as cnt, ROUND(AVG(signal_score), 2) as avg_score, MIN(signal_score) as min_score, MAX(signal_score) as max_score FROM entries GROUP BY source """) sources = {r["source"]: dict(r) for r in cur.fetchall()} cur.execute("SELECT COUNT(*) FROM entries WHERE summary IS NOT NULL") summarized = cur.fetchone()[0] cur.execute("SELECT COUNT(*) FROM entries WHERE summary IS NULL") pending = cur.fetchone()[0] # Bucket distribution try: cur.execute(""" SELECT bucket, COUNT(*) as cnt FROM entries WHERE bucket IS NOT NULL GROUP BY bucket ORDER BY cnt DESC """) buckets = {r["bucket"]: r["cnt"] for r in cur.fetchall()} except Exception: buckets = {} return { "total_entries": total, "sources": sources, "summarized": summarized, "pending_summary": pending, "buckets": buckets, } def query_top(conn: sqlite3.Connection, n: int = 10, source: Optional[str] = None, min_score: float = 0) -> list[dict]: """Get top N entries by signal score.""" cur = conn.cursor() where_parts = [] params = [] if min_score > 0: where_parts.append("signal_score >= ?") params.append(min_score) if source: where_parts.append("source = ?") params.append(source) where = (" AND " + " AND ".join(where_parts)) if where_parts else "" cur.execute( f"SELECT * FROM entries {where} ORDER BY signal_score DESC LIMIT ?", params + [n], ) return [dict(r) for r in cur.fetchall()] def query_recent(conn: sqlite3.Connection, hours: int = 24) -> list[dict]: """Get entries from the last N hours.""" from datetime import datetime, timezone, timedelta cutoff = (datetime.now(timezone.utc) - timedelta(hours=hours)).strftime( "%Y-%m-%dT%H:%M:%SZ" ) cur = conn.cursor() cur.execute( "SELECT * FROM entries WHERE first_seen >= ? ORDER BY first_seen DESC", (cutoff,), ) return [dict(r) for r in cur.fetchall()] def query_search(conn: sqlite3.Connection, q: str, limit: int = 20) -> list[dict]: """Search entries by title, summary, and key technical point.""" cur = conn.cursor() pattern = f"%{q}%" cur.execute(""" SELECT * FROM entries WHERE title LIKE ? OR json_extract(summary,'$.one_liner') LIKE ? OR json_extract(summary,'$.key_technical_point') LIKE ? ORDER BY signal_score DESC LIMIT ? """, (pattern, pattern, pattern, limit)) return [dict(r) for r in cur.fetchall()] def query_by_tag(conn: sqlite3.Connection, tag: str, limit: int = 20) -> list[dict]: """Get entries matching a category tag.""" cur = conn.cursor() cur.execute(""" SELECT * FROM entries WHERE json_extract(category_tags,'$') LIKE ? ORDER BY signal_score DESC LIMIT ? """, (f'%"{tag}"%', limit)) return [dict(r) for r in cur.fetchall()]