"""
anti_cluster_loss.py — V2 (Mattia 04/08/2026)
==============================================
Modulo anti-cluster-loss: pausa simbolo per 4h dopo 2 SL consecutivi stesso symbol.

Logica:
- Mantiene stato persistente: per ogni symbol conta gli SL consecutivi recenti
- Quando count >= 2: marca symbol "paused_until" = now + 4h
- Tutte le nuove entries su quel symbol sono bloccate fino a scadenza pausa
- Stato persistito in anti_cluster_state.json (atomic write)

Usage:
    from anti_cluster_loss import should_pause, record_sl, record_win
    if should_pause(sym):
        log(f"ANTI-CLUSTER: {sym} paused after 2 SL, skip")
        return
    # ... apri posizione ...
    if closed_in_sl:
        record_sl(sym)
    else:
        record_win(sym)
"""
import json
import os
import time
from datetime import datetime, timezone
from pathlib import Path

STATE_FILE = Path("/opt/charter-live/live_deploy_v2/logs/anti_cluster_state.json")
LOG_FILE = Path("/opt/charter-live/live_deploy_v2/logs/anti_cluster_loss.log")

MAX_CONSECUTIVE_SL = 2          # 2 SL consecutivi = pausa
PAUSE_DURATION_SEC = 4 * 3600   # 4 ore di pausa
SL_LOOKBACK_SEC = 24 * 3600     # considera solo SL nelle ultime 24h


def _log(msg):
    ts = datetime.now(timezone.utc).isoformat()
    line = f"[{ts}] {msg}"
    print(line, flush=True)
    try:
        os.makedirs(LOG_FILE.parent, exist_ok=True)
        with open(LOG_FILE, "a", encoding="utf-8") as f:
            f.write(line + "\n")
    except Exception:
        pass


def _load_state():
    try:
        if STATE_FILE.exists():
            with open(STATE_FILE, "r", encoding="utf-8") as f:
                return json.load(f)
    except Exception as e:
        _log(f"load_state err: {e}")
    return {"symbols": {}}


def _save_state(state):
    try:
        os.makedirs(STATE_FILE.parent, exist_ok=True)
        tmp = STATE_FILE.with_suffix(".tmp")
        with open(tmp, "w", encoding="utf-8") as f:
            json.dump(state, f, indent=2)
        tmp.replace(STATE_FILE)
    except Exception as e:
        _log(f"save_state err: {e}")


def _prune_old_sl(symbol_state, now):
    """Rimuovi SL piu' vecchi di SL_LOOKBACK_SEC. Resetta counter a 0 se troppi vecchi."""
    sl_times = symbol_state.get("sl_times", [])
    cutoff = now - SL_LOOKBACK_SEC
    fresh_sl_times = [t for t in sl_times if t > cutoff]
    symbol_state["sl_times"] = fresh_sl_times


def should_pause(symbol: str) -> bool:
    """Ritorna True se il symbol e' in pausa (2+ SL consecutivi nelle ultime 24h)."""
    state = _load_state()
    sym_state = state.get("symbols", {}).get(symbol)
    if not sym_state:
        return False
    paused_until = sym_state.get("paused_until", 0)
    return paused_until > time.time()


def record_sl(symbol: str):
    """Registra un SL su symbol. Se raggiunge MAX_CONSECUTIVE_SL, marca pausa."""
    state = _load_state()
    now = time.time()
    sym_state = state.setdefault("symbols", {}).setdefault(symbol, {"sl_times": [], "paused_until": 0})
    _prune_old_sl(sym_state, now)
    sym_state["sl_times"].append(now)
    sym_state["last_sl_ts"] = now
    # Check pause trigger
    if len(sym_state["sl_times"]) >= MAX_CONSECUTIVE_SL:
        sym_state["paused_until"] = now + PAUSE_DURATION_SEC
        _log(f"PAUSE TRIGGERED: {symbol} 2 SL in <24h, paused until {datetime.fromtimestamp(sym_state['paused_until'], tz=timezone.utc).isoformat()}")
    _save_state(state)


def record_win(symbol: str):
    """Registra un WIN/ TP su symbol. Resetta counter SL consecutivi e sblocca eventuale pausa."""
    state = _load_state()
    sym_state = state.get("symbols", {}).get(symbol)
    if sym_state:
        sym_state["sl_times"] = []  # reset counter
        sym_state["paused_until"] = 0  # FIX: WIN resetta anche la pausa
        sym_state["last_win_ts"] = time.time()
        _save_state(state)


def get_state():
    return _load_state()


if __name__ == "__main__":
    import sys
    if len(sys.argv) > 1:
        cmd = sys.argv[1]
        if cmd == "pause":
            sym = sys.argv[2] if len(sys.argv) > 2 else "BTCUSDT"
            record_sl(sym)
        elif cmd == "win":
            sym = sys.argv[2] if len(sys.argv) > 2 else "BTCUSDT"
            record_win(sym)
    print(json.dumps(get_state(), indent=2, ensure_ascii=False))
