"""
simulate.py -- trade1 backtester.

Replays saved trade-by-trade market data through the REAL bot: it builds
monitor.SessionOrchestrator (the same class that trades live) with a
replay clock, a simulated broker and a replayed data stream, so every
buy/sell decision comes from the same code path and the same rule files
(breakout_rules.py, reversal_rules.py, exit_rules.py) as live trading.

    python simulate.py --days all                 every saved day
    python simulate.py --date 2026-09-22          one day
    python simulate.py --days all --entry breakout_rules --exit exit_rules
    python simulate.py --days all --rules-dir /some/dir --entry my_test_rules
        (--rules-dir: also import rule files from that folder -- for trying
         a variant without touching the real rule files)
    python simulate.py --fixed-entries sim/entries/j_20day.json --exit exit_rules
        EXIT test: buys exactly the listed entries (sim/rules/fixed_entries.py),
        only on their days and symbols, no slot limit -- compare exits on the
        same entries
    python simulate.py ... --set exit_rules.swing_minutes=3 --set exit_rules.trigger=touch
        override any config value for this run only (value parsed as JSON)

Data:
  candidates  sim/candidates/<date>.json (the 9:28 top-30 lists; 8/27-9/24
              are the rebuilt lists from the 20-day study, 9/25 the bot's
              real list), else data/candidates/<date>_scanner.json
  ticks       simulator.cache_dir in config.json (shared with screener/trade:
              <cache_dir>/<date>/<SYMBOL>.jsonl); missing symbols are
              downloaded from Alpaca on first use
Output:
  data/simulations/<name>.json   every trade + per-day summary
  data/simulations/<name>/decisions_<date>.jsonl   rule decisions (on change)
  logs/sim_*.log                 (kept apart from the live bot's log)

Simplifications (same as screener/trade's simulator): fills at the last
trade price (no slippage), fixed equity for sizing, no intraday rescans
(each day's 9:28 list is watched all day), 5-second poll grid.
"""
import argparse
import gc
import importlib
import json
import sys
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from datetime import datetime, timedelta, timezone
from pathlib import Path
from zoneinfo import ZoneInfo

BASE_DIR = Path(__file__).resolve().parent
sys.path.insert(0, str(BASE_DIR))

from config_loader import get_config

# keep simulator logging out of the live bot's log file -- must run before
# anything calls get_logger()
get_config().setdefault("logging", {})["filename_prefix"] = "sim"

import volatility
from stream import SymbolBuffer
import monitor

UTC = timezone.utc
ET = ZoneInfo(get_config()["schedule"]["timezone"])


def _et_dt(date_str: str, hhmmss: str) -> datetime:
    h, m, s = (int(x) for x in hhmmss.split(":"))
    d = datetime.fromisoformat(date_str)
    return datetime(d.year, d.month, d.day, h, m, s, tzinfo=ET)


# ----------------------------------------------------------------------
# tick data: cache + reader
# ----------------------------------------------------------------------
def _chunks(start, end, minutes=20):
    cur = start
    while cur < end:
        nxt = min(cur + timedelta(minutes=minutes), end)
        yield cur, nxt
        cur = nxt


def _cache_symbol(client, symbol, start, end, path: Path):
    """One 20-min chunk at a time -- this host has ~2 GB RAM."""
    tmp = path.with_suffix(".jsonl.tmp")
    n = 0
    with open(tmp, "w") as f:
        for a, b in _chunks(start, end):
            trades = client.get_historical_trades(symbol, a, b)
            quotes = client.get_historical_quotes(symbol, a, b)
            ev = [(t["t"], "trade", t["p"], t["s"]) for t in trades] + \
                 [(q["t"], "quote", q["b"], q["a"]) for q in quotes]
            ev.sort(key=lambda e: e[0])
            for ts, kind, x, y in ev:
                f.write(json.dumps([ts.isoformat(), kind, x, y]) + "\n")
            n += len(ev)
            del trades, quotes, ev
            gc.collect()
    tmp.rename(path)
    return n


def ensure_cache(client_fn, symbols, date_str, cache_dir: Path):
    cache_dir.mkdir(parents=True, exist_ok=True)
    todo = [s for s in symbols if not (cache_dir / f"{s}.jsonl").exists()]
    if not todo:
        return
    sched = get_config()["schedule"]
    start = _et_dt(date_str, sched["market_open_time"]).astimezone(UTC)
    end = _et_dt(date_str, sched["market_close_time"]).astimezone(UTC)
    client = client_fn()
    print(f"  downloading ticks for {len(todo)} symbols ...", file=sys.stderr)
    with ThreadPoolExecutor(max_workers=3) as pool:
        futs = {pool.submit(_cache_symbol, client, s, start, end, cache_dir / f"{s}.jsonl"): s for s in todo}
        for fut in as_completed(futs):
            try:
                print(f"    {futs[fut]}: {fut.result()} events", file=sys.stderr)
            except Exception as e:
                print(f"    {futs[fut]}: FAILED {e}", file=sys.stderr)


