Download data/storage.py from raghava4u/Trading-Bot-M20: direct link, hf CLI and curl.
- Browser
- Download file 20.1 kB
-
https://huggingface.co/raghava4u/Trading-Bot-M20/resolve/main/data/storage.py
- Command line
-
hf download hf://raghava4u/Trading-Bot-M20/data/storage.py
-
curl -L -o storage.py https://huggingface.co/raghava4u/Trading-Bot-M20/resolve/main/data/storage.py
20.1 kB
| """ | |
| 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 | |