# ============================================================
# MT5 Trade Stream Server - Database Layer (v2)
# Supports: trade_log + account_snapshot + open_positions
# ============================================================

import sqlite3
import threading
from datetime import datetime, timedelta
from config import DB_PATH, DATA_RETENTION_DAYS

_local = threading.local()


def get_db() -> sqlite3.Connection:
    if not hasattr(_local, "conn") or _local.conn is None:
        _local.conn = sqlite3.connect(DB_PATH, check_same_thread=False)
        _local.conn.row_factory = sqlite3.Row
        _local.conn.execute("PRAGMA journal_mode=WAL")
        _local.conn.execute("PRAGMA synchronous=NORMAL")
    return _local.conn


def init_db():
    conn = get_db()
    conn.executescript("""
        CREATE TABLE IF NOT EXISTS trade_log (
            id          INTEGER PRIMARY KEY AUTOINCREMENT,
            account_id  TEXT    NOT NULL,
            ticket      INTEGER NOT NULL,
            symbol      TEXT    NOT NULL,
            type        TEXT    NOT NULL,
            action      TEXT    NOT NULL,
            lots        REAL    NOT NULL,
            price       REAL    NOT NULL,
            sl          REAL    DEFAULT 0,
            tp          REAL    DEFAULT 0,
            profit      REAL    DEFAULT 0,
            comment     TEXT    DEFAULT '',
            trade_time  TEXT    NOT NULL,
            created_at  TEXT    DEFAULT (datetime('now','localtime'))
        );

        CREATE TABLE IF NOT EXISTS account_snapshot (
            account_id      TEXT    PRIMARY KEY,
            account_name    TEXT    DEFAULT '',
            broker          TEXT    DEFAULT '',
            balance         REAL    DEFAULT 0,
            equity          REAL    DEFAULT 0,
            margin          REAL    DEFAULT 0,
            free_margin     REAL    DEFAULT 0,
            margin_level    REAL    DEFAULT 0,
            floating_profit REAL    DEFAULT 0,
            total_profit    REAL    DEFAULT 0,
            net_deposit     REAL    DEFAULT 0,
            open_positions  INTEGER DEFAULT 0,
            updated_at      TEXT    DEFAULT (datetime('now','localtime'))
        );

        CREATE TABLE IF NOT EXISTS open_position (
            id              INTEGER PRIMARY KEY AUTOINCREMENT,
            account_id      TEXT    NOT NULL,
            ticket          INTEGER NOT NULL,
            symbol          TEXT    NOT NULL,
            type            TEXT    NOT NULL,
            lots            REAL    NOT NULL,
            open_price      REAL    NOT NULL,
            current_price   REAL    DEFAULT 0,
            sl              REAL    DEFAULT 0,
            tp              REAL    DEFAULT 0,
            profit          REAL    DEFAULT 0,
            swap            REAL    DEFAULT 0,
            comment         TEXT    DEFAULT '',
            open_time       TEXT    NOT NULL,
            updated_at      TEXT    DEFAULT (datetime('now','localtime'))
        );

        CREATE INDEX IF NOT EXISTS idx_trade_account   ON trade_log(account_id);
        CREATE INDEX IF NOT EXISTS idx_trade_created    ON trade_log(created_at);
        CREATE INDEX IF NOT EXISTS idx_trade_action     ON trade_log(action);
        CREATE INDEX IF NOT EXISTS idx_pos_account      ON open_position(account_id);
    """)
    conn.commit()


# ── Trade Log ────────────────────────────────────────────────
def insert_trade(data: dict) -> int:
    conn = get_db()
    cur = conn.execute("""
        INSERT INTO trade_log
            (account_id, ticket, symbol, type, action, lots, price, sl, tp, profit, comment, trade_time)
        VALUES
            (:account_id, :ticket, :symbol, :type, :action, :lots, :price, :sl, :tp, :profit, :comment, :trade_time)
    """, data)
    conn.commit()
    return cur.lastrowid


def get_trades(since_id: int = 0, account: str = None, limit: int = 200) -> list[dict]:
    conn = get_db()
    query = "SELECT * FROM trade_log WHERE id > ?"
    params: list = [since_id]
    if account:
        query += " AND account_id = ?"
        params.append(account)
    query += " ORDER BY id DESC LIMIT ?"
    params.append(limit)
    return [dict(r) for r in conn.execute(query, params).fetchall()]


