from __future__ import annotations

import argparse
import json
import os
import time
from datetime import datetime, timezone
from decimal import Decimal
from pathlib import Path
from typing import Any

from src.market.context import DEFAULT_CONTEXT_UNIVERSE
from src.tools.trader_desk_report import build_report
from src.trader_desk.finance_bridge import SAFE_DEFAULT_TARGET, export_finance_snapshot

RUNTIME_DIR = Path(os.getenv("CTB_TRADER_DESK_SHADOW_RUNTIME_DIR", "runtime/experiments/trader_desk_shadow"))
STATE_PATH = RUNTIME_DIR / "state.json"
DECISIONS = RUNTIME_DIR / "decision_journal.jsonl"
OUTCOMES = RUNTIME_DIR / "outcome_journal.jsonl"
HEALTH = RUNTIME_DIR / "runtime_health.jsonl"
LATEST_REPORT = RUNTIME_DIR / "trader_desk_latest.json"
PID_PATH = RUNTIME_DIR / "runner.pid"

SHADOW_ACTIONS = {"watch_shadow", "tiny_live_candidate"}


def D(value: Any, default: str = "0") -> Decimal:
    try:
        return Decimal(str(value))
    except Exception:
        return Decimal(default)


def _now() -> str:
    return datetime.now(timezone.utc).isoformat()


def _json_default(value: Any) -> Any:
    if isinstance(value, Decimal):
        return str(value)
    if isinstance(value, datetime):
        return value.isoformat()
    return str(value)


