"""
Experimental script to test Japanese word segmentation approaches.

Compares different segmentation libraries for improving LLM domain extraction
on Japanese text in the Airflow document processing pipeline.
"""

import sys
import time
from typing import List, Dict, Any
sys.path.insert(0, '/home/neil/Documents/textiq-doc-extraction')

from src.extraction_v2.domain_term_extractor import DomainTermExtractor


# Sample Japanese test cases
TEST_CASES = [
    {
        "name": "Finance Document",
        "text": """
        2023年度第4四半期財務結果
        売上高は125億円で前年同期比23%増加しました。
        EBITDAマージンは32%に達し、営業キャッシュフローは大幅に改善されました。
        純利益は予想を15%上回りました。
        """,
        "expected_category": "finance",
        "expected_terms": ["売上", "EBITDA", "利益", "四半期", "財務"]
    },
    {
        "name": "HR Document",
        "text": """
        第3四半期採用レポート
        新規採用社員数: エンジニア15名
        総従業員数: 200名
        平均給与: 12万ドル
        福利厚生パッケージには株式報酬が含まれます。
        研修プログラムを拡充しました。
        """,
        "expected_category": "hr",
        "expected_terms": ["採用", "社員", "従業員", "給与", "研修"]
    },
    {
        "name": "Mixed Content",
        "text": """
        株式会社テキストIQは、データ分析とAI技術を活用した
        ドキュメント処理システムの開発を行っています。
        売上高、収益率、顧客満足度の向上を目指しています。
        """,
        "expected_category": "general",
        "expected_terms": ["株式会社", "データ分析", "AI", "売上", "収益"]
    }
]


def segment_with_mecab(text: str) -> str:
    """Segment Japanese text using MeCab."""
    try:
        import MeCab
        tagger = MeCab.Tagger("-Owakati")
        return tagger.parse(text).strip()
    except ImportError:
        return None
    except Exception as e:
        print(f"MeCab error: {e}")
        return None


def segment_with_janome(text: str) -> str:
    """Segment Japanese text using Janome."""
    try:
        from janome.tokenizer import Tokenizer
        tokenizer = Tokenizer()
        tokens = [token.surface for token in tokenizer.tokenize(text)]
        return " ".join(tokens)
    except ImportError:
        return None
    except Exception as e:
        print(f"Janome error: {e}")
        return None


def segment_with_sudachi(text: str) -> str:
    """Segment Japanese text using SudachiPy."""
    try:
        from sudachipy import tokenizer, dictionary
        tokenizer_obj = dictionary.Dictionary().create()
        mode = tokenizer.Tokenizer.SplitMode.C
        tokens = [m.surface() for m in tokenizer_obj.tokenize(text, mode)]
        return " ".join(tokens)
    except ImportError:
        return None
    except Exception as e:
        print(f"SudachiPy error: {e}")
        return None


def test_baseline(test_case: Dict[str, Any]) -> Dict[str, Any]:
    """Test baseline (no segmentation) approach."""
    extractor = DomainTermExtractor()

    start = time.time()
    result = extractor.extract_from_chunk(test_case["text"])
    elapsed = time.time() - start

    return {
        "method": "Baseline (No Segmentation)",
        "terms": result['domain_terms'],
        "category": result['domain_category'],
        "extraction_method": result['extraction_method'],
        "term_count": len(result['domain_terms']),
        "elapsed_ms": int(elapsed * 1000),
        "success": len(result['domain_terms']) > 0
    }


def test_segmented(test_case: Dict[str, Any], segmenter_name: str, segmenter_func) -> Dict[str, Any]:
    """Test with word segmentation preprocessing."""
    segmented_text = segmenter_func(test_case["text"])

    if segmented_text is None:
        return {
            "method": segmenter_name,
            "error": "Library not installed or failed",
            "success": False
        }

    extractor = DomainTermExtractor()

    start = time.time()
    result = extractor.extract_from_chunk(segmented_text)
    elapsed = time.time() - start

    return {
        "method": segmenter_name,
        "segmented_text": segmented_text[:200] + "..." if len(segmented_text) > 200 else segmented_text,
        "terms": result['domain_terms'],
        "category": result['domain_category'],
        "extraction_method": result['extraction_method'],
        "term_count": len(result['domain_terms']),
        "elapsed_ms": int(elapsed * 1000),
        "success": len(result['domain_terms']) > 0
    }