class EventReader:
    """Streams one symbol's cached ticks line by line."""

    def __init__(self, path: Path):
        self._f = open(path) if path.exists() else None
        self._next = None
        self._advance()

    def _advance(self):
        line = self._f.readline() if self._f else ""
        if not line:
            self._next = None
            return
        ts, kind, a, b = json.loads(line)
        self._next = (datetime.fromisoformat(ts), kind, a, b)

    def drain_up_to(self, t, on_trade, on_quote):
        while self._next is not None and self._next[0] <= t:
            ts, kind, a, b = self._next
            (on_trade if kind == "trade" else on_quote)(ts, a, b)
            self._advance()

    def close(self):
        if self._f:
            self._f.close()


def _daily_ref(client_fn, symbols, date_str, cache_dir):
    """Daily ATR and 5-day reference as of the morning of date_str
    (cached next to the ticks, same files screener/trade's simulator uses)."""
    day = _et_dt(date_str, "00:00:00").astimezone(UTC)
    out = {}
    for name, days_back, fn in (("daily_atr", 45, volatility.daily_atr),
                                ("ref5d", 15, volatility.five_day_reference)):
        p = cache_dir / f"{name}.json"
        data = json.loads(p.read_text()) if p.exists() else {}
        missing = [s for s in symbols if s not in data]
        if missing and not p.exists():
            bars = client_fn().get_daily_bars_bulk(missing, day - timedelta(days=days_back), day - timedelta(seconds=1))
            for s, bl in bars.items():
                v = fn([b for b in bl if b["t"] < day])
                if v:
                    data[s] = v
            p.write_text(json.dumps(data))
        out[name] = data
    return out


BAR_DELAY = timedelta(seconds=2)   # live: Alpaca's official bar lands ~1-2 s after the minute


def _minute_bars(client_fn, symbols, date_str, cache_dir, prefix="bars1m"):
    """Alpaca's OFFICIAL 1-min bars (what the live bar stream delivers),
    cached next to the ticks. Used for the watched symbols (the live stream
    upserts them into the buffer and forwards them to the rule modules) and
    for config streaming.benchmark_symbols (SPY/IWM, bars only)."""
    out = {}
    for s in symbols:
        p = cache_dir / f"{prefix}_{s}.json"
        if not p.exists():
            a = _et_dt(date_str, "09:30:00")
            raw = client_fn().get_minute_bars(s, start=a, end=_et_dt(date_str, "16:00:00"), limit=1000)
            p.write_text(json.dumps([[b.timestamp.isoformat(), float(b.open), float(b.high), float(b.low),
                                      float(b.close), float(b.volume)] for b in raw]))
        out[s] = [{"t": datetime.fromisoformat(r[0]), "o": r[1], "h": r[2], "l": r[3], "c": r[4], "v": r[5]}
                  for r in json.loads(p.read_text())]
    return out


