You cannot select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
274 lines
12 KiB
Python
274 lines
12 KiB
Python
"""SQLite state management for TGTC paper trading."""
|
|
from __future__ import annotations
|
|
|
|
import os
|
|
import sqlite3
|
|
import uuid
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
|
|
from apps.tgtc_trader.models import (
|
|
CREATE_TABLES_SQL,
|
|
TGTCCandidateRow,
|
|
TGTCDailySnapshotRow,
|
|
TGTCPositionRow,
|
|
TGTCSessionRow,
|
|
TGTCSnapshotRow,
|
|
TGTCTradeRow,
|
|
)
|
|
|
|
_DEFAULT_DB = "data/paper/tgtc.db"
|
|
|
|
|
|
class TGTCStateManager:
|
|
"""SQLite CRUD layer for TGTC paper trading state."""
|
|
|
|
def __init__(self, db_path: str | None = None) -> None:
|
|
self._db_path = db_path or os.environ.get("TGTC_TRADER_DB", _DEFAULT_DB)
|
|
Path(self._db_path).parent.mkdir(parents=True, exist_ok=True)
|
|
self._init_db()
|
|
|
|
def _connect(self) -> sqlite3.Connection:
|
|
conn = sqlite3.connect(self._db_path)
|
|
conn.row_factory = sqlite3.Row
|
|
conn.execute("PRAGMA journal_mode=WAL")
|
|
conn.execute("PRAGMA foreign_keys=ON")
|
|
return conn
|
|
|
|
def _init_db(self) -> None:
|
|
with self._connect() as conn:
|
|
conn.executescript(CREATE_TABLES_SQL)
|
|
|
|
# ── Sessions ──────────────────────────────────────────────────────────────
|
|
|
|
def create_session(self, name: str, config_path: str, initial_equity: float) -> str:
|
|
session_id = str(uuid.uuid4())[:8]
|
|
now = datetime.now(timezone.utc).isoformat()
|
|
with self._connect() as conn:
|
|
conn.execute(
|
|
"""INSERT INTO tgtc_sessions (session_id, session_name, config_path, initial_equity, created_at)
|
|
VALUES (?, ?, ?, ?, ?)""",
|
|
(session_id, name, config_path, initial_equity, now),
|
|
)
|
|
return session_id
|
|
|
|
def get_session(self, name_or_id: str) -> TGTCSessionRow | None:
|
|
with self._connect() as conn:
|
|
row = conn.execute(
|
|
"SELECT * FROM tgtc_sessions WHERE session_id=? OR session_name=?",
|
|
(name_or_id, name_or_id),
|
|
).fetchone()
|
|
return TGTCSessionRow(**dict(row)) if row else None
|
|
|
|
def list_sessions(self) -> list[TGTCSessionRow]:
|
|
with self._connect() as conn:
|
|
rows = conn.execute("SELECT * FROM tgtc_sessions ORDER BY created_at DESC").fetchall()
|
|
return [TGTCSessionRow(**dict(r)) for r in rows]
|
|
|
|
def set_session_status(self, session_id: str, status: str) -> None:
|
|
with self._connect() as conn:
|
|
conn.execute("UPDATE tgtc_sessions SET status=? WHERE session_id=?", (status, session_id))
|
|
|
|
def delete_session(self, session_id: str) -> None:
|
|
with self._connect() as conn:
|
|
for tbl in ("tgtc_snapshots", "tgtc_candidates", "tgtc_positions",
|
|
"tgtc_trades", "tgtc_daily_snapshots"):
|
|
conn.execute(f"DELETE FROM {tbl} WHERE session_id=?", (session_id,))
|
|
conn.execute("DELETE FROM tgtc_sessions WHERE session_id=?", (session_id,))
|
|
|
|
def get_equity(self, session_id: str) -> float | None:
|
|
with self._connect() as conn:
|
|
row = conn.execute(
|
|
"""SELECT equity FROM tgtc_daily_snapshots
|
|
WHERE session_id=? ORDER BY date DESC LIMIT 1""",
|
|
(session_id,),
|
|
).fetchone()
|
|
return float(row["equity"]) if row else None
|
|
|
|
# ── Snapshots ─────────────────────────────────────────────────────────────
|
|
|
|
def save_snapshot_batch(self, rows: list[TGTCSnapshotRow]) -> None:
|
|
with self._connect() as conn:
|
|
conn.executemany(
|
|
"""INSERT OR IGNORE INTO tgtc_snapshots
|
|
(session_id, date, captured_at, symbol, rank, price, pct_change, volume, market_cap)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
|
[
|
|
(r.session_id, r.date, r.captured_at, r.symbol, r.rank,
|
|
r.price, r.pct_change, r.volume, r.market_cap)
|
|
for r in rows
|
|
],
|
|
)
|
|
|
|
def get_latest_snapshots(self, session_id: str, date: str,
|
|
limit_per_tick: int = 100) -> list[dict]:
|
|
"""Return latest tick snapshots sorted by rank."""
|
|
with self._connect() as conn:
|
|
latest_tick = conn.execute(
|
|
"""SELECT MAX(captured_at) FROM tgtc_snapshots
|
|
WHERE session_id=? AND date=?""",
|
|
(session_id, date),
|
|
).fetchone()
|
|
if not latest_tick or not latest_tick[0]:
|
|
return []
|
|
rows = conn.execute(
|
|
"""SELECT * FROM tgtc_snapshots
|
|
WHERE session_id=? AND date=? AND captured_at=?
|
|
ORDER BY rank ASC LIMIT ?""",
|
|
(session_id, date, latest_tick[0], limit_per_tick),
|
|
).fetchall()
|
|
return [dict(r) for r in rows]
|
|
|
|
# ── Candidates ────────────────────────────────────────────────────────────
|
|
|
|
def save_candidates(self, rows: list[TGTCCandidateRow]) -> None:
|
|
with self._connect() as conn:
|
|
conn.executemany(
|
|
"""INSERT OR REPLACE INTO tgtc_candidates
|
|
(session_id, date, symbol, score, rank_persistence, rank_velocity,
|
|
price_structure, volume_quality, relative_strength,
|
|
pct_change_at_10, price_at_10, vwap_at_10, above_vwap, decided_at, status)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
|
[
|
|
(r.session_id, r.date, r.symbol, r.score, r.rank_persistence,
|
|
r.rank_velocity, r.price_structure, r.volume_quality,
|
|
r.relative_strength, r.pct_change_at_10, r.price_at_10,
|
|
r.vwap_at_10, int(r.above_vwap), r.decided_at, r.status)
|
|
for r in rows
|
|
],
|
|
)
|
|
|
|
def get_candidates(self, session_id: str, date: str) -> list[dict]:
|
|
with self._connect() as conn:
|
|
rows = conn.execute(
|
|
"""SELECT * FROM tgtc_candidates
|
|
WHERE session_id=? AND date=? ORDER BY score DESC""",
|
|
(session_id, date),
|
|
).fetchall()
|
|
return [dict(r) for r in rows]
|
|
|
|
def update_candidate_status(self, session_id: str, date: str,
|
|
symbol: str, status: str) -> None:
|
|
with self._connect() as conn:
|
|
conn.execute(
|
|
"""UPDATE tgtc_candidates SET status=?
|
|
WHERE session_id=? AND date=? AND symbol=?""",
|
|
(status, session_id, date, symbol),
|
|
)
|
|
|
|
# ── Positions ─────────────────────────────────────────────────────────────
|
|
|
|
def save_position(self, pos: TGTCPositionRow) -> None:
|
|
with self._connect() as conn:
|
|
conn.execute(
|
|
"""INSERT OR REPLACE INTO tgtc_positions
|
|
(session_id, date, symbol, entry_signal, entry_price, stop_price,
|
|
current_stop, shares, entered_at, peak_price, partial_taken,
|
|
be_stop_active, exit_price, exit_reason, exited_at, pnl, r_multiple,
|
|
is_dry_run, status)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
|
(
|
|
pos.session_id, pos.date, pos.symbol, pos.entry_signal,
|
|
pos.entry_price, pos.stop_price, pos.current_stop, pos.shares,
|
|
pos.entered_at, pos.peak_price, int(pos.partial_taken),
|
|
int(pos.be_stop_active), pos.exit_price, pos.exit_reason,
|
|
pos.exited_at, pos.pnl, pos.r_multiple, int(pos.is_dry_run),
|
|
pos.status,
|
|
),
|
|
)
|
|
|
|
def get_open_positions(self, session_id: str, date: str) -> list[TGTCPositionRow]:
|
|
with self._connect() as conn:
|
|
rows = conn.execute(
|
|
"""SELECT * FROM tgtc_positions
|
|
WHERE session_id=? AND date=? AND status='open'""",
|
|
(session_id, date),
|
|
).fetchall()
|
|
return [_row_to_position(r) for r in rows]
|
|
|
|
def get_all_positions(self, session_id: str) -> list[dict]:
|
|
with self._connect() as conn:
|
|
rows = conn.execute(
|
|
"""SELECT * FROM tgtc_positions WHERE session_id=?
|
|
ORDER BY entered_at DESC LIMIT 200""",
|
|
(session_id,),
|
|
).fetchall()
|
|
return [dict(r) for r in rows]
|
|
|
|
# ── Trades ────────────────────────────────────────────────────────────────
|
|
|
|
def save_trade(self, trade: TGTCTradeRow) -> None:
|
|
with self._connect() as conn:
|
|
conn.execute(
|
|
"""INSERT OR IGNORE INTO tgtc_trades
|
|
(trade_id, session_id, date, symbol, entry_signal, entry_price,
|
|
exit_price, entered_at, exited_at, shares, pnl, r_multiple,
|
|
exit_reason, is_dry_run)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
|
(
|
|
trade.trade_id, trade.session_id, trade.date, trade.symbol,
|
|
trade.entry_signal, trade.entry_price, trade.exit_price,
|
|
trade.entered_at, trade.exited_at, trade.shares, trade.pnl,
|
|
trade.r_multiple, trade.exit_reason, int(trade.is_dry_run),
|
|
),
|
|
)
|
|
|
|
def get_trades(self, session_id: str, limit: int = 200) -> list[dict]:
|
|
with self._connect() as conn:
|
|
rows = conn.execute(
|
|
"""SELECT * FROM tgtc_trades WHERE session_id=?
|
|
ORDER BY exited_at DESC LIMIT ?""",
|
|
(session_id, limit),
|
|
).fetchall()
|
|
return [dict(r) for r in rows]
|
|
|
|
# ── Daily snapshots ───────────────────────────────────────────────────────
|
|
|
|
def save_daily_snapshot(self, snap: TGTCDailySnapshotRow) -> None:
|
|
with self._connect() as conn:
|
|
conn.execute(
|
|
"""INSERT OR REPLACE INTO tgtc_daily_snapshots
|
|
(session_id, date, equity, daily_pnl, total_pnl, trades_taken, phase)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?)""",
|
|
(snap.session_id, snap.date, snap.equity, snap.daily_pnl,
|
|
snap.total_pnl, snap.trades_taken, snap.phase),
|
|
)
|
|
|
|
def get_daily_snapshots(self, session_id: str, limit: int = 100) -> list[dict]:
|
|
with self._connect() as conn:
|
|
rows = conn.execute(
|
|
"""SELECT * FROM tgtc_daily_snapshots WHERE session_id=?
|
|
ORDER BY date ASC LIMIT ?""",
|
|
(session_id, limit),
|
|
).fetchall()
|
|
return [dict(r) for r in rows]
|
|
|
|
def get_daily_state(self, session_id: str, date: str) -> TGTCDailySnapshotRow | None:
|
|
with self._connect() as conn:
|
|
row = conn.execute(
|
|
"SELECT * FROM tgtc_daily_snapshots WHERE session_id=? AND date=?",
|
|
(session_id, date),
|
|
).fetchone()
|
|
if row is None:
|
|
return None
|
|
return TGTCDailySnapshotRow(**dict(row))
|
|
|
|
def update_phase(self, session_id: str, date: str, phase: str) -> None:
|
|
now_eq = self.get_equity(session_id) or 0.0
|
|
with self._connect() as conn:
|
|
conn.execute(
|
|
"""INSERT INTO tgtc_daily_snapshots (session_id, date, equity, phase)
|
|
VALUES (?, ?, ?, ?)
|
|
ON CONFLICT(session_id, date) DO UPDATE SET phase=excluded.phase""",
|
|
(session_id, date, now_eq, phase),
|
|
)
|
|
|
|
|
|
def _row_to_position(row: sqlite3.Row) -> TGTCPositionRow:
|
|
d = dict(row)
|
|
d["partial_taken"] = bool(d.get("partial_taken", 0))
|
|
d["be_stop_active"] = bool(d.get("be_stop_active", 0))
|
|
d["is_dry_run"] = bool(d.get("is_dry_run", 1))
|
|
d.pop("id", None)
|
|
return TGTCPositionRow(**d)
|