"""Test script for domain extraction."""
import sys
sys.path.insert(0, '/home/neil/Documents/textiq-doc-extraction')

from src.extraction_v2.domain_term_extractor import DomainTermExtractor

def test_finance_extraction():
    """Test extraction on finance content."""
    extractor = DomainTermExtractor()

    finance_text = """
    Q4 2023 Financial Results
    Revenue: $12.5M (23% YoY growth)
    EBITDA: $4.0M (32% margin)
    Operating cash flow increased significantly
    Net income exceeded forecast by 15%
    """

    result = extractor.extract_from_chunk(finance_text)
    print("=" * 60)
    print("FINANCE CONTENT TEST")
    print("=" * 60)
    print(f"Terms extracted: {result['domain_terms']}")
    print(f"Category: {result['domain_category']}")
    print(f"Extraction method: {result['extraction_method']}")
    print(f"Number of terms: {len(result['domain_terms'])}")

    # Verify finance category
    assert result['domain_category'] == 'finance', f"Expected 'finance', got '{result['domain_category']}'"
    assert 'ebitda' in result['domain_terms'] or 'revenue' in result['domain_terms'], "Should extract finance terms"
    print("✓ Finance test passed")
    return result


def test_hr_extraction():
    """Test extraction on HR content."""
    extractor = DomainTermExtractor()

    hr_text = """
    Q3 Hiring Report
    New employees hired: 15 engineers
    Total headcount: 200 employees
    Average salary: $120K
    Benefits package includes equity compensation
    Training programs expanded
    """

    result = extractor.extract_from_chunk(hr_text)
    print("\n" + "=" * 60)
    print("HR CONTENT TEST")
    print("=" * 60)
    print(f"Terms extracted: {result['domain_terms']}")
    print(f"Category: {result['domain_category']}")
    print(f"Extraction method: {result['extraction_method']}")
    print(f"Number of terms: {len(result['domain_terms'])}")

    # Verify HR category
    assert result['domain_category'] == 'hr', f"Expected 'hr', got '{result['domain_category']}'"
    assert any(term in result['domain_terms'] for term in ['employee', 'salary', 'hire']), "Should extract HR terms"
    print("✓ HR test passed")
    return result


def test_excel_extraction():
    """Test extraction on Excel record."""
    extractor = DomainTermExtractor()

    headers = ["Employee ID", "Name", "Department", "Salary", "Hire Date"]
    row_data = {
        "Employee ID": "E12345",
        "Name": "John Doe",
        "Department": "Engineering",
        "Salary": 120000,
        "Hire Date": "2023-01-15"
    }

    result = extractor.extract_from_excel_record(headers, row_data, sheet_name="Employee Data")
    print("\n" + "=" * 60)
    print("EXCEL RECORD TEST")
    print("=" * 60)
    print(f"Terms extracted: {result['domain_terms']}")
    print(f"Category: {result['domain_category']}")
    print(f"Extraction method: {result['extraction_method']}")

    # Should extract employee/salary terms
    assert any(term in result['domain_terms'] for term in ['employee', 'salary', 'engineering']), \
        "Should extract employee-related terms"
    print("✓ Excel extraction test passed")
    return result


def test_empty_content():
    """Test extraction on empty content."""
    extractor = DomainTermExtractor()

    result = extractor.extract_from_chunk("")
    print("\n" + "=" * 60)
    print("EMPTY CONTENT TEST")
    print("=" * 60)
    print(f"Terms extracted: {result['domain_terms']}")
    print(f"Category: {result['domain_category']}")

    assert result['domain_terms'] == [], "Should return empty list for empty content"
    assert result['domain_category'] == 'general', "Should return 'general' category"
    print("✓ Empty content test passed")
    return result


if __name__ == "__main__":
    print("Testing Domain Term Extractor\n")

    try:
        test_finance_extraction()
        test_hr_extraction()
        test_excel_extraction()
        test_empty_content()

        print("\n" + "=" * 60)
        print("ALL TESTS PASSED ✓")
        print("=" * 60)

    except Exception as e:
        print(f"\n✗ TEST FAILED: {e}")
        import traceback
        traceback.print_exc()
        sys.exit(1)
