"""
alpaca_client.py

Thin wrapper around alpaca-py. Isolates every direct Alpaca SDK call in
one place so the rest of the system (scanner, entry engine, position
manager) never talks to the SDK directly. This makes it possible to
swap data providers later without touching strategy code, and makes
simulation mode trivial (just don't call the order methods).

Uses:
- alpaca.data.historical for premarket snapshots / bars
- alpaca.data.live for SIP streaming
- alpaca.trading for account, positions, and order submission (paper by default)
"""

from datetime import datetime, timedelta
from typing import Optional

from config_loader import get_env, get_config
from logger_setup import get_logger

log = get_logger("alpaca_client")

import json
import re

from alpaca.data.historical import StockHistoricalDataClient
from alpaca.data.live import StockDataStream
from alpaca.data.requests import (
    StockBarsRequest,
    StockLatestQuoteRequest,
    StockLatestTradeRequest,
    StockSnapshotRequest,
    StockTradesRequest,
    StockQuotesRequest,
)
from alpaca.data.timeframe import TimeFrame, TimeFrameUnit
from alpaca.trading.client import TradingClient
from alpaca.trading.requests import (
    MarketOrderRequest,
    GetAssetsRequest,
    GetOrdersRequest,
)
from alpaca.trading.enums import OrderSide, TimeInForce, AssetClass, AssetStatus, QueryOrderStatus
from alpaca.data.enums import DataFeed


# Alpaca's "potential wash trade" rejection code -- returned when we try
# to submit a second order for a symbol while an earlier, still-open
# order (ours or otherwise) exists on the opposite side. See
# parse_wash_trade_error() below; this is the error PositionManager was
# previously discarding instead of acting on. [BUGFIX 2026-08-18]
WASH_TRADE_ERROR_CODE = 40310000

# [BUGFIX 2026-08-24] Alpaca's "position not found" rejection code --
# returned when close_position() is called for a symbol the broker
# doesn't (yet) recognize as holding shares of. This is a DIFFERENT
# error from the wash-trade one above and was NOT covered by the
# 2026-08-18 fix, even though it produces the identical symptom: real
# session logs (2026-08-24) show CRML/FSM/EXK/XPON each hitting this
# error 4-8 times in a row, 20-30 seconds apart, for 4-6 minutes with
# no working stop protection the whole time, before the broker-side
# state finally caught up (at which point a resubmitted close usually
# collided with a resting order and surfaced AS a wash-trade error
# instead, which the existing fix does handle). Root cause: a market
# BUY order can take longer to settle at the broker than the 5-second
# poll interval, but enter_position() marks the position "open" locally
# immediately at submission time (see position_manager.py's own
# docstring) -- so a stop touch moments later tries to close a position
# the broker hasn't finished registering yet.
POSITION_NOT_FOUND_ERROR_CODE = 40410000


def parse_position_not_found_error(exc: Exception) -> bool:
    """
    Returns True if exc is Alpaca's "position not found" rejection
    (code 40410000), False otherwise. Same best-effort JSON-then-regex
    parsing strategy as parse_wash_trade_error() below, for the same
    reason: a minor SDK error-formatting change shouldn't silently
    disable this check.
    """
    text = str(exc)
    payload = None
    try:
        payload = json.loads(text)
    except (json.JSONDecodeError, TypeError):
        start = text.find("{")
        end = text.rfind("}")
        if start != -1 and end != -1 and end > start:
            try:
                payload = json.loads(text[start:end + 1])
            except json.JSONDecodeError:
                payload = None

    if isinstance(payload, dict) and payload.get("code") == POSITION_NOT_FOUND_ERROR_CODE:
        return True

    return "position not found" in text.lower()


def parse_wash_trade_error(exc: Exception):
    """
    Alpaca's SDK raises an APIError whose string form is (or embeds) the
    raw JSON error body, e.g.:
        {"code":40310000,"existing_order_id":"...","message":"potential
         wash trade detected. use complex orders","reject_reason":
         "opposite side market/stop order exists"}

    Returns the existing_order_id (str) if this exception is a wash-trade
    rejection that names a conflicting order, otherwise None. Best-effort:
    tries json.loads first, falls back to a regex so a minor SDK
    formatting difference doesn't silently disable the guard.
    """
    text = str(exc)
    payload = None
    try:
        payload = json.loads(text)
    except (json.JSONDecodeError, TypeError):
        start = text.find("{")
        end = text.rfind("}")
        if start != -1 and end != -1 and end > start:
            try:
                payload = json.loads(text[start:end + 1])
            except json.JSONDecodeError:
                payload = None

    if isinstance(payload, dict) and payload.get("code") == WASH_TRADE_ERROR_CODE:
        return payload.get("existing_order_id")

    if "wash trade" in text.lower():
        m = re.search(r'"existing_order_id"\s*:\s*"([^"]+)"', text)
        if m:
            return m.group(1)

    return None


