from __future__ import annotations

import argparse
import json
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from statistics import mean
from typing import Any

from config import BotConfig, RuntimePaths


@dataclass(frozen=True)
class TradeContextMatch:
    coin: str
    entry_ts: datetime
    exit_ts: datetime
    context_ts: datetime | None
    context_tags: tuple[str, ...]
    context_reliability: float | None
    realized_pnl_usd: float
    exit_reason: str
    time_in_trade_minutes: int


@dataclass(frozen=True)
class TagStats:
    tag: str
    sample_size: int
    win_rate: float
    avg_pnl_usd: float
    total_pnl_usd: float
    avg_time_in_trade_minutes: float


@dataclass(frozen=True)
class CorrelationResult:
    matches: list[TradeContextMatch]
    tag_stats: dict[str, TagStats]
    closed_trade_count: int
    open_entry_count: int


def _parse_ts(value: Any) -> datetime:
    dt = datetime.fromisoformat(str(value))
    if dt.tzinfo is None:
        return dt.replace(tzinfo=timezone.utc)
    return dt.astimezone(timezone.utc)


def _read_jsonl(path: str | Path) -> list[dict[str, Any]]:
    p = Path(path).expanduser()
    if not p.exists():
        return []
    rows = []
    for line in p.read_text(encoding="utf-8").splitlines():
        line = line.strip()
        if line:
            rows.append(json.loads(line))
    return rows


def _best_context(entry: dict[str, Any], contexts: list[dict[str, Any]]) -> dict[str, Any] | None:
    coin = str(entry.get("coin", "")).upper()
    entry_ts = _parse_ts(entry.get("ts"))
    candidates = []
    for ctx in contexts:
        if str(ctx.get("coin", "")).upper() != coin:
            continue
        try:
            ctx_ts = _parse_ts(ctx.get("ts"))
        except Exception:
            continue
        if ctx_ts <= entry_ts:
            candidates.append((ctx_ts, ctx))
    if not candidates:
        return None
    return max(candidates, key=lambda x: x[0])[1]


def _pair_closed_trades(events: list[dict[str, Any]]) -> tuple[list[tuple[dict[str, Any], dict[str, Any]]], int]:
    open_entries: dict[str, list[dict[str, Any]]] = {}
    pairs: list[tuple[dict[str, Any], dict[str, Any]]] = []
    sorted_events = sorted(events, key=lambda e: str(e.get("ts", "")))
    for event in sorted_events:
        if event.get("dry_run") is True:
            continue
        coin = str(event.get("coin", "")).upper()
        if event.get("event_type") == "entry":
            open_entries.setdefault(coin, []).append(event)
        elif event.get("event_type") == "exit":
            entries = open_entries.get(coin) or []
            if entries:
                entry = entries.pop(0)
                pairs.append((entry, event))
    open_count = sum(len(v) for v in open_entries.values())
    return pairs, open_count


def _aggregate_tag_stats(matches: list[TradeContextMatch]) -> dict[str, TagStats]:
    buckets: dict[str, list[TradeContextMatch]] = {}
    for match in matches:
        for tag in match.context_tags:
            buckets.setdefault(tag, []).append(match)
    stats: dict[str, TagStats] = {}
    for tag, rows in buckets.items():
        pnls = [r.realized_pnl_usd for r in rows]
        stats[tag] = TagStats(
            tag=tag,
            sample_size=len(rows),
            win_rate=round(sum(1 for p in pnls if p > 0) / len(pnls), 4) if pnls else 0.0,
            avg_pnl_usd=round(mean(pnls), 8) if pnls else 0.0,
            total_pnl_usd=round(sum(pnls), 8),
            avg_time_in_trade_minutes=round(mean(r.time_in_trade_minutes for r in rows), 4) if rows else 0.0,
        )
    return dict(sorted(stats.items(), key=lambda kv: (kv[1].sample_size, kv[1].total_pnl_usd), reverse=True))


def correlate_trades_with_context(journal_path: str | Path, market_context_path: str | Path) -> CorrelationResult:
    events = _read_jsonl(journal_path)
    contexts = _read_jsonl(market_context_path)
    pairs, open_count = _pair_closed_trades(events)
    matches: list[TradeContextMatch] = []
    for entry, exit_event in pairs:
        context = _best_context(entry, contexts)
        entry_ts = _parse_ts(entry.get("ts"))
        exit_ts = _parse_ts(exit_event.get("ts"))
        context_ts = _parse_ts(context["ts"]) if context else None
        tags = tuple(str(x) for x in (context or {}).get("tags", []))
        pnl = float(exit_event.get("realized_pnl_usd", 0.0))
        minutes = int((exit_ts - entry_ts).total_seconds() // 60)
        reliability_value = context.get("reliability_score") if context else None
        matches.append(TradeContextMatch(
            coin=str(entry.get("coin", "")).upper(),
            entry_ts=entry_ts,
            exit_ts=exit_ts,
            context_ts=context_ts,
            context_tags=tags,
            context_reliability=float(reliability_value) if reliability_value is not None else None,
            realized_pnl_usd=pnl,
            exit_reason=str(exit_event.get("reason", "")),
            time_in_trade_minutes=minutes,
        ))
    return CorrelationResult(
        matches=matches,
        tag_stats=_aggregate_tag_stats(matches),
        closed_trade_count=len(pairs),
        open_entry_count=open_count,
    )


def render_correlation_report(result: CorrelationResult, *, min_samples: int = 5) -> str:
    total_pnl = sum(m.realized_pnl_usd for m in result.matches)
    lines = [
        "Signal Correlation Report",
        f"Closed trades: {result.closed_trade_count}",
        f"Open entries: {result.open_entry_count}",
        f"Total PnL: {total_pnl:+.4f} USDC",
    ]
    if result.closed_trade_count < min_samples:
        lines.append(f"Observation only: Sample Size < {min_samples}; keine Strategieänderung aus diesen Daten ableiten.")
    if result.tag_stats:
        lines.append("Tag stats:")
        for stat in result.tag_stats.values():
            lines.append(
                f"- {stat.tag}: n={stat.sample_size}, win={stat.win_rate:.2f}, avg={stat.avg_pnl_usd:+.4f}, total={stat.total_pnl_usd:+.4f}, avg_time={stat.avg_time_in_trade_minutes:.1f}m"
            )
    else:
        lines.append("Tag stats: keine geschlossenen Trades mit Kontext gefunden.")
    return "\n".join(lines)


def main(argv: list[str] | None = None) -> int:
    cfg = BotConfig.from_file()
    paths = RuntimePaths.from_config(cfg)
    parser = argparse.ArgumentParser(description="Correlate paper/live trade outcomes with collected market context.")
    parser.add_argument("--journal", default=str(paths.trade_journal))
    parser.add_argument("--market-context", default=str(paths.market_context))
    parser.add_argument("--min-samples", type=int, default=5)
    args = parser.parse_args(argv)
    result = correlate_trades_with_context(args.journal, args.market_context)
    print(render_correlation_report(result, min_samples=args.min_samples))
    return 0


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