#!/usr/bin/env python3
"""
Strategia del Quadrato — Implementazione Python per bot Bybit V5.

Logica:
- Ogni sessione (9:00 Europe/London) calcola box_top e box_bottom dal giorno prima
- Aspetta che il prezzo testi il top o il bottom
- Aspetta il retest
- Cerca pattern Doji Hammer con conferma volume
- Entra nella direzione opposta al breakout (mean reversion)
- Gestisce stop loss, take profit, trailing stop

Uso:
    from square_strategy import SquareStrategy
    from bybit_demo_client import BybitDemoClient

    bot = SquareStrategy(
        bybit_client=bybit,
        symbol="ETHUSDT",
        timeframe="15",
        on_signal=on_strategy_signal  # callback per inviare a Bybit
    )
    bot.run_forever()
"""

import time
import logging
from datetime import datetime, timezone, timedelta
from dataclasses import dataclass
from typing import Callable, Optional

log = logging.getLogger("square_strategy")


@dataclass
class Box:
    """Rettangolo del giorno precedente."""
    top: float
    bottom: float
    mid: float
    height: float
    no_trade_top: float
    no_trade_bot: float
    date: str  # ISO date della sessione a cui si applica

    @property
    def is_in_top_zone(self, price: float) -> bool:
        return price > self.mid

    @property
    def is_in_bottom_zone(self, price: float) -> bool:
        return price < self.mid

    @property
    def is_in_no_trade_zone(self, price: float) -> bool:
        return self.no_trade_bot <= price <= self.no_trade_top


@dataclass
class Candle:
    """Candela OHLCV normalizzata."""
    timestamp: int       # unix seconds
    open: float
    high: float
    low: float
    close: float
    volume: float

    @property
    def body(self) -> float:
        return abs(self.close - self.open)

    @property
    def range(self) -> float:
        return self.high - self.low

    @property
    def is_red(self) -> bool:
        return self.close < self.open

    @property
    def is_green(self) -> bool:
        return self.close > self.open

    @property
    def upper_wick(self) -> float:
        return self.high - max(self.close, self.open)

    @property
    def lower_wick(self) -> float:
        return min(self.close, self.open) - self.low


