"""Results analyzer with Rich table visualizations."""

import json
from collections import defaultdict
from pathlib import Path

from rich.console import Console
from rich.table import Table


def analyze_run(run_dir: Path | None = None) -> None:
    """
    Analyze benchmark run results and display tables.

    Args:
        run_dir: Path to run directory, or None to find most recent
    """
    console = Console()

    # Find run directory if not specified
    if run_dir is None:
        logs_dir = Path("logs")
        if not logs_dir.exists():
            console.print("[red]Error:[/red] No logs directory found")
            return

        # Find most recent run
        runs = sorted(logs_dir.glob("benchmark_run_*"), reverse=True)
        if not runs:
            console.print("[red]Error:[/red] No benchmark runs found in logs/")
            return

        run_dir = runs[0]

    run_dir = Path(run_dir)
    if not run_dir.exists():
        console.print(f"[red]Error:[/red] Run directory not found: {run_dir}")
        return

    # Load summary.json
    summary_file = run_dir / "summary.json"
    if not summary_file.exists():
        console.print(f"[red]Error:[/red] No summary.json found in {run_dir}")
        return

    try:
        with open(summary_file) as f:
            summary = json.load(f)
    except (OSError, json.JSONDecodeError) as e:
        console.print(f"[red]Error:[/red] Failed to load summary.json: {e}")
        return

    console.print(f"\n[bold]Analyzing results from:[/bold] {run_dir.name}\n")

    # 1. Overview Summary
    _print_overview_table(console, summary)

    # 2. Success by Level
    _print_success_by_level_table(console, summary["results"])

    # 3. Success by Tag
    _print_success_by_tag_table(console, summary["results"])

    # 4. Cost and Duration Analysis
    _print_cost_duration_table(console, summary["results"])

    # 5. Failures Detail
    _print_failures_table(console, summary["results"])


def _print_overview_table(console: Console, summary: dict) -> None:
    """Print overview summary table."""
    table = Table(title="Overview Summary", show_header=False)
    table.add_column("Metric", style="cyan")
    table.add_column("Value", style="green")

    table.add_row("Total Benchmarks", str(summary["total_benchmarks"]))
    table.add_row(
        "Completed",
        f"{summary['completed']} ({summary['completion_rate']:.1f}%)"
    )
    table.add_row("With Flags", str(summary.get("with_flags", "N/A")))
    table.add_row("Timeout", str(summary["timeout"]))
    table.add_row("Error", str(summary["error"]))
    table.add_row("Total Duration", f"{summary['total_duration_seconds'] / 3600:.2f}h")
    table.add_row("Average Duration", f"{summary['average_duration_seconds'] / 60:.1f}m")
    table.add_row("Total Cost", f"${summary['total_cost_usd']:.2f}")
    table.add_row("Average Cost", f"${summary['average_cost_usd']:.2f}")

    console.print(table)
    console.print()


def _print_success_by_level_table(console: Console, results: list[dict]) -> None:
    """Print completion stats by difficulty level."""
    # Group by level
    by_level = defaultdict(lambda: {"total": 0, "completed": 0, "with_flags": 0, "timeout": 0, "error": 0})

    for r in results:
        level = r["level"]
        by_level[level]["total"] += 1

        if r["status"] == "COMPLETED":
            by_level[level]["completed"] += 1
            if len(r.get("found_flags", [])) > 0:
                by_level[level]["with_flags"] += 1
        elif r["status"] == "TIMEOUT":
            by_level[level]["timeout"] += 1
        elif r["status"] == "ERROR":
            by_level[level]["error"] += 1

    table = Table(title="Results by Difficulty Level")
    table.add_column("Level", style="cyan", justify="center")
    table.add_column("Total", justify="right")
    table.add_column("Completed", style="green", justify="right")
    table.add_column("With Flags", style="bright_green", justify="right")
    table.add_column("Timeout", style="yellow", justify="right")
    table.add_column("Error", style="magenta", justify="right")
    table.add_column("Completion %", justify="right")

    for level in sorted(by_level.keys()):
        stats = by_level[level]
        completion_rate = (stats["completed"] / stats["total"] * 100) if stats["total"] > 0 else 0

        table.add_row(
            str(level),
            str(stats["total"]),
            str(stats["completed"]),
            str(stats["with_flags"]),
            str(stats["timeout"]),
            str(stats["error"]),
            f"{completion_rate:.1f}%"
        )

    console.print(table)
    console.print()


