from __future__ import annotations

import math
from statistics import pstdev

from investly.domain import Bar


def sma(values: list[float], n: int) -> float | None:
    return sum(values[-n:]) / n if len(values) >= n else None


def momentum(values: list[float], n: int) -> float | None:
    if len(values) <= n or values[-n - 1] == 0:
        return None
    return values[-1] / values[-n - 1] - 1


def atr(bars: list[Bar], n: int = 14) -> float | None:
    if len(bars) < n + 1:
        return None
    trs: list[float] = []
    for prev, cur in zip(bars[-n - 1 : -1], bars[-n:], strict=True):
        trs.append(max(cur.high - cur.low, abs(cur.high - prev.close), abs(cur.low - prev.close)))
    return sum(trs) / len(trs)


def annualized_volatility(values: list[float], n: int = 60) -> float | None:
    if len(values) < n + 1:
        return None
    rets = [
        math.log(values[i] / values[i - 1])
        for i in range(len(values) - n, len(values))
        if values[i - 1] > 0 and values[i] > 0
    ]
    return pstdev(rets) * math.sqrt(252) if len(rets) >= 2 else None


def feature_set(bars: list[Bar]) -> dict[str, float | None]:
    closes = [b.close for b in bars]
    last = closes[-1]
    s50 = sma(closes, 50)
    s200 = sma(closes, 200)
    high52 = max((b.high for b in bars[-252:]), default=last)
    low52 = min((b.low for b in bars[-252:]), default=last)
    range_pos = (last - low52) / (high52 - low52) if high52 > low52 else 0.5
    a = atr(bars)
    return {
        "price": last,
        "sma50": s50,
        "sma200": s200,
        "dist_sma50": (last / s50 - 1) if s50 else None,
        "dist_sma200": (last / s200 - 1) if s200 else None,
        "mom_1m": momentum(closes, 21),
        "mom_3m": momentum(closes, 63),
        "mom_6m": momentum(closes, 126),
        "mom_12m": momentum(closes, 252),
        "atr14": a,
        "atr_pct": a / last if a and last else None,
        "vol60": annualized_volatility(closes, 60),
        "range52": range_pos,
    }
