#!/usr/bin/env python3
"""External token refresh cron for the cTrader vault.

Fetches encrypted tokens from the forwarder, decrypts with
``CTRADER_TOKEN_SECRET``, calls the cTrader OAuth refresh endpoint,
re-encrypts with the same secret, and pushes the refreshed payload back to
the forwarder.

Usage::

    # Manual single run
    python scripts/ctrader_refresh.py

    # Dry-run (show what would be refreshed)
    python scripts/ctrader_refresh.py --dry-run

    # Refresh only tokens expiring within 2 hours
    python scripts/ctrader_refresh.py --threshold 7200

    # Scheduled cron (every 30 min via crontab):
    # */30 * * * * cd /path/to/project && python scripts/ctrader_refresh.py >> refresh.log 2>&1

Environment (from .env or env vars):
    CTRADER_API_BASE       — Forwarder URL (default https://quill.nx.kg/f)
    WEB_PIN                — Dashboard PIN
    CTRADER_TOKEN_SECRET   — Shared secret used to encrypt/decrypt tokens
    CTRADER_CLIENT_ID      — cTrader App Client ID
    CTRADER_CLIENT_SECRET  — cTrader App Client Secret
"""

from __future__ import annotations

import argparse
import json
import time
import urllib.error
import urllib.parse
import urllib.request
from pathlib import Path
from typing import Any

from common import load_env, ForwarderClient


def _load_secret(env: dict[str, str]) -> str:
    """Return the CTRADER_TOKEN_SECRET from env or raise."""
    secret = env.get("CTRADER_TOKEN_SECRET", "")
    if not secret:
        raise SystemExit(
            "CTRADER_TOKEN_SECRET not found in .env. "
            "Generate one with: python dev_generate_keys.py"
        )
    return secret


def _do_refresh(
    client_id: str, client_secret: str, refresh_token: str
) -> dict[str, Any] | None:
    """Call cTrader OAuth refresh endpoint. Returns dict or None on failure."""
    data = urllib.parse.urlencode({
        "grant_type": "refresh_token",
        "refresh_token": refresh_token,
        "client_id": client_id,
        "client_secret": client_secret,
    }).encode("utf-8")

    req = urllib.request.Request(
        "https://openapi.ctrader.com/apps/token",
        data=data,
        headers={
            "Content-Type": "application/x-www-form-urlencoded",
            "Accept": "application/json",
        },
        method="POST",
    )

    try:
        with urllib.request.urlopen(req, timeout=30) as resp:
            body = json.loads(resp.read().decode("utf-8"))
    except urllib.error.HTTPError as exc:
        body = exc.read().decode("utf-8", errors="ignore")
        print(f"    cTrader refresh HTTP {exc.code}: {body}")
        return None

    access = body.get("accessToken") or body.get("access_token")
    if not access:
        print(f"    No access_token in response: {body}")
        return None

    refresh = body.get("refreshToken") or body.get("refresh_token", refresh_token)
    expires_in = body.get("expiresIn") or body.get("expires_in", 2_628_000)
    return {
        "access_token": access,
        "refresh_token": refresh,
        "expires_at": time.time() + expires_in,
    }


def main(argv: list[str] | None = None) -> int:
    parser = argparse.ArgumentParser(description="Refresh cTrader OAuth tokens externally")
    parser.add_argument("--env", default=".env", help="Path to .env file")
    parser.add_argument(
        "--threshold",
        type=int,
        default=3600,
        help="Refresh tokens expiring within N seconds (default 3600 = 1h)",
    )
    parser.add_argument("--dry-run", action="store_true", help="Do not push changes")
    args = parser.parse_args(argv)

    env = load_env(Path(args.env))
    base_url = env.get("CTRADER_API_BASE", "https://quill.nx.kg/f")
    pin = env.get("WEB_PIN", "")
    client_id = env.get("CTRADER_CLIENT_ID", "")
    client_secret = env.get("CTRADER_CLIENT_SECRET", "")

    if not pin:
        raise SystemExit("WEB_PIN not found in .env")
    if not client_id or not client_secret:
        raise SystemExit("CTRADER_CLIENT_ID and CTRADER_CLIENT_SECRET required in .env")

    secret = _load_secret(env)

    from crypto_vault import decrypt_payload, encrypt_payload

    print(f"Connecting to {base_url} …")
    client = ForwarderClient(base_url, pin)
    accounts = client.list_accounts()
    print(f"Fetched {len(accounts)} account(s)\n")

    now = time.time()
    refreshed = 0
    skipped = 0
    failed = 0

    for acct in accounts:
        account_id = acct.get("account_id")
        enc = acct.get("encrypted_payload")
        broker = acct.get("broker_title") or acct.get("broker_name") or "Unknown"

        if not enc or not account_id:
            print(f"  [!] {broker} (ID {account_id}) — missing payload")
            failed += 1
            continue

        # Decrypt
        try:
            payload = decrypt_payload(secret, enc)
        except Exception as exc:
            print(f"  [X] {broker} (ID {account_id}) — decrypt failed: {exc}")
            failed += 1
            continue

        expires_at = payload.get("expires_at", 0)
        remaining = expires_at - now

        print(f"  {broker} (ID {account_id}) — expires in {remaining/60:.0f} min")

        if remaining > args.threshold:
            print(f"    → skip (threshold {args.threshold}s)")
            skipped += 1
            continue

        refresh_token = payload.get("refresh_token", "")
        if not refresh_token:
            print("    → no refresh_token available; cannot refresh")
            failed += 1
            continue

        # Call cTrader refresh
        new_tokens = _do_refresh(client_id, client_secret, refresh_token)
        if not new_tokens:
            print("    → cTrader refresh failed")
            failed += 1
            continue

        # Preserve any extra metadata from the original payload
        new_payload = {**payload, **new_tokens}

        # Re-encrypt with the same secret
        try:
            new_enc = encrypt_payload(secret, new_payload)
        except Exception as exc:
            print(f"    → re-encrypt failed: {exc}")
            failed += 1
            continue

        if args.dry_run:
            print(
                f"    → [DRY-RUN] new access={new_tokens['access_token'][:24]}… "
                f"expires_in {new_tokens['expires_at'] - now:.0f}s"
            )
            refreshed += 1
            continue

        # Push back
        try:
            res = client.push_refresh(account_id, new_enc)
            if res.get("status") == "ok":
                print(
                    f"    → refreshed OK (new access={new_tokens['access_token'][:24]}… "
                    f"expires in {(new_tokens['expires_at'] - now)/60:.0f} min)"
                )
                refreshed += 1
            else:
                print(f"    → push back failed: {res}")
                failed += 1
        except Exception as exc:
            print(f"    → push back error: {exc}")
            failed += 1

    print(f"\nDone: {refreshed} refreshed, {skipped} skipped, {failed} failed")
    return 1 if failed else 0


if __name__ == "__main__":
    raise SystemExit(main())
