"""BTC/ETH 收益增强策略的单服务、多实例 Web 应用。"""

import argparse
import json
import logging
import os
import sys
import tempfile
import threading
import time as pytime

from flask import Flask, jsonify, make_response, request
from flask_sock import Sock


BASE_DIR = os.path.dirname(__file__)
ENV_FILE = os.path.join(BASE_DIR, ".env")


def _load_env_file():
    """先加载本地 .env，再导入依赖环境变量的策略模块。"""
    if not os.path.exists(ENV_FILE):
        return
    with open(ENV_FILE, "r", encoding="utf-8") as handle:
        for raw_line in handle:
            line = raw_line.strip()
            if line and not line.startswith("#") and "=" in line:
                key, value = line.split("=", 1)
                os.environ.setdefault(key.strip(), value.strip())


_load_env_file()

from deribit_api import DeribitClient
from strategy_engine import StrategyEngine


logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
    stream=sys.stdout,
    force=True,
)
logger = logging.getLogger(__name__)

DERIBIT_CLIENT_ID = os.environ.get("DERIBIT_ID", "")
DERIBIT_CLIENT_SECRET = os.environ.get("DERIBIT_SECRET", "")
USE_TESTNET = os.environ.get("DERIBIT_TESTNET", "1") == "1"
STATE_DIR = os.environ.get("STRAT_STATE_DIR", os.path.join(BASE_DIR, "runtime"))

SYMBOL_CONFIGS = {
    "BTC_USDC": {"instrument_name": "BTC_USDC", "index_name": "btc_usdc"},
    "ETH_USDC": {"instrument_name": "ETH_USDC", "index_name": "eth_usdc"},
}


def _configured_symbols():
    raw = os.environ.get("ENABLED_SYMBOLS", "BTC_USDC,ETH_USDC")
    symbols = [
        value.strip().upper()
        for value in raw.split(",")
        if value.strip().upper() in SYMBOL_CONFIGS
    ]
    return tuple(dict.fromkeys(symbols)) or ("BTC_USDC", "ETH_USDC")


ENABLED_SYMBOLS = _configured_symbols()
DEFAULT_SYMBOL = os.environ.get("DEFAULT_SYMBOL", "BTC_USDC").strip().upper()
if DEFAULT_SYMBOL not in ENABLED_SYMBOLS:
    DEFAULT_SYMBOL = ENABLED_SYMBOLS[0]

TOTAL_CAPITAL_USDC = float(os.environ.get("TOTAL_CAPITAL_USDC", "500"))
DEFAULT_ALLOCATION = TOTAL_CAPITAL_USDC / len(ENABLED_SYMBOLS)
SYMBOL_ALLOCATIONS = {
    symbol: float(
        os.environ.get(
            f"{symbol.split('_', 1)[0]}_ALLOCATION_USDC",
            str(DEFAULT_ALLOCATION),
        )
    )
    for symbol in ENABLED_SYMBOLS
}
if sum(SYMBOL_ALLOCATIONS.values()) > TOTAL_CAPITAL_USDC + 1e-9:
    raise RuntimeError(
        "BTC_ALLOCATION_USDC + ETH_ALLOCATION_USDC 不能超过 TOTAL_CAPITAL_USDC"
    )
TRADE_SIZE_USDC = float(os.environ.get("TRADE_SIZE_USDC", "25"))
if TOTAL_CAPITAL_USDC <= 0 or any(value <= 0 for value in SYMBOL_ALLOCATIONS.values()):
    raise RuntimeError("本金和每个标的额度必须大于 0")
if any(TRADE_SIZE_USDC > value for value in SYMBOL_ALLOCATIONS.values()):
    raise RuntimeError("TRADE_SIZE_USDC 不能超过任一标的额度")

app = Flask(__name__, static_url_path="/static", static_folder="static")
sock = Sock(app)

REQUIRE_CF_ACCESS = os.environ.get("REQUIRE_CF_ACCESS", "0") == "1"
CF_ACCESS_ALLOWED_EMAILS = {
    email.strip().lower()
    for email in os.environ.get("CF_ACCESS_ALLOWED_EMAILS", "").split(",")
    if email.strip()
}

engines: dict[str, StrategyEngine] = {}
shared_api_clients: dict[bool, DeribitClient] = {}
engine_lock = threading.RLock()
ws_clients = set()
ws_clients_lock = threading.Lock()


