"""
monitor.py

Bot CORE orchestrator for screener/trade1. Contains NO entry or exit
strategy: every trade decision comes from the rule modules named in
config.json "rules" (breakout_rules.py, reversal_rules.py,
exit_rules.py -- contract in rules_api.py). Work on entries/exits
happens in those files only.

    START -> wait for schedule.premarket_scan_time
          -> universe.py + scanner.py (today's candidate list)
          -> wait for market open, then subscribe the stream (only after
             the open, so premarket prints never land in session bars)
          -> every schedule.poll_interval_seconds:
               apply a finished background rescan (scanner.rescan_*)
               each OPEN position   -> exit module.evaluate() -> sell?
               each WATCHED symbol  -> entry modules in order  -> buy?
                  (while a slot is free, before no_new_entries_after)
          -> force_liquidate_time: sell everything, write the summary

What stays in the core (not strategy -- plumbing and safety):
  scanning/rescans, streaming, order placement, position sizing
  (trading.account_risk_pct_per_trade to the rule's stop, notional cap),
  max_positions, end-of-day liquidation, and a fallback that sells at
  the entry stop ONLY if the exit module can't load or raises.

Run directly:  python monitor.py
A singleton lock (data_store.acquire_singleton_lock) stops two copies
running from this folder. NOTE: screener/trade uses the SAME Alpaca
paper account -- never run both bots at the same time.
"""

import importlib
import os
import sys
import time
import signal
import threading
from datetime import datetime, timezone

from config_loader import get_config, get_env
from logger_setup import get_logger
import market_time
import data_store
import universe
import scanner
from position_manager import PositionManager
from stream import StreamManager
from alpaca_client import get_client
from rules_api import MarketView, EntryDecision, ExitDecision

log = get_logger("monitor")

BASE_DIR = os.path.dirname(os.path.abspath(__file__))
PID_PATH = os.path.join(BASE_DIR, "monitor.pid")


class RuleModule:
    """One loaded rule file. Wraps every call so a rule error is logged
    (once per symbol+method, to keep the log readable) and never raises
    into the core. Optionally re-imports the file when it changes."""

    def __init__(self, name: str, hot_reload: bool):
        self.name = name
        self.hot_reload = hot_reload
        self.mod = None
        self._mtime = None
        self._errors_logged = set()
        self.load()

    def _path(self):
        f = getattr(self.mod, "__file__", None)
        return f or os.path.join(BASE_DIR, f"{self.name}.py")

    def load(self) -> bool:
        try:
            if self.mod is None:
                self.mod = importlib.import_module(self.name)
            else:
                self.mod = importlib.reload(self.mod)
            self._mtime = os.path.getmtime(self._path())
            self._errors_logged.clear()
            log.info(f"[RULES] loaded {self.name}")
            return True
        except Exception:
            log.exception(f"[RULES] failed to load {self.name}"
                          + (" -- keeping the previous version" if self.mod is not None else ""))
            return False

    def maybe_reload(self):
        if not self.hot_reload:
            return
        try:
            mtime = os.path.getmtime(self._path())
        except OSError:
            return
        if mtime != self._mtime:
            self._mtime = mtime   # a broken file is not retried every poll
            log.info(f"[RULES] {self.name}.py changed on disk -- reloading")
            self.load()

    def has(self, method: str) -> bool:
        return self.mod is not None and callable(getattr(self.mod, method, None))

    def call(self, method: str, *args, symbol: str = "", default=None):
        if not self.has(method):
            return default
        try:
            return getattr(self.mod, method)(*args)
        except Exception:
            key = (method, symbol)
            if key not in self._errors_logged:
                self._errors_logged.add(key)
                log.exception(f"[RULES] {self.name}.{method}({symbol}) raised -- ignored "
                              f"(logged once per symbol)")
            return default

    def cfg(self) -> dict:
        return get_config().get(self.name, {})

    def enabled(self) -> bool:
        return self.mod is not None and self.cfg().get("enabled", True)


class _TickFanout:
    """Stream listener that forwards raw ticks to whichever rule modules
    define on_trade/on_quote/on_bar/seed_bars (looked up per call, so a
    hot reload takes effect immediately)."""

    def __init__(self, modules):
        self.modules = modules

    def _fan(self, method, *args):
        for m in self.modules:
            if m.has(method):
                m.call(method, *args, symbol=str(args[0]) if args else "")

    def on_trade(self, *a):
        self._fan("on_trade", *a)

    def on_quote(self, *a):
        self._fan("on_quote", *a)

    def on_bar(self, *a):
        self._fan("on_bar", *a)

    def seed_bars(self, *a):
        self._fan("seed_bars", *a)