# ----------------------------------------------------------------------
# simulated broker + stream (same interfaces monitor.py uses live)
# ----------------------------------------------------------------------
class SimPositionManager:
    def __init__(self, equity, trading_cfg, max_positions, clock):
        self.equity = equity
        self.cfg = trading_cfg
        self.max_positions = max_positions
        self.clock = clock
        self.positions = {}
        self.trades = []

    def has_available_slot(self):
        return len(self.positions) < self.max_positions

    def is_symbol_open(self, symbol):
        return symbol in self.positions

    def get_open_symbols(self):
        return list(self.positions)

    def closed_trades_today(self):
        return list(self.trades)

    def calculate_qty(self, entry_price, stop_price):
        # identical to position_manager.PositionManager.calculate_qty
        risk_per_share = entry_price - stop_price
        if risk_per_share <= 0 or not entry_price:
            return 0
        qty = int(self.equity * self.cfg["account_risk_pct_per_trade"] / 100.0 / risk_per_share)
        qty = min(qty, int(self.equity * self.cfg["max_position_notional_pct_of_equity"] / 100.0 / entry_price))
        return qty if qty >= self.cfg["min_shares"] else 0

    def enter_position(self, symbol, entry_price, stop_price, reasons=None, setup=None, plan=None):
        if self.is_symbol_open(symbol) or not self.has_available_slot():
            return False
        qty = self.calculate_qty(entry_price, stop_price)
        if qty <= 0:
            return False
        self.positions[symbol] = {
            "symbol": symbol, "qty": qty, "entry_price": entry_price, "stop_price": stop_price,
            "entry_time": self.clock().isoformat(), "reasons": reasons or [],
            "setup": setup, "plan": plan or {}, "peak_price": entry_price, "low_price": entry_price,
        }
        return True

    def mark(self, symbol, price):
        p = self.positions.get(symbol)
        if p:
            p["peak_price"] = max(p["peak_price"], price)
            p["low_price"] = min(p["low_price"], price)

    def exit_position(self, symbol, exit_price, reason):
        p = self.positions.pop(symbol, None)
        if not p:
            return
        peak, low = max(p["peak_price"], exit_price), min(p["low_price"], exit_price)
        pl = (exit_price - p["entry_price"]) * p["qty"]
        self.trades.append({
            **p, "exit_price": exit_price, "exit_time": self.clock().isoformat(), "exit_reason": reason,
            "status": "closed", "peak_price": peak, "low_price": low,
            "pl_pct": round((exit_price / p["entry_price"] - 1) * 100, 3), "pl_dollars": round(pl, 2),
            "mfe_dollars": round((peak - p["entry_price"]) * p["qty"], 2),
            "giveback_dollars": round((peak - p["entry_price"]) * p["qty"] - pl, 2),
        })


class SimStream:
    """Stand-in for stream.StreamManager: the same SymbolBuffer class the
    live stream uses, fed from saved ticks; forwards ticks and completed
    1-min bars to listeners exactly like the live stream does."""

    def __init__(self, symbols):
        sc = get_config().get("streaming", {})
        sub_s = sc.get("sub_minute_bucket_seconds", 30)
        sub_n = max(1, int(sc.get("sub_minute_buffer_minutes", 15) * 60 / sub_s))
        self.buffers = {s: SymbolBuffer(sub_bucket_seconds=sub_s, sub_maxlen=sub_n) for s in symbols}
        self.listeners = []
        self.bars_only = set()
        self._subscribed = set(symbols)
        self._last_bar = {}
        self.bench = {}          # benchmark symbol -> today's official 1-min bars
        self.official = {}       # watched symbol -> today's official 1-min bars
        self._released = {}      # symbol -> official bars released so far

    def _release(self, s, bl, t, buf=None):
        i = self._released.get(s, 0)
        while i < len(bl) and bl[i]["t"] + timedelta(minutes=1) + BAR_DELAY <= t:
            b = bl[i]
            if buf is not None:   # live _on_bar: upsert into the buffer, then notify
                buf.on_bar(b["o"], b["h"], b["l"], b["c"], b["v"], b["t"])
            self._notify("on_bar", s, dict(b))
            i += 1
        self._released[s] = i

    def feed_bench(self, t):
        for s, bl in self.bench.items():
            self._release(s, bl, t)

    def _notify(self, method, *args):
        for lst in self.listeners:
            fn = getattr(lst, method, None)
            if fn:
                fn(*args)

    def feed(self, symbol, reader: EventReader, t):
        buf = self.buffers[symbol]

        def on_trade(ts, price, size):
            buf.on_trade(price, size, ts)
            self._notify("on_trade", symbol, ts, price, size)

        def on_quote(ts, bid, ask):
            buf.on_quote(bid, ask, ts)
            self._notify("on_quote", symbol, ts, bid, ask, 0.0, 0.0)

        reader.drain_up_to(t, on_trade, on_quote)
        # official 1-min bars, as Alpaca's bar stream delivers them live
        self._release(symbol, self.official.get(symbol, []), t, buf)

    def get_bars(self, symbol):
        b = self.buffers.get(symbol)
        if b is None and symbol in self.bench:
            return self.bench[symbol][:self._released.get(symbol, 0)]
        return b.get_bars() if b else []

    def get_bars_sub(self, symbol):
        b = self.buffers.get(symbol)
        return b.get_bars_sub() if b else []

    def get_quote(self, symbol):
        b = self.buffers.get(symbol)
        return b.latest_quote if b else None

    def get_trade_imbalance(self, symbol, window_seconds=20.0):
        b = self.buffers.get(symbol)
        return b.get_trade_imbalance(window_seconds) if b else None

    def add_symbols(self, symbols):
        pass

    def remove_symbols(self, symbols):
        pass

    def stop(self):
        pass