@app.before_request
def require_cloudflare_access_identity():
    """可选的源站纵深校验；实际登录层仍由 Cloudflare Access 提供。"""
    if not REQUIRE_CF_ACCESS:
        return None
    email = request.headers.get(
        "Cf-Access-Authenticated-User-Email", ""
    ).strip().lower()
    if not email or email not in CF_ACCESS_ALLOWED_EMAILS:
        return jsonify({"error": "Cloudflare Access authentication required"}), 403
    return None


def _save_env(client_id, client_secret, testnet):
    """只更新本地凭证相关键，原子替换 .env 并保留其他资金配置。"""
    replacements = {
        "DERIBIT_ID": client_id,
        "DERIBIT_SECRET": client_secret,
        "DERIBIT_TESTNET": "1" if testnet else "0",
    }
    existing = []
    if os.path.exists(ENV_FILE):
        with open(ENV_FILE, "r", encoding="utf-8") as handle:
            existing = handle.readlines()

    found = set()
    output = []
    for line in existing:
        key = line.split("=", 1)[0].strip() if "=" in line else ""
        if key in replacements:
            output.append(f"{key}={replacements[key]}\n")
            found.add(key)
        else:
            output.append(line)
    for key, value in replacements.items():
        if key not in found:
            output.append(f"{key}={value}\n")

    temp = tempfile.NamedTemporaryFile(
        mode="w",
        encoding="utf-8",
        dir=BASE_DIR,
        prefix=".env.",
        suffix=".tmp",
        delete=False,
    )
    try:
        temp.writelines(output)
        temp.flush()
        os.fsync(temp.fileno())
        temp.close()
        os.chmod(temp.name, 0o600)
        os.replace(temp.name, ENV_FILE)
    except Exception:
        try:
            os.unlink(temp.name)
        except OSError:
            pass
        raise


def _normalize_symbol(value):
    symbol = (value or DEFAULT_SYMBOL).strip().upper()
    if symbol not in ENABLED_SYMBOLS:
        raise ValueError(f"Unsupported symbol: {symbol}")
    return symbol


def _request_symbol(data=None):
    value = request.args.get("symbol")
    if not value and isinstance(data, dict):
        value = data.get("symbol")
    return _normalize_symbol(value)


def _symbol_error(exc):
    return jsonify({"success": False, "message": str(exc)}), 400


def _wait_for_engine_ready(target, timeout=15):
    """等待初始化真正完成，避免后台线程刚启动就保存空状态。"""
    deadline = pytime.time() + timeout
    while pytime.time() < deadline:
        if target.status in ("ready", "running"):
            return True
        if target.status == "error" or not target._running:
            return False
        pytime.sleep(0.25)
    return False


def _make_engine(symbol, testnet=None):
    use_testnet = USE_TESTNET if testnet is None else bool(testnet)
    config = {
        **SYMBOL_CONFIGS[symbol],
        "allocation_usdc": SYMBOL_ALLOCATIONS[symbol],
        "trade_size_usdc": TRADE_SIZE_USDC,
    }
    return StrategyEngine(
        DERIBIT_CLIENT_ID,
        DERIBIT_CLIENT_SECRET,
        config=config,
        testnet=use_testnet,
        state_callback=lambda state, selected=symbol: on_state_update(selected, state),
        state_dir=STATE_DIR,
        api_client=_get_shared_api_client(use_testnet),
    )


def _get_shared_api_client(testnet):
    """同一账户/网络的所有标的共享 REST token，避免重复认证和刷新竞争。"""
    key = bool(testnet)
    client = shared_api_clients.get(key)
    if client is None:
        client = DeribitClient(
            DERIBIT_CLIENT_ID,
            DERIBIT_CLIENT_SECRET,
            testnet=key,
        )
        shared_api_clients[key] = client
    return client


def _get_or_create_engine(symbol, testnet=None):
    target = engines.get(symbol)
    if target is None:
        target = _make_engine(symbol, testnet=testnet)
        engines[symbol] = target
    return target


def _empty_state(symbol):
    return {
        "status": "stopped",
        "symbol": symbol,
        "spot_currency": symbol.split("_", 1)[0],
        "trading_enabled": False,
        "api_connected": False,
        "config": {
            **SYMBOL_CONFIGS[symbol],
            "allocation_usdc": SYMBOL_ALLOCATIONS[symbol],
            "trade_size_usdc": TRADE_SIZE_USDC,
        },
    }


def _state_envelope(symbol, state):
    return {"type": "state", "symbol": symbol, "state": state}


