#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Benchmark script for the SSE /candles/dump and /candles/dump/raw endpoints.
Calculates the number of candles received per second over a given duration.
Supports running multiple iterations per chunk size to get more precise results.
Uses aiosseclient for asynchronous SSE processing.
"""
import argparse
import json
import re
import time
import statistics
import asyncio
from aiosseclient import aiosseclient
from tqdm import tqdm
from rich.console import Console
from rich.table import Table
import urllib.parse

console = Console()


async def run_single_test(test_url, duration, endpoint_type):
    """Run a single benchmark test and return the results."""
    total_candles = 0
    total_events = 0
    expected_total = None
    start_time = time.time()

    # Create a progress bar for the current test
    test_pbar = tqdm(
        desc="Receiving events", unit="candles", position=1, leave=False, ncols=80
    )

    try:
        async for event in aiosseclient(test_url, headers={"Accept": "text/event-stream"}):
            now = time.time()
            elapsed = now - start_time

            if event.event == "summary":
                summary = json.loads(event.data)
                expected_total = summary.get("count")
                if expected_total:
                    test_pbar.reset(total=expected_total)
                    tqdm.write(
                        f"Summary: Range from {summary.get('min')} to {summary.get('max')}, {expected_total} candles total"
                    )

            elif event.event == "candles" and endpoint_type == "dump":
                # Count based on finding opening brackets of inner arrays
                batch_size = len(re.findall(r"\[\d", event.data))
                total_candles += batch_size
                total_events += 1
                test_pbar.update(batch_size)

            elif event.event == "raw_chunk" and endpoint_type == "raw":
                # Count based on the number of items in the JSON array
                try:
                    data_list = json.loads(event.data)
                    batch_size = len(data_list) if isinstance(data_list, list) else 0
                except json.JSONDecodeError:
                    tqdm.write(f"Warning: Could not decode raw_chunk data: {event.data[:100]}...")
                    batch_size = 0 # Or estimate based on string length?
                total_candles += batch_size
                total_events += 1
                test_pbar.update(batch_size)

            # Update progress description regardless of event type if candles were received
            if event.event in ["candles", "raw_chunk"]:
                 test_pbar.set_description(
                    f"Received {total_events} events, {total_candles} candles"
                )

            # Check if we've run long enough
            if elapsed >= duration:
                break

    except KeyboardInterrupt:
        tqdm.write("Benchmark interrupted!")
        return None
    except Exception as e:
        tqdm.write(f"Error: {str(e)}")
        return None
    finally:
        elapsed = time.time() - start_time
        rate = total_candles / elapsed if elapsed > 0 else 0
        events_per_sec = total_events / elapsed if elapsed > 0 else 0
        avg_chunk_size = total_candles / total_events if total_events > 0 else 0

        test_pbar.close()

        return {
            "total_candles": total_candles,
            "total_events": total_events,
            "elapsed_seconds": elapsed,
            "candles_per_second": rate,
            "events_per_second": events_per_sec,
            "avg_chunk_size": avg_chunk_size,
        }


def calculate_average_results(results_list, exclude_percent=20):
    """
    Calculate average values from multiple test runs, excluding specified percent of extreme values.
    Excludes the top and bottom (exclude_percent/2) percent of the results.
    """
    if not results_list or len(results_list) <= 2:
        return results_list[0] if results_list else None

    # Sort the results by candles_per_second
    sorted_results = sorted(results_list, key=lambda x: x["candles_per_second"])

    # Calculate how many results to exclude from each end
    exclude_count = int(len(sorted_results) * (exclude_percent / 200))

    # Trim the extreme values
    trimmed_results = sorted_results[exclude_count:len(sorted_results)-exclude_count]

    if not trimmed_results: # Handle case where too many results were excluded
        trimmed_results = sorted_results # Fallback to using all results

    # Calculate averages from the trimmed results
    avg_result = {
        "total_candles": statistics.mean(r["total_candles"] for r in trimmed_results),
        "total_events": statistics.mean(r["total_events"] for r in trimmed_results),
        "elapsed_seconds": statistics.mean(r["elapsed_seconds"] for r in trimmed_results),
        "candles_per_second": statistics.mean(r["candles_per_second"] for r in trimmed_results),
        "events_per_second": statistics.mean(r["events_per_second"] for r in trimmed_results),
        "avg_chunk_size": statistics.mean(r["avg_chunk_size"] for r in trimmed_results),
        "iterations": len(results_list),
        "used_iterations": len(trimmed_results),
    }

    return avg_result


async def main_async():
    parser = argparse.ArgumentParser(
        description="Benchmark SSE endpoint for candles stream"
    )
    parser.add_argument(
        "--base-url",
        default="https://candles.macrofinder.flolep.fr",
        help="Base URL of the server (e.g., http://localhost:3000)",
    )
    parser.add_argument(
        "--ticker",
        default="BINANCE:BTCUSDT",
        help="Ticker symbol for the request",
    )
    parser.add_argument(
        "--timeframe",
        default="1",
        help="Timeframe for the request",
    )
    parser.add_argument(
        "--endpoint-type",
        choices=["dump", "raw"],
        default="dump",
        help="Which dump endpoint to test ('dump' or 'raw')",
    )
    parser.add_argument(
        "--chunk-size", type=int, help="Single chunkSize parameter for the endpoint"
    )
    parser.add_argument(
        "--chunk-sizes", help="Comma-separated list of chunk sizes to test"
    )
    parser.add_argument(
        "--duration", type=int, default=20, help="Seconds to run each benchmark"
    )
    parser.add_argument(
        "--iterations", type=int, default=3, help="Number of iterations to run for each chunk size"
    )
    parser.add_argument(
        "--exclude-percent", type=float, default=20,
        help="Percentage of extreme values to exclude from the average (from 0 to 100)"
    )
    args = parser.parse_args()

    # Construct the base path based on endpoint type
    endpoint_path = f"/candles/dump/raw" if args.endpoint_type == "raw" else "/candles/dump"
    base_query = f"ticker={urllib.parse.quote(args.ticker)}&timeframe={urllib.parse.quote(args.timeframe)}"
    base_url_with_path = f"{args.base_url.rstrip('/')}{endpoint_path}?{base_query}"

    # Determine chunk sizes to benchmark
    if args.chunk_sizes:
        chunk_sizes = [int(x) for x in args.chunk_sizes.split(",") if x]
    elif args.chunk_size:
        chunk_sizes = [args.chunk_size]
    else:
        # Default range depends slightly on endpoint type (raw might handle larger chunks better)
        chunk_sizes = [100,300,500,800,1000,1500]

    aggregated_results = []

    # Create the main progress bar with tqdm
    main_pbar = tqdm(
        total=len(chunk_sizes),
        desc=f"Benchmarking '{args.endpoint_type}' endpoint",
        position=0,
        leave=True,
        ncols=80,
    )

    for chunk_idx, chunk_size in enumerate(chunk_sizes):
        main_pbar.set_description(f"Testing chunkSize={chunk_size}")
        tqdm.write(f"\n[Test {chunk_idx+1}/{len(chunk_sizes)}] Benchmark for chunkSize={chunk_size} ({args.endpoint_type} endpoint)")

        # Store results for all iterations of this chunk size
        chunk_results = []

        # Progress bar for iterations
        iter_pbar = tqdm(
            total=args.iterations,
            desc=f"Running iterations",
            position=1,
            leave=False,
            ncols=80,
        )

        for i in range(args.iterations):
            test_url = f"{base_url_with_path}&chunkSize={chunk_size}"
            tqdm.write(f"\nIteration {i+1}/{args.iterations} - Testing: {test_url}")

            # Run a single test, passing the endpoint type
            result = await run_single_test(test_url, args.duration, args.endpoint_type)

            if result is None:  # Test was interrupted or failed
                break

            # Add chunk_size to the result
            result["chunk_size"] = chunk_size
            chunk_results.append(result)

            # Show summary for this iteration
            tqdm.write(f"Iteration result: {result['candles_per_second']:,.2f} candles/sec")
            iter_pbar.update(1)

        iter_pbar.close()

        if chunk_results:
            # Calculate the average results after excluding extremes
            avg_result = calculate_average_results(chunk_results, args.exclude_percent)
            avg_result["chunk_size"] = chunk_size
            aggregated_results.append(avg_result)

            # Show summary for this chunk size
            tqdm.write(f"\nChunk Size {chunk_size} Summary:")
            tqdm.write(f"Ran {avg_result['iterations']} iterations, used {avg_result['used_iterations']} for averaging")
            tqdm.write(f"Average Rate: {avg_result['candles_per_second']:,.2f} candles/sec")
            tqdm.write(f"Average Events Rate: {avg_result['events_per_second']:.2f} events/sec")

        # Update the main progress bar
        main_pbar.update(1)

    # Close the main progress bar
    main_pbar.close()

    # Print final results table using rich
    console.print(f"\n[bold magenta]===== BENCHMARK RESULTS ({args.endpoint_type.upper()} Endpoint) =====[/bold magenta]")
    table = Table(show_header=True, header_style="bold")
    table.add_column("Chunk Size", justify="right", style="cyan")
    table.add_column("Candles/sec", justify="right", style="green")
    table.add_column("Events/sec", justify="right")
    table.add_column("Avg. Event Size", justify="right")
    table.add_column("Total Candles", justify="right")
    table.add_column("Duration (s)", justify="right")
    table.add_column("Iterations", justify="center")

    # Sort results by chunk size for consistent display
    aggregated_results.sort(key=lambda x: x["chunk_size"])

    for result in aggregated_results:
        table.add_row(
            f"{result['chunk_size']:,}",
            f"{result['candles_per_second']:,.2f}",
            f"{result['events_per_second']:.2f}",
            f"{result['avg_chunk_size']:.2f}",
            f"{result['total_candles']:,.0f}",
            f"{result['elapsed_seconds']:.2f}",
            f"{result['used_iterations']}/{result['iterations']}",
        )

    console.print(table)

    # Print the best result
    if aggregated_results:
        best_result = max(aggregated_results, key=lambda x: x["candles_per_second"])
        console.print(
            f"\n[bold green]Best performance:[/bold green] [cyan]chunkSize={best_result['chunk_size']}[/cyan] with [bold]{best_result['candles_per_second']:,.2f}[/bold] candles/sec"
        )
        console.print(f"Based on {best_result['used_iterations']} iterations (excluding {args.exclude_percent}% extreme values)")


def main():
    """Entry point that runs the async main function"""
    asyncio.run(main_async())


if __name__ == "__main__":
    main()