"""
stream.py

Manages the Alpaca SIP real-time stream for a set of symbols. Buffers
incoming trades/quotes/bars into an in-memory per-symbol rolling window
that entry_engine.py and position_manager.py read from.

Runs the stream in a background thread so monitor.py's main loop stays
simple synchronous polling logic (poll the buffer, not the network).

Handles:
- reconnect with backoff (config: streaming.reconnect_backoff_seconds)
- stale data detection (config: streaming.stale_data_seconds)
- symbol subscribe/unsubscribe as the watchlist changes through the day
"""

import asyncio
import threading
import time
from collections import defaultdict, deque, OrderedDict
from datetime import datetime, timezone

from config_loader import get_config
from logger_setup import get_logger
from alpaca_client import get_client

log = get_logger("stream")


class SymbolBuffer:
    """Rolling window of 1-min bars + latest quote for one symbol, plus a
    second, parallel sub-minute rolling window (bars_sub / get_bars_sub())
    at a configurable bucket width for finer-grained stream-feature work.
    See config.json's streaming._note_sub_minute."""

    def __init__(self, maxlen=390, sub_bucket_seconds=30, sub_maxlen=30):
        # [BUGFIX 2026-09-03] Was a plain deque that both on_bar() and
        # _accumulate_bar() (fed by on_trade()) appended to independently
        # -- since _subscribe_async()/`_run_forever()` subscribe to BOTH
        # subscribe_trades AND subscribe_bars for every symbol, every
        # single minute could end up with TWO entries in `bars`: Alpaca's
        # own aggregated bar (on_bar) and a second bar built here from
        # raw trade prints (on_trade -> _accumulate_bar), with nothing
        # deduplicating them by minute. Every consumer of get_bars() --
        # vwap(), atr(), rsi(), the slope/structure functions in
        # indicators.py -- assumes one bar per minute; double-booking a
        # minute silently double-weights its volume in vwap()/relative_
        # volume() and adds a near-duplicate, independently-sourced OHLC
        # point into every regression/structure calculation, which is
        # exactly the kind of noise that would make price_slope/momentum_
        # slope flip on single poll cycles (confirmed in the 2026-09-03
        # session log: ACHR's momentum_slope flipped negative->positive
        # inside one 5-second window at entry, and the same symbol logged
        # 79 separate stop-touch/grace-reset cycles across ~110 minutes
        # of otherwise unremarkable chop). Fix: key the window by minute
        # bucket so on_bar() and _accumulate_bar()'s finalization both
        # UPSERT the same slot for a given minute instead of appending a
        # second one -- whichever source writes a minute last wins for
        # that minute, and neither can be double-counted.
        self.maxlen = maxlen
        self.bars = OrderedDict()  # minute_bucket -> bar dict, oldest first
        self.latest_quote = None  # (bid, ask, timestamp)
        self.latest_trade_price = None
        self.last_update = None
        self._current_bar = None

        # [SIMULATION-ONLY 2026-09-05] Second, parallel rolling buffer at
        # finer-than-1-minute resolution -- see config.json's
        # streaming._note_sub_minute. Fed by the exact same on_trade()
        # ticks as _accumulate_bar() above; does not touch self.bars,
        # self._current_bar, or on_bar() in any way, so every existing
        # consumer of get_bars() (entry_engine, intraday_health,
        # position_manager, risk_manager) is completely unaffected.
        # Nothing currently calls get_bars_sub() -- inert until a
        # stream-features consumer is built.
        self.sub_bucket_seconds = sub_bucket_seconds
        self.sub_maxlen = sub_maxlen
        self.bars_sub = OrderedDict()  # sub_bucket -> bar dict, oldest first
        self._current_bar_sub = None

    def _upsert_bar(self, bar: dict):
        key = bar["t"]
        self.bars[key] = bar
        self.bars.move_to_end(key)
        while len(self.bars) > self.maxlen:
            self.bars.popitem(last=False)

    def _upsert_bar_sub(self, bar: dict):
        key = bar["t"]
        self.bars_sub[key] = bar
        self.bars_sub.move_to_end(key)
        while len(self.bars_sub) > self.sub_maxlen:
            self.bars_sub.popitem(last=False)

    def _sub_bucket_for(self, ts):
        width = self.sub_bucket_seconds
        bucket_second = (ts.second // width) * width
        return ts.replace(second=bucket_second, microsecond=0)

    def _accumulate_bar_sub(self, price, size, ts, prev_price):
        """
        [SIMULATION-ONLY 2026-09-05] prev_price is the trade price
        immediately before this one (captured by on_trade() BEFORE it
        overwrites self.latest_trade_price) -- compared against the new
        price for tick direction (upticks/downticks), deliberately
        continuous across bucket boundaries rather than resetting
        direction tracking to "unknown" at the start of every new bucket,
        so a bucket's uptick_ratio isn't artificially diluted by treating
        its first tick as directionless. tick_count/upticks/downticks are
        plain O(1) counters -- no raw ticks are stored, so this adds no
        unbounded memory regardless of how many trades print per bucket.
        """
        bucket = self._sub_bucket_for(ts)
        is_new_bucket = self._current_bar_sub is None or self._current_bar_sub["t"] != bucket
        if is_new_bucket:
            if self._current_bar_sub is not None:
                self._upsert_bar_sub(self._current_bar_sub)
            self._current_bar_sub = {
                "t": bucket, "o": price, "h": price, "l": price, "c": price, "v": size,
                "tick_count": 0, "upticks": 0, "downticks": 0,
            }
        b = self._current_bar_sub
        if not is_new_bucket:
            b["h"] = max(b["h"], price)
            b["l"] = min(b["l"], price)
            b["c"] = price
            b["v"] += size
        b["tick_count"] += 1
        if prev_price is not None:
            if price > prev_price:
                b["upticks"] += 1
            elif price < prev_price:
                b["downticks"] += 1

    def get_bars_sub(self, include_forming=True) -> list:
        bars = list(self.bars_sub.values())
        if include_forming and self._current_bar_sub is not None:
            bars = bars + [self._current_bar_sub]
        return bars

    def on_trade(self, price: float, size: float, ts):
        prev_price = self.latest_trade_price
        self.latest_trade_price = price
        self.last_update = ts
        self._accumulate_bar(price, size, ts)
        self._accumulate_bar_sub(price, size, ts, prev_price)

    def on_quote(self, bid: float, ask: float, ts):
        self.latest_quote = (bid, ask, ts)
        self.last_update = ts

    def on_bar(self, o, h, l, c, v, ts):
        """Direct minute-bar ingestion (Alpaca also streams aggregated bars)."""
        minute_bucket = ts.replace(second=0, microsecond=0) if hasattr(ts, "replace") else ts
        self._upsert_bar({"t": minute_bucket, "o": o, "h": h, "l": l, "c": c, "v": v})
        self.last_update = ts
        # The official bar for this minute just landed -- if a synthetic
        # trade-built bar for the SAME minute is still forming, drop it
        # rather than let get_bars() append it again on top of the
        # official one a moment later.
        if self._current_bar is not None and self._current_bar["t"] == minute_bucket:
            self._current_bar = None

    def _accumulate_bar(self, price, size, ts):
        minute_bucket = ts.replace(second=0, microsecond=0)
        if self._current_bar is None or self._current_bar["t"] != minute_bucket:
            if self._current_bar is not None:
                self._upsert_bar(self._current_bar)
            self._current_bar = {
                "t": minute_bucket, "o": price, "h": price, "l": price, "c": price, "v": size
            }
        else:
            b = self._current_bar
            b["h"] = max(b["h"], price)
            b["l"] = min(b["l"], price)
            b["c"] = price
            b["v"] += size

    def get_bars(self, include_forming=True) -> list:
        bars = list(self.bars.values())
        if include_forming and self._current_bar is not None:
            bars = bars + [self._current_bar]
        return bars

    def is_stale(self, stale_seconds: int) -> bool:
        if self.last_update is None:
            return True
        now = datetime.now(timezone.utc)
        last = self.last_update
        if last.tzinfo is None:
            last = last.replace(tzinfo=timezone.utc)
        return (now - last).total_seconds() > stale_seconds


class StreamManager:
    def __init__(self):
        self.cfg = get_config()["streaming"]
        self.client = get_client()
        sub_bucket_seconds = self.cfg.get("sub_minute_bucket_seconds", 30)
        sub_buffer_minutes = self.cfg.get("sub_minute_buffer_minutes", 15)
        sub_maxlen = max(1, int(sub_buffer_minutes * 60 / sub_bucket_seconds))
        self.buffers: dict[str, SymbolBuffer] = defaultdict(
            lambda: SymbolBuffer(sub_bucket_seconds=sub_bucket_seconds, sub_maxlen=sub_maxlen)
        )
        self._subscribed = set()
        self._stream = None
        self._thread = None
        self._loop = None
        self._stop_event = threading.Event()
        self._connected = threading.Event()

    # ------------------------------------------------------------------
    def start(self, symbols: list):
        self._thread = threading.Thread(target=self._run_forever, args=(symbols,), daemon=True)
        self._thread.start()
        # Give the connection a moment to establish before returning
        self._connected.wait(timeout=10)

    def stop(self):
        self._stop_event.set()
        if self._loop and self._stream:
            try:
                asyncio.run_coroutine_threadsafe(self._stream.stop_ws(), self._loop)
            except Exception:
                pass

    def add_symbols(self, symbols: list):
        new_syms = [s for s in symbols if s not in self._subscribed]
        if not new_syms or self._stream is None or self._loop is None:
            return
        self._subscribed.update(new_syms)
        try:
            asyncio.run_coroutine_threadsafe(
                self._subscribe_async(new_syms), self._loop
            )
        except Exception as e:
            log.warning(f"add_symbols failed: {e}")

    async def _subscribe_async(self, symbols):
        self._stream.subscribe_trades(self._on_trade, *symbols)
        self._stream.subscribe_quotes(self._on_quote, *symbols)
        self._stream.subscribe_bars(self._on_bar, *symbols)

    # ------------------------------------------------------------------
    async def _on_trade(self, trade):
        buf = self.buffers[trade.symbol]
        buf.on_trade(float(trade.price), float(trade.size), trade.timestamp)

    async def _on_quote(self, quote):
        buf = self.buffers[quote.symbol]
        if quote.bid_price and quote.ask_price:
            buf.on_quote(float(quote.bid_price), float(quote.ask_price), quote.timestamp)

    async def _on_bar(self, bar):
        buf = self.buffers[bar.symbol]
        buf.on_bar(float(bar.open), float(bar.high), float(bar.low),
                   float(bar.close), float(bar.volume), bar.timestamp)

    # ------------------------------------------------------------------
    def _run_forever(self, symbols: list):
        backoffs = self.cfg["reconnect_backoff_seconds"]
        max_attempts = self.cfg["max_reconnect_attempts"]
        attempt = 0

        while not self._stop_event.is_set() and attempt < max_attempts:
            try:
                self._loop = asyncio.new_event_loop()
                asyncio.set_event_loop(self._loop)
                self._stream = self.client.new_stream()
                self._subscribed = set(symbols)
                self._stream.subscribe_trades(self._on_trade, *symbols)
                self._stream.subscribe_quotes(self._on_quote, *symbols)
                self._stream.subscribe_bars(self._on_bar, *symbols)

                log.info(f"[STREAM] Connecting SIP stream for {len(symbols)} symbols")
                self._connected.set()
                attempt = 0  # reset backoff counter on a successful connection
                self._stream.run()  # blocks until stopped or disconnected
            except Exception as e:
                if self._stop_event.is_set():
                    break
                delay = backoffs[min(attempt, len(backoffs) - 1)]
                log.warning(f"[STREAM] Disconnected ({e}); reconnecting in {delay}s "
                            f"(attempt {attempt + 1}/{max_attempts})")
                time.sleep(delay)
                attempt += 1

        if attempt >= max_attempts:
            log.error("[STREAM] Max reconnect attempts exceeded. Stream is DOWN. "
                      "monitor.py should fall back to REST polling.")

    # ------------------------------------------------------------------
    def get_bars(self, symbol: str) -> list:
        return self.buffers[symbol].get_bars()

    def get_bars_sub(self, symbol: str) -> list:
        """Sub-minute (configurable bucket width, default 30s) rolling
        bars -- see SymbolBuffer.get_bars_sub() / config.json's
        streaming._note_sub_minute. Not consumed by anything yet."""
        return self.buffers[symbol].get_bars_sub()

    def get_quote(self, symbol: str):
        return self.buffers[symbol].latest_quote

    def is_symbol_stale(self, symbol: str) -> bool:
        return self.buffers[symbol].is_stale(self.cfg["stale_data_seconds"])

    def is_healthy(self) -> bool:
        return self._connected.is_set() and not self._stop_event.is_set()