def broadcast_json(payload):
    encoded = json.dumps(payload, ensure_ascii=False, default=str)
    with ws_clients_lock:
        clients = list(ws_clients)
    dead = []
    for client in clients:
        try:
            client.send(encoded)
        except Exception:
            dead.append(client)
    if dead:
        with ws_clients_lock:
            for client in dead:
                ws_clients.discard(client)


def on_state_update(symbol, state):
    broadcast_json(_state_envelope(symbol, state))


@sock.route("/ws")
def ws_handler(ws):
    with ws_clients_lock:
        ws_clients.add(ws)
    try:
        ws.send(
            json.dumps(
                {
                    "type": "symbols",
                    "symbols": list(ENABLED_SYMBOLS),
                    "default_symbol": DEFAULT_SYMBOL,
                }
            )
        )
        with engine_lock:
            snapshots = {
                symbol: engines[symbol].get_state()
                if symbol in engines
                else _empty_state(symbol)
                for symbol in ENABLED_SYMBOLS
            }
        for symbol, state in snapshots.items():
            ws.send(
                json.dumps(
                    _state_envelope(symbol, state),
                    ensure_ascii=False,
                    default=str,
                )
            )
        while ws.receive() is not None:
            pass
    except Exception:
        pass
    finally:
        with ws_clients_lock:
            ws_clients.discard(ws)


@app.route("/")
def index():
    html_path = os.path.join(BASE_DIR, "static", "dashboard.html")
    with open(html_path, "r", encoding="utf-8") as handle:
        response = make_response(handle.read())
    response.headers["Content-Type"] = "text/html; charset=utf-8"
    response.headers["Cache-Control"] = "no-store, no-cache, must-revalidate, max-age=0"
    return response


@app.route("/api/symbols")
def api_symbols():
    with engine_lock:
        items = []
        for symbol in ENABLED_SYMBOLS:
            target = engines.get(symbol)
            items.append(
                {
                    "symbol": symbol,
                    "currency": symbol.split("_", 1)[0],
                    "allocation_usdc": SYMBOL_ALLOCATIONS[symbol],
                    "status": target.status if target else "stopped",
                    "trading_enabled": bool(target and target._trading_enabled),
                }
            )
    return jsonify(
        {
            "symbols": items,
            "default_symbol": DEFAULT_SYMBOL,
            "total_capital_usdc": TOTAL_CAPITAL_USDC,
            "trade_size_usdc": TRADE_SIZE_USDC,
        }
    )


@app.route("/api/status")
def api_status():
    try:
        symbol = _request_symbol()
    except ValueError as exc:
        return _symbol_error(exc)
    with engine_lock:
        target = engines.get(symbol)
        state = target.get_state() if target else _empty_state(symbol)
    return jsonify(state)


@app.route("/api/init", methods=["POST"])
def api_init():
    if not DERIBIT_CLIENT_ID or not DERIBIT_CLIENT_SECRET:
        return jsonify({"success": False, "message": "请先配置 Deribit API 凭证"}), 400
    data = request.get_json(silent=True) or {}
    try:
        symbol = _request_symbol(data)
    except ValueError as exc:
        return _symbol_error(exc)
    with engine_lock:
        target = _get_or_create_engine(symbol, testnet=data.get("testnet"))
        if target._running:
            return jsonify(
                {
                    "success": True,
                    "message": "Engine already running",
                    "symbol": symbol,
                    "status": target.status,
                }
            )
        if not target.initialize():
            return jsonify({"success": False, "message": "Initialization failed"}), 500
    if not _wait_for_engine_ready(target):
        return jsonify(
            {
                "success": False,
                "message": "Initialization failed",
                "symbol": symbol,
                "status": target.status,
            }
        ), 502
    return jsonify(
        {
            "success": True,
            "message": "Engine initialized",
            "symbol": symbol,
            "status": target.status,
        }
    )


@app.route("/api/start", methods=["POST"])
def api_start():
    if not DERIBIT_CLIENT_ID or not DERIBIT_CLIENT_SECRET:
        return jsonify({"success": False, "message": "请先配置 Deribit API 凭证"}), 400
    data = request.get_json(silent=True) or {}
    try:
        symbol = _request_symbol(data)
    except ValueError as exc:
        return _symbol_error(exc)

    with engine_lock:
        target = _get_or_create_engine(symbol)
        needs_init = not target._running
        if needs_init and not target.initialize():
            return jsonify({"success": False, "message": "Initialization failed"}), 500
    if needs_init and not _wait_for_engine_ready(target):
        return jsonify({"success": False, "message": "Initialization failed"}), 502
    with engine_lock:
        if target._trading_enabled:
            return jsonify({"success": False, "message": "Already trading"}), 409
        success = target.start()
    return jsonify(
        {
            "success": success,
            "symbol": symbol,
            "message": "Trading started" if success else "Failed to start trading",
        }
    )


