"""Aggressive Concurrency & Stress Harness for FastMCP Mem0 Service.

Adversarial Verification for Milestone 3 Challenger:
- 50 Rapid Sequential Queries (p50, p95, p99 latency measurement)
- 30 Concurrent Parallel SSE Sessions (burst traffic)
- Mixed Load (read queries, status probes, edge cases)
- Continuous Process Health Monitoring (verifying PID stability, zero socket drops, zero locking)
"""

import asyncio
import json
import statistics
import sys
import time
from typing import Any, Dict, List

from mcp.client.session import ClientSession
from mcp.client.sse import sse_client

if hasattr(sys.stdout, "reconfigure"):
    sys.stdout.reconfigure(encoding="utf-8", errors="replace")
if hasattr(sys.stderr, "reconfigure"):
    sys.stderr.reconfigure(encoding="utf-8", errors="replace")

MEM0_SSE_URL = "http://127.0.0.1:8765/sse"
USER_ID = "developer"

BENCHMARK_QUERIES = [
    "FreeExile combat",
    "Feral NPC dialogue HUD",
    "Savage Infinite Atlas",
    "developer preferred language",
    "preferred backend language",
    "VoLamWeb cloud budget",
    "0 VND Google Cloud",
    "DSCons rule",
    "frontend theme",
    "backend architecture",
    "AEC pricing norms",
    "CAD takeoff geometry",
    "database precision",
    "testing workflow",
    "MCP tools protocol",
]


async def single_search(query: str, session_id: int = 0) -> Dict[str, Any]:
    """Execute a single search query over a dedicated SSE connection."""
    t0 = time.perf_counter()
    try:
        async with asyncio.timeout(15.0):
            async with (
                sse_client(MEM0_SSE_URL) as (read, write),
                ClientSession(read, write) as session,
            ):
                await session.initialize()
                res = await session.call_tool(
                    "mem0_search",
                    arguments={"query": query, "user_id": USER_ID, "limit": 3},
                )
                elapsed = time.perf_counter() - t0
                raw = res.content[0].text
                data = json.loads(raw)
                hits = data.get("memories", {}).get("results", [])
                top1_score = float(hits[0].get("score", 0.0)) if hits else 0.0
                return {
                    "success": True,
                    "session_id": session_id,
                    "query": query,
                    "elapsed": elapsed,
                    "hits_count": len(hits),
                    "top1_score": top1_score,
                    "error": None,
                }
    except Exception as e:
        elapsed = time.perf_counter() - t0
        return {
            "success": False,
            "session_id": session_id,
            "query": query,
            "elapsed": elapsed,
            "hits_count": 0,
            "top1_score": 0.0,
            "error": str(e),
        }


async def run_rapid_sequential(count: int = 50) -> Dict[str, Any]:
    """Execute rapid sequential queries to test connection churn and latency."""
    print(f"\n--- [Phase 1] Rapid Sequential Stress Test ({count} iterations) ---")
    latencies = []
    errors = []
    t_start = time.perf_counter()

    for i in range(count):
        q = BENCHMARK_QUERIES[i % len(BENCHMARK_QUERIES)]
        res = await single_search(q, session_id=i)
        if res["success"]:
            latencies.append(res["elapsed"])
        else:
            errors.append(res)

    total_time = time.perf_counter() - t_start
    p50 = statistics.median(latencies) if latencies else 0.0
    p95 = statistics.quantiles(latencies, n=20)[18] if len(latencies) >= 20 else max(latencies, default=0.0)
    p99 = statistics.quantiles(latencies, n=100)[98] if len(latencies) >= 100 else max(latencies, default=0.0)
    avg_lat = statistics.mean(latencies) if latencies else 0.0

    report = {
        "total_requests": count,
        "successful_requests": len(latencies),
        "failed_requests": len(errors),
        "total_duration_sec": total_time,
        "throughput_rps": count / total_time if total_time > 0 else 0.0,
        "latency_p50_ms": p50 * 1000,
        "latency_p95_ms": p95 * 1000,
        "latency_p99_ms": p99 * 1000,
        "latency_avg_ms": avg_lat * 1000,
        "errors": errors,
    }

    print(f"Total time: {total_time:.2f}s | Throughput: {report['throughput_rps']:.2f} req/s")
    print(f"Latency: Avg={report['latency_avg_ms']:.1f}ms, p50={report['latency_p50_ms']:.1f}ms, p95={report['latency_p95_ms']:.1f}ms")
    print(f"Success rate: {len(latencies)}/{count} ({len(latencies)/count*100:.1f}%) | Errors: {len(errors)}")
    return report