# ----------------------------------------------------------------------
def load_candidates(date_str, override=None):
    for p in ([Path(override)] if override else []) + [
            BASE_DIR / "sim" / "candidates" / f"{date_str}.json",
            BASE_DIR / "data" / "candidates" / f"{date_str}_scanner.json"]:
        if p.exists():
            return json.loads(p.read_text()), p
    raise FileNotFoundError(f"no candidate list for {date_str}")


def all_days():
    return sorted(p.stem for p in (BASE_DIR / "sim" / "candidates").glob("*.json"))


def run_day(date_str, args, equity, client_fn, out_dir):
    cfg = get_config()
    sched = cfg["schedule"]
    cands, cpath = load_candidates(date_str, args.candidates_file)
    if args._fixed_syms is not None:
        want = args._fixed_syms.get(date_str, set())
        cands = [c for c in cands if c["symbol"] in want] + \
                [{"symbol": x, "metrics": {}} for x in sorted(want - {c["symbol"] for c in cands})]
    symbols = [c["symbol"] for c in cands]
    cache_dir = Path(cfg.get("simulator", {}).get("cache_dir", BASE_DIR / "data" / "cache")) / date_str
    ensure_cache(client_fn, symbols, date_str, cache_dir)
    ref = _daily_ref(client_fn, symbols, date_str, cache_dir)

    now = {"t": _et_dt(date_str, sched["market_open_time"]).astimezone(UTC)}
    clock = lambda: now["t"]
    pm = SimPositionManager(equity, cfg["trading"], args.max_positions or cfg["trading"]["max_positions"], clock)
    stream = SimStream(symbols)
    stream.bench = _minute_bars(client_fn, [b for b in cfg.get("streaming", {}).get("benchmark_symbols", [])
                                            if b not in symbols], date_str, cache_dir, prefix="bench")
    stream.official = _minute_bars(client_fn, symbols, date_str, cache_dir)
    dec_path = out_dir / f"decisions_{date_str}.jsonl"
    dec_f = open(dec_path, "w")

    def sink(rec):
        dec_f.write(json.dumps({"timestamp": now["t"].isoformat(), **rec}, default=str) + "\n")

    o = monitor.SessionOrchestrator(client=object(), position_mgr=pm, stream=stream, clock=clock,
                                    decision_sink=sink, live=False)
    for m in o.all_modules:
        m.load()   # fresh module state each simulated day
    o._apply_scan({
        "candidates": cands,
        "baselines": {c["symbol"]: c["metrics"].get("avg_daily_volume") for c in cands},
        "daily_atr": ref["daily_atr"], "ref5d": ref["ref5d"],
    })

    readers = {s: EventReader(cache_dir / f"{s}.jsonl") for s in symbols}
    t = now["t"]
    cutoff = _et_dt(date_str, sched["no_new_entries_after"]).astimezone(UTC)
    eod = _et_dt(date_str, sched["force_liquidate_time"]).astimezone(UTC)
    step = timedelta(seconds=sched["poll_interval_seconds"])
    while t <= eod:
        now["t"] = t
        stream.feed_bench(t)
        for s in symbols:
            stream.feed(s, readers[s], t)
        for s in pm.get_open_symbols():
            bars = stream.get_bars(s)
            if bars:
                pm.mark(s, bars[-1]["c"])
        o._update_open_positions()
        if pm.has_available_slot() and t < cutoff:
            o._scan_for_entries()
        t += step

    now["t"] = eod
    for m in o.all_modules:
        m.call("on_session_end")
    for s in pm.get_open_symbols():
        bars = stream.get_bars(s)
        o._sell(s, bars[-1]["c"] if bars else pm.positions[s]["entry_price"], "END_OF_DAY")
    for r in readers.values():
        r.close()
    dec_f.close()
    del stream, readers
    gc.collect()
    return pm.trades, str(cpath)


