#!/usr/bin/env python3
"""Shared helpers for all deployment and maintenance scripts.

This module consolidates the common patterns used across the scripts/
directory:

  - **load_env**: parse a ``.env`` file into a dict (no external deps).
  - **ssh helpers**: ``ssh_run``, ``ssh_upload``, ``ssh_check`` wrap
    ``subprocess.run`` with consistent BatchMode + capture behaviour.
  - **RemoteDB**: query the cPanel MariaDB over SSH without leaking
    credentials on the command line.
  - **ForwarderClient**: tiny HTTP client with cookie/CSRF session
    management, shared by ``ctrader_refresh.py`` and
    ``ctrader_consumer.py``.

Usage from a sibling script::

    from common import load_env, ssh_run, RemoteDB, ForwarderClient
"""

from __future__ import annotations

import json
import subprocess
import sys
import urllib.error
import urllib.request
from http.cookiejar import CookieJar
from pathlib import Path
from typing import Any, Optional


# ── .env parsing ────────────────────────────────────────────────────────────


def load_env(path: Path) -> dict[str, str]:
    """Parse a ``.env`` file into a flat ``{key: value}`` dict.

    Skips blank lines, comments, and lines without ``=``.  Strips
    surrounding quotes from values.
    """
    vals: dict[str, str] = {}
    if not path.exists():
        return vals
    for line in path.read_text("utf-8").splitlines():
        line = line.strip()
        if not line or line.startswith("#") or "=" not in line:
            continue
        k, v = line.split("=", 1)
        vals[k.strip()] = v.strip().strip('"').strip("'")
    return vals


# ── SSH helpers ─────────────────────────────────────────────────────────────

_SSH_FLAGS = ["-q", "-o", "BatchMode=yes"]


def ssh_run(
    host: str,
    *args: str,
    input: Optional[bytes] = None,
) -> subprocess.CompletedProcess:
    """Run a command over SSH and return the ``CompletedProcess``."""
    return subprocess.run(
        ["ssh", *_SSH_FLAGS, host, *args],
        input=input,
        capture_output=True,
    )


def ssh_upload(host: str, remote_path: str, data: bytes, mode: str = "600") -> subprocess.CompletedProcess:
    """Upload bytes to a remote path via SSH stdin (avoids temp files)."""
    return ssh_run(host, f"cat > {remote_path} && chmod {mode} {remote_path}", input=data)


def ssh_check(host: str, timeout: int = 8) -> None:
    """Verify SSH connectivity or raise ``SystemExit``."""
    res = subprocess.run(
        ["ssh", *_SSH_FLAGS, "-o", f"ConnectTimeout={timeout}", host, "true"],
        capture_output=True,
    )
    if res.returncode != 0:
        raise SystemExit(f"Cannot SSH to {host}")


# ── RemoteDB: query cPanel MariaDB over SSH ────────────────────────────────


class RemoteDB:
    """Query the remote cPanel MySQL database over SSH.

    Credentials are passed via a temporary ``.my.cnf`` file uploaded to
    ``/tmp``, so they never appear in shell arguments or process lists.
    """

    def __init__(self, ssh_host: str, env: dict[str, str]):
        self.ssh_host = ssh_host
        self.env = env
        self._cnf_uploaded = False

    def _ensure_cnf(self) -> None:
        if self._cnf_uploaded:
            return
        for key in ("MYSQL_HOST", "MYSQL_USER", "MYSQL_PASSWORD", "MYSQL_DB"):
            if not self.env.get(key):
                raise SystemExit(f"Missing {key} in .env")
        cnf = (
            "[client]\n"
            f"host={self.env['MYSQL_HOST']}\n"
            f"user={self.env['MYSQL_USER']}\n"
            f"password={self.env['MYSQL_PASSWORD']}\n"
            f"database={self.env['MYSQL_DB']}\n"
        ).encode()
        r = ssh_upload(self.ssh_host, "/tmp/remote_db.cnf", cnf)
        if r.returncode != 0:
            print(r.stderr.decode("utf-8", errors="ignore"), file=sys.stderr)
            sys.exit(1)
        self._cnf_uploaded = True

    def query(self, sql: str, table: bool = True) -> str:
        """Run a SQL query and return the output as a string."""
        self._ensure_cnf()
        r_upload = ssh_upload(self.ssh_host, "/tmp/remote_db_query.sql", sql.encode("utf-8"))
        if r_upload.returncode != 0:
            print(r_upload.stderr.decode("utf-8", errors="ignore"), file=sys.stderr)
            sys.exit(1)

        args = "mysql --defaults-file=/tmp/remote_db.cnf"
        if table:
            args += " --table"
        args += " < /tmp/remote_db_query.sql"

        r = ssh_run(self.ssh_host, args)
        stdout = r.stdout.decode("utf-8", errors="ignore")
        stderr = r.stderr.decode("utf-8", errors="ignore").strip()
        if stderr:
            print(stderr, file=sys.stderr)
        if r.returncode != 0:
            sys.exit(1)
        return stdout

    def cleanup(self) -> None:
        if self._cnf_uploaded:
            ssh_run(self.ssh_host, "rm -f /tmp/remote_db.cnf /tmp/remote_db_query.sql")
            self._cnf_uploaded = False


# ── ForwarderClient: HTTP client with session/CSRF ─────────────────────────


class ForwarderClient:
    """Minimal HTTP client for the forwarder API.

    Manages Flask session cookies and CSRF tokens automatically.  Used by
    ``ctrader_refresh.py`` and ``ctrader_consumer.py``.
    """

    def __init__(self, base_url: str, pin: str, timeout: int = 30):
        self.base = base_url.rstrip("/")
        self.pin = pin
        self.timeout = timeout
        self._cookie_jar = CookieJar()
        self._csrf: str | None = None
        self._login()

    def _request(
        self, path: str, method: str = "GET", data: dict[str, Any] | None = None
    ) -> dict:
        url = f"{self.base}{path}"
        headers: dict[str, str] = {}
        body: bytes | None = None
        if data is not None:
            body = json.dumps(data).encode("utf-8")
            headers["Content-Type"] = "application/json"
        if self._csrf and method in ("POST", "PATCH", "DELETE"):
            headers["X-CSRF-Token"] = self._csrf

        req = urllib.request.Request(url, data=body, headers=headers, method=method)
        opener = urllib.request.build_opener(
            urllib.request.HTTPCookieProcessor(self._cookie_jar)
        )
        with opener.open(req, timeout=self.timeout) as resp:
            return json.loads(resp.read().decode("utf-8"))

    def _login(self) -> None:
        status = self._request("/api/auth/status")
        if not status.get("pin_required"):
            self._csrf = status.get("csrf_token")
            return
        login = self._request("/api/auth/login", "POST", {"pin": self.pin})
        if login.get("error"):
            raise SystemExit(f"PIN login failed: {login['error']}")
        self._csrf = login.get("csrf_token")
        if not self._csrf:
            raise SystemExit("Login succeeded but no CSRF token returned")

    def list_accounts(self, encrypted: bool = True) -> list[dict]:
        flag = "?encrypted=true" if encrypted else ""
        return self._request(f"/api/ctrader/accounts{flag}").get("accounts", [])

    def push_refresh(self, account_id: int, encrypted_payload: str) -> dict:
        return self._request(
            f"/api/ctrader/accounts/{account_id}/refresh",
            "POST",
            {"encrypted_payload": encrypted_payload},
        )
