"""Benchmark script for domain-aware retrieval performance."""
import sys
import asyncio
import time
sys.path.insert(0, '/home/neil/Documents/textiq-doc-extraction')

from src.knowledge.domain_aware_retriever import DomainAwareRetriever


async def benchmark():
    """Benchmark retrieval performance."""
    retriever = DomainAwareRetriever()

    queries = [
        ("Finance", "What is the quarterly EBITDA and profit margin?"),
        ("Finance", "Show me revenue projections for Q4"),
        ("HR", "How many employees were hired last month?"),
        ("HR", "What is the average salary for engineers?"),
        ("Operations", "What are the inventory levels?"),
        ("General", "Tell me about the business"),
    ]

    print("="*80)
    print("RETRIEVAL PERFORMANCE BENCHMARK")
    print("="*80)
    print(f"\nTesting {len(queries)} queries with top_k=10\n")

    total_time = 0
    results_summary = []

    for category, query in queries:
        start = time.perf_counter()

        try:
            results = await retriever.retrieve(query, top_k=10)
            latency = (time.perf_counter() - start) * 1000
            total_time += latency

            # Analyze results
            categories_found = {}
            for r in results:
                cat = r.get('domain_category', 'unknown')
                categories_found[cat] = categories_found.get(cat, 0) + 1

            print(f"[{category:12}] Query: {query[:50]}")
            print(f"               Latency: {latency:6.2f}ms | Results: {len(results):2} | Categories: {categories_found}")

            results_summary.append({
                'category': category,
                'latency_ms': latency,
                'result_count': len(results),
                'categories': categories_found
            })

        except Exception as e:
            print(f"[{category:12}] ✗ FAILED: {e}")
            results_summary.append({
                'category': category,
                'latency_ms': 0,
                'result_count': 0,
                'error': str(e)
            })

    # Calculate statistics
    avg_latency = total_time / len(queries)
    latencies = [r['latency_ms'] for r in results_summary if r['latency_ms'] > 0]
    min_latency = min(latencies) if latencies else 0
    max_latency = max(latencies) if latencies else 0

    print("\n" + "="*80)
    print("PERFORMANCE SUMMARY")
    print("="*80)
    print(f"Total Queries:     {len(queries)}")
    print(f"Average Latency:   {avg_latency:.2f}ms")
    print(f"Min Latency:       {min_latency:.2f}ms")
    print(f"Max Latency:       {max_latency:.2f}ms")
    print(f"Target:            < 100ms")
    print(f"Status:            {'✓ PASS' if avg_latency < 100 else '✗ FAIL'}")

    # Expected improvement analysis
    print("\n" + "="*80)
    print("EXPECTED VS ACTUAL")
    print("="*80)
    print("Expected Improvement: 5-10x faster for domain-specific queries")
    print(f"Actual Performance:   {avg_latency:.2f}ms average")

    if avg_latency < 100:
        print("✓ Performance target achieved!")
    else:
        print("⚠ Performance target not met (may need optimization)")

    print("\n" + "="*80)


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