from __future__ import annotations

from dataclasses import asdict, dataclass
from datetime import datetime, timezone
from math import sqrt
from statistics import mean, pstdev
from typing import Iterable

ENGINE_VERSION = "rules-v1"
HORIZONS = {"15m": 900, "1h": 3600, "4h": 14400}
LEVERAGES = (1, 5, 10, 25)

@dataclass(frozen=True)
class Candle:
    timestamp: int
    open: float
    high: float
    low: float
    close: float
    volume: float | None = None

@dataclass(frozen=True)
class Signal:
    symbol_id: int
    symbol: str
    horizon: str
    direction: str
    entry: float
    stop_loss: float | None
    take_profit: float | None
    confidence: float
    expected_move: float
    expected_volatility: float
    max_adverse_move: float
    risk_reward: float
    score: float
    regime: str
    recommended_leverage: int
    reasons: list[str]
    feature_snapshot: dict
    engine_version: str
    candle_timestamp: int
    generated_at: str
    expires_at: str

    def payload(self) -> dict:
        return asdict(self)

def ema(values: list[float], period: int) -> float:
    if not values: return 0.0
    alpha, value = 2 / (period + 1), values[0]
    for point in values[1:]: value = alpha * point + (1 - alpha) * value
    return value

def rsi(values: list[float], period: int = 14) -> float:
    changes = [b-a for a,b in zip(values[-period-1:-1], values[-period:])]
    if not changes: return 50.0
    gains, losses = mean([max(x,0) for x in changes]), mean([max(-x,0) for x in changes])
    if losses == 0: return 100.0
    return 100 - 100 / (1 + gains/losses)

def atr(candles: list[Candle], period: int = 14) -> float:
    rows = candles[-period-1:]
    tr = [max(cur.high-cur.low, abs(cur.high-prev.close), abs(cur.low-prev.close)) for prev,cur in zip(rows, rows[1:])]
    return mean(tr) if tr else 0.0

def adx_proxy(candles: list[Candle], period: int = 14) -> float:
    rows = candles[-period-1:]
    moves = [abs(b.close-a.close) for a,b in zip(rows, rows[1:])]
    path = sum(moves)
    return min(100.0, (abs(rows[-1].close-rows[0].close) / path * 100)) if path else 0.0

def features(candles: list[Candle]) -> dict:
    if len(candles) < 210: raise ValueError("At least 210 closed candles are required")
    closes = [c.close for c in candles]
    returns = [(b/a)-1 for a,b in zip(closes[-21:-1], closes[-20:]) if a]
    e20,e50,e100,e200 = (ema(closes, p) for p in (20,50,100,200))
    macd = ema(closes,12)-ema(closes,26)
    a = atr(candles)
    bb_mid, bb_std = mean(closes[-20:]), pstdev(closes[-20:])
    volumes = [c.volume for c in candles[-20:] if c.volume is not None]
    support, resistance = min(c.low for c in candles[-50:]), max(c.high for c in candles[-50:])
    return {
        "return_1": closes[-1]/closes[-2]-1, "momentum_10": closes[-1]/closes[-11]-1,
        "roc_20": closes[-1]/closes[-21]-1, "price_acceleration": (closes[-1]-closes[-2])-(closes[-2]-closes[-3]),
        "ema20": e20, "ema50": e50, "ema100": e100, "ema200": e200, "rsi14": rsi(closes),
        "macd": macd, "atr14": a, "adx14": adx_proxy(candles),
        "bollinger_width": (4*bb_std/bb_mid) if bb_mid else 0, "realized_volatility": pstdev(returns)*sqrt(20) if returns else 0,
        "vwap20": sum(c.close*(c.volume or 0) for c in candles[-20:])/sum(volumes) if volumes and sum(volumes) else closes[-1],
        "volume_acceleration": (volumes[-1]/mean(volumes[:-1])-1) if len(volumes)>1 and mean(volumes[:-1]) else 0,
        "support": support, "resistance": resistance,
        "support_distance": (closes[-1]-support)/closes[-1], "resistance_distance": (resistance-closes[-1])/closes[-1],
        "maximum_adverse_movement": max((abs(x) for x in returns if x < 0),default=0),
    }

def generate(symbol_id: int, symbol: str, horizon: str, candles: Iterable[Candle], now: datetime | None = None) -> Signal:
    rows, now = list(candles), now or datetime.now(timezone.utc)
    f, price = features(rows), rows[-1].close
    bull = 0.0
    bull += 18 if f["ema20"] > f["ema50"] else -18
    bull += 12 if f["ema50"] > f["ema200"] else -12
    bull += max(-15,min(15,f["momentum_10"]*500))
    bull += 10 if f["macd"] > 0 else -10
    bull += 8 if 52 <= f["rsi14"] <= 72 else (-8 if 28 <= f["rsi14"] <= 48 else 0)
    bull += 8 if price >= f["resistance"]*.999 else (-8 if price <= f["support"]*1.001 else 0)
    vol_pct = f["atr14"]/price if price else 0
    if f["adx14"] >= 45: regime = "bullish trend" if bull > 0 else "bearish trend"
    elif f["bollinger_width"] > .06: regime = "high volatility"
    elif price >= f["resistance"]*.999 or price <= f["support"]*1.001: regime = "breakout"
    elif f["bollinger_width"] < .015: regime = "low volatility"
    else: regime = "range"
    directional = abs(bull)
    confidence = round(min(90,max(50,50+directional*.42)),2)
    score = round(min(100,directional*.9 + confidence*.5 - min(15,vol_pct*400)),2)
    direction = "LONG" if bull > 0 else "SHORT"
    stop_distance = max(f["atr14"]*1.5, price*.0025)
    stop = price-stop_distance if direction=="LONG" else price+stop_distance
    target = price+stop_distance*1.75 if direction=="LONG" else price-stop_distance*1.75
    rr = 1.75
    if score < 70 or confidence < 65 or rr < 1.5 or regime == "range": direction,stop,target = "NO_TRADE",None,None
    volatility_cap = max(1,min(25,int(1/max(vol_pct*4,.04))))
    quality_cap = 25 if confidence >= 78 and score >= 88 else (10 if confidence >= 72 and score >= 80 else 5)
    max_leverage = min(volatility_cap,quality_cap) if direction != "NO_TRADE" else 1
    leverage = max(x for x in LEVERAGES if x <= max_leverage)
    seconds = HORIZONS[horizon]
    reasons = [f"EMA trend is {'bullish' if bull>0 else 'bearish'}", f"RSI is {f['rsi14']:.1f}", f"Regime: {regime}"]
    return Signal(symbol_id,symbol,horizon,direction,price,stop,target,confidence,abs(f["momentum_10"]),f["realized_volatility"],f["maximum_adverse_movement"],rr,score,regime,leverage,reasons,f,ENGINE_VERSION,rows[-1].timestamp,now.isoformat(),datetime.fromtimestamp(now.timestamp()+seconds,tz=timezone.utc).isoformat())

def resample(candles: list[Candle], seconds: int) -> list[Candle]:
    buckets: dict[int,list[Candle]] = {}
    for candle in sorted(candles,key=lambda c:c.timestamp): buckets.setdefault(candle.timestamp//seconds*seconds,[]).append(candle)
    return [Candle(ts,rows[0].open,max(x.high for x in rows),min(x.low for x in rows),rows[-1].close,sum(x.volume or 0 for x in rows) or None) for ts,rows in buckets.items()]
