import base64
from contextlib import contextmanager
import hashlib
import os
from typing import Any, Dict, List, Optional

import pymysql
from cryptography.fernet import Fernet
from sqlalchemy.pool import QueuePool

import config


def _create_connection():
    return pymysql.connect(
        host=config.MYSQL_HOST,
        user=config.MYSQL_USER,
        password=config.MYSQL_PASSWORD,
        database=config.MYSQL_DB,
        charset="utf8mb4",
        cursorclass=pymysql.cursors.DictCursor,
        autocommit=True,
        # Localhost traffic gains nothing from TLS, and this host's MariaDB
        # produces intermittent "DECRYPTION_FAILED_OR_BAD_RECORD_MAC" errors
        # over SSL — keep the connection plain and stable.
        ssl_disabled=True,
        # Fail fast on dead sockets instead of hanging for the TCP timeout.
        read_timeout=10,
        write_timeout=10,
        connect_timeout=5,
    )


_pool = QueuePool(
    _create_connection,
    pool_size=5,
    max_overflow=3,
    # Recycle well under shared-hosting wait_timeout (~300 s) so pooled
    # connections are never handed out after the server killed them.
    recycle=120,
)


def _derive_fernet() -> Fernet:
    """Derive a Fernet key from FLASK_SECRET_KEY (deterministic)."""
    raw = config.FLASK_SECRET_KEY.encode("utf-8")
    key = base64.urlsafe_b64encode(hashlib.sha256(raw).digest())
    return Fernet(key)