@app.route("/api/stop", methods=["POST"])
def api_stop():
    data = request.get_json(silent=True) or {}
    try:
        symbol = _request_symbol(data)
    except ValueError as exc:
        return _symbol_error(exc)
    with engine_lock:
        target = engines.get(symbol)
    if target is None:
        return jsonify({"success": False, "message": "No strategy running"}), 400
    target.stop()
    return jsonify({"success": True, "symbol": symbol, "message": "Strategy stopped"})


@app.route("/api/credentials", methods=["GET", "POST"])
def api_credentials():
    global DERIBIT_CLIENT_ID, DERIBIT_CLIENT_SECRET, USE_TESTNET

    if request.method == "GET":
        masked = (
            DERIBIT_CLIENT_ID[:4] + "****"
            if len(DERIBIT_CLIENT_ID) > 4
            else "****"
        )
        return jsonify({"client_id_masked": masked, "testnet": USE_TESTNET})

    data = request.get_json(silent=True) or {}
    new_id = str(data.get("client_id", "")).strip()
    new_secret = str(data.get("client_secret", "")).strip()
    new_testnet = bool(data.get("testnet", True))
    if not new_id or not new_secret:
        return jsonify({"success": False, "message": "ID 和 Secret 不能为空"}), 400
    try:
        symbol = _request_symbol(data)
    except ValueError as exc:
        return _symbol_error(exc)

    with engine_lock:
        old_engines = list(engines.values())
        engines.clear()
        shared_api_clients.clear()
    for target in old_engines:
        target.stop()

    DERIBIT_CLIENT_ID = new_id
    DERIBIT_CLIENT_SECRET = new_secret
    USE_TESTNET = new_testnet
    try:
        _save_env(new_id, new_secret, new_testnet)
    except Exception as exc:
        logger.exception("Failed to update .env")
        return jsonify({"success": False, "message": f".env 保存失败: {exc}"}), 500

    with engine_lock:
        target = _get_or_create_engine(symbol)
        if not target.initialize():
            return jsonify({"success": False, "message": "新凭证连接失败"}), 500
    if not _wait_for_engine_ready(target):
        return jsonify({"success": False, "message": "新凭证连接失败"}), 502
    return jsonify({"success": True, "message": "凭证已保存，当前标的已重新连接"})


@app.route("/api/params", methods=["GET", "POST"])
def api_params():
    data = request.get_json(silent=True) or {} if request.method == "POST" else {}
    try:
        symbol = _request_symbol(data)
    except ValueError as exc:
        return _symbol_error(exc)

    with engine_lock:
        target = engines.get(symbol)
        if target is None:
            target = _get_or_create_engine(symbol)

        if request.method == "GET":
            cfg = target.cfg
            return jsonify(
                {
                    "editable": True,
                    "symbol": symbol,
                    "anchor_price": target.anchor_price,
                    "allocation_usdc": cfg["allocation_usdc"],
                    "trade_size_usdc": cfg["trade_size_usdc"],
                    "rv_min": cfg["rv_min"],
                    "rv_max": cfg["rv_max"],
                    "rv_update_interval_minutes": cfg.get(
                        "rv_update_interval_minutes", 15
                    ),
                    "poll_interval": cfg["poll_interval"],
                    "cooldown_seconds": cfg.get("cooldown_seconds", 180),
                    "min_poll_balance_usdc": cfg["min_poll_balance_usdc"],
                }
            )

        if not data:
            return jsonify({"success": False, "message": "No data"}), 400

        changed = []
        int_ranges = {
            "poll_interval": (5, 300),
            "cooldown_seconds": (10, 600),
            "rv_update_interval_minutes": (5, 1440),
        }
        float_ranges = {
            "trade_size_usdc": (10, target.cfg["allocation_usdc"]),
            "rv_min": (0.0001, 0.05),
            "rv_max": (0.001, 0.1),
            "min_poll_balance_usdc": (0, target.cfg["allocation_usdc"]),
        }
        try:
            for key, (minimum, maximum) in int_ranges.items():
                if key in data:
                    value = int(data[key])
                    if not minimum <= value <= maximum:
                        raise ValueError(f"{key} 超出范围")
                    target.cfg[key] = value
                    changed.append(f"{key}={value}")
            for key, (minimum, maximum) in float_ranges.items():
                if key in data:
                    value = float(data[key])
                    if not minimum <= value <= maximum:
                        raise ValueError(f"{key} 超出范围")
                    target.cfg[key] = value
                    changed.append(f"{key}={value}")
            if target.cfg["rv_min"] > target.cfg["rv_max"]:
                raise ValueError("rv_min 不能大于 rv_max")
            if "anchor_price" in data:
                value = float(data["anchor_price"])
                if value <= 0:
                    raise ValueError("anchor_price 必须大于 0")
                target.anchor_price = value
                target._recalc_thresholds()
                changed.append(f"anchor=${value:.2f}")
        except (TypeError, ValueError) as exc:
            return jsonify({"success": False, "message": str(exc)}), 400

        if changed:
            target._save_state()
            on_state_update(symbol, target.get_state())
    return jsonify({"success": True, "symbol": symbol, "changed": changed})