def main():
    ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--date", help="one day, YYYY-MM-DD")
    ap.add_argument("--days", help="'all' (every file in sim/candidates) or a comma list of dates")
    ap.add_argument("--candidates-file", help="use this candidate list (single --date only)")
    ap.add_argument("--entry", help="comma list of entry modules (default: config rules.entry_modules)")
    ap.add_argument("--exit", help="exit module (default: config rules.exit_module)")
    ap.add_argument("--rules-dir", help="extra folder to import rule files from")
    ap.add_argument("--max-positions", type=int, help="default: config trading.max_positions")
    ap.add_argument("--equity", type=float, help="sizing equity (default: the paper account's current equity)")
    ap.add_argument("--name", help="results name (default: from the modules and time)")
    ap.add_argument("--fixed-entries", help="exit test: buy exactly the entries in this file")
    ap.add_argument("--set", action="append", default=[], metavar="SECTION.KEY=VALUE",
                    help="override a config value for this run (repeatable)")
    args = ap.parse_args()

    for kv in args.set:
        k, v = kv.split("=", 1)
        try:
            v = json.loads(v)
        except ValueError:
            pass
        node = get_config()
        parts = k.split(".")
        for p in parts[:-1]:
            node = node.setdefault(p, {})
        node[parts[-1]] = v
    args._fixed_syms = None
    if args.fixed_entries:
        fe = Path(args.fixed_entries).resolve()
        sys.path.insert(0, str(BASE_DIR / "sim" / "rules"))
        get_config()["rules"]["entry_modules"] = ["fixed_entries"]
        get_config().setdefault("fixed_entries", {})["file"] = str(fe)
        args._fixed_syms = {}
        for e in json.load(open(fe))["entries"]:
            args._fixed_syms.setdefault(e["date"], set()).add(e["symbol"])
        if not args.max_positions:
            args.max_positions = 99
        if not args.date and not args.days:
            args.days = ",".join(sorted(args._fixed_syms))

    if args.rules_dir:
        sys.path.insert(0, str(Path(args.rules_dir).resolve()))
    rc = get_config().setdefault("rules", {})
    rc["hot_reload"] = False
    if args.entry and not args.fixed_entries:
        rc["entry_modules"] = [x.strip() for x in args.entry.split(",") if x.strip()]
    if args.exit:
        rc["exit_module"] = args.exit.strip()

    days = [args.date] if args.date else (all_days() if args.days == "all" else
                                           [d.strip() for d in (args.days or "").split(",") if d.strip()])
    if not days:
        ap.error("give --date or --days")

    _client = {}

    def client_fn():
        if "c" not in _client:
            from alpaca_client import get_client
            _client["c"] = get_client()
        return _client["c"]

    equity = args.equity or float(client_fn().get_account().equity)
    name = args.name or ("sim_" + "+".join(rc.get("entry_modules", [])) + "__" + str(rc.get("exit_module"))
                         + "_" + datetime.now().strftime("%Y%m%d_%H%M%S"))
    out_dir = BASE_DIR / "data" / "simulations" / name
    out_dir.mkdir(parents=True, exist_ok=True)
    print(f"trade1 simulator | entry={rc.get('entry_modules')} exit={rc.get('exit_module')} | "
          f"equity=${equity:,.2f} | {len(days)} day(s) -> {out_dir}", file=sys.stderr)

    results = {"name": name, "entry_modules": rc.get("entry_modules"), "exit_module": rc.get("exit_module"),
               "equity": equity, "days": {}}
    for d in days:
        t0 = time.time()
        trades, src = run_day(d, args, equity, client_fn, out_dir)
        results["days"][d] = {"candidates": src, "trades": trades}
        pl = sum(t["pl_dollars"] for t in trades)
        print(f"{d}: {len(trades)} trades, {sum(t['pl_dollars'] > 0 for t in trades)} winners, "
              f"${pl:+.2f}  [{time.time() - t0:.0f}s]", flush=True)
        for t in trades:
            print(f"    {t['symbol']:6} [{t['setup']}] {t['entry_time'][11:19]}Z @ {t['entry_price']:.2f} -> "
                  f"{t['exit_time'][11:19]}Z @ {t['exit_price']:.2f}  ${t['pl_dollars']:+.2f}  ({t['exit_reason'][:60]})",
                  flush=True)
        (BASE_DIR / "data" / "simulations" / f"{name}.json").write_text(json.dumps(results, indent=1, default=str))

    allt = [t for v in results["days"].values() for t in v["trades"]]
    wins = [t for t in allt if t["pl_dollars"] > 0]
    print(f"\nTOTAL {len(days)} days: {len(allt)} trades, {len(wins)} winners "
          f"(${sum(t['pl_dollars'] for t in wins):+.2f}), {len(allt) - len(wins)} losers "
          f"(${sum(t['pl_dollars'] for t in allt if t['pl_dollars'] <= 0):+.2f}), "
          f"net ${sum(t['pl_dollars'] for t in allt):+.2f}")


if __name__ == "__main__":
    main()
