""" data/storage.py — Database persistence layer (SQLite WAL or DuckDB). Thread-safe via ConnectionManager with write lock. Implements get_bars() per contracts.py. """ from __future__ import annotations import datetime import logging import sqlite3 import threading from pathlib import Path from typing import Literal import pandas as pd import config from contracts import InsufficientDataError, validate_bar_dataframe logger = logging.getLogger("trading_system.storage") DB_DIR = Path(config.BASE_DIR) / "data" DB_DIR.mkdir(parents=True, exist_ok=True) SQLITE_PATH = DB_DIR / "market_data.db" TIMEFRAME_TABLE_MAP = { "5Min": "bars_5m", "5m": "bars_5m", "1Hour": "bars_1h", "1H": "bars_1h", "1h": "bars_1h", "1Day": "bars_1d", "1D": "bars_1d", "1d": "bars_1d", } # ── SQL DDL ────────────────────────────────────────────────────────────────── CREATE_BARS_SQL = """ CREATE TABLE IF NOT EXISTS {table} ( symbol TEXT NOT NULL, timestamp TEXT NOT NULL, open REAL NOT NULL, high REAL NOT NULL, low REAL NOT NULL, close REAL NOT NULL, volume INTEGER NOT NULL, vwap REAL, adjusted INTEGER NOT NULL DEFAULT 0, PRIMARY KEY (symbol, timestamp) ); """ CREATE_CORPORATE_ACTIONS_SQL = """ CREATE TABLE IF NOT EXISTS corporate_actions ( symbol TEXT NOT NULL, date TEXT NOT NULL, action_type TEXT NOT NULL, ratio REAL NOT NULL, PRIMARY KEY (symbol, date, action_type) ); """ CREATE_SENTIMENT_CACHE_SQL = """ CREATE TABLE IF NOT EXISTS sentiment_cache ( symbol TEXT PRIMARY KEY, score REAL NOT NULL, cached_at TEXT NOT NULL, source_count INTEGER NOT NULL DEFAULT 0 ); """ CREATE_PDT_TRADES_SQL = """ CREATE TABLE IF NOT EXISTS pdt_trades ( id INTEGER PRIMARY KEY AUTOINCREMENT, symbol TEXT NOT NULL, open_date TEXT NOT NULL, close_date TEXT NOT NULL, settlement_date TEXT NOT NULL, side TEXT NOT NULL, qty REAL NOT NULL ); """ CREATE_LINKED_ORDER_GROUPS_SQL = """ CREATE TABLE IF NOT EXISTS linked_order_groups ( entry_order_id TEXT PRIMARY KEY, symbol TEXT NOT NULL, stop_order_id TEXT, tp_order_id TEXT, entry_filled INTEGER NOT NULL DEFAULT 0, stop_submitted INTEGER NOT NULL DEFAULT 0, tp_submitted INTEGER NOT NULL DEFAULT 0, orphaned INTEGER NOT NULL DEFAULT 0, is_fractional INTEGER NOT NULL DEFAULT 0, entry_price REAL, created_at TEXT NOT NULL ); """ CREATE_CAPITAL_TRACKER_SQL = """ CREATE TABLE IF NOT EXISTS capital_tracker ( date TEXT PRIMARY KEY, daily_pnl REAL NOT NULL DEFAULT 0.0, cumulative_pnl REAL NOT NULL DEFAULT 0.0, starting_capital REAL NOT NULL DEFAULT 0.0 ); """ class ConnectionManager: """Thread-safe SQLite connection manager with WAL mode.""" def __init__(self, db_path: Path | None = None): self._db_path = str(db_path or SQLITE_PATH) self._write_lock = threading.Lock() self._local = threading.local() self._init_db() def _get_conn(self) -> sqlite3.Connection: if not hasattr(self._local, "conn") or self._local.conn is None: conn = sqlite3.connect(self._db_path, timeout=30) conn.execute("PRAGMA journal_mode=WAL;") conn.execute("PRAGMA busy_timeout=30000;") conn.execute("PRAGMA wal_autocheckpoint=1000;") conn.row_factory = sqlite3.Row self._local.conn = conn return self._local.conn def _init_db(self): conn = self._get_conn() with self._write_lock: for table in ("bars_5m", "bars_1h", "bars_1d"): conn.execute(CREATE_BARS_SQL.format(table=table)) conn.execute(CREATE_CORPORATE_ACTIONS_SQL) conn.execute(CREATE_SENTIMENT_CACHE_SQL) conn.execute(CREATE_PDT_TRADES_SQL) conn.execute(CREATE_LINKED_ORDER_GROUPS_SQL) conn.execute(CREATE_CAPITAL_TRACKER_SQL) conn.commit() logger.info("Database initialized at %s (WAL mode)", self._db_path) def execute_write(self, sql: str, params: tuple = ()) -> sqlite3.Cursor: import time as _time conn = self._get_conn() for attempt in range(5): try: with self._write_lock: cursor = conn.execute(sql, params) conn.commit() return cursor except sqlite3.OperationalError as e: if "locked" in str(e) and attempt < 4: logger.warning("DB locked on write (attempt %d), retrying...", attempt + 1) _time.sleep(0.5 * (attempt + 1)) else: raise def execute_many_write(self, sql: str, params_list: list[tuple]) -> None: import time as _time conn = self._get_conn() for attempt in range(5): try: with self._write_lock: conn.executemany(sql, params_list) conn.commit() return except sqlite3.OperationalError as e: if "locked" in str(e) and attempt < 4: logger.warning("DB locked on batch write (attempt %d), retrying...", attempt + 1) _time.sleep(0.5 * (attempt + 1)) else: raise def execute_read(self, sql: str, params: tuple = ()) -> list[sqlite3.Row]: conn = self._get_conn() return conn.execute(sql, params).fetchall() def read_dataframe(self, sql: str, params: tuple = ()) -> pd.DataFrame: conn = self._get_conn() return pd.read_sql_query(sql, conn, params=params) def close(self): if hasattr(self._local, "conn") and self._local.conn: self._local.conn.close() self._local.conn = None # ── Module-level singleton ─────────────────────────────────────────────────── _db: ConnectionManager | None = None def get_db() -> ConnectionManager: global _db if _db is None: _db = ConnectionManager() return _db def close_db(): global _db if _db is not None: _db.close() _db = None # ── Bar operations ─────────────────────────────────────────────────────────── def _resolve_table(timeframe: str) -> str: table = TIMEFRAME_TABLE_MAP.get(timeframe) if table is None: raise ValueError(f"Unknown timeframe: {timeframe}") return table def insert_bars(symbol: str, timeframe: str, bars: list[dict]) -> int: """Insert bars, skipping duplicates. Returns count inserted.""" table = _resolve_table(timeframe) db = get_db() sql = f""" INSERT OR IGNORE INTO {table} (symbol, timestamp, open, high, low, close, volume, vwap, adjusted) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) """ rows = [ ( symbol, str(b["timestamp"]), float(b["open"]), float(b["high"]), float(b["low"]), float(b["close"]), int(b["volume"]), float(b["vwap"]) if b.get("vwap") is not None else None, 1 if b.get("adjusted", False) else 0, ) for b in bars ] db.execute_many_write(sql, rows) return len(rows) def get_first_timestamp(symbol: str, timeframe: str) -> datetime.datetime | None: """Get the oldest stored timestamp for a symbol/timeframe.""" table = _resolve_table(timeframe) db = get_db() rows = db.execute_read( f"SELECT MIN(timestamp) as ts FROM {table} WHERE symbol = ?", (symbol,), ) if rows and rows[0]["ts"]: return pd.Timestamp(rows[0]["ts"]).to_pydatetime().replace( tzinfo=datetime.timezone.utc ) return None def get_last_timestamp(symbol: str, timeframe: str) -> datetime.datetime | None: """Get the most recent stored timestamp for a symbol/timeframe.""" table = _resolve_table(timeframe) db = get_db() rows = db.execute_read( f"SELECT MAX(timestamp) as ts FROM {table} WHERE symbol = ?", (symbol,), ) if rows and rows[0]["ts"]: return pd.Timestamp(rows[0]["ts"]).to_pydatetime().replace( tzinfo=datetime.timezone.utc ) return None def get_bars( symbol: str, timeframe: str, n_bars: int, adjusted: bool = True, ) -> pd.DataFrame: """Fetch bars from DB matching contracts.py specification. Returns DataFrame with UTC timestamp index and columns: open, high, low, close, volume, vwap, adjusted """ table = _resolve_table(timeframe) db = get_db() adj_clause = "AND adjusted = 1" if adjusted else "" sql = f""" SELECT timestamp, open, high, low, close, volume, vwap, adjusted FROM {table} WHERE symbol = ? {adj_clause} ORDER BY timestamp DESC LIMIT ? """ df = db.read_dataframe(sql, (symbol, n_bars)) if df.empty: raise InsufficientDataError(f"{symbol}: no bars in {table}") df["timestamp"] = pd.to_datetime(df["timestamp"], utc=True) df = df.set_index("timestamp").sort_index() df["adjusted"] = df["adjusted"].astype(bool) return validate_bar_dataframe(df, symbol, n_bars) def get_all_bars( symbol: str, timeframe: str, adjusted: bool = True, ) -> pd.DataFrame: """Fetch ALL bars for a symbol (no minimum count enforced).""" table = _resolve_table(timeframe) db = get_db() adj_clause = "AND adjusted = 1" if adjusted else "" sql = f""" SELECT timestamp, open, high, low, close, volume, vwap, adjusted FROM {table} WHERE symbol = ? {adj_clause} ORDER BY timestamp ASC """ df = db.read_dataframe(sql, (symbol,)) if df.empty: return df df["timestamp"] = pd.to_datetime(df["timestamp"], utc=True) df = df.set_index("timestamp").sort_index() df["adjusted"] = df["adjusted"].astype(bool) return df def bar_exists(symbol: str, timeframe: str, timestamp: str) -> bool: """Check if a specific bar already exists (dedup for WebSocket).""" table = _resolve_table(timeframe) db = get_db() rows = db.execute_read( f"SELECT 1 FROM {table} WHERE symbol = ? AND timestamp = ? LIMIT 1", (symbol, timestamp), ) return len(rows) > 0 def count_bars(symbol: str, timeframe: str) -> int: """Count total bars stored for a symbol/timeframe.""" table = _resolve_table(timeframe) db = get_db() rows = db.execute_read( f"SELECT COUNT(*) as cnt FROM {table} WHERE symbol = ?", (symbol,), ) return rows[0]["cnt"] if rows else 0 # ── Corporate actions ──────────────────────────────────────────────────────── def insert_corporate_action(symbol: str, date: str, action_type: str, ratio: float): db = get_db() db.execute_write( "INSERT OR IGNORE INTO corporate_actions (symbol, date, action_type, ratio) VALUES (?, ?, ?, ?)", (symbol, date, action_type, ratio), ) def get_corporate_actions(symbol: str) -> list[dict]: db = get_db() rows = db.execute_read( "SELECT date, action_type, ratio FROM corporate_actions WHERE symbol = ? ORDER BY date", (symbol,), ) return [dict(r) for r in rows] def apply_split_adjustment(symbol: str, split_date: str, ratio: float): """Apply split ratio to all historical bars before split date.""" db = get_db() for table in ("bars_5m", "bars_1h", "bars_1d"): db.execute_write( f""" UPDATE {table} SET open = open * ?, high = high * ?, low = low * ?, close = close * ?, volume = CAST(volume / ? AS INTEGER), adjusted = 1 WHERE symbol = ? AND timestamp < ? """, (ratio, ratio, ratio, ratio, ratio, symbol, split_date), ) logger.critical("Applied split adjustment for %s: ratio=%.4f, date=%s", symbol, ratio, split_date) # ── Sentiment cache ────────────────────────────────────────────────────────── def upsert_sentiment(symbol: str, score: float, source_count: int): db = get_db() now = datetime.datetime.now(datetime.timezone.utc).isoformat() db.execute_write( """INSERT OR REPLACE INTO sentiment_cache (symbol, score, cached_at, source_count) VALUES (?, ?, ?, ?)""", (symbol, score, now, source_count), ) def get_cached_sentiment(symbol: str) -> dict | None: db = get_db() rows = db.execute_read( "SELECT score, cached_at, source_count FROM sentiment_cache WHERE symbol = ?", (symbol,), ) if rows: return dict(rows[0]) return None # ── Capital Tracker ────────────────────────────────────────────────────────── def get_cumulative_pnl() -> float: """Get the total cumulative P&L from all previous trading days.""" db = get_db() rows = db.execute_read( "SELECT cumulative_pnl FROM capital_tracker ORDER BY date DESC LIMIT 1" ) if rows: return float(rows[0]["cumulative_pnl"]) return 0.0 def record_daily_pnl(date: str, daily_pnl: float, starting_capital: float): """Record today's P&L and update cumulative total.""" db = get_db() prev_cumulative = get_cumulative_pnl() new_cumulative = prev_cumulative + daily_pnl db.execute_write( """INSERT OR REPLACE INTO capital_tracker (date, daily_pnl, cumulative_pnl, starting_capital) VALUES (?, ?, ?, ?)""", (date, daily_pnl, new_cumulative, starting_capital), ) logger.info( "Capital tracker: date=%s, daily_pnl=$%.2f, cumulative=$%.2f, capital=$%.2f", date, daily_pnl, new_cumulative, starting_capital, ) def get_capital_history() -> list[dict]: """Get full capital tracking history.""" db = get_db() rows = db.execute_read( "SELECT date, daily_pnl, cumulative_pnl, starting_capital FROM capital_tracker ORDER BY date" ) return [dict(r) for r in rows] # ── PDT trades ─────────────────────────────────────────────────────────────── def insert_pdt_trade(symbol: str, open_date: str, close_date: str, settlement_date: str, side: str, qty: float): db = get_db() db.execute_write( "INSERT INTO pdt_trades (symbol, open_date, close_date, settlement_date, side, qty) VALUES (?, ?, ?, ?, ?, ?)", (symbol, open_date, close_date, settlement_date, side, qty), ) def count_pdt_trades_in_window(start_date: str, end_date: str) -> int: db = get_db() rows = db.execute_read( "SELECT COUNT(*) as cnt FROM pdt_trades WHERE settlement_date >= ? AND settlement_date <= ?", (start_date, end_date), ) return rows[0]["cnt"] if rows else 0 # ── Linked order groups ───────────────────────────────────────────────────── def insert_linked_order_group(group: dict): db = get_db() db.execute_write( """INSERT INTO linked_order_groups (entry_order_id, symbol, stop_order_id, tp_order_id, entry_filled, stop_submitted, tp_submitted, orphaned, is_fractional, entry_price, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", ( group["entry_order_id"], group["symbol"], group.get("stop_order_id"), group.get("tp_order_id"), int(group.get("entry_filled", False)), int(group.get("stop_submitted", False)), int(group.get("tp_submitted", False)), int(group.get("orphaned", False)), int(group.get("is_fractional", False)), group.get("entry_price"), group.get("created_at", datetime.datetime.now(datetime.timezone.utc).isoformat()), ), ) def update_linked_order_group(entry_order_id: str, updates: dict): db = get_db() set_clauses = [] params = [] for key, val in updates.items(): set_clauses.append(f"{key} = ?") params.append(int(val) if isinstance(val, bool) else val) params.append(entry_order_id) db.execute_write( f"UPDATE linked_order_groups SET {', '.join(set_clauses)} WHERE entry_order_id = ?", tuple(params), ) def get_orphan_candidates() -> list[dict]: """Find groups where entry filled but stop/TP not submitted and not yet orphaned.""" db = get_db() rows = db.execute_read( """SELECT * FROM linked_order_groups WHERE orphaned = 0 AND entry_filled = 1 AND (stop_submitted = 0 OR tp_submitted = 0)""" ) return [dict(r) for r in rows] def get_all_linked_groups() -> list[dict]: db = get_db() rows = db.execute_read("SELECT * FROM linked_order_groups ORDER BY created_at DESC") return [dict(r) for r in rows] def get_kv(key: str) -> str | None: """Get a persistent key-value pair.""" db = get_db() conn = db._get_conn() conn.execute( "CREATE TABLE IF NOT EXISTS kv_store (key TEXT PRIMARY KEY, value TEXT)" ) row = conn.execute("SELECT value FROM kv_store WHERE key = ?", (key,)).fetchone() return row[0] if row else None def set_kv(key: str, value: str): """Set a persistent key-value pair.""" db = get_db() conn = db._get_conn() conn.execute( "CREATE TABLE IF NOT EXISTS kv_store (key TEXT PRIMARY KEY, value TEXT)" ) with db._write_lock: conn.execute( "INSERT OR REPLACE INTO kv_store (key, value) VALUES (?, ?)", (key, value), ) conn.commit() def count_trades_today(date_str: str) -> int: """Count trade entries recorded today.""" db = get_db() # Assume pdt_trades or linked_order_groups could be used? # Wait, the user said "trade_log", but in main risk it uses daily_trade_count. # User's snippet: SELECT COUNT(*) FROM trade_log WHERE date(entry_time) = ? # Wait, is there a trade_log table? Let's use the user's exact snippet adapted to sqlite/ConnectionManager conn = db._get_conn() try: row = conn.execute( "SELECT COUNT(*) FROM pdt_trades WHERE date(timestamp) = ?", (date_str,), ).fetchone() return row[0] if row else 0 except sqlite3.OperationalError: # If the user literally meant trade_log, I'll provide both or just exactly what user asked. try: row = conn.execute( "SELECT COUNT(*) FROM trade_log WHERE date(entry_time) = ?", (date_str,), ).fetchone() return row[0] if row else 0 except sqlite3.OperationalError: return 0