#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Search Quality Benchmark - Compare optimized vs standard BM25

Usage:
    python3 benchmark_search.py [--test-cases <file>] [--output <file>]
"""

import argparse
import json
from pathlib import Path
from core import search, CSV_CONFIG, AVAILABLE_DOMAINS


# ============ TEST CASES ============
TEST_CASES = [
    # 同义词扩展测试
    {
        "name": "同义词扩展 - btn vs button",
        "query": "btn primary",
        "domain": "component",
        "description": "测试 'btn' 能否通过同义词扩展匹配到 'button' 相关结果"
    },
    {
        "name": "同义词扩展 - 中文查询",
        "query": "表单输入",
        "domain": "form",
        "description": "测试中文查询能否通过同义词扩展匹配"
    },
    {
        "name": "同义词扩展 - hover cursor",
        "query": "hover cursor",
        "domain": "component",
        "description": "测试多词查询的同义词扩展效果"
    },
    
    # 字段权重测试
    {
        "name": "字段权重 - Keywords 优先",
        "query": "primary button",
        "domain": "component",
        "description": "测试 Keywords 字段权重是否让关键词匹配结果优先显示"
    },
    {
        "name": "字段权重 - 精确匹配",
        "query": "icon hover",
        "domain": "component",
        "description": "测试关键词字段匹配的结果是否排在前面"
    },
    
    # 综合测试
    {
        "name": "综合测试 - button group",
        "query": "button group",
        "domain": "component",
        "description": "综合测试：同义词 + 字段权重"
    },
    {
        "name": "综合测试 - form input",
        "query": "form input required",
        "domain": "form",
        "description": "综合测试：多词查询 + 同义词 + 字段权重"
    },
    {
        "name": "综合测试 - table datatable",
        "query": "table datatable",
        "domain": "table",
        "description": "综合测试：同义词扩展（table = 表格 = datatable）"
    },
    
    # 边界情况
    {
        "name": "边界情况 - 单字查询",
        "query": "btn",
        "domain": "component",
        "description": "测试单字查询的同义词扩展"
    },
    {
        "name": "边界情况 - 空结果",
        "query": "nonexistent_component_xyz",
        "domain": "component",
        "description": "测试无匹配结果的查询"
    },
]


# ============ BENCHMARK FUNCTIONS ============
def run_search_variant(query, domain, use_weights, use_cache, expand_synonyms, max_results=5):
    """运行一个搜索变体"""
    try:
        result = search(
            query,
            domain=domain,
            max_results=max_results,
            use_weights=use_weights,
            use_cache=use_cache,
            expand_synonyms=expand_synonyms
        )
        return {
            "success": True,
            "count": result.get("count", 0),
            "results": result.get("results", []),
            "error": result.get("error")
        }
    except Exception as e:
        return {
            "success": False,
            "count": 0,
            "results": [],
            "error": str(e)
        }


def compare_variants(query, domain, max_results=5):
    """对比不同配置的搜索结果"""
    variants = {
        "优化前 (标准)": {
            "use_weights": False,
            "use_cache": False,  # 关闭缓存以便公平对比
            "expand_synonyms": False
        },
        "优化后 (完整)": {
            "use_weights": True,
            "use_cache": False,  # 关闭缓存以便公平对比
            "expand_synonyms": True
        },
        "仅同义词扩展": {
            "use_weights": False,
            "use_cache": False,
            "expand_synonyms": True
        },
        "仅字段权重": {
            "use_weights": True,
            "use_cache": False,
            "expand_synonyms": False
        }
    }
    
    results = {}
    for variant_name, config in variants.items():
        result = run_search_variant(
            query, domain, max_results=max_results,
            **config
        )
        results[variant_name] = result
    
    return results


def analyze_results(comparison_results):
    """分析对比结果"""
    baseline = comparison_results.get("优化前 (标准)", {})
    optimized = comparison_results.get("优化后 (完整)", {})
    
    analysis = {
        "结果数量对比": {
            "优化前": baseline.get("count", 0),
            "优化后": optimized.get("count", 0),
            "差异": optimized.get("count", 0) - baseline.get("count", 0),
            "提升": "✅ 有提升" if optimized.get("count", 0) > baseline.get("count", 0) else "❌ 无提升"
        },
        "结果质量": {
            "优化前Top3": [r.get("Component") or r.get("Element") or r.get("Pattern", "N/A") for r in baseline.get("results", [])[:3]],
            "优化后Top3": [r.get("Component") or r.get("Element") or r.get("Pattern", "N/A") for r in optimized.get("results", [])[:3]],
        },
        "各优化功能效果": {
            "同义词扩展": {
                "结果数": comparison_results.get("仅同义词扩展", {}).get("count", 0),
                "vs 标准": comparison_results.get("仅同义词扩展", {}).get("count", 0) - baseline.get("count", 0)
            },
            "字段权重": {
                "结果数": comparison_results.get("仅字段权重", {}).get("count", 0),
                "vs 标准": comparison_results.get("仅字段权重", {}).get("count", 0) - baseline.get("count", 0)
            }
        }
    }
    
    return analysis


def format_comparison_report(test_case, comparison_results, analysis):
    """格式化对比报告"""
    lines = []
    lines.append("=" * 80)
    lines.append(f"测试用例: {test_case['name']}")
    lines.append(f"查询: {test_case['query']}")
    lines.append(f"Domain: {test_case['domain']}")
    lines.append(f"说明: {test_case['description']}")
    lines.append("=" * 80)
    lines.append("")
    
    # 结果数量对比
    lines.append("## 📊 结果数量对比")
    for variant_name, result in comparison_results.items():
        status = "✅" if result.get("success") else "❌"
        count = result.get("count", 0)
        error = result.get("error", "")
        lines.append(f"- {status} **{variant_name}**: {count} 个结果" + (f" (错误: {error})" if error else ""))
    lines.append("")
    
    # 分析结果
    lines.append("## 📈 分析结果")
    count_analysis = analysis["结果数量对比"]
    lines.append(f"- **优化前结果数**: {count_analysis['优化前']}")
    lines.append(f"- **优化后结果数**: {count_analysis['优化后']}")
    lines.append(f"- **差异**: {count_analysis['差异']:+d}")
    lines.append(f"- **评估**: {count_analysis['提升']}")
    lines.append("")
    
    # 各功能效果
    lines.append("## 🔍 各优化功能效果")
    func_analysis = analysis["各优化功能效果"]
    lines.append(f"- **仅同义词扩展**: {func_analysis['同义词扩展']['结果数']} 个结果 (vs 标准: {func_analysis['同义词扩展']['vs 标准']:+d})")
    lines.append(f"- **仅字段权重**: {func_analysis['字段权重']['结果数']} 个结果 (vs 标准: {func_analysis['字段权重']['vs 标准']:+d})")
    lines.append("")
    
    # 结果质量对比
    lines.append("## 🎯 Top 3 结果对比")
    quality = analysis["结果质量"]
    lines.append(f"- **优化前**: {', '.join(quality['优化前Top3']) or '无结果'}")
    lines.append(f"- **优化后**: {', '.join(quality['优化后Top3']) or '无结果'}")
    lines.append("")
    
    # 详细结果对比
    lines.append("## 📝 详细结果对比")
    for variant_name, result in comparison_results.items():
        lines.append(f"\n### {variant_name}")
        if result.get("success") and result.get("results"):
            for i, res in enumerate(result.get("results", [])[:3], 1):
                title = res.get("Component") or res.get("Element") or res.get("Pattern") or res.get("Icon Name", "N/A")
                keywords = res.get("Keywords", "N/A")
                lines.append(f"  {i}. {title} (Keywords: {keywords})")
        elif result.get("error"):
            lines.append(f"  错误: {result.get('error')}")
        else:
            lines.append("  无结果")
    
    lines.append("")
    lines.append("-" * 80)
    lines.append("")
    
    return "\n".join(lines)


def run_benchmark(test_cases=None, output_file=None):
    """运行完整的基准测试"""
    if test_cases is None:
        test_cases = TEST_CASES
    
    print(f"🚀 开始运行基准测试，共 {len(test_cases)} 个测试用例...")
    print()
    
    all_reports = []
    summary = {
        "总测试数": len(test_cases),
        "优化提升": 0,
        "无提升": 0,
        "同义词扩展有效": 0,
        "字段权重有效": 0
    }
    
    for i, test_case in enumerate(test_cases, 1):
        print(f"[{i}/{len(test_cases)}] 运行测试: {test_case['name']}...")
        
        # 运行对比
        comparison_results = compare_variants(
            test_case["query"],
            test_case["domain"]
        )
        
        # 分析结果
        analysis = analyze_results(comparison_results)
        
        # 生成报告
        report = format_comparison_report(test_case, comparison_results, analysis)
        all_reports.append(report)
        
        # 更新统计
        if analysis["结果数量对比"]["差异"] > 0:
            summary["优化提升"] += 1
        else:
            summary["无提升"] += 1
        
        if analysis["各优化功能效果"]["同义词扩展"]["vs 标准"] > 0:
            summary["同义词扩展有效"] += 1
        
        if analysis["各优化功能效果"]["字段权重"]["vs 标准"] >= 0:
            summary["字段权重有效"] += 1
        
        print(f"  ✅ 完成: 优化前={analysis['结果数量对比']['优化前']}, 优化后={analysis['结果数量对比']['优化后']}")
    
    # 生成总结报告
    summary_report = []
    summary_report.append("=" * 80)
    summary_report.append("# 📊 基准测试总结报告")
    summary_report.append("=" * 80)
    summary_report.append("")
    summary_report.append(f"- **总测试数**: {summary['总测试数']}")
    summary_report.append(f"- **优化有提升**: {summary['优化提升']} ({summary['优化提升']/summary['总测试数']*100:.1f}%)")
    summary_report.append(f"- **无提升**: {summary['无提升']} ({summary['无提升']/summary['总测试数']*100:.1f}%)")
    summary_report.append(f"- **同义词扩展有效**: {summary['同义词扩展有效']} ({summary['同义词扩展有效']/summary['总测试数']*100:.1f}%)")
    summary_report.append(f"- **字段权重有效**: {summary['字段权重有效']} ({summary['字段权重有效']/summary['总测试数']*100:.1f}%)")
    summary_report.append("")
    summary_report.append("=" * 80)
    summary_report.append("")
    
    # 合并所有报告
    final_report = "\n".join(summary_report) + "\n".join(all_reports)
    
    # 输出到文件或控制台
    if output_file:
        with open(output_file, 'w', encoding='utf-8') as f:
            f.write(final_report)
        print(f"\n📄 完整报告已保存到: {output_file}")
    else:
        print("\n" + final_report)
    
    return {
        "summary": summary,
        "reports": all_reports,
        "full_report": final_report
    }


# ============ MAIN ============
if __name__ == "__main__":
    parser = argparse.ArgumentParser(
        description="Search Quality Benchmark - Compare optimized vs standard BM25",
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog="""
Examples:
  python3 benchmark_search.py
  python3 benchmark_search.py --output benchmark_report.md
  python3 benchmark_search.py --test-cases custom_cases.json
        """
    )
    parser.add_argument("--test-cases", help="Path to custom test cases JSON file")
    parser.add_argument("--output", "-o", help="Output file path (default: print to console)")
    
    args = parser.parse_args()
    
    # 加载测试用例
    test_cases = TEST_CASES
    if args.test_cases:
        with open(args.test_cases, 'r', encoding='utf-8') as f:
            test_cases = json.load(f)
    
    # 运行基准测试
    result = run_benchmark(test_cases, args.output)
    
    # 输出总结
    print("\n" + "=" * 80)
    print("✅ 基准测试完成！")
    print("=" * 80)