class Database:
    """MariaDB wrapper using a SQLAlchemy QueuePool over pymysql."""

    def __init__(self):
        self._fernet = _derive_fernet()

    @contextmanager
    def _conn(self):
        conn = _pool.connect()
        try:
            # pymysql-level liveness check with transparent reconnect —
            # raw QueuePool has no dialect, so pool pre_ping can't be used.
            conn.ping(reconnect=True)
        except Exception:
            conn.invalidate()
            raise
        try:
            yield conn
        finally:
            conn.close()

    def dispose(self):
        _pool.dispose()

    def acquire_core_lock(self):
        """Try to become the Telegram-core leader across Passenger workers.

        Passenger spawns several lswsgi workers; if two of them connect to
        Telegram with the same session auth key, the DC drops one worker's
        pending requests (the stale worker then times out and gives up).
        MySQL GET_LOCK is a connection-scoped cross-process mutex: the caller
        must HOLD the returned connection for the app's lifetime (MySQL
        auto-releases the lock when the connection closes or the process
        dies). Returns the held connection, or None if another worker owns
        the lock.
        """
        conn = _create_connection()
        try:
            with conn.cursor() as cur:
                cur.execute("SELECT GET_LOCK('quill_tg_core', 0) AS locked")
                row = cur.fetchone()
            if row and row.get("locked") == 1:
                return conn
            conn.close()
            return None
        except Exception:
            try:
                conn.close()
            except Exception:
                pass
            return None

    def release_core_lock(self, conn) -> None:
        """Release the core lock (called on graceful shutdown only)."""
        try:
            with conn.cursor() as cur:
                cur.execute("SELECT RELEASE_LOCK('quill_tg_core')")
                cur.fetchone()
        except Exception:
            pass
        finally:
            try:
                conn.close()
            except Exception:
                pass

    def get_config(self, key: str) -> Optional[str]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT var_value FROM config_vars WHERE var_key=%s", (key,)
                )
                row = cur.fetchone()
                return row["var_value"] if row else None

    def set_config(self, key: str, value: str):
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "INSERT INTO config_vars (var_key, var_value) VALUES (%s, %s) "
                    "ON DUPLICATE KEY UPDATE var_value=%s",
                    (key, value, value),
                )

    # ── Secret config (Fernet-encrypted at rest) ────────────────────────────
    # Used for values that must never be readable in plaintext from the DB
    # (e.g. the LLM API key). Encrypted values are prefixed with "enc:"
    # so legacy/plain values keep working and can be detected.

    _SECRET_PREFIX = "enc:"

    def set_secret_config(self, key: str, value: str):
        """Store ``value`` Fernet-encrypted under ``key`` (empty clears it)."""
        if value:
            blob = self._fernet.encrypt(value.encode("utf-8")).decode("utf-8")
            stored = self._SECRET_PREFIX + blob
        else:
            stored = ""
        self.set_config(key, stored)

    def get_secret_config(self, key: str) -> Optional[str]:
        """Read and decrypt a secret config value. Returns plaintext or None."""
        raw = self.get_config(key)
        if not raw:
            return None
        if not raw.startswith(self._SECRET_PREFIX):
            # Not encrypted (legacy/plain) — return as-is.
            return raw
        try:
            return self._fernet.decrypt(
                raw[len(self._SECRET_PREFIX) :].encode("utf-8")
            ).decode("utf-8")
        except Exception:
            return None

    def get_all_config(self) -> dict:
        """Return all config_vars as a flat dict."""
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute("SELECT var_key, var_value FROM config_vars")
                return {row["var_key"]: row["var_value"] for row in cur.fetchall()}

    def delete_config(self, key: str):
        """Remove a config_vars entry (resets to env default)."""
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute("DELETE FROM config_vars WHERE var_key=%s", (key,))

    # ── Auth state (survives restart) ─────────────────────────────────────────

    def get_auth_state(self) -> dict:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT phone, code_hash, status FROM auth_state WHERE id=1"
                )
                row = cur.fetchone()
                return dict(row) if row else {}

    def set_auth_state(
        self,
        phone: Optional[str] = None,
        code_hash: Optional[str] = None,
        status: Optional[str] = None,
    ):
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "INSERT INTO auth_state (id, phone, code_hash, status) VALUES (1, %s, %s, %s) "
                    "ON DUPLICATE KEY UPDATE phone=VALUES(phone), code_hash=VALUES(code_hash), status=VALUES(status)",
                    (phone, code_hash, status),
                )

    def clear_auth_state(self):
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute("DELETE FROM auth_state WHERE id=1")

    # ── Channel pairs ─────────────────────────────────────────────────────────

    def get_pairs(self, only_enabled: bool = False) -> List[dict]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                sql = (
                    "SELECT cp.*, "
                    "ta.name AS account_name, "
                    "bt.bot_username AS bot_username "
                    "FROM channel_pairs cp "
                    "LEFT JOIN tg_accounts ta ON cp.account_id = ta.id "
                    "LEFT JOIN bot_tokens bt ON cp.bot_token_id = bt.id"
                )
                if only_enabled:
                    sql += " WHERE cp.enabled=TRUE"
                sql += " ORDER BY cp.id"
                cur.execute(sql)
                return list(cur.fetchall())

    def get_pair_by_id(self, pair_id: int) -> Optional[dict]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute("SELECT * FROM channel_pairs WHERE id=%s", (pair_id,))
                return cur.fetchone()

    def save_pair(
        self,
        source_chat_id: int,
        dest_chat_id: int,
        source_title: str = "",
        dest_title: str = "",
        enabled: bool = True,
        include_text: bool = True,
        include_media: bool = True,
        skip_standalone_media: bool = False,
        forward_as_link: bool = False,
        price_augment: bool = False,
        forward_via_bot: bool = False,
        augment_symbol: Optional[str] = None,
        filter_type: str = "none",
        llm_enabled: bool = False,
        account_id: Optional[int] = None,
        bot_token_id: Optional[int] = None,
    ) -> int:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "INSERT INTO channel_pairs "
                    "(source_chat_id, dest_chat_id, source_title, dest_title, enabled, "
                    "include_text, include_media, skip_standalone_media, forward_as_link, price_augment, "
                    "forward_via_bot, augment_symbol, filter_type, llm_enabled, account_id, bot_token_id) "
                    "VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) "
                    "ON DUPLICATE KEY UPDATE "
                    "source_title=%s, dest_title=%s, enabled=%s, "
                    "include_text=%s, include_media=%s, skip_standalone_media=%s, "
                    "forward_as_link=%s, price_augment=%s, forward_via_bot=%s, augment_symbol=%s, filter_type=%s, "
                    "llm_enabled=%s, account_id=%s, bot_token_id=%s",
                    (
                        source_chat_id,
                        dest_chat_id,
                        source_title,
                        dest_title,
                        enabled,
                        include_text,
                        include_media,
                        skip_standalone_media,
                        forward_as_link,
                        price_augment,
                        forward_via_bot,
                        augment_symbol,
                        filter_type,
                        llm_enabled,
                        account_id,
                        bot_token_id,
                        source_title,
                        dest_title,
                        enabled,
                        include_text,
                        include_media,
                        skip_standalone_media,
                        forward_as_link,
                        price_augment,
                        forward_via_bot,
                        augment_symbol,
                        filter_type,
                        llm_enabled,
                        account_id,
                        bot_token_id,
                    ),
                )
                if cur.lastrowid:
                    return cur.lastrowid
                cur.execute(
                    "SELECT id FROM channel_pairs WHERE source_chat_id=%s AND dest_chat_id=%s",
                    (source_chat_id, dest_chat_id),
                )
                row = cur.fetchone()
                return row["id"] if row else 0

    def update_pair(self, pair_id: int, **fields) -> bool:
        allowed = {
            "source_title",
            "dest_title",
            "enabled",
            "include_text",
            "include_media",
            "skip_standalone_media",
            "forward_as_link",
            "price_augment",
            "forward_via_bot",
            "augment_symbol",
            "filter_type",
            "llm_enabled",
            "account_id",
            "bot_token_id",
        }
        updates = {k: v for k, v in fields.items() if k in allowed}
        if not updates:
            return False
        with self._conn() as conn:
            with conn.cursor() as cur:
                set_clause = ", ".join(f"{k}=%s" for k in updates)
                values = list(updates.values()) + [pair_id]
                cur.execute(
                    f"UPDATE channel_pairs SET {set_clause} WHERE id=%s", values
                )
                return cur.rowcount > 0

    def ensure_llm_schema(self) -> None:
        """Idempotently add LLM columns to existing databases.

        deploy.sh only applies schema.sql (CREATE IF NOT EXISTS), which never
        alters existing tables, so new columns must be added explicitly. This
        is safe to run on every startup (MariaDB 10.5+ supports IF NOT EXISTS).
        """
        with self._conn() as conn:
            with conn.cursor() as cur:
                try:
                    cur.execute(
                        "ALTER TABLE channel_pairs "
                        "ADD COLUMN IF NOT EXISTS llm_enabled BOOLEAN DEFAULT FALSE"
                    )
                except Exception:
                    # Older MariaDB without IF NOT EXISTS: probe then add.
                    cur.execute(
                        "SELECT COUNT(*) AS c FROM information_schema.columns "
                        "WHERE table_schema=DATABASE() AND table_name='channel_pairs' "
                        "AND column_name='llm_enabled'"
                    )
                    if (cur.fetchone() or {}).get("c", 0) == 0:
                        cur.execute(
                            "ALTER TABLE channel_pairs "
                            "ADD COLUMN llm_enabled BOOLEAN DEFAULT FALSE"
                        )

    def ensure_pair_schema(self) -> None:
        """Idempotently add account_id / bot_token_id to existing databases.

        Mirrors ensure_llm_schema: schema.sql's CREATE IF NOT EXISTS never
        alters existing tables, so new columns are added explicitly at startup.
        """
        with self._conn() as conn:
            with conn.cursor() as cur:
                for col, ddl in (
                    ("account_id", "ADD COLUMN IF NOT EXISTS account_id INT DEFAULT NULL"),
                    ("bot_token_id", "ADD COLUMN IF NOT EXISTS bot_token_id INT DEFAULT NULL"),
                ):
                    try:
                        cur.execute(f"ALTER TABLE channel_pairs {ddl}")
                    except Exception:
                        # Older MariaDB without IF NOT EXISTS: probe then add.
                        cur.execute(
                            "SELECT COUNT(*) AS c FROM information_schema.columns "
                            "WHERE table_schema=DATABASE() AND table_name='channel_pairs' "
                            "AND column_name=%s",
                            (col,),
                        )
                        if (cur.fetchone() or {}).get("c", 0) == 0:
                            cur.execute(f"ALTER TABLE channel_pairs ADD COLUMN {col} INT DEFAULT NULL")

    def delete_pair(self, pair_id: int) -> bool:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute("DELETE FROM channel_pairs WHERE id=%s", (pair_id,))
                return cur.rowcount > 0

    def toggle_pair(self, pair_id: int) -> bool:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "UPDATE channel_pairs SET enabled=NOT enabled WHERE id=%s",
                    (pair_id,),
                )
                return cur.rowcount > 0

    # ── Message map ───────────────────────────────────────────────────────────

    def save_message_mapping(
        self,
        source_chat_id: int,
        source_msg_id: int,
        dest_chat_id: int,
        dest_msg_id: int,
        via_bot: bool = False,
    ):
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "INSERT INTO message_map (source_chat_id, source_msg_id, dest_chat_id, dest_msg_id, via_bot) "
                    "VALUES (%s, %s, %s, %s, %s) "
                    "ON DUPLICATE KEY UPDATE dest_msg_id=%s, via_bot=%s",
                    (
                        source_chat_id,
                        source_msg_id,
                        dest_chat_id,
                        dest_msg_id,
                        bool(via_bot),
                        dest_msg_id,
                        bool(via_bot),
                    ),
                )

    def get_mapped_message(
        self, source_chat_id: int, dest_chat_id: int, source_msg_id: int
    ) -> Optional[int]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT dest_msg_id FROM message_map "
                    "WHERE source_chat_id=%s AND dest_chat_id=%s AND source_msg_id=%s",
                    (source_chat_id, dest_chat_id, source_msg_id),
                )
                row = cur.fetchone()
                return row["dest_msg_id"] if row else None

    def get_message_mapping(
        self, source_chat_id: int, dest_chat_id: int, source_msg_id: int
    ) -> Optional[dict]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT dest_msg_id, via_bot FROM message_map "
                    "WHERE source_chat_id=%s AND dest_chat_id=%s AND source_msg_id=%s",
                    (source_chat_id, dest_chat_id, source_msg_id),
                )
                return cur.fetchone()

    def prune_message_map(self, days: int) -> int:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "DELETE FROM message_map WHERE created_at < NOW() - INTERVAL %s DAY",
                    (days,),
                )
                return cur.rowcount

    # ── Signal price snapshots ────────────────────────────────────────────────

    def save_signal_snapshot(self, snapshot: dict) -> int:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "INSERT INTO signal_log "
                    "(source_chat_id, source_message_id, signal_text, signal_date, "
                    "reply_to_message_id, received_at_ms, price_fetched_at_ms, "
                    "symbol, symbol_id, bid, ask, spread, tick_timestamp_ms, price_available) "
                    "VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) "
                    "ON DUPLICATE KEY UPDATE "
                    "signal_text=VALUES(signal_text), signal_date=VALUES(signal_date), "
                    "received_at_ms=VALUES(received_at_ms), price_fetched_at_ms=VALUES(price_fetched_at_ms), "
                    "symbol=VALUES(symbol), symbol_id=VALUES(symbol_id), bid=VALUES(bid), "
                    "ask=VALUES(ask), spread=VALUES(spread), tick_timestamp_ms=VALUES(tick_timestamp_ms), "
                    "price_available=VALUES(price_available)",
                    (
                        snapshot.get("source_chat_id"),
                        snapshot.get("source_message_id"),
                        snapshot.get("signal_text", ""),
                        snapshot.get("signal_date"),
                        snapshot.get("reply_to_message_id"),
                        snapshot.get("received_at_ms"),
                        snapshot.get("price_fetched_at_ms"),
                        snapshot.get("symbol", ""),
                        snapshot.get("symbol_id"),
                        snapshot.get("bid"),
                        snapshot.get("ask"),
                        snapshot.get("spread"),
                        snapshot.get("tick_timestamp_ms"),
                        bool(snapshot.get("price_available", False)),
                    ),
                )
                return cur.lastrowid or 0

    def get_signal_snapshots(
        self,
        source_chat_id: int,
        limit: int = 50,
        since_id: Optional[int] = None,
    ) -> List[dict]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                sql = "SELECT * FROM signal_log WHERE source_chat_id=%s "
                params: list = [source_chat_id]
                if since_id is not None:
                    sql += "AND id > %s "
                    params.append(since_id)
                sql += "ORDER BY id DESC LIMIT %s"
                params.append(limit)
                cur.execute(sql, params)
                return list(cur.fetchall())

    def get_all_signal_snapshots(
        self,
        limit: int = 50,
        since_id: Optional[int] = None,
    ) -> List[dict]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                sql = "SELECT * FROM signal_log "
                params: list = []
                if since_id is not None:
                    sql += "WHERE id > %s "
                    params.append(since_id)
                sql += "ORDER BY id DESC LIMIT %s"
                params.append(limit)
                cur.execute(sql, params)
                return list(cur.fetchall())

    # ── Telegram accounts (multi-account) ────────────────────────────────────

    def get_accounts(self) -> List[dict]:
        """List accounts for the UI. Never exposes session credentials."""
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT id, name, phone, is_primary, is_authorized, status, "
                    "phone_code_hash, created_at, updated_at "
                    "FROM tg_accounts ORDER BY is_primary DESC, id"
                )
                return list(cur.fetchall())

    def get_account(self, account_id: int) -> Optional[dict]:
        """Account metadata for the UI. Never exposes session credentials."""
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT id, name, phone, is_primary, is_authorized, status, "
                    "qr_token, qr_expires_at, phone_code_hash, created_at, updated_at "
                    "FROM tg_accounts WHERE id=%s",
                    (account_id,),
                )
                return cur.fetchone()

    def get_account_session(self, account_id: int) -> Optional[str]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT session_string FROM tg_accounts WHERE id=%s",
                    (account_id,),
                )
                row = cur.fetchone()
                return row["session_string"] if row else None

    def get_account_pending_session(self, account_id: int) -> Optional[str]:
        """Pending (half-authenticated) session, used to finish sign-in after a worker recycle."""
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT pending_session FROM tg_accounts WHERE id=%s",
                    (account_id,),
                )
                row = cur.fetchone()
                return row["pending_session"] if row else None

    def save_account(
        self,
        name: str,
        phone: str = "",
        is_primary: bool = False,
    ) -> int:
        with self._conn() as conn:
            with conn.cursor() as cur:
                if is_primary:
                    cur.execute("UPDATE tg_accounts SET is_primary=FALSE")
                cur.execute(
                    "INSERT INTO tg_accounts (name, phone, is_primary) VALUES (%s, %s, %s)",
                    (name, phone, is_primary),
                )
                return cur.lastrowid or 0

    def update_account(self, account_id: int, **fields) -> bool:
        allowed = {
            "name",
            "phone",
            "is_primary",
            "is_authorized",
            "status",
            "session_string",
            "qr_token",
            "qr_expires_at",
            "phone_code_hash",
            "pending_session",
        }
        updates = {k: v for k, v in fields.items() if k in allowed}
        if not updates:
            return False
        with self._conn() as conn:
            with conn.cursor() as cur:
                if updates.get("is_primary"):
                    cur.execute("UPDATE tg_accounts SET is_primary=FALSE")
                set_clause = ", ".join(f"{k}=%s" for k in updates)
                values = list(updates.values()) + [account_id]
                cur.execute(
                    f"UPDATE tg_accounts SET {set_clause} WHERE id=%s", values
                )
                return cur.rowcount > 0

    def delete_account(self, account_id: int) -> bool:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute("DELETE FROM tg_accounts WHERE id=%s", (account_id,))
                return cur.rowcount > 0

    def reset_account(self, account_id: int) -> bool:
        """Clear all auth state for an account so it can start fresh.

        Clears: pending_session, phone_code_hash, qr_token, qr_expires_at,
        and resets status to 'pending'. Does NOT delete the account row.
        """
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "UPDATE tg_accounts SET "
                    "pending_session=NULL, phone_code_hash='', "
                    "qr_token=NULL, qr_expires_at=NULL, "
                    "status='pending', is_authorized=FALSE "
                    "WHERE id=%s",
                    (account_id,),
                )
                return cur.rowcount > 0

    def reset_all_stale(self) -> int:
        """Reset all accounts that have stale pending auth state.

        Returns the number of accounts reset.
        """
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "UPDATE tg_accounts SET "
                    "pending_session=NULL, phone_code_hash='', "
                    "qr_token=NULL, qr_expires_at=NULL, "
                    "status='pending' "
                    "WHERE is_authorized=FALSE "
                    "AND (pending_session IS NOT NULL OR phone_code_hash != '' "
                    "OR qr_token IS NOT NULL)"
                )
                return cur.rowcount

    def get_primary_account(self) -> Optional[dict]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT id, name, phone, is_authorized, status "
                    "FROM tg_accounts WHERE is_primary=TRUE LIMIT 1"
                )
                return cur.fetchone()

    def get_authorized_accounts(self) -> List[dict]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT id, name, phone, is_primary FROM tg_accounts "
                    "WHERE is_authorized=TRUE ORDER BY is_primary DESC, id"
                )
                return list(cur.fetchall())

    def migrate_legacy_session(self) -> bool:
        """Migrate legacy session_string from config_vars to tg_accounts (one-time)."""
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT var_value FROM config_vars WHERE var_key='session_string'"
                )
                row = cur.fetchone()
                if not row or not row["var_value"]:
                    return False
                session_str = row["var_value"]
                cur.execute("SELECT COUNT(*) AS cnt FROM tg_accounts")
                if cur.fetchone()["cnt"] > 0:
                    return False
                cur.execute(
                    "INSERT INTO tg_accounts (name, phone, session_string, is_primary, is_authorized, status) "
                    "VALUES (%s, %s, %s, TRUE, TRUE, 'active')",
                    ("Default Account", config.TG_PHONE or "", session_str),
                )
                return True

    # ── Bot tokens ───────────────────────────────────────────────────────────

    def _encrypt_token(self, token: str) -> str:
        return self._fernet.encrypt(token.encode("utf-8")).decode("utf-8")

    def _decrypt_token(self, encrypted: str) -> str:
        return self._fernet.decrypt(encrypted.encode("utf-8")).decode("utf-8")

    def get_bot_tokens(self) -> List[dict]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT id, token_hash, bot_name, bot_username, is_active, verified, "
                    "created_at, updated_at FROM bot_tokens ORDER BY id"
                )
                return list(cur.fetchall())

    def get_bot_token(self, bot_id: int) -> Optional[dict]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT id, token_hash, bot_name, bot_username, is_active, verified, "
                    "created_at, updated_at FROM bot_tokens WHERE id=%s",
                    (bot_id,),
                )
                return cur.fetchone()

    def get_decrypted_bot_token(self, bot_id: int) -> Optional[str]:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT token_encrypted FROM bot_tokens WHERE id=%s AND is_active=TRUE",
                    (bot_id,),
                )
                row = cur.fetchone()
                if not row:
                    return None
                return self._decrypt_token(row["token_encrypted"])

    def save_bot_token(
        self,
        token: str,
        bot_name: str = "",
        bot_username: str = "",
    ) -> int:
        token_hash = hashlib.sha256(token.encode()).hexdigest()
        encrypted = self._encrypt_token(token)
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "INSERT INTO bot_tokens (token_hash, token_encrypted, bot_name, bot_username) "
                    "VALUES (%s, %s, %s, %s) "
                    "ON DUPLICATE KEY UPDATE token_encrypted=%s, bot_name=%s, bot_username=%s, is_active=TRUE",
                    (
                        token_hash,
                        encrypted,
                        bot_name,
                        bot_username,
                        encrypted,
                        bot_name,
                        bot_username,
                    ),
                )
                if cur.lastrowid:
                    return cur.lastrowid
                cur.execute(
                    "SELECT id FROM bot_tokens WHERE token_hash=%s", (token_hash,)
                )
                row = cur.fetchone()
                return row["id"] if row else 0

    def update_bot_token(self, bot_id: int, **fields) -> bool:
        allowed = {"bot_name", "bot_username", "is_active", "verified"}
        updates = {k: v for k, v in fields.items() if k in allowed}
        if not updates:
            return False
        with self._conn() as conn:
            with conn.cursor() as cur:
                set_clause = ", ".join(f"{k}=%s" for k in updates)
                values = list(updates.values()) + [bot_id]
                cur.execute(
                    f"UPDATE bot_tokens SET {set_clause} WHERE id=%s", values
                )
                return cur.rowcount > 0

    def delete_bot_token(self, bot_id: int) -> bool:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute("DELETE FROM bot_tokens WHERE id=%s", (bot_id,))
                return cur.rowcount > 0

    # ── cTrader encrypted token vault ───────────────────────────────────────

    def list_ctrader_accounts(self) -> List[dict]:
        """Return metadata for all stored cTrader accounts (no encrypted payload)."""
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT id, account_id, user_id, broker_name, broker_title, "
                    "account_number, is_live, deposit_currency, balance_cents, "
                    "money_digits, leverage, account_type, account_status, "
                    "created_at, updated_at "
                    "FROM ctrader_accounts ORDER BY created_at DESC"
                )
                return list(cur.fetchall())

    def get_ctrader_account(self, account_id: int) -> Optional[dict]:
        """Return full row including encrypted_payload for a single account."""
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "SELECT * FROM ctrader_accounts WHERE account_id=%s", (account_id,)
                )
                return cur.fetchone()

    def save_ctrader_account(
        self,
        account_id: int,
        encrypted_payload: str,
        user_id: Optional[int] = None,
        broker_name: str = "",
        broker_title: str = "",
        account_number: Optional[int] = None,
        is_live: bool = False,
        deposit_currency: str = "",
        balance_cents: Optional[int] = None,
        money_digits: int = 2,
        leverage: Optional[int] = None,
        account_type: str = "",
        account_status: str = "",
    ) -> int:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "INSERT INTO ctrader_accounts "
                    "(account_id, user_id, broker_name, broker_title, account_number, "
                    "is_live, deposit_currency, encrypted_payload, balance_cents, "
                    "money_digits, leverage, account_type, account_status) "
                    "VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) "
                    "ON DUPLICATE KEY UPDATE "
                    "encrypted_payload=%s, broker_name=%s, broker_title=%s, "
                    "account_number=%s, is_live=%s, deposit_currency=%s, "
                    "balance_cents=%s, money_digits=%s, leverage=%s, "
                    "account_type=%s, account_status=%s",
                    (
                        account_id, user_id, broker_name, broker_title, account_number,
                        is_live, deposit_currency, encrypted_payload,
                        balance_cents, money_digits, leverage,
                        account_type, account_status,
                        encrypted_payload, broker_name, broker_title,
                        account_number, is_live, deposit_currency,
                        balance_cents, money_digits, leverage,
                        account_type, account_status,
                    ),
                )
                return cur.lastrowid or 0

    def update_ctrader_account_snapshot(
        self,
        account_id: int,
        *,
        broker_name: str = "",
        broker_title: str = "",
        account_number: Optional[int] = None,
        is_live: bool = False,
        deposit_currency: str = "",
        balance_cents: Optional[int] = None,
        money_digits: int = 2,
        leverage: Optional[int] = None,
        account_type: str = "",
        account_status: str = "",
    ) -> None:
        """Refresh display fields from a live Spotware fetch (tokens untouched)."""
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "UPDATE ctrader_accounts SET broker_name=%s, broker_title=%s, "
                    "account_number=%s, is_live=%s, deposit_currency=%s, "
                    "balance_cents=%s, money_digits=%s, leverage=%s, "
                    "account_type=%s, account_status=%s "
                    "WHERE account_id=%s",
                    (
                        broker_name, broker_title, account_number, is_live,
                        deposit_currency, balance_cents, money_digits, leverage,
                        account_type, account_status, account_id,
                    ),
                )

    def update_ctrader_account_payload(
        self, account_id: int, encrypted_payload: str
    ) -> bool:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    "UPDATE ctrader_accounts SET encrypted_payload=%s WHERE account_id=%s",
                    (encrypted_payload, account_id),
                )
                return cur.rowcount > 0

    def delete_ctrader_account(self, account_id: int) -> bool:
        with self._conn() as conn:
            with conn.cursor() as cur:
                cur.execute("DELETE FROM ctrader_accounts WHERE account_id=%s", (account_id,))
                return cur.rowcount > 0

    # ── Utility ─────────────────────────────────────────────────────────────

    def _dump_table(self, table: str) -> Optional[List[dict]]:
        """Dump all rows from a table. Returns None if table doesn't exist."""
        try:
            with self._conn() as conn:
                with conn.cursor() as cur:
                    cur.execute(f"SELECT * FROM {table}")
                    return list(cur.fetchall())
        except Exception:
            return None