@app.route("/api/kline")
def api_kline():
    try:
        symbol = _request_symbol()
    except ValueError as exc:
        return _symbol_error(exc)
    try:
        import requests as requests_module

        end = int(pytime.time() * 1000)
        payload = {
            "jsonrpc": "2.0",
            "id": 1,
            "method": "public/get_tradingview_chart_data",
            "params": {
                "instrument_name": symbol,
                "start_timestamp": end - 7 * 86400 * 1000,
                "end_timestamp": end,
                "resolution": "5",
            },
        }
        response = requests_module.post(
            "https://www.deribit.com/api/v2/", json=payload, timeout=15
        )
        response.raise_for_status()
        result = response.json().get("result")
        return jsonify(result or {"error": "no data"})
    except Exception as exc:
        return jsonify({"error": str(exc)}), 500


@app.route("/api/test-connection")
def api_test_connection():
    if not DERIBIT_CLIENT_ID or not DERIBIT_CLIENT_SECRET:
        return jsonify({"connected": False, "auth_error": "请先配置 Deribit API 凭证"}), 400
    try:
        symbol = _request_symbol()
    except ValueError as exc:
        return _symbol_error(exc)

    index_name = SYMBOL_CONFIGS[symbol]["index_name"]
    spot_currency = symbol.split("_", 1)[0]
    results = {}
    for label, testnet in (("mainnet", False), ("testnet", True)):
        client = DeribitClient(
            DERIBIT_CLIENT_ID, DERIBIT_CLIENT_SECRET, testnet=testnet
        )
        info = client.check_connection()
        if info["connected"]:
            info["btc_index_price"] = client.get_index_price(index_name)
            info["spot_currency"] = spot_currency
            try:
                usdc = client.get_account_summary(currency="USDC")
                if usdc:
                    info["usdc_balance"] = usdc.get("balance", 0)
                spot = client.get_account_summary(currency=spot_currency)
                if spot:
                    info["btc_balance"] = spot.get("balance", 0)
            except Exception:
                pass
        results[label] = info
    return jsonify(results)


@app.route("/api/config")
def api_config():
    """兼容旧接口：配置只读，标的不可在运行实例内切换。"""
    try:
        symbol = _request_symbol()
    except ValueError as exc:
        return _symbol_error(exc)
    with engine_lock:
        target = engines.get(symbol)
        config = target.cfg if target else _empty_state(symbol)["config"]
    return jsonify(config)


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description="BTC/ETH 收益增强策略")
    parser.add_argument("--port", type=int, default=5050, help="Web 端口")
    args = parser.parse_args()

    print("=" * 60)
    print("  BTC/ETH 收益增强策略 - 单服务多实例 Dashboard")
    print(f"  http://localhost:{args.port}")
    print(
        f"  总本金 ${TOTAL_CAPITAL_USDC:.2f} · "
        + " · ".join(
            f"{symbol.split('_', 1)[0]} ${SYMBOL_ALLOCATIONS[symbol]:.2f}"
            for symbol in ENABLED_SYMBOLS
        )
    )
    print(f"  默认单笔 ${TRADE_SIZE_USDC:.2f}")
    print("=" * 60)
    app.run(
        host=os.environ.get("APP_HOST", "127.0.0.1"),
        port=args.port,
        debug=False,
    )
