"""Main orchestrator for benchmark execution."""

import signal
import sys
from datetime import datetime

from pentestgpt.benchmark.docker import load_benchmarks, start_benchmark, stop_benchmark
from pentestgpt.benchmark.executor import PentestGPTExecutor
from pentestgpt.benchmark.models import BenchmarkConfig, BenchmarkResult
from pentestgpt.benchmark.parser import OutputParser
from pentestgpt.benchmark.registry import BenchmarkInfo
from pentestgpt.benchmark.reporter import Reporter
from pentestgpt.benchmark.state import StateManager


class BenchmarkRunner:
    """Main orchestrator coordinating all benchmark execution."""

    def __init__(self, config: BenchmarkConfig):
        """
        Initialize runner.

        Args:
            config: Benchmark configuration
        """
        self.config = config
        self.executor = PentestGPTExecutor(model=config.model)
        self.parser = OutputParser()
        self.reporter = Reporter(config.output_dir)
        self.state = StateManager(config.state_file)

        self.interrupted = False
        self.current_benchmark_path = None

        self._setup_signal_handlers()

    def _setup_signal_handlers(self):
        """Setup signal handlers for graceful shutdown."""
        def handle_interrupt(signum, frame):
            print("\n\nInterrupt received. Cleaning up...")
            self.interrupted = True

            # Stop current benchmark if any
            if self.current_benchmark_path:
                print("Stopping current benchmark...")
                stop_benchmark(self.current_benchmark_path)

            # Save state
            self.state.save()

            print("Cleanup complete.")
            print("Resume with: --resume")
            sys.exit(130)

        signal.signal(signal.SIGINT, handle_interrupt)
        signal.signal(signal.SIGTERM, handle_interrupt)

    async def run_all(self) -> dict:
        """
        Run all selected benchmarks.

        Returns:
            Summary dictionary
        """
        start_time = datetime.now()

        # Load all benchmarks
        print(f"Loading benchmarks from {self.config.benchmarks_dir}...")
        all_benchmarks = load_benchmarks(self.config.benchmarks_dir)
        print(f"Found {len(all_benchmarks)} benchmarks")

        # Filter by selected IDs
        selected_benchmarks = []
        for bench_id in self.config.benchmark_ids:
            # Convert ID to benchmark name format
            bench_name = f"XBEN-{bench_id:03d}-24"
            if bench_name in all_benchmarks:
                selected_benchmarks.append(all_benchmarks[bench_name])
            else:
                print(f"Warning: Benchmark {bench_name} not found")

        if not selected_benchmarks:
            print("Error: No valid benchmarks selected")
            return {}

        # Filter out completed if resuming
        if self.config.resume:
            remaining_ids = self.state.get_remaining([b.id for b in selected_benchmarks])
            selected_benchmarks = [b for b in selected_benchmarks if b.id in remaining_ids]
            print(f"Resuming: {len(selected_benchmarks)} benchmarks remaining")

        total = len(selected_benchmarks)
        print(f"\nStarting benchmark run: {total} benchmarks")
        print(f"Timeout: {self.config.timeout_seconds}s per benchmark")
        print("=" * 60)

        # Run each benchmark
        results = []
        for index, info in enumerate(selected_benchmarks, 1):
            if self.interrupted:
                break

            self.reporter.log_start(info.id, index, total)

            result = await self.run_single_benchmark(info)
            results.append(result)

            # Mark in state
            self.state.mark_completed(info.id, result.success)

            # Log result
            self.reporter.log_result(result)

        end_time = datetime.now()

        # Generate summary
        if results:
            self.reporter.generate_summary(results, start_time, end_time)

        return {"total": total, "completed": len(results)}

    async def run_single_benchmark(self, info: BenchmarkInfo) -> BenchmarkResult:
        """
        Execute single benchmark end-to-end.

        Args:
            info: Benchmark information

        Returns:
            Benchmark result
        """
        start_time = datetime.now()
        self.current_benchmark_path = info.path

        try:
            # 1. Start Docker containers
            docker_result = start_benchmark(info.path)
            if not docker_result["success"]:
                return BenchmarkResult(
                    benchmark_id=info.id,
                    benchmark_name=info.name,
                    level=info.level,
                    tags=info.tags,
                    status="ERROR",
                    success=False,
                    expected_flag=info.expected_flag,
                    found_flags=[],
                    correct_flag=False,
                    duration_seconds=(datetime.now() - start_time).total_seconds(),
                    cost_usd=0.0,
                    session_id="",
                    error_message=f"Docker start failed: {docker_result['message']}",
                    timeout_occurred=False,
                    start_time=start_time,
                    end_time=datetime.now()
                )

            target_url = docker_result["target_url"]

            # 2. Execute PentestGPT
            output_file = self.reporter.get_benchmark_log_path(info.id)
            exec_result = await self.executor.execute(
                target_url,
                info.id,
                self.config.timeout_seconds,
                output_file
            )

            # 3. Parse output
            parsed = self.parser.parse_output(exec_result["output_lines"])

            # 4. Extract flags found
            found_flags = parsed["flags"]

            # 5. Determine objective status (not judging success/failure)
            if exec_result["timed_out"]:
                status = "TIMEOUT"
            elif parsed["has_error"]:
                status = "ERROR"
            else:
                status = "COMPLETED"

            # For backward compatibility, set success based on flags found
            # (but status now focuses on completion state, not correctness)
            success = len(found_flags) > 0 and status == "COMPLETED"

            # 6. Build result
            return BenchmarkResult(
                benchmark_id=info.id,
                benchmark_name=info.name,
                level=info.level,
                tags=info.tags,
                status=status,
                success=success,
                expected_flag=info.expected_flag,
                found_flags=found_flags,
                correct_flag=False,  # Not validating correctness anymore
                duration_seconds=(datetime.now() - start_time).total_seconds(),
                cost_usd=parsed.get("cost", 0.0),
                session_id=parsed.get("session_id", ""),
                error_message=parsed.get("error_msg"),
                timeout_occurred=exec_result["timed_out"],
                start_time=start_time,
                end_time=datetime.now()
            )

        except Exception as e:
            # Unexpected error
            return BenchmarkResult(
                benchmark_id=info.id,
                benchmark_name=info.name,
                level=info.level,
                tags=info.tags,
                status="ERROR",
                success=False,
                expected_flag=info.expected_flag,
                found_flags=[],
                correct_flag=False,
                duration_seconds=(datetime.now() - start_time).total_seconds(),
                cost_usd=0.0,
                session_id="",
                error_message=f"Unexpected error: {e!s}",
                timeout_occurred=False,
                start_time=start_time,
                end_time=datetime.now()
            )

        finally:
            # ALWAYS cleanup Docker containers
            try:
                stop_benchmark(info.path)
            except Exception as e:
                print(f"  Warning: Error stopping containers: {e}")

            self.current_benchmark_path = None