class AlpacaClient:
    def __init__(self):
        env = get_env()
        env.validate()
        self.env = env
        self.cfg = get_config()

        self.trading = TradingClient(
            env.api_key, env.secret_key, paper="paper" in env.base_url
        )
        self.hist_data = StockHistoricalDataClient(env.api_key, env.secret_key)

        feed = self.cfg.get("streaming", {}).get("feed", "sip")
        try:
            self._feed = DataFeed(feed)
        except ValueError:
            log.warning(f"Unknown feed '{feed}' in config; defaulting to DataFeed.SIP")
            self._feed = DataFeed.SIP

    # ------------------------------------------------------------------
    # Streaming
    # ------------------------------------------------------------------

    def new_stream(self) -> StockDataStream:
        """Returns a fresh StockDataStream configured for SIP feed."""
        return StockDataStream(
            self.env.api_key, self.env.secret_key, feed=self._feed
        )

    # ------------------------------------------------------------------
    # Historical / snapshot data (used by scanner)
    # ------------------------------------------------------------------

    def get_tradable_assets(self):
        req = GetAssetsRequest(asset_class=AssetClass.US_EQUITY, status=AssetStatus.ACTIVE)
        return self.trading.get_all_assets(req)

    def get_snapshots(self, symbols: list):
        """Bulk snapshot (latest trade/quote/min bar/day bar) for a symbol list."""
        if not symbols:
            return {}
        req = StockSnapshotRequest(symbol_or_symbols=symbols, feed=self._feed)
        try:
            return self.hist_data.get_stock_snapshot(req)
        except Exception as e:
            log.warning(f"get_snapshots failed: {e}")
            return {}

    def get_minute_bars(self, symbol: str, start: datetime, end: datetime, limit=500):
        req = StockBarsRequest(
            symbol_or_symbols=symbol,
            timeframe=TimeFrame(1, TimeFrameUnit.Minute),
            start=start,
            end=end,
            limit=limit,
            feed=self._feed,
        )
        try:
            bars = self.hist_data.get_stock_bars(req)
            return bars[symbol] if symbol in bars.data else []
        except Exception as e:
            log.warning(f"get_minute_bars({symbol}) failed: {e}")
            return []

    def get_daily_bars_bulk(self, symbols: list, start: datetime, end: datetime,
                             limit: int = 60, chunk_size: int = 200) -> dict:
        """
        [FEATURE 2026-09-15] Bulk multi-symbol daily bars, chunked to keep
        each request URL/param size reasonable -- used by
        premarket_scanner.py's volatility-floor filter (see its
        prefilter_by_snapshot() docstring), which needs each symbol's
        real day-to-day range history, not just today's intraday bars.
        One bulk call per chunk instead of one REST call per symbol,
        same efficiency reasoning as get_snapshots() already uses.

        Returns {symbol: [bar dicts]}, oldest first; a symbol with no
        data (delisted, too new, or the request failed) is simply absent
        from the result rather than raising -- callers must treat a
        missing symbol as "unknown," never as "confirmed low volatility."

        [BUGFIX 2026-09-15] `limit` on a MULTI-symbol StockBarsRequest is
        a combined cap across the whole response, not per-symbol --
        confirmed live: requesting 7 symbols with limit=60 silently
        returned 49 bars for the first alphabetically and 11 for the
        second, and NOTHING for the other five (49+11=60 exactly). Every
        chunk request below multiplies `limit` by the chunk size so each
        symbol still gets its full requested history regardless of how
        many other symbols share the request.
        """
        out = {}
        if not symbols:
            return out
        for i in range(0, len(symbols), chunk_size):
            chunk = symbols[i:i + chunk_size]
            req = StockBarsRequest(
                symbol_or_symbols=chunk,
                timeframe=TimeFrame(1, TimeFrameUnit.Day),
                start=start,
                end=end,
                limit=limit * len(chunk),
                feed=self._feed,
            )
            try:
                bars = self.hist_data.get_stock_bars(req)
            except Exception as e:
                log.warning(f"get_daily_bars_bulk() chunk failed: {e}")
                continue
            for symbol in chunk:
                if symbol in bars.data:
                    out[symbol] = [
                        {"t": b.timestamp, "o": float(b.open), "h": float(b.high),
                         "l": float(b.low), "c": float(b.close), "v": float(b.volume)}
                        for b in bars.data[symbol]
                    ]
        return out

    def get_historical_trades(self, symbol: str, start: datetime, end: datetime) -> list:
        """
        [SIMULATION-ONLY 2026-09-18] Raw tick trades for one symbol over
        a window, for replaying a session through stream.py's SymbolBuffer
        offline (see simulate.py) -- NOT used by the live path. Returns
        plain dicts (not SDK objects) so callers don't hold onto the
        SDK's own trade objects any longer than this call.
        """
        req = StockTradesRequest(symbol_or_symbols=symbol, start=start, end=end,
                                  feed=self._feed, limit=None)
        try:
            resp = self.hist_data.get_stock_trades(req)
            trades = resp.data.get(symbol, [])
        except Exception as e:
            log.warning(f"get_historical_trades({symbol}) failed: {e}")
            return []
        return [{"t": t.timestamp, "p": float(t.price), "s": float(t.size)} for t in trades]

    def get_historical_quotes(self, symbol: str, start: datetime, end: datetime) -> list:
        """[SIMULATION-ONLY 2026-09-18] Raw NBBO quotes for one symbol over
        a window -- same reasoning as get_historical_trades() above."""
        req = StockQuotesRequest(symbol_or_symbols=symbol, start=start, end=end,
                                  feed=self._feed, limit=None)
        try:
            resp = self.hist_data.get_stock_quotes(req)
            quotes = resp.data.get(symbol, [])
        except Exception as e:
            log.warning(f"get_historical_quotes({symbol}) failed: {e}")
            return []
        return [{"t": q.timestamp, "b": float(q.bid_price), "a": float(q.ask_price)}
                for q in quotes if q.bid_price and q.ask_price]

    def get_latest_quote(self, symbol: str):
        req = StockLatestQuoteRequest(symbol_or_symbols=symbol, feed=self._feed)
        try:
            q = self.hist_data.get_stock_latest_quote(req)
            return q.get(symbol)
        except Exception as e:
            log.warning(f"get_latest_quote({symbol}) failed: {e}")
            return None

    # ------------------------------------------------------------------
    # Account / trading
    # ------------------------------------------------------------------

    def get_asset(self, symbol: str):
        try:
            return self.trading.get_asset(symbol)
        except Exception as e:
            log.warning(f"get_asset({symbol}) failed: {e}")
            return None

    def get_account(self):
        return self.trading.get_account()

    def get_open_positions(self):
        try:
            return self.trading.get_all_positions()
        except Exception as e:
            log.warning(f"get_open_positions failed: {e}")
            return []

    def submit_market_order(self, symbol: str, qty: int, side: str):
        order_side = OrderSide.BUY if side.lower() == "buy" else OrderSide.SELL
        req = MarketOrderRequest(
            symbol=symbol,
            qty=qty,
            side=order_side,
            time_in_force=TimeInForce.DAY,
        )
        return self.trading.submit_order(req)

    def close_position(self, symbol: str):
        try:
            return self.trading.close_position(symbol)
        except Exception as e:
            log.error(f"close_position({symbol}) failed: {e}")
            raise

    def get_order(self, order_id: str):
        """Fetch a single order by id -- used to poll a pending close's
        fill status instead of blindly resubmitting it."""
        try:
            return self.trading.get_order_by_id(order_id)
        except Exception as e:
            log.warning(f"get_order({order_id}) failed: {e}")
            return None

    def get_open_orders(self, symbol: str):
        """Open (unfilled) orders for a single symbol, most recent first.
        Used to recover the order_id Alpaca is referencing in a wash-trade
        rejection so we can track/poll it instead of resubmitting."""
        req = GetOrdersRequest(status=QueryOrderStatus.OPEN, symbols=[symbol])
        try:
            return self.trading.get_orders(req)
        except Exception as e:
            log.warning(f"get_open_orders({symbol}) failed: {e}")
            return []

    def cancel_order(self, order_id: str) -> bool:
        try:
            self.trading.cancel_order_by_id(order_id)
            return True
        except Exception as e:
            log.warning(f"cancel_order({order_id}) failed: {e}")
            return False

    def get_position_qty(self, symbol: str):
        """Shares the broker actually holds for symbol: 0 if none, None if
        the lookup itself failed (caller must not treat that as flat)."""
        try:
            return float(self.trading.get_open_position(symbol).qty)
        except Exception as e:
            if parse_position_not_found_error(e):
                return 0.0
            log.warning(f"get_position_qty({symbol}) failed: {e}")
            return None

    def close_all_positions(self):
        try:
            return self.trading.close_all_positions(cancel_orders=True)
        except Exception as e:
            log.error(f"close_all_positions failed: {e}")
            raise


_client_singleton: Optional[AlpacaClient] = None


def get_client() -> AlpacaClient:
    global _client_singleton
    if _client_singleton is None:
        _client_singleton = AlpacaClient()
    return _client_singleton