def print_result(result: Dict[str, Any], test_name: str):
    """Pretty print a single test result."""
    print(f"\n{'─' * 70}")
    print(f"Method: {result['method']}")
    print(f"{'─' * 70}")

    if result.get("error"):
        print(f"❌ ERROR: {result['error']}")
        return

    if result.get("segmented_text"):
        print(f"Segmented Text: {result['segmented_text']}")

    print(f"Terms Found: {result['terms']}")
    print(f"Term Count: {result['term_count']}")
    print(f"Category: {result['category']}")
    print(f"Extraction Method: {result['extraction_method']}")
    print(f"Time: {result['elapsed_ms']}ms")
    print(f"Success: {'✓' if result['success'] else '✗'}")


def compare_all_methods():
    """Run comparison across all test cases and segmentation methods."""
    print("=" * 70)
    print("JAPANESE WORD SEGMENTATION EXPERIMENT")
    print("=" * 70)
    print("Testing different segmentation libraries for Japanese domain extraction")
    print()

    segmenters = [
        ("MeCab", segment_with_mecab),
        ("Janome", segment_with_janome),
        ("SudachiPy", segment_with_sudachi),
    ]

    # Check which libraries are available
    print("Checking library availability...")
    available_libs = []
    for name, func in segmenters:
        test_result = func("テスト")
        if test_result is not None:
            print(f"  ✓ {name} available")
            available_libs.append((name, func))
        else:
            print(f"  ✗ {name} not installed")

    if not available_libs:
        print("\n❌ No segmentation libraries installed!")
        print("\nInstall at least one:")
        print("  pip install mecab-python3 unidic-lite")
        print("  pip install janome")
        print("  pip install sudachipy sudachidict_core")
        return

    print()

    # Run tests for each test case
    all_results = []

    for test_case in TEST_CASES:
        print("\n" + "=" * 70)
        print(f"TEST CASE: {test_case['name']}")
        print("=" * 70)
        print(f"Original Text:\n{test_case['text'][:150]}...")
        print(f"\nExpected Category: {test_case['expected_category']}")
        print(f"Expected Terms: {test_case['expected_terms']}")

        # Test baseline
        baseline_result = test_baseline(test_case)
        print_result(baseline_result, test_case['name'])
        all_results.append({
            "test_case": test_case['name'],
            "baseline": baseline_result
        })

        # Test each segmenter
        segmented_results = {}
        for name, func in available_libs:
            result = test_segmented(test_case, name, func)
            print_result(result, test_case['name'])
            segmented_results[name] = result

        all_results[-1]["segmented"] = segmented_results

    # Summary comparison
    print("\n\n" + "=" * 70)
    print("SUMMARY COMPARISON")
    print("=" * 70)

    for result_set in all_results:
        print(f"\n{result_set['test_case']}:")
        print(f"  Baseline: {result_set['baseline']['term_count']} terms, "
              f"{result_set['baseline']['elapsed_ms']}ms, "
              f"method={result_set['baseline']['extraction_method']}")

        for lib_name, lib_result in result_set['segmented'].items():
            if not lib_result.get('error'):
                print(f"  {lib_name}: {lib_result['term_count']} terms, "
                      f"{lib_result['elapsed_ms']}ms, "
                      f"method={lib_result['extraction_method']}")

    # Recommendations
    print("\n\n" + "=" * 70)
    print("RECOMMENDATIONS")
    print("=" * 70)

    # Calculate average term counts for each method
    avg_counts = {"baseline": 0}
    for lib_name, _ in available_libs:
        avg_counts[lib_name] = 0

    for result_set in all_results:
        avg_counts["baseline"] += result_set['baseline']['term_count']
        for lib_name, lib_result in result_set['segmented'].items():
            if not lib_result.get('error'):
                avg_counts[lib_name] += lib_result['term_count']

    num_tests = len(all_results)
    print("\nAverage term extraction:")
    for method, total in avg_counts.items():
        avg = total / num_tests
        print(f"  {method}: {avg:.1f} terms/document")

    # Find best performer
    best_method = max(avg_counts.items(), key=lambda x: x[1])
    print(f"\n✓ Best performer: {best_method[0]} ({best_method[1]/num_tests:.1f} avg terms)")

    print("\n" + "=" * 70)
    print("NEXT STEPS")
    print("=" * 70)
    print("1. Choose the best-performing segmentation library")
    print("2. Integrate into DomainTermExtractor._llm_extract() method")
    print("3. Add language detection to only segment CJK text")
    print("4. Test on real customer documents")
    print("=" * 70)


if __name__ == "__main__":
    compare_all_methods()