def append_jsonl(path: Path, row: dict[str, Any]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    row.setdefault("timestamp", _now())
    path.open("a", encoding="utf-8").write(json.dumps(row, sort_keys=True, default=_json_default) + "\n")


def write_json(path: Path, payload: dict[str, Any]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    tmp = path.with_suffix(path.suffix + ".tmp")
    tmp.write_text(json.dumps(payload, indent=2, sort_keys=True, default=_json_default), encoding="utf-8")
    tmp.replace(path)


def load_state() -> dict[str, Any]:
    if STATE_PATH.exists():
        return json.loads(STATE_PATH.read_text(encoding="utf-8"))
    return {
        "schema_version": "trader_desk_shadow_state.v1",
        "positions": {},
        "closed_hypotheses": 0,
        "wins": 0,
        "losses": 0,
        "scratch": 0,
        "realized_r": "0",
        "live_order_allowed": False,
        "mainnet_signed_action": False,
        "paper_only": True,
    }


def save_state(state: dict[str, Any]) -> None:
    state["live_order_allowed"] = False
    state["mainnet_signed_action"] = False
    state["paper_only"] = True
    write_json(STATE_PATH, state)


def _mid_from_entry_zone(plan: dict[str, Any]) -> Decimal:
    zone = plan.get("entry_zone") or {}
    low = D(zone.get("low"))
    high = D(zone.get("high"))
    if low > 0 and high > 0:
        return (low + high) / Decimal("2")
    if low > 0:
        return low
    if high > 0:
        return high
    return Decimal("0")


def _tp1(plan_or_position: dict[str, Any]) -> Decimal:
    take_profit = plan_or_position.get("take_profit") or {}
    return D(take_profit.get("tp1"))


def _eligible_for_shadow(plan: dict[str, Any]) -> bool:
    return (
        str(plan.get("action")) in SHADOW_ACTIONS
        and str(plan.get("direction")) in {"long", "short"}
        and _mid_from_entry_zone(plan) > 0
        and D(plan.get("stop_loss")) > 0
        and _tp1(plan) > 0
    )


def _r_multiple(position: dict[str, Any], mark: Decimal) -> Decimal:
    entry = D(position.get("entry"))
    stop = D(position.get("stop_loss"))
    side = str(position.get("side"))
    risk = abs(entry - stop)
    if risk <= 0:
        return Decimal("0")
    pnl = mark - entry if side == "long" else entry - mark
    return pnl / risk


def _close_position(state: dict[str, Any], coin: str, position: dict[str, Any], *, mark: Decimal, reason: str, report_timestamp: str) -> dict[str, Any]:
    r_mult = _r_multiple(position, mark)
    state.setdefault("positions", {}).pop(coin, None)
    state["closed_hypotheses"] = int(state.get("closed_hypotheses", 0)) + 1
    if r_mult > Decimal("0.05"):
        state["wins"] = int(state.get("wins", 0)) + 1
        bucket = "win"
    elif r_mult < Decimal("-0.05"):
        state["losses"] = int(state.get("losses", 0)) + 1
        bucket = "loss"
    else:
        state["scratch"] = int(state.get("scratch", 0)) + 1
        bucket = "scratch"
    state["realized_r"] = str(D(state.get("realized_r")) + r_mult)
    row = {
        "event": "shadow_close",
        "coin": coin,
        "side": position.get("side"),
        "entry": position.get("entry"),
        "mark": str(mark),
        "stop_loss": position.get("stop_loss"),
        "tp1": position.get("take_profit", {}).get("tp1"),
        "r_multiple": str(r_mult.quantize(Decimal("0.0001"))),
        "bucket": bucket,
        "reason": reason,
        "setup": position.get("setup"),
        "score": position.get("score"),
        "source_action": position.get("action"),
        "opened_at": position.get("opened_at"),
        "report_timestamp": report_timestamp,
        "paper_only": True,
        "live_order_allowed": False,
        "mainnet_signed_action": False,
    }
    append_jsonl(OUTCOMES, row)
    return row


def reconcile_shadow_positions(state: dict[str, Any], plans: list[dict[str, Any]], *, report_timestamp: str, max_age_seconds: int = 24 * 3600) -> list[dict[str, Any]]:
    outcomes: list[dict[str, Any]] = []
    plans_by_coin = {str(plan.get("coin") or "").upper(): plan for plan in plans}
    positions = state.setdefault("positions", {})
    now_ts = time.time()
    for coin in list(positions):
        position = positions[coin]
        plan = plans_by_coin.get(coin)
        mark = _mid_from_entry_zone(plan or {})
        if mark <= 0:
            continue
        side = str(position.get("side"))
        stop = D(position.get("stop_loss"))
        tp1 = _tp1(position)
        opened_ts = float(position.get("opened_ts") or now_ts)
        reason: str | None = None
        if side == "long" and mark <= stop:
            reason = "shadow_stop_hit"
        elif side == "short" and mark >= stop:
            reason = "shadow_stop_hit"
        elif side == "long" and tp1 > 0 and mark >= tp1:
            reason = "shadow_tp1_hit"
        elif side == "short" and tp1 > 0 and mark <= tp1:
            reason = "shadow_tp1_hit"
        elif plan and str(plan.get("action")) == "no_trade":
            reason = "signal_invalidated"
        elif now_ts - opened_ts > max_age_seconds:
            reason = "max_shadow_age"
        if reason:
            outcomes.append(_close_position(state, coin, position, mark=mark, reason=reason, report_timestamp=report_timestamp))
    return outcomes


def open_shadow_positions(
    state: dict[str, Any],
    plans: list[dict[str, Any]],
    *,
    report_timestamp: str,
    max_open: int = 8,
    skip_coins: set[str] | None = None,
) -> list[dict[str, Any]]:
    opened: list[dict[str, Any]] = []
    positions = state.setdefault("positions", {})
    skip_coins = skip_coins or set()
    for plan in plans:
        coin = str(plan.get("coin") or "").upper()
        if not coin or coin in skip_coins or coin in positions or len(positions) >= max_open or not _eligible_for_shadow(plan):
            continue
        entry = _mid_from_entry_zone(plan)
        row = {
            "coin": coin,
            "side": str(plan.get("direction")),
            "entry": str(entry),
            "stop_loss": str(plan.get("stop_loss")),
            "take_profit": plan.get("take_profit") or {},
            "score": str(plan.get("score") or "0.00"),
            "action": str(plan.get("action")),
            "setup": plan.get("setup"),
            "opened_at": _now(),
            "opened_ts": time.time(),
            "report_timestamp": report_timestamp,
            "paper_only": True,
            "live_order_allowed": False,
            "mainnet_signed_action": False,
        }
        positions[coin] = row
        event = {"event": "shadow_open", **row}
        append_jsonl(OUTCOMES, event)
        opened.append(event)
    return opened


def _summarize_state(state: dict[str, Any]) -> dict[str, Any]:
    closed = int(state.get("closed_hypotheses", 0))
    wins = int(state.get("wins", 0))
    losses = int(state.get("losses", 0))
    return {
        "open_shadow_positions": len(state.get("positions") or {}),
        "closed_hypotheses": closed,
        "wins": wins,
        "losses": losses,
        "scratch": int(state.get("scratch", 0)),
        "win_rate_pct": str((Decimal(wins) / Decimal(closed) * Decimal("100")).quantize(Decimal("0.01"))) if closed else "0.00",
        "realized_r": str(D(state.get("realized_r")).quantize(Decimal("0.0001"))),
        "paper_only": True,
        "live_order_allowed": False,
        "mainnet_signed_action": False,
    }


def run_once(
    *,
    source_path: str | None = None,
    env: str = "mainnet",
    coins: tuple[str, ...] = DEFAULT_CONTEXT_UNIVERSE,
    max_plans: int = 8,
    export_finance: bool = True,
    finance_output: str | Path = SAFE_DEFAULT_TARGET,
) -> dict[str, Any]:
    if os.getenv("CTB_LIVE_TRADING_ALLOWED", "").lower() == "true":
        raise PermissionError("trader desk shadow runner refuses CTB_LIVE_TRADING_ALLOWED=true")
    payload = build_report(source_path=source_path, env=env, coins=coins, max_plans=max_plans)
    payload["live_order_allowed"] = False
    payload["mainnet_signed_action"] = False
    report = payload.get("report") or {}
    report_timestamp = str(report.get("timestamp") or _now())
    plans = list(report.get("plans") or [])
    write_json(LATEST_REPORT, payload)
    for plan in plans:
        append_jsonl(DECISIONS, {
            "event": "desk_decision",
            "report_timestamp": report_timestamp,
            "coin": plan.get("coin"),
            "action": plan.get("action"),
            "direction": plan.get("direction"),
            "score": plan.get("score"),
            "setup": plan.get("setup"),
            "blockers": plan.get("blockers") or [],
            "requires_preflight": bool(plan.get("requires_preflight", True)),
            "paper_only": True,
            "live_order_allowed": False,
            "mainnet_signed_action": False,
        })
    state = load_state()
    closed = reconcile_shadow_positions(state, plans, report_timestamp=report_timestamp)
    closed_coins = {str(row.get("coin") or "").upper() for row in closed}
    opened = open_shadow_positions(state, plans, report_timestamp=report_timestamp, skip_coins=closed_coins)
    save_state(state)
    finance_result = None
    if export_finance:
        finance_result = export_finance_snapshot(payload, target=finance_output)
    result = {
        "status": "ok",
        "runtime_dir": str(RUNTIME_DIR),
        "latest_report": str(LATEST_REPORT),
        "report_timestamp": report_timestamp,
        "decisions": len(plans),
        "opened": len(opened),
        "closed": len(closed),
        "summary": report.get("summary") or {},
        "shadow_performance": _summarize_state(state),
        "finance_export": {k: v for k, v in (finance_result or {}).items() if k != "snapshot"} if finance_result else None,
        "paper_only": True,
        "live_order_allowed": False,
        "mainnet_signed_action": False,
    }
    append_jsonl(HEALTH, {"event": "shadow_tick", **result})
    return result


def main(argv: list[str] | None = None) -> int:
    parser = argparse.ArgumentParser(description="Run the read-only JARVIS Trader Desk shadow journaler.")
    parser.add_argument("--from-report", help="Existing market_confluence_report JSON path. If omitted, build a fresh read-only report.")
    parser.add_argument("--env", choices=["mainnet", "testnet"], default="mainnet")
    parser.add_argument("--coins", default=",".join(DEFAULT_CONTEXT_UNIVERSE))
    parser.add_argument("--max-plans", type=int, default=8)
    parser.add_argument("--iterations", type=int, default=1)
    parser.add_argument("--interval-seconds", type=int, default=300)
    parser.add_argument("--no-export-finance", action="store_true")
    parser.add_argument("--finance-output", default=str(SAFE_DEFAULT_TARGET))
    parser.add_argument("--json", action="store_true")
    args = parser.parse_args(argv)

    RUNTIME_DIR.mkdir(parents=True, exist_ok=True)
    PID_PATH.write_text(str(os.getpid()), encoding="utf-8")
    coins = tuple(c.strip().upper() for c in args.coins.split(",") if c.strip())
    results: list[dict[str, Any]] = []
    for i in range(args.iterations):
        try:
            results.append(run_once(
                source_path=args.from_report,
                env=args.env,
                coins=coins,
                max_plans=args.max_plans,
                export_finance=not args.no_export_finance,
                finance_output=args.finance_output,
            ))
        except Exception as exc:
            row = {
                "event": "shadow_tick_degraded",
                "status": "degraded",
                "error_type": type(exc).__name__,
                "message": str(exc)[:300],
                "entries_blocked": True,
                "paper_only": True,
                "live_order_allowed": False,
                "mainnet_signed_action": False,
            }
            append_jsonl(HEALTH, row)
            results.append(row)
        if i < args.iterations - 1:
            time.sleep(args.interval_seconds)
    payload = {
        "status": "degraded" if any(r.get("status") == "degraded" for r in results) else "ok",
        "runtime_dir": str(RUNTIME_DIR),
        "results": results,
        "paper_only": True,
        "live_order_allowed": False,
        "mainnet_signed_action": False,
    }
    if args.json:
        print(json.dumps(payload, indent=2, sort_keys=True, default=_json_default))
    else:
        latest = results[-1] if results else {}
        print(
            "JARVIS Trader Desk Shadow Runner: "
            f"status={payload['status']}, decisions={latest.get('decisions', 0)}, "
            f"opened={latest.get('opened', 0)}, closed={latest.get('closed', 0)}, "
            "paper_only=true, live=false, signed=false"
        )
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