class SessionOrchestrator:
    """Live by default. simulate.py builds the SAME class with a replay
    clock, a simulated broker (position_mgr) and a replayed stream, so a
    backtest runs exactly the decision code that trades live."""

    def __init__(self, client=None, position_mgr=None, stream=None, clock=None,
                 decision_sink=None, live=True):
        self.cfg = get_config()
        self.live = live
        if live:
            get_env().validate()

        self.client = client or (get_client() if live else None)
        self.position_mgr = position_mgr or PositionManager()
        self.stream = stream or StreamManager()
        self._now = clock or (lambda: datetime.now(timezone.utc))
        self._decision_sink = decision_sink or data_store.append_decision_record

        rc = self.cfg.get("rules", {})
        hot = bool(rc.get("hot_reload", False))
        self.entry_modules = [RuleModule(n, hot) for n in rc.get("entry_modules", [])]
        self.exit_module = RuleModule(rc["exit_module"], hot) if rc.get("exit_module") else None
        all_mods = self.entry_modules + ([self.exit_module] if self.exit_module else [])
        self.stream.listeners.append(_TickFanout(all_mods))
        self.all_modules = all_mods

        self.candidates = []           # today's ranked scanner output
        self._scan_by_symbol = {}      # symbol -> scanner record
        self._levels = {}              # symbol -> {label: level}
        self._avg_vol_baseline = {}    # symbol -> prior-session volume
        # rule-module state, owned by the modules, kept here between polls
        self._entry_state = {}         # (module, symbol) -> dict
        self._exit_state = {}          # symbol -> dict (per open position)
        self._last_logged = {}         # (kind, symbol) -> (state, reason): decision log on change only
        self._shutdown = False
        # Intraday rescans run on a background thread (15-40 s of REST
        # calls would stall exits); applied on the main thread next poll.
        self._last_scan_started = None
        self._early_rescan_done = False
        self._rescan_thread = None
        self._pending_scan = None
        self._pending_lock = threading.Lock()

        if live:
            signal.signal(signal.SIGTERM, self._handle_signal)
            signal.signal(signal.SIGINT, self._handle_signal)

    # ------------------------------------------------------------------
    def _handle_signal(self, signum, frame):
        log.info(f"[SHUTDOWN] Received signal {signum}, shutting down gracefully")
        self._shutdown = True

    # ------------------------------------------------------------------
    def run(self):
        log.info(f"[START] monitor.py (trade1 core) starting | mode={self.cfg['mode']['execution_mode']} | "
                 f"entry modules={[m.name for m in self.entry_modules]} | "
                 f"exit module={self.exit_module.name if self.exit_module else None}")
        if self.exit_module is None or self.exit_module.mod is None:
            log.error("[RULES] NO working exit module -- positions will only be sold at their entry stop "
                      "(core fallback) or at end of day")

        if not market_time.is_weekday():
            log.info("[START] Not a weekday, exiting")
            return

        self._wait_for_premarket_scan()
        if self._shutdown:
            return

        self._run_scan()
        self._wait_for_market_open()
        self._start_streaming()
        self._main_trading_loop()
        self._end_of_day_liquidation()

        log.info("[SHUTDOWN] Session complete")

    # ------------------------------------------------------------------
    def _wait_for_premarket_scan(self):
        while not market_time.is_premarket_scan_time() and not self._shutdown:
            time.sleep(5)

    def _compute_scan(self, candidates_file_suffix: str = "") -> dict:
        """Pure scan -- returns results without touching live state, so it
        can run on the rescan thread."""
        all_symbols = universe.get_universe_symbols(self.client)
        log.info(f"[SCAN] Universe size after asset filtering: {len(all_symbols)}")

        baselines, prev_day_highs, prev_closes = {}, {}, {}
        prefiltered = universe.prefilter_by_snapshot(
            self.client, all_symbols,
            baseline_out=baselines,
            prev_day_high_out=prev_day_highs,
            prev_close_out=prev_closes,
        )
        log.info(f"[SCAN] Prefiltered to {len(prefiltered)} symbols")

        daily_atr, ref5d = {}, {}
        range_20d_highs = scanner.compute_range_20d_high(self.client, prefiltered, daily_atr_out=daily_atr,
                                                         ref5d_out=ref5d)
        candidates = scanner.scan(
            self.client, prefiltered=prefiltered,
            volume_baselines=baselines,
            prev_day_highs=prev_day_highs, prev_closes=prev_closes,
            range_20d_highs=range_20d_highs,
            candidates_file_suffix=candidates_file_suffix,
        )
        return {"candidates": candidates, "baselines": baselines, "daily_atr": daily_atr, "ref5d": ref5d}

    def _apply_scan(self, result: dict):
        """Main-thread only. Swaps in the new candidate list; keeps each
        already-known symbol's levels from the first scan it appeared in
        (a mid-session scan's 'premarket_high' is really the last-6h high)."""
        self._avg_vol_baseline.update(result["baselines"])
        for c in result["candidates"]:
            m = c["metrics"]
            lv = self._levels.setdefault(c["symbol"], {
                "premarket_high": m.get("premarket_high"),
                "prev_day_high": m.get("previous_day_high"),
                "range_20d_high": m.get("range_20d_high"),
            })
            lv["daily_atr"] = result.get("daily_atr", {}).get(c["symbol"])
            lv.update(result.get("ref5d", {}).get(c["symbol"], {}))
            self._scan_by_symbol[c["symbol"]] = c
        old_syms = [c["symbol"] for c in self.candidates]
        self.candidates = result["candidates"]
        log.info(f"[SCAN] {len(self.candidates)} candidates selected: "
                 f"{[c['symbol'] for c in self.candidates]}")
        return old_syms

    def _run_scan(self):
        self._last_scan_started = time.time()
        self._apply_scan(self._compute_scan())

    def _maybe_rescan(self):
        """Refresh the candidate list every scanner.rescan_interval_minutes
        (0 = off) until no_new_entries_after. Dropped symbols are
        unsubscribed unless a position is open in them; new ones are
        backfilled and subscribed."""
        interval = self.cfg["scanner"].get("rescan_interval_minutes", 0)
        if not interval:
            return

        with self._pending_lock:
            result, self._pending_scan = self._pending_scan, None
        if result is not None:
            old_syms = self._apply_scan(result)
            new_syms = [c["symbol"] for c in self.candidates]
            open_syms = set(self.position_mgr.get_open_symbols())
            added = [s for s in new_syms if s not in old_syms]
            dropped = [s for s in self.stream._subscribed
                       if s not in new_syms and s not in open_syms and s not in self._benchmarks()]
            log.info(f"[RESCAN] applied: +{len(added)} {added} / -{len(dropped)} {dropped}")
            self.stream.add_symbols(added)
            self.stream.remove_symbols(dropped)
            for key in [k for k in self._entry_state if k[1] in dropped]:
                self._entry_state.pop(key, None)

        if market_time.is_new_entries_cutoff():
            return
        if self._rescan_thread is not None and self._rescan_thread.is_alive():
            return
        # first_rescan_after_open_minutes forces one rescan that many minutes
        # after the open; rescan_interval_first_hour_minutes replaces the
        # interval during the first hour. 0 = off.
        sc = self.cfg["scanner"]
        mins_open = market_time.minutes_since_open()
        early = sc.get("first_rescan_after_open_minutes", 0)
        first_hour = sc.get("rescan_interval_first_hour_minutes", 0)
        if first_hour and 0 <= mins_open < 60:
            interval = min(interval, first_hour)
        due_early = early and not self._early_rescan_done and mins_open >= early
        if due_early:
            self._early_rescan_done = True
        elif self._last_scan_started and time.time() - self._last_scan_started < interval * 60:
            return

        self._last_scan_started = time.time()
        suffix = "_" + market_time.now_et().strftime("%H%M")

        def worker():
            try:
                res = self._compute_scan(candidates_file_suffix=suffix)
            except Exception:
                log.exception("[RESCAN] scan failed; keeping current candidate list")
                return
            with self._pending_lock:
                self._pending_scan = res

        log.info(f"[RESCAN] starting background rescan ({interval}-min interval)")
        self._rescan_thread = threading.Thread(target=worker, daemon=True)
        self._rescan_thread.start()

    def _benchmarks(self) -> list:
        return list(self.cfg.get("streaming", {}).get("benchmark_symbols", []))

    def _start_streaming(self):
        symbols = [c["symbol"] for c in self.candidates]
        symbols += [b for b in self._benchmarks() if b not in symbols]
        self.stream.bars_only = set(self._benchmarks()) - {c["symbol"] for c in self.candidates}
        if symbols:
            self.stream.start(symbols)

    def _wait_for_market_open(self):
        while not market_time.is_past_market_open() and not self._shutdown:
            time.sleep(2)

    # ------------------------------------------------------------------
    def _main_trading_loop(self):
        interval = self.cfg["schedule"]["poll_interval_seconds"]

        while not self._shutdown and not market_time.is_force_liquidate_time():
            # one bad poll cycle must never kill the bot with positions open
            try:
                for m in self.all_modules:
                    m.maybe_reload()
                self._maybe_rescan()
                self._update_open_positions()

                if self.position_mgr.has_available_slot() and not market_time.is_new_entries_cutoff():
                    self._scan_for_entries()
            except Exception:
                log.exception("[LOOP] poll cycle failed; continuing")

            time.sleep(interval)

    # ------------------------------------------------------------------
    def _minutes_since_open(self) -> float:
        now_et = self._now().astimezone(market_time._tz())
        h, m, sec = (int(x) for x in self.cfg["schedule"]["market_open_time"].split(":"))
        return (now_et - now_et.replace(hour=h, minute=m, second=sec, microsecond=0)).total_seconds() / 60.0

    def _view(self, symbol: str, bars: list, module: RuleModule, last_trade=None) -> MarketView:
        levels = dict(self._levels.get(symbol, {}))
        levels["session_high"] = max(b["h"] for b in bars)
        q = self.stream.get_quote(symbol)
        return MarketView(
            symbol=symbol,
            now=self._now(),
            price=bars[-1]["c"],
            bars=bars,
            bars_sub=self.stream.get_bars_sub(symbol),
            quote=(q[0], q[1]) if q else None,
            levels=levels,
            volume_baseline=self._avg_vol_baseline.get(symbol),
            scan=self._scan_by_symbol.get(symbol, {}),
            minutes_since_open=self._minutes_since_open(),
            slots_free=self.position_mgr.has_available_slot(),
            last_trade=last_trade,
            cfg=module.cfg() if module else {},
            benchmarks={b: self.stream.get_bars(b) for b in self._benchmarks()},
            _imbalance=lambda w, s=symbol: self.stream.get_trade_imbalance(s, window_seconds=w),
        )

    def _log_decision(self, kind: str, symbol: str, state: str, reason: str, record: dict, force=False):
        """Decision log (data/decisions/<date>.jsonl), written when a
        symbol's (state, reason) changes -- not every 5-second poll."""
        key = (kind, symbol)
        sig = (state, reason)
        if not force and self._last_logged.get(key) == sig:
            return
        self._last_logged[key] = sig
        self._decision_sink({"symbol": symbol, "kind": kind, "state": state,
                             "reason": reason, **record})

    def _sell(self, symbol: str, price: float, reason: str):
        position = dict(self.position_mgr.positions.get(symbol, {}))
        self.position_mgr.exit_position(symbol, price, reason)
        if not self.position_mgr.is_symbol_open(symbol):
            self._exit_state.pop(symbol, None)
            for m in self.all_modules:
                m.call("on_exit", symbol, position, reason, symbol=symbol)

    def _update_open_positions(self):
        for symbol in self.position_mgr.get_open_symbols():
            bars = self.stream.get_bars(symbol)
            if not bars:
                continue
            p = self.position_mgr.positions[symbol]
            price = bars[-1]["c"]

            em = self.exit_module
            decision = None
            if em is not None and em.enabled():
                view = self._view(symbol, bars, em)
                decision = em.call("evaluate", view, dict(p), self._exit_state.setdefault(symbol, {}),
                                   symbol=symbol)
                if decision is not None and not isinstance(decision, ExitDecision):
                    log.error(f"[RULES] {em.name}.evaluate returned {type(decision).__name__}, "
                              f"not ExitDecision -- ignored")
                    decision = None

            if decision is None:
                # SAFETY NET only: the exit module is missing/disabled/broken.
                stop = p.get("stop_price")
                if stop is not None and price <= stop:
                    log.warning(f"[EXIT] {symbol} core fallback: no exit decision from the exit module, "
                                f"price ${price:.2f} <= entry stop ${stop:.2f}")
                    self._sell(symbol, price, f"core fallback stop (exit module unavailable) ${stop:.2f}")
                continue

            self._log_decision("exit", symbol, decision.state, decision.reason,
                               {"metrics": decision.metrics}, force=decision.should_exit)
            if decision.should_exit:
                self._sell(symbol, price, decision.reason or decision.state)

    def _scan_for_entries(self):
        last_trade = {}
        for t in self.position_mgr.closed_trades_today():
            last_trade[t["symbol"]] = t

        for c in self.candidates:
            symbol = c["symbol"]
            if self.position_mgr.is_symbol_open(symbol):
                continue
            if not self.position_mgr.has_available_slot():
                break
            bars = self.stream.get_bars(symbol)
            if not bars:
                continue

            for m in self.entry_modules:
                if not m.enabled():
                    continue
                state = self._entry_state.setdefault((m.name, symbol), {})
                view = self._view(symbol, bars, m, last_trade=last_trade.get(symbol))
                d = m.call("evaluate", view, state, symbol=symbol)
                if d is None:
                    continue
                if not isinstance(d, EntryDecision):
                    log.error(f"[RULES] {m.name}.evaluate returned {type(d).__name__}, not EntryDecision -- ignored")
                    continue
                self._log_decision(f"entry:{m.name}", symbol, d.state, d.reason,
                                   {"reasons": d.reasons, "plan": d.plan, "metrics": d.metrics},
                                   force=d.should_enter)
                if not d.should_enter:
                    continue
                if d.stop is None or not d.stop < view.price:
                    log.warning(f"[ENTRY] {symbol} {m.name} wanted to buy with an invalid stop "
                                f"({d.stop} vs price {view.price}) -- skipped")
                    continue
                setup = getattr(m.mod, "NAME", m.name)
                entered = self.position_mgr.enter_position(
                    symbol, view.price, d.stop, reasons=d.reasons or [d.reason], setup=setup, plan=d.plan)
                if entered:
                    for mm in self.entry_modules:
                        self._entry_state.pop((mm.name, symbol), None)
                    self._exit_state[symbol] = {}
                    pos = dict(self.position_mgr.positions.get(symbol, {}))
                    for mm in self.all_modules:
                        mm.call("on_entry", symbol, pos, symbol=symbol)
                break   # one decision per symbol per poll

    # ------------------------------------------------------------------
    def _end_of_day_liquidation(self):
        log.info("[EOD] Force-liquidating all open positions")
        for m in self.all_modules:
            m.call("on_session_end")
        max_wait = self.cfg["schedule"]["eod_liquidation_max_wait_seconds"]
        poll_interval = self.cfg["schedule"]["poll_interval_seconds"]
        deadline = time.time() + max_wait

        while time.time() < deadline:
            open_symbols = self.position_mgr.get_open_symbols()
            if not open_symbols:
                break
            for symbol in open_symbols:
                bars = self.stream.get_bars(symbol)
                current_price = bars[-1]["c"] if bars else self.position_mgr.positions[symbol]["entry_price"]
                self._sell(symbol, current_price, "END_OF_DAY")
            time.sleep(poll_interval)
        else:
            still_open = self.position_mgr.get_open_symbols()
            if still_open:
                log.warning(f"[EOD] {len(still_open)} position(s) still not resolved after "
                            f"{max_wait}s: {sorted(still_open)}")

        self.stream.stop()

        trades = data_store.load_today_trades()
        data_store.write_trades_summary(trades)
        total_pl = sum(t.get("pl_dollars", 0) for t in trades if t.get("status") == "closed")
        wins = sum(1 for t in trades if t.get("status") == "closed" and t.get("pl_dollars", 0) > 0)
        losses = sum(1 for t in trades if t.get("status") == "closed" and t.get("pl_dollars", 0) <= 0)
        log.info(f"[EOD] Session summary: trades={len(trades)} wins={wins} losses={losses} "
                 f"total_P/L=${total_pl:.2f}")


def main():
    try:
        with data_store.acquire_singleton_lock():
            data_store.ensure_dirs()
            with open(PID_PATH, "w") as f:
                f.write(str(os.getpid()))
            try:
                orchestrator = SessionOrchestrator()
                orchestrator.run()
            finally:
                try:
                    os.remove(PID_PATH)
                except OSError:
                    pass
    except RuntimeError as e:
        log.error(str(e))
        sys.exit(1)


if __name__ == "__main__":
    main()
