"""
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 fast_pipeline.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, timedelta

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

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

        # [FEATURE 2026-09-17] Trade-flow imbalance -- classifies every
        # trade print as buyer- or seller-initiated against the latest
        # known NBBO quote (Lee-Ready style: at/above ask -> buy, at/below
        # bid -> sell, otherwise a midpoint tiebreak), kept as a short
        # rolling window of raw (ts, side, size) for get_trade_imbalance().
        # Both smart_engine.py's Stage 3 AND exit.py's imbalance-decline
        # gate read this same live rolling window now (exit.py originally
        # had its own separate per-completed-minute tally, checked once
        # per bar roll -- removed 2026-09-18 after real-data testing
        # showed that cadence completely missed a real reversal while
        # this continuous one caught it within 5 seconds of the actual
        # price low).
        self._trade_log = deque()  # (ts, side, size), trimmed to _TRADE_LOG_MAX_AGE_SECONDS

        # [BUGFIX 2026-09-22] The stream thread writes these buffers while
        # the monitor thread reads them; iterating _trade_log mid-append
        # raised "deque mutated during iteration" and killed the bot at
        # 10:49 ET with positions still open. Every read and write of
        # bars / bars_sub / _trade_log / the forming bars goes through it.
        self._lock = threading.RLock()

    _TRADE_LOG_MAX_AGE_SECONDS = 90  # comfortably covers Stage 3's ~20s window

    def _classify_trade(self, price: float) -> int:
        """+1 buyer-initiated, -1 seller-initiated, 0 unknown/no quote yet."""
        if self.latest_quote is None:
            return 0
        bid, ask, _ = self.latest_quote
        if not bid or not ask:
            return 0
        if price >= ask:
            return 1
        if price <= bid:
            return -1
        mid = (bid + ask) / 2
        if price > mid:
            return 1
        if price < mid:
            return -1
        return 0

    def _trim_trade_log(self, now_ts):
        cutoff = now_ts - timedelta(seconds=self._TRADE_LOG_MAX_AGE_SECONDS)
        while self._trade_log and self._trade_log[0][0] < cutoff:
            self._trade_log.popleft()

    def get_trade_imbalance(self, window_seconds: float = 20.0):
        """(buy_vol - sell_vol) / (buy_vol + sell_vol) over the last
        window_seconds of classified trades, or None if there's nothing
        to compute it from (fails open, same convention as spread_pct
        when a quote isn't available)."""
        with self._lock:
            if not self._trade_log:
                return None
            now_ts = self._trade_log[-1][0]
            cutoff = now_ts - timedelta(seconds=window_seconds)
            buy_vol = sell_vol = 0.0
            for ts, side, size in reversed(self._trade_log):
                if ts < cutoff:
                    break
                if side > 0:
                    buy_vol += size
                elif side < 0:
                    sell_vol += size
        total = buy_vol + sell_vol
        return (buy_vol - sell_vol) / total if total else None

    def _upsert_bar(self, bar: dict):
        with self._lock:
            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:
        with self._lock:
            bars = list(self.bars_sub.values())
            if include_forming and self._current_bar_sub is not None:
                bars = bars + [dict(self._current_bar_sub)]
        return bars

    def on_trade(self, price: float, size: float, ts):
        with self._lock:
            self._on_trade_locked(price, size, ts)

    def _on_trade_locked(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)
        side = self._classify_trade(price)
        self._trade_log.append((ts, side, size))
        self._trim_trade_log(ts)

    def on_quote(self, bid: float, ask: float, ts):
        with self._lock:
            self._on_quote_locked(bid, ask, ts)

    def _on_quote_locked(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)."""
        with self._lock:
            self._on_bar_locked(o, h, l, c, v, ts)

    def _on_bar_locked(self, o, h, l, c, v, ts):
        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:
        with self._lock:
            bars = list(self.bars.values())
            if include_forming and self._current_bar is not None:
                bars = bars + [dict(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()
        # [2026-09-26] optional entry_rules.ShadowEntryRules, set by monitor.py
        # (shadow-only recorder; its methods never raise)
        self.shadow = None
        # symbols subscribed to 1-min BARS only (no trades/quotes) -- the shadow
        # recorder's SPY/IWM benchmarks: their tick volume would flood this
        # thread, and relative strength only needs bars.
        self.bars_only = set()
        self._stream = None
        self._thread = None
        self._loop = None
        self._stop_event = threading.Event()
        self._connected = threading.Event()

    # ------------------------------------------------------------------
    def _backfill(self, symbols: list):
        """
        [BUGFIX 2026-09-09] stream_features.compute_features() treats
        bars[0] as "the session's first 1-min bar" (session_open in its
        extension_from_open_pct calc) on the assumption a symbol's
        buffer always spans the whole session from 09:30. That's only
        true if the symbol has been subscribed since before the open --
        false for (a) any symbol added mid-day via add_symbols() (no
        history, subscribe_trades/bars only deliver ticks going
        forward), and (b) EVERY symbol after any intraday process
        restart, since SymbolBuffer is in-memory only and starts empty
        again. Confirmed root cause of a real miss: SRAD, added to the
        pool at 12:13 on 2026-09-09 (well after its 11:55 reversal low
        and post a 12:27 restart besides), had its "session open" read
        as ~$12.65 (the price when the buffer happened to start) instead
        of the real $12.69 09:30 open -- silently understating how
        extended it actually was, undermining the whole extension gate.
        Fix: before subscribing a symbol to the live stream, pull its
        real 1-min bars since market_open_time via REST and pre-load the
        buffer, so bars[0] is always genuinely the session's first bar
        regardless of when/why the subscription started. No-op before
        the open (nothing to backfill yet -- the live stream will
        deliver the real bar 0 itself, same as today).
        """
        open_dt = session_open_dt()
        now = datetime.now(timezone.utc)
        if now <= open_dt.astimezone(timezone.utc):
            return
        for symbol in symbols:
            try:
                bars = self.client.get_minute_bars(symbol, start=open_dt, end=now)
            except Exception as e:
                log.warning(f"[STREAM] Backfill failed for {symbol}: {e}")
                continue
            if not bars:
                continue
            buf = self.buffers[symbol]
            for bar in bars:
                minute_bucket = bar.timestamp.replace(second=0, microsecond=0)
                buf._upsert_bar({
                    "t": minute_bucket, "o": float(bar.open), "h": float(bar.high),
                    "l": float(bar.low), "c": float(bar.close), "v": float(bar.volume),
                })
            if self.shadow is not None:
                try:
                    from entry_rules import Bar as _SBar
                    self.shadow.seed_bars(symbol, [_SBar(b.timestamp, float(b.open), float(b.high), float(b.low),
                                                         float(b.close), float(b.volume)) for b in bars])
                except Exception:
                    pass  # shadow recorder must never affect the stream
            log.info(f"[STREAM] Backfilled {symbol} with {len(bars)} historical bars "
                      f"(session open ${float(bars[0].open):.4f}) before live subscription")

    # ------------------------------------------------------------------
    def start(self, symbols: list):
        self._backfill(symbols)
        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):
        """[FIX 2026-09-23] Used to schedule _subscribe_async() onto
        self._loop -- but self._stream.run() runs on alpaca-py's OWN loop
        (asyncio.run inside run()), so self._loop never ran and the
        coroutine never executed: added symbols were silently never
        subscribed. alpaca-py's subscribe_*/unsubscribe_* are themselves
        thread-safe (they post onto the stream's running loop and wait),
        so they're called directly from the caller's thread here."""
        # a bars-only benchmark (SPY/IWM) that the scanner now picks needs
        # its trades/quotes too
        promote = [s for s in symbols if s in self.bars_only and s in self._subscribed]
        if promote and self._stream is not None:
            self.bars_only -= set(promote)
            try:
                self._stream.subscribe_trades(self._on_trade, *promote)
                self._stream.subscribe_quotes(self._on_quote, *promote)
            except Exception as e:
                log.warning(f"[STREAM] tick subscribe failed for {promote}: {e}")
        new_syms = [s for s in symbols if s not in self._subscribed]
        if not new_syms or self._stream is None:
            return
        self._backfill(new_syms)
        self._subscribed.update(new_syms)
        try:
            self._stream.subscribe_trades(self._on_trade, *new_syms)
            self._stream.subscribe_quotes(self._on_quote, *new_syms)
            self._stream.subscribe_bars(self._on_bar, *new_syms)
            log.info(f"[STREAM] Subscribed {len(new_syms)} added symbols: {new_syms}")
        except Exception as e:
            log.warning(f"[STREAM] add_symbols failed for {new_syms}: {e}")

    def remove_symbols(self, symbols: list):
        """Unsubscribes and drops the buffer for symbols no longer
        watched, so the subscription set stays ~one scanner list wide
        instead of growing with every rescan."""
        gone = [s for s in symbols if s in self._subscribed]
        if not gone:
            return
        self._subscribed.difference_update(gone)
        if self._stream is not None:
            try:
                self._stream.unsubscribe_trades(*gone)
                self._stream.unsubscribe_quotes(*gone)
                self._stream.unsubscribe_bars(*gone)
            except Exception as e:
                log.warning(f"[STREAM] remove_symbols failed for {gone}: {e}")
        for sym in gone:
            self.buffers.pop(sym, None)
        log.info(f"[STREAM] Unsubscribed {len(gone)} dropped symbols: {gone}")

    # ------------------------------------------------------------------
    async def _on_trade(self, trade):
        buf = self.buffers[trade.symbol]
        buf.on_trade(float(trade.price), float(trade.size), trade.timestamp)
        if self.shadow is not None:
            try:
                self.shadow.on_trade(trade.symbol, trade.timestamp, float(trade.price), float(trade.size))
            except Exception:
                pass

    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)
            if self.shadow is not None:
                try:
                    self.shadow.on_quote(quote.symbol, quote.timestamp, float(quote.bid_price),
                                         float(quote.ask_price), float(quote.bid_size or 0), float(quote.ask_size or 0))
                except Exception:
                    pass

    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)
        if self.shadow is not None:
            try:
                from entry_rules import Bar as _SBar
                self.shadow.on_bar(bar.symbol, _SBar(bar.timestamp, float(bar.open), float(bar.high),
                                                     float(bar.low), float(bar.close), float(bar.volume)))
            except Exception:
                pass

    # ------------------------------------------------------------------
    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()
                # [FIX 2026-09-23] On a reconnect, resubscribe whatever is
                # watched NOW (add_symbols/remove_symbols keep _subscribed
                # current), not the list start() was first called with.
                if not self._subscribed:
                    self._subscribed = set(symbols)
                current = sorted(self._subscribed)
                ticks = [x for x in current if x not in self.bars_only]
                if ticks:
                    self._stream.subscribe_trades(self._on_trade, *ticks)
                    self._stream.subscribe_quotes(self._on_quote, *ticks)
                self._stream.subscribe_bars(self._on_bar, *current)

                log.info(f"[STREAM] Connecting SIP stream for {len(current)} 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 get_trade_imbalance(self, symbol: str, window_seconds: float = 20.0):
        return self.buffers[symbol].get_trade_imbalance(window_seconds)

    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()