async def run_concurrent_burst(concurrency: int = 30) -> Dict[str, Any]:
    """Execute burst concurrent queries across parallel SSE sessions."""
    print(f"\n--- [Phase 2] Concurrent Burst Stress Test ({concurrency} parallel sessions) ---")
    tasks = []
    for i in range(concurrency):
        q = BENCHMARK_QUERIES[i % len(BENCHMARK_QUERIES)]
        tasks.append(single_search(q, session_id=i))

    t_start = time.perf_counter()
    results = await asyncio.gather(*tasks)
    total_time = time.perf_counter() - t_start

    successful = [r for r in results if r["success"]]
    failed = [r for r in results if not r["success"]]
    latencies = [r["elapsed"] for r in successful]

    avg_lat = statistics.mean(latencies) if latencies else 0.0
    max_lat = max(latencies) if latencies else 0.0
    min_lat = min(latencies) if latencies else 0.0

    report = {
        "concurrency": concurrency,
        "total_requests": concurrency,
        "successful_requests": len(successful),
        "failed_requests": len(failed),
        "total_duration_sec": total_time,
        "throughput_rps": concurrency / total_time if total_time > 0 else 0.0,
        "latency_avg_ms": avg_lat * 1000,
        "latency_min_ms": min_lat * 1000,
        "latency_max_ms": max_lat * 1000,
        "errors": failed,
    }

    print(f"Total time: {total_time:.2f}s | Concurrent Throughput: {report['throughput_rps']:.2f} req/s")
    print(f"Latency: Avg={report['latency_avg_ms']:.1f}ms, Min={report['latency_min_ms']:.1f}ms, Max={report['latency_max_ms']:.1f}ms")
    print(f"Success rate: {len(successful)}/{concurrency} ({len(successful)/concurrency*100:.1f}%) | Errors: {len(failed)}")
    return report


async def run_mixed_stress_workload() -> Dict[str, Any]:
    """Execute mixed workload: rapid searches, status requests, and edge queries."""
    print("\n--- [Phase 3] Mixed Stress Workload (Search + Status + Boundary) ---")
    async def status_check(i: int):
        t0 = time.perf_counter()
        async with (
            sse_client(MEM0_SSE_URL) as (read, write),
            ClientSession(read, write) as session,
        ):
            await session.initialize()
            res = await session.call_tool("mem0_status", arguments={})
            elapsed = time.perf_counter() - t0
            data = json.loads(res.content[0].text)
            return {"type": "status", "success": data.get("status") == "online", "elapsed": elapsed}

    search_tasks = [single_search(BENCHMARK_QUERIES[i % len(BENCHMARK_QUERIES)], session_id=100+i) for i in range(15)]
    status_tasks = [status_check(i) for i in range(5)]

    t_start = time.perf_counter()
    all_results = await asyncio.gather(*(search_tasks + status_tasks), return_exceptions=True)
    total_time = time.perf_counter() - t_start

    errors = [r for r in all_results if isinstance(r, Exception) or (isinstance(r, dict) and not r.get("success", False))]
    print(f"Completed {len(all_results)} mixed operations in {total_time:.2f}s with {len(errors)} errors.")
    return {
        "total_operations": len(all_results),
        "total_time_sec": total_time,
        "error_count": len(errors),
        "errors": [str(e) for e in errors],
    }


async def main():
    print("=" * 70)
    print("   MEM0 FASTMCP SERVICE CONCURRENCY & STRESS TEST HARNESS")
    print("   Endpoint: http://127.0.0.1:8765/sse")
    print("=" * 70)

    # Initial Health Check
    print("\n[Pre-test Health Check]")
    pre_status = await single_search("DSCons rule", session_id=-1)
    assert pre_status["success"], f"Initial check failed: {pre_status['error']}"
    print("Service connection healthy. Starting stress phases...")

    # Phase 1: 50 Rapid Sequential
    seq_report = await run_rapid_sequential(count=50)

    # Phase 2: 30 Concurrent Burst
    burst_report = await run_concurrent_burst(concurrency=30)

    # Phase 3: Mixed Workload
    mixed_report = await run_mixed_stress_workload()

    # Post-test Health Check
    print("\n[Post-test Final Health Check]")
    async with (
        sse_client(MEM0_SSE_URL) as (read, write),
        ClientSession(read, write) as session,
    ):
        await session.initialize()
        res = await session.call_tool("mem0_status", arguments={})
        post_status = json.loads(res.content[0].text)
        print(f"Final Status: {json.dumps(post_status, indent=2)}")

    total_failures = seq_report["failed_requests"] + burst_report["failed_requests"] + mixed_report["error_count"]
    print("\n" + "=" * 70)
    if total_failures == 0 and post_status.get("status") == "online":
        print(f" STRESS TEST VERDICT: [PASS] Zero socket drops, zero locking errors, 100% success rate!")
    else:
        print(f" STRESS TEST VERDICT: [FAIL] Encountered {total_failures} failures.")
    print("=" * 70)

    sys.exit(0 if total_failures == 0 else 1)


if __name__ == "__main__":
    asyncio.run(main())
