Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
82 changes: 76 additions & 6 deletions adapters/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,62 @@ def connect_db(
attempt += 1


def is_sqlite(conn: Any) -> bool:
"""Return whether a DB connection is SQLite-backed."""
return isinstance(conn, sqlite3.Connection)


def is_postgres(conn: Any) -> bool:
"""Return whether a DB connection is Postgres-backed."""
return not is_sqlite(conn) and hasattr(conn, "execute")


def get_placeholder(conn: Any) -> str:
"""Return the parameter placeholder for the active DB connection."""
return "?" if is_sqlite(conn) else "%s"


def table_exists(conn: Any, table_name: str) -> bool:
"""Return whether a table exists on SQLite or Postgres."""
if is_sqlite(conn):
row = conn.execute(
"SELECT 1 FROM sqlite_master WHERE type='table' AND name = ?",
(table_name,),
).fetchone()
return row is not None
row = conn.execute("SELECT to_regclass(%s)", (table_name,)).fetchone()
return bool(row and row[0])


def get_table_columns(conn: Any, table_name: str) -> set[str]:
"""Return table column names for SQLite or Postgres."""
try:
if is_sqlite(conn):
escaped_table = table_name.replace("'", "''")
rows = conn.execute(f"SELECT name FROM pragma_table_info('{escaped_table}')").fetchall()
return {str(row[0]) for row in rows}
rows = conn.execute(
"SELECT column_name FROM information_schema.columns "
"WHERE table_schema = current_schema() AND table_name = %s",
(table_name,),
).fetchall()
return {str(row[0]) for row in rows}
except Exception:
return set()


def manager_id_column(conn: Any, *, require_table: bool = False) -> str | None:
"""Return the manager primary-key column used by the active schema."""
if require_table and not table_exists(conn, "managers"):
return None
columns = get_table_columns(conn, "managers")
if "manager_id" in columns:
return "manager_id"
if "id" in columns:
return "id"
return None


def _ensure_sqlite_usage_schema(conn: sqlite3.Connection) -> None:
conn.execute("""CREATE TABLE IF NOT EXISTS api_usage (
id INTEGER PRIMARY KEY,
Expand All @@ -113,6 +169,16 @@ def _ensure_sqlite_usage_schema(conn: sqlite3.Connection) -> None:


def _ensure_postgres_usage_schema(conn: Any) -> None:
conn.execute("""CREATE TABLE IF NOT EXISTS api_usage (
id BIGSERIAL PRIMARY KEY,
ts TIMESTAMPTZ DEFAULT now(),
source TEXT,
endpoint TEXT,
status INT,
bytes INT,
latency_ms INT,
cost_usd NUMERIC(10,4)
)""")
conn.execute("SELECT to_regclass('api_usage')")
conn.execute("""
DO $$
Expand All @@ -137,6 +203,14 @@ def _ensure_postgres_usage_schema(conn: Any) -> None:
""")


def ensure_api_usage_schema(conn: Any) -> None:
"""Ensure the api_usage table and monthly_usage view exist for the active dialect."""
if is_sqlite(conn):
_ensure_sqlite_usage_schema(conn)
return
_ensure_postgres_usage_schema(conn)


@asynccontextmanager
async def tracked_call(
source: str,
Expand Down Expand Up @@ -192,12 +266,8 @@ def _store(resp: Any) -> None:
else:
computed_cost = 0.0
conn = connect_db(db_path)
if isinstance(conn, sqlite3.Connection):
_ensure_sqlite_usage_schema(conn)
placeholder = "?"
else: # Postgres
_ensure_postgres_usage_schema(conn)
placeholder = "%s"
ensure_api_usage_schema(conn)
placeholder = get_placeholder(conn)
values_clause = ",".join([placeholder] * 6)
sql = (
"INSERT INTO api_usage(source, endpoint, status, bytes, latency_ms, cost_usd)"
Expand Down
46 changes: 13 additions & 33 deletions api/activism.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@

from __future__ import annotations

import sqlite3
from datetime import date, datetime
from typing import Any

Expand All @@ -14,7 +13,7 @@

APIRouter, BaseModel, Field, Query = offline_api_imports()

from adapters.base import connect_db
from adapters.base import connect_db, get_placeholder, is_sqlite, table_exists

router = APIRouter()

Expand Down Expand Up @@ -61,25 +60,6 @@ class ActiveCampaignResponse(BaseModel):
latest_event_type: str | None


def _is_sqlite(conn: Any) -> bool:
return isinstance(conn, sqlite3.Connection)


def _placeholder(conn: Any) -> str:
return "?" if _is_sqlite(conn) else "%s"


def _table_exists(conn: Any, table_name: str) -> bool:
if _is_sqlite(conn):
row = conn.execute(
"SELECT 1 FROM sqlite_master WHERE type='table' AND name = ?",
(table_name,),
).fetchone()
return row is not None
row = conn.execute("SELECT to_regclass(%s)", (table_name,)).fetchone()
return bool(row and row[0])


def _to_float(value: Any) -> float | None:
if value is None:
return None
Expand Down Expand Up @@ -125,10 +105,10 @@ def query_activism_filings(
since: date | None = None,
limit: int = 100,
) -> list[ActivismFilingResponse]:
if not _table_exists(conn, "activism_filings"):
if not table_exists(conn, "activism_filings"):
return []

ph = _placeholder(conn)
ph = get_placeholder(conn)
filters: list[str] = []
params: list[Any] = []

Expand Down Expand Up @@ -182,10 +162,10 @@ def query_activism_events(
since: date | None = None,
limit: int = 100,
) -> list[ActivismEventResponse]:
if not _table_exists(conn, "activism_events"):
if not table_exists(conn, "activism_events"):
return []

ph = _placeholder(conn)
ph = get_placeholder(conn)
filters: list[str] = []
params: list[Any] = []

Expand All @@ -199,7 +179,7 @@ def query_activism_events(
filters.append(f"upper(COALESCE(ae.subject_cusip, '')) = upper({ph})")
params.append(cusip)
if since is not None:
if _is_sqlite(conn):
if is_sqlite(conn):
filters.append(f"date(ae.detected_at) >= date({ph})")
else:
filters.append(f"ae.detected_at::date >= {ph}")
Expand Down Expand Up @@ -234,14 +214,14 @@ def query_activism_events(


def query_activism_timeline(conn: Any, manager_id: int) -> list[ActivismTimelineEntry]:
if not _table_exists(conn, "activism_filings"):
if not table_exists(conn, "activism_filings"):
return []

ph = _placeholder(conn)
ph = get_placeholder(conn)
event_cte = (
", event_entries AS ("
" SELECT "
+ ("date(ae.detected_at)" if _is_sqlite(conn) else "ae.detected_at::date")
+ ("date(ae.detected_at)" if is_sqlite(conn) else "ae.detected_at::date")
+ " AS entry_date, 'event' AS entry_type, "
" ae.event_type || ' on ' || ae.subject_company || "
" CASE WHEN ae.threshold_crossed IS NOT NULL THEN ' (threshold ' || ae.threshold_crossed || '%)' ELSE '' END AS description, "
Expand All @@ -252,7 +232,7 @@ def query_activism_timeline(conn: Any, manager_id: int) -> list[ActivismTimeline
)
union_source = "SELECT * FROM filing_entries"
params: tuple[Any, ...]
if _table_exists(conn, "activism_events"):
if table_exists(conn, "activism_events"):
union_source = "SELECT * FROM filing_entries UNION ALL SELECT * FROM event_entries"
params = (manager_id, manager_id)
else:
Expand Down Expand Up @@ -290,11 +270,11 @@ def query_active_campaigns(
min_ownership_pct: float = 5.0,
limit: int = 100,
) -> list[ActiveCampaignResponse]:
if not _table_exists(conn, "activism_filings"):
if not table_exists(conn, "activism_filings"):
return []

ph = _placeholder(conn)
if _table_exists(conn, "activism_events"):
ph = get_placeholder(conn)
if table_exists(conn, "activism_events"):
event_count_sql = (
"COALESCE((SELECT COUNT(*) FROM activism_events ae "
"WHERE ae.manager_id = ranked.manager_id "
Expand Down
39 changes: 7 additions & 32 deletions api/chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,6 @@
import logging
import math
import os
import sqlite3
import time
import uuid
from collections import defaultdict, deque
Expand All @@ -29,7 +28,7 @@
from prometheus_client import CONTENT_TYPE_LATEST, Histogram, generate_latest
from pydantic import BaseModel, ConfigDict, Field

from adapters.base import connect_db
from adapters.base import connect_db, get_placeholder, is_sqlite, manager_id_column, table_exists
from api.activism import router as activism_router
from api.alerts import router as alerts_router
from api.data import router as data_router
Expand Down Expand Up @@ -365,7 +364,6 @@ async def search_api(
"nl_query": ("chains.nl_query", "NLQueryChain"),
"rag_search": ("chains.rag_search", "RAGSearchChain"),
}
SQLITE_TABLE_INFO_SQL = "SELECT name FROM pragma_table_info(?)"


class ChainUnavailableError(RuntimeError):
Expand Down Expand Up @@ -570,22 +568,8 @@ def _chat_zone_disabled() -> bool:
return os.getenv("LLM_ZONE", "").strip().lower() == "disabled"


def _is_sqlite_connection(conn: Any) -> bool:
return isinstance(conn, sqlite3.Connection)


def _manager_id_column(conn: Any) -> str:
if _is_sqlite_connection(conn):
rows = conn.execute(SQLITE_TABLE_INFO_SQL, ("managers",)).fetchall()
columns = {str(row[0]) for row in rows}
else:
rows = conn.execute(
"SELECT column_name FROM information_schema.columns "
"WHERE table_schema = current_schema() AND table_name = %s",
("managers",),
).fetchall()
columns = {str(row[0]) for row in rows}
return "manager_id" if "manager_id" in columns else "id"
return manager_id_column(conn, require_table=True) or "id"


def _normalize_manager_ids(value: Any) -> list[int]:
Expand Down Expand Up @@ -630,7 +614,7 @@ def _manager_ids_for_name(conn: Any, manager_name: str) -> list[int]:
if not normalized_name:
return []
manager_id_col = _manager_id_column(conn)
placeholder = "?" if _is_sqlite_connection(conn) else "%s"
placeholder = get_placeholder(conn)
rows = conn.execute(
f"SELECT {manager_id_col} FROM managers WHERE lower(name) = lower({placeholder})",
(normalized_name,),
Expand Down Expand Up @@ -882,17 +866,8 @@ def _response_id_from_trace_url(trace_url: str | None) -> str | None:
return None


def _placeholder(conn: Any) -> str:
return "?" if isinstance(conn, sqlite3.Connection) else "%s"


def _postgres_table_exists(conn: Any, table_name: str) -> bool:
row = conn.execute("SELECT to_regclass(%s)", (f"public.{table_name}",)).fetchone()
return bool(row and row[0])


def _ensure_chat_feedback_table(conn: Any) -> None:
if isinstance(conn, sqlite3.Connection):
if is_sqlite(conn):
conn.execute("""CREATE TABLE IF NOT EXISTS chat_feedback (
feedback_id INTEGER PRIMARY KEY,
response_id TEXT NOT NULL,
Expand All @@ -902,16 +877,16 @@ def _ensure_chat_feedback_table(conn: Any) -> None:
)""")
return

if not _postgres_table_exists(conn, "chat_feedback"):
if not table_exists(conn, "chat_feedback"):
raise RuntimeError("chat_feedback table missing; run migrations before accepting feedback")


def _store_feedback(feedback: FeedbackRequest) -> int:
conn = connect_db()
try:
_ensure_chat_feedback_table(conn)
ph = _placeholder(conn)
if isinstance(conn, sqlite3.Connection):
ph = get_placeholder(conn)
if is_sqlite(conn):
cursor = conn.execute(
f"INSERT INTO chat_feedback(response_id, rating, comment) VALUES ({ph}, {ph}, {ph})",
(feedback.response_id, feedback.rating, feedback.comment),
Expand Down
18 changes: 6 additions & 12 deletions api/managers.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
from pydantic import BaseModel, ConfigDict, Field, ValidationError

from adapters.base import connect_db
from adapters.base import manager_id_column as shared_manager_id_column
from api.cache import cache_query, invalidate_cache_prefix
from api.models import (
BulkImportFailure,
Expand All @@ -29,6 +30,7 @@
ManagerStatsResponse,
UniverseImportResponse,
)
from utils.identifiers import normalize_cik

router = APIRouter()
logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -159,7 +161,9 @@ def _ensure_manager_table(conn) -> None:

def _manager_id_column(conn) -> str:
"""Return the manager primary-key column for the active database backend."""
return "id" if isinstance(conn, sqlite3.Connection) else "manager_id"
if not isinstance(conn, sqlite3.Connection):
return "manager_id"
return shared_manager_id_column(conn) or "id"


def _json_array(raw: object) -> list[str]:
Expand Down Expand Up @@ -242,16 +246,6 @@ def _to_manager_response(row: tuple[object, ...]) -> ManagerResponse:
)


def _normalize_cik(raw: Any) -> str:
cik = "" if raw is None else str(raw).strip()
if not cik:
return ""
digits = "".join(ch for ch in cik if ch.isdigit())
if not digits:
return ""
return digits.zfill(10)


def _ensure_universe_schema(conn: Any) -> None:
"""Ensure managers table has the columns/index needed for universe imports."""
_ensure_manager_table(conn)
Expand Down Expand Up @@ -1209,7 +1203,7 @@ async def import_manager_universe(
logger.warning("Universe import skipped record %s: record must be an object", index)
continue
name = str(record.get("name", "")).strip()
cik = _normalize_cik(record.get("cik"))
cik = normalize_cik(record.get("cik"))
jurisdiction = str(record.get("jurisdiction", "")).strip().lower()
if not name or not cik or not jurisdiction:
skipped += 1
Expand Down
Loading
Loading