def _print_success_by_tag_table(console: Console, results: list[dict]) -> None:
    """Print completion stats by vulnerability tag (top 10)."""
    # Group by tag
    by_tag = defaultdict(lambda: {"total": 0, "completed": 0, "with_flags": 0})

    for r in results:
        for tag in r["tags"]:
            by_tag[tag]["total"] += 1
            if r["status"] == "COMPLETED":
                by_tag[tag]["completed"] += 1
                if len(r.get("found_flags", [])) > 0:
                    by_tag[tag]["with_flags"] += 1

    # Sort by total count, take top 10
    sorted_tags = sorted(by_tag.items(), key=lambda x: x[1]["total"], reverse=True)[:10]

    table = Table(title="Results by Vulnerability Type (Top 10)")
    table.add_column("Tag", style="cyan")
    table.add_column("Total", justify="right")
    table.add_column("Completed", style="green", justify="right")
    table.add_column("With Flags", style="bright_green", justify="right")
    table.add_column("Completion %", justify="right")

    for tag, stats in sorted_tags:
        completion_rate = (stats["completed"] / stats["total"] * 100) if stats["total"] > 0 else 0

        table.add_row(
            tag,
            str(stats["total"]),
            str(stats["completed"]),
            str(stats["with_flags"]),
            f"{completion_rate:.1f}%"
        )

    console.print(table)
    console.print()


def _print_cost_duration_table(console: Console, results: list[dict]) -> None:
    """Print cost and duration analysis."""
    if not results:
        return

    durations = [r["duration_seconds"] / 60 for r in results]  # Convert to minutes
    costs = [r["cost_usd"] for r in results]

    # Find min/max/avg/median
    durations_sorted = sorted(durations)
    costs_sorted = sorted(costs)

    duration_min = min(durations)
    duration_max = max(durations)
    duration_avg = sum(durations) / len(durations)
    duration_median = durations_sorted[len(durations_sorted) // 2]

    cost_min = min(costs)
    cost_max = max(costs)
    cost_avg = sum(costs) / len(costs)
    cost_median = costs_sorted[len(costs_sorted) // 2]

    table = Table(title="Cost and Duration Analysis")
    table.add_column("Metric", style="cyan")
    table.add_column("Min", justify="right")
    table.add_column("Max", justify="right")
    table.add_column("Avg", justify="right")
    table.add_column("Median", justify="right")

    table.add_row(
        "Duration (min)",
        f"{duration_min:.1f}",
        f"{duration_max:.1f}",
        f"{duration_avg:.1f}",
        f"{duration_median:.1f}"
    )
    table.add_row(
        "Cost (USD)",
        f"${cost_min:.2f}",
        f"${cost_max:.2f}",
        f"${cost_avg:.2f}",
        f"${cost_median:.2f}"
    )

    console.print(table)
    console.print()


def _print_failures_table(console: Console, results: list[dict]) -> None:
    """Print details of failed/timeout/error benchmarks."""
    # Filter non-success results
    failures = [r for r in results if not r["success"]]

    if not failures:
        console.print("[green]No failures![/green]\n")
        return

    table = Table(title="Failures Detail")
    table.add_column("ID", style="cyan")
    table.add_column("Status", style="yellow")
    table.add_column("Duration", justify="right")
    table.add_column("Cost", justify="right")
    table.add_column("Reason", style="red")

    for r in failures:
        duration_m = int(r["duration_seconds"] // 60)
        duration_s = int(r["duration_seconds"] % 60)

        # Determine reason
        if r["timeout_occurred"]:
            reason = "Timeout"
        elif r["error_message"]:
            reason = r["error_message"][:50]  # Truncate long messages
        elif not r["found_flags"]:
            reason = "No flags found"
        else:
            reason = "Incorrect flag"

        table.add_row(
            r["benchmark_id"],
            r["status"],
            f"{duration_m}m {duration_s}s",
            f"${r['cost_usd']:.2f}",
            reason
        )

    console.print(table)
    console.print()