class SquareStrategy:
    # Parametri (con defaults dalla spec)
    NO_TRADE_ZONE_PCT = 0.15       # ±15% del box height = zona centrale
    DOJI_BODY_RATIO = 0.30
    DOJI_WICK_RATIO = 0.40
    VOLUME_MULT = 1.5
    STOP_BUFFER_PCT = 0.2
    MIN_MINUTES_AFTER_OPEN = 30    # SPEC §7.3: niente entry primi 30min dopo 9:00 London
    MIN_ATR_PCT = 0.003            # SPEC §7.1: ATR/price >= 0.3%
    ATR_PERIOD = 14                # SPEC §7.1: ATR(14)
    PREV_WICK_RATIO = 0.30         # SPEC §3.2: candela precedente deve avere ombra >= 30% del range nella direzione del rifiuto

    # Stati della state machine
    STATE_IDLE = 0
    STATE_FIRST_TEST_TOP = 1
    STATE_FIRST_TEST_BOTTOM = 2
    STATE_RETEST_TOP = 3
    STATE_RETEST_BOTTOM = 4
    STATE_IN_TRADE = 5

    def __init__(
        self,
        bybit_client,
        symbol: str,
        timeframe: str = "15",
        risk_per_trade_pct: float = 1.0,
        leverage: int = 3,
        on_signal: Optional[Callable] = None,
    ):
        self.bybit = bybit_client
        self.symbol = symbol
        self.timeframe = timeframe
        self.risk_pct = risk_per_trade_pct / 100.0
        self.leverage = leverage
        self.on_signal = on_signal or self._default_signal_handler

        # Stato runtime
        self.box: Optional[Box] = None
        self.state = self.STATE_IDLE
        self.candles: list[Candle] = []
        self.last_test_bar_idx: Optional[int] = None
        self.in_position: bool = False
        self.position_side: Optional[str] = None  # "long" | "short"
        self.entry_price: Optional[float] = None
        self.stop_price: Optional[float] = None
        self.tp_price: Optional[float] = None
        # SPEC §7.3: time filter 30min dopo session open (9:00 London)
        # Settato automaticamente quando update_box() viene chiamato la prima volta.
        self.session_open_time: Optional[int] = None  # unix seconds, 9:00 London del giorno corrente
        self._last_box_date: Optional[str] = None

    # =========================================================================
    # BOX CALCULATION
    # =========================================================================
    def update_box(self, yesterday_high: float, yesterday_low: float, date: str,
                   session_open_unix: Optional[int] = None):
        """Aggiorna il rettangolo del giorno precedente.
        SPEC §1: range = high/low sessione precedente (9:00 London -> 9:00 London).
        SPEC §7.3: time filter parte da session_open_unix (9:00 London del giorno corrente).
        Se session_open_unix non e' passato, prova a calcolarlo dalla data (UTC midnight, semplificato).
        Se la box e' aggiornata per un nuovo giorno, resetta lo state machine.
        """
        h = yesterday_high - yesterday_low
        m = (yesterday_high + yesterday_low) / 2.0
        ntt = m + h * self.NO_TRADE_ZONE_PCT / 2.0
        ntb = m - h * self.NO_TRADE_ZONE_PCT / 2.0
        self.box = Box(
            top=yesterday_high,
            bottom=yesterday_low,
            mid=m,
            height=h,
            no_trade_top=ntt,
            no_trade_bot=ntb,
            date=date,
        )
        # Se cambia il giorno, resetta state machine e registra session_open_time
        if self._last_box_date != date:
            self._last_box_date = date
            self.state = self.STATE_IDLE
            if session_open_unix is not None:
                self.session_open_time = session_open_unix
            else:
                # Fallback: se non passato, prova a dedurlo correttamente.
                # SPEC §1: session open = 9:00 Europe/London. zoneinfo gestisce BST/GMT automatico.
                # Il bot runtime dovrebbe passarlo esplicitamente per evitare problemi timezone.
                try:
                    from datetime import datetime as _dt
                    from zoneinfo import ZoneInfo
                    d = _dt.strptime(date, "%Y-%m-%d")
                    # 9:00 Europe/London (BST = UTC+1 estate, GMT = UTC+0 inverno, auto)
                    london_tz = ZoneInfo("Europe/London")
                    session_dt = _dt(d.year, d.month, d.day, 9, 0, 0, tzinfo=london_tz)
                    self.session_open_time = int(session_dt.timestamp())
                except Exception:
                    self.session_open_time = None
        log.info(
            "Box updated [%s]: top=%.4f bottom=%.4f mid=%.4f no_trade=[%.4f, %.4f] session_open=%s",
            self.symbol, self.box.top, self.box.bottom, self.box.mid,
            self.box.no_trade_bot, self.box.no_trade_top,
            self.session_open_time
        )

    # =========================================================================
    # PATTERN DETECTION
    # =========================================================================
    def _is_doji_hammer_short(self, c: Candle) -> bool:
        """Doji Hammer ROSSA (per entry SHORT)."""
        if not c.is_red or c.range == 0:
            return False
        return (
            c.body <= c.range * self.DOJI_BODY_RATIO
            and c.upper_wick >= c.range * self.DOJI_WICK_RATIO
        )

    def _is_doji_hammer_long(self, c: Candle) -> bool:
        """Doji Hammer VERDE (per entry LONG)."""
        if not c.is_green or c.range == 0:
            return False
        return (
            c.body <= c.range * self.DOJI_BODY_RATIO
            and c.lower_wick >= c.range * self.DOJI_WICK_RATIO
        )

    def _volume_ok(self, c: Candle) -> bool:
        """Verifica che il volume sia sopra la media di almeno VOLUME_MULT."""
        if len(self.candles) < 20:
            return False
        recent = self.candles[-20:]
        avg = sum(x.volume for x in recent) / 20.0
        return c.volume >= avg * self.VOLUME_MULT

    def _prev_candle_confirms_rejection(self, c: Candle, side: str) -> bool:
        """SPEC §3.2 pattern 2 candele: la candela precedente deve confermare
        il rifiuto del bordo con ombra nella direzione giusta.
        - side='short' (test TOP): candela precedente GREEN con upper_wick >= 30% del range
          (il prezzo ha tentato di salire oltre il top ma e' stato rifiutato)
        - side='long' (test BOTTOM): candela precedente RED con lower_wick >= 30% del range
          (il prezzo ha tentato di scendere sotto il bottom ma e' stato rifiutato)
        Ritorna True se conferma, False altrimenti. Se non c'è candela precedente, False.
        """
        if len(self.candles) < 2:
            return False
        # candles[-1] = candela corrente (gia' appendata in on_new_candle PRIMA di _check_entry)
        # candles[-2] = candela precedente
        prev = self.candles[-2]
        if prev.range == 0:
            return False
        if side == "short":
            # candela precedente GREEN con ombra superiore >= 30% del range
            if not prev.is_green:
                return False
            return prev.upper_wick >= prev.range * self.PREV_WICK_RATIO
        else:  # long
            # candela precedente RED con ombra inferiore >= 30% del range
            if not prev.is_red:
                return False
            return prev.lower_wick >= prev.range * self.PREV_WICK_RATIO

    def _pattern_in_correct_zone(self, c: Candle, side: str) -> bool:
        """SPEC §1/§2.1: il pattern (Doji Hammer) deve essere FUORI dalla no-trade zone.
        - side='short': il BODY della Doji deve essere nella ZONA ALTA (sopra no_trade_top)
        - side='long': il BODY della Doji deve essere nella ZONA BASSA (sotto no_trade_bot)
        La condizione precedente (c.high >= mid) era lasca — controllava solo l'estremo high.
        """
        if self.box is None:
            return False
        body_center = (c.open + c.close) / 2.0
        if side == "short":
            return body_center > self.box.no_trade_top
        else:  # long
            return body_center < self.box.no_trade_bot

    def _atr_filter_ok(self) -> bool:
        """SPEC §7.1: ATR(14) / prezzo corrente deve essere >= MIN_ATR_PCT (0.3%).
        Se troppo basso, mercato piatto = niente setup. Calcola su candele 1-min o timeframe scelto.
        """
        if len(self.candles) < self.ATR_PERIOD + 1:
            return False
        recent = self.candles[-(self.ATR_PERIOD + 1):]
        # True range per ogni candela (escludo l'ultima che e' la corrente, gia' calcolata sopra)
        trs = []
        for i in range(1, len(recent)):
            cur = recent[i]
            prev_close = recent[i - 1].close
            tr = max(cur.high - cur.low, abs(cur.high - prev_close), abs(cur.low - prev_close))
            trs.append(tr)
        if not trs:
            return False
        atr = sum(trs) / len(trs)
        # Usa l'ultimo close come riferimento prezzo
        last_close = self.candles[-1].close
        if last_close <= 0:
            return False
        return (atr / last_close) >= self.MIN_ATR_PCT

    def _minutes_since_session_open(self) -> Optional[float]:
        """Ritorna i minuti trascorsi dall'apertura della sessione (9:00 London).
        La sessione e' 9:00 London = 8:00 UTC (winter) o 7:00 UTC (DST) - semplifichiamo
        assumendo UTC+0 (London winter) per il calcolo.
        Se non e' stata ancora impostata una session_open_time, ritorna None.
        """
        if self.session_open_time is None or not self.candles:
            return None
        now_ts = self.candles[-1].timestamp
        return (now_ts - self.session_open_time) / 60.0

    # =========================================================================
    # STATE MACHINE
    # =========================================================================
    def _detect_first_test(self, c: Candle) -> int:
        """Rileva il primo test di top/bottom. Ritorna nuovo state."""
        if self.box is None or self.in_position:
            return self.state
        # Test del top: wick sopra box_top di almeno 0.1%
        if c.high >= self.box.top * 1.001:
            return self.STATE_FIRST_TEST_TOP
        # Test del bottom: wick sotto box_bottom di almeno 0.1%
        if c.low <= self.box.bottom * 0.999:
            return self.STATE_FIRST_TEST_BOTTOM
        return self.state

    def _detect_retest(self, c: Candle) -> int:
        """Rileva retest dopo il primo test."""
        if self.box is None or self.in_position:
            return self.state
        if self.state == self.STATE_FIRST_TEST_TOP:
            # retest top: high vicino al top (entro 0.5%)
            if abs(c.high - self.box.top) <= self.box.top * 0.005:
                return self.STATE_RETEST_TOP
        elif self.state == self.STATE_FIRST_TEST_BOTTOM:
            # retest bottom: low vicino al bottom (entro 0.5%)
            if abs(c.low - self.box.bottom) <= self.box.bottom * 0.005:
                return self.STATE_RETEST_BOTTOM
        return self.state

    # =========================================================================
    # SIGNAL GENERATION
    # =========================================================================
    def _check_entry(self, c: Candle) -> Optional[dict]:
        """Verifica se c'è un segnale di entry. Ritorna dict o None.

        SPEC §2/§3: prima applica TUTTI i filtri (no-trade zone, time, ATR),
        poi verifica pattern 2 candele (Doji Hammer corrente + candela precedente
        con ombra di rifiuto), poi calcola entry trigger / SL / TP.
        """
        if self.box is None or self.in_position:
            return None

        # === FILTRO 1: Time filter (SPEC §7.3) ===
        # Skip se siamo nei primi 30 min dopo 9:00 London
        mins = self._minutes_since_session_open()
        if mins is not None and mins < self.MIN_MINUTES_AFTER_OPEN:
            return None  # troppo presto, zona chaos

        # === FILTRO 2: ATR (SPEC §7.1) ===
        # Skip se volatilita' troppo bassa
        if not self._atr_filter_ok():
            return None  # mercato piatto, niente setup

        # === SHORT: test del TOP, cerca Doji Hammer RED ===
        if self.state == self.STATE_RETEST_TOP:
            if not self._is_doji_hammer_short(c):
                return None
            if not self._volume_ok(c):
                return None
            # SPEC §1/§2.1: pattern deve essere FUORI dalla no-trade zone (ZONA ALTA)
            if not self._pattern_in_correct_zone(c, "short"):
                return None
            # SPEC §3.2 pattern 2 candele: candela precedente deve confermare rifiuto del top
            if not self._prev_candle_confirms_rejection(c, "short"):
                return None
            return {
                "side": "short",
                "entry_trigger": c.low,  # entry alla rottura del low della Doji
                "stop": c.high * (1.0 + self.STOP_BUFFER_PCT / 100.0),
                "tp": self.box.bottom,  # target al bordo opposto
                "reason": "SHORT: Doji Hammer RED + prev GREEN upper-wick rejection (retest top)",
            }

        # === LONG: test del BOTTOM, cerca Doji Hammer GREEN ===
        if self.state == self.STATE_RETEST_BOTTOM:
            if not self._is_doji_hammer_long(c):
                return None
            if not self._volume_ok(c):
                return None
            # SPEC §1/§2.1: pattern deve essere FUORI dalla no-trade zone (ZONA BASSA)
            if not self._pattern_in_correct_zone(c, "long"):
                return None
            # SPEC §3.2 pattern 2 candele: candela precedente deve confermare rifiuto del bottom
            if not self._prev_candle_confirms_rejection(c, "long"):
                return None
            return {
                "side": "long",
                "entry_trigger": c.high,  # entry alla rottura del high della Doji
                "stop": c.low * (1.0 - self.STOP_BUFFER_PCT / 100.0),
                "tp": self.box.top,  # target al bordo opposto
                "reason": "LONG: Doji Hammer GREEN + prev RED lower-wick rejection (retest bottom)",
            }

        return None

    # =========================================================================
    # POSITION SIZING
    # =========================================================================
    def _calc_qty(self, entry: float, stop: float) -> float:
        """Calcola la qty in base al rischio % del capitale."""
        try:
            balance = float(self.bybit.get_wallet_balance())
        except Exception as e:
            log.error("Cannot fetch balance: %s — fallback qty=0.01", e)
            return 0.01

        risk_amount = balance * self.risk_pct
        risk_per_unit = abs(entry - stop)
        if risk_per_unit <= 0:
            return 0.01
        qty = risk_amount / risk_per_unit
        # arrotonda per difetto (conservativo)
        return round(qty, 4)

    # =========================================================================
    # MAIN LOOP
    # =========================================================================
    def on_new_candle(self, c: Candle):
        """Callback per ogni nuova candela chiusa."""
        self.candles.append(c)
        if len(self.candles) > 100:
            self.candles.pop(0)

        if self.box is None:
            log.debug("Box not yet set, skipping candle")
            return

        # Aggiorna state machine
        if self.state in (self.STATE_IDLE,):
            new_state = self._detect_first_test(c)
            if new_state != self.state:
                log.info("State: IDLE -> %s (first test detected)", new_state)
                self.state = new_state
        elif self.state in (self.STATE_FIRST_TEST_TOP, self.STATE_FIRST_TEST_BOTTOM):
            new_state = self._detect_retest(c)
            if new_state != self.state:
                log.info("State: %s -> %s (retest detected)", self.state, new_state)
                self.state = new_state

        # Check entry
        signal = self._check_entry(c)
        if signal:
            log.info("ENTRY SIGNAL: %s @ trigger=%.4f stop=%.4f tp=%.4f",
                     signal["side"], signal["entry_trigger"], signal["stop"], signal["tp"])
            self.on_signal(signal)
            self.in_position = True
            self.position_side = signal["side"]
            self.stop_price = signal["stop"]
            self.tp_price = signal["tp"]

    def on_price_update(self, price: float):
        """Callback per ogni tick di prezzo (per gestione trailing stop)."""
        if not self.in_position or self.entry_price is None:
            return

        # Move to breakeven: quando il prezzo si muove di 1% a nostro favore
        if self.position_side == "short" and price <= self.entry_price * 0.99:
            new_stop = self.entry_price * 1.001
            if new_stop < self.stop_price:
                log.info("Move SL to breakeven: %.4f", new_stop)
                self.stop_price = new_stop
                # NB: implementare la chiamata bybit.update_sl(...) qui
        elif self.position_side == "long" and price >= self.entry_price * 1.01:
            new_stop = self.entry_price * 0.999
            if new_stop > self.stop_price:
                log.info("Move SL to breakeven: %.4f", new_stop)
                self.stop_price = new_stop
                # NB: implementare la chiamata bybit.update_sl(...) qui

    # =========================================================================
    # SIGNAL CALLBACK
    # =========================================================================
    def _default_signal_handler(self, signal: dict):
        """Default: stampa solo. Sostituisci con on_signal reale per invio ordini."""
        log.info("[SIGNAL] %s | reason=%s | entry_trigger=%.4f | stop=%.4f | tp=%.4f",
                 signal["side"].upper(), signal["reason"],
                 signal["entry_trigger"], signal["stop"], signal["tp"])


# =========================================================================
# ESEMPIO D'USO
# =========================================================================
if __name__ == "__main__":
    logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")

    # Mock client per test
    class MockBybit:
        def get_wallet_balance(self): return 10000.0

    bot = SquareStrategy(
        bybit_client=MockBybit(),
        symbol="ETHUSDT",
        timeframe="15",
        risk_per_trade_pct=1.0,
    )

    # Setup box con i dati di ieri
    bot.update_box(yesterday_high=3000, yesterday_low=2900, date="2026-07-15")

    # Simula candele
    test_candles = [
        Candle(timestamp=1, open=2950, high=3010, low=2945, close=2955, volume=100),  # test top
        Candle(timestamp=2, open=2955, high=2960, low=2940, close=2945, volume=90),
        Candle(timestamp=3, open=2945, high=2995, low=2940, close=2985, volume=110),  # retest top
        Candle(timestamp=4, open=2985, high=3005, low=2975, close=2980, volume=200),  # Doji Hammer RED?
        Candle(timestamp=5, open=2980, high=2982, low=2970, close=2972, volume=180),  # trigger entry SHORT
    ]
    for c in test_candles:
        bot.on_new_candle(c)