def get_recent_trades(limit: int = 100) -> list[dict]:
    conn = get_db()
    return [dict(r) for r in conn.execute(
        "SELECT * FROM trade_log ORDER BY id DESC LIMIT ?", (limit,)
    ).fetchall()]


# ── Account Snapshot ─────────────────────────────────────────
def upsert_account_snapshot(data: dict):
    """Insert or update account snapshot (UPSERT)."""
    conn = get_db()
    conn.execute("""
        INSERT INTO account_snapshot
            (account_id, account_name, broker, balance, equity, margin, free_margin,
             margin_level, floating_profit, total_profit, net_deposit, open_positions, updated_at)
        VALUES
            (:account_id, :account_name, :broker, :balance, :equity, :margin, :free_margin,
             :margin_level, :floating_profit, :total_profit, :net_deposit, :open_positions, datetime('now','localtime'))
        ON CONFLICT(account_id) DO UPDATE SET
            account_name    = excluded.account_name,
            broker          = excluded.broker,
            balance         = excluded.balance,
            equity          = excluded.equity,
            margin          = excluded.margin,
            free_margin     = excluded.free_margin,
            margin_level    = excluded.margin_level,
            floating_profit = excluded.floating_profit,
            total_profit    = excluded.total_profit,
            net_deposit     = excluded.net_deposit,
            open_positions  = excluded.open_positions,
            updated_at      = datetime('now','localtime')
    """, data)
    conn.commit()


def get_all_accounts() -> list[dict]:
    """Return only accounts updated within the last 5 minutes."""
    conn = get_db()
    return [dict(r) for r in conn.execute(
        "SELECT * FROM account_snapshot WHERE updated_at >= datetime('now','localtime','-5 minutes') ORDER BY account_id"
    ).fetchall()]


# ── Open Positions ───────────────────────────────────────────
def replace_open_positions(account_id: str, positions: list[dict]):
    """Replace all open positions for an account (delete + insert)."""
    conn = get_db()
    conn.execute("DELETE FROM open_position WHERE account_id = ?", (account_id,))
    for pos in positions:
        pos["account_id"] = account_id
        conn.execute("""
            INSERT INTO open_position
                (account_id, ticket, symbol, type, lots, open_price, current_price,
                 sl, tp, profit, swap, comment, open_time)
            VALUES
                (:account_id, :ticket, :symbol, :type, :lots, :open_price, :current_price,
                 :sl, :tp, :profit, :swap, :comment, :open_time)
        """, pos)
    conn.commit()


def get_open_positions(account_id: str = None) -> list[dict]:
    conn = get_db()
    if account_id:
        return [dict(r) for r in conn.execute(
            "SELECT * FROM open_position WHERE account_id = ? ORDER BY open_time DESC",
            (account_id,)
        ).fetchall()]
    else:
        return [dict(r) for r in conn.execute(
            "SELECT * FROM open_position ORDER BY account_id, open_time DESC"
        ).fetchall()]


# ── Stats ────────────────────────────────────────────────────
def get_stats() -> dict:
    conn = get_db()
    accounts = get_all_accounts()
    total_balance = sum(a["balance"] for a in accounts)
    total_equity = sum(a["equity"] for a in accounts)
    total_floating = sum(a["floating_profit"] for a in accounts)
    total_profit = sum(a["total_profit"] for a in accounts)
    total_net_deposit = sum(a["net_deposit"] for a in accounts)
    total_positions = sum(a["open_positions"] for a in accounts)
    total_events = conn.execute("SELECT COUNT(*) as c FROM trade_log").fetchone()["c"]

    return {
        "total_accounts": len(accounts),
        "total_balance": round(total_balance, 2),
        "total_equity": round(total_equity, 2),
        "total_floating_profit": round(total_floating, 2),
        "total_profit": round(total_profit, 2),
        "total_net_deposit": round(total_net_deposit, 2),
        "total_open_positions": total_positions,
        "total_events": total_events,
        "accounts": accounts,
    }


# ── Cleanup ──────────────────────────────────────────────────
def cleanup_old_data() -> int:
    conn = get_db()
    cutoff = (datetime.now() - timedelta(days=DATA_RETENTION_DAYS)).strftime("%Y-%m-%d %H:%M:%S")
    cur = conn.execute("DELETE FROM trade_log WHERE created_at < ?", (cutoff,))
    conn.commit()
    return cur.rowcount
