"""Unit tests for Japanese segmentation in DomainTermExtractor."""

import os
import sys
import threading
from unittest.mock import MagicMock, patch

import pytest

from src.extraction_v2.domain_term_extractor import DomainTermExtractor


class TestJapaneseLanguageDetection:
    """Tests for _is_japanese() method."""

    def test_japanese_hiragana_detected(self):
        """Should detect Hiragana text as Japanese."""
        extractor = DomainTermExtractor()
        assert extractor._is_japanese("こんにちは") is True

    def test_japanese_katakana_detected(self):
        """Should detect Katakana text as Japanese."""
        extractor = DomainTermExtractor()
        assert extractor._is_japanese("カタカナ") is True

    def test_japanese_kanji_detected(self):
        """Should detect Kanji text as Japanese."""
        extractor = DomainTermExtractor()
        assert extractor._is_japanese("日本語") is True

    def test_english_not_detected(self):
        """Should not detect English text as Japanese."""
        extractor = DomainTermExtractor()
        assert extractor._is_japanese("Hello World") is False

    def test_mixed_japanese_english_detected(self):
        """Should detect mixed Japanese/English text as Japanese."""
        extractor = DomainTermExtractor()
        # Text with Hiragana is unambiguously Japanese
        assert extractor._is_japanese("Hello 日本です") is True

    def test_empty_string_not_detected(self):
        """Should not detect empty string as Japanese."""
        extractor = DomainTermExtractor()
        assert extractor._is_japanese("") is False

    def test_whitespace_only_not_detected(self):
        """Should not detect whitespace-only text as Japanese."""
        extractor = DomainTermExtractor()
        assert extractor._is_japanese("   ") is False

    def test_chinese_detected_as_japanese(self):
        """Should detect Chinese text as Japanese (known limitation - Kanji overlap)."""
        extractor = DomainTermExtractor()
        # Chinese "你好" contains Kanji range characters
        # This is expected behavior - segmentation will fail gracefully and fall back
        assert extractor._is_japanese("你好") is True


class TestJapaneseSegmentation:
    """Tests for _segment_japanese() method."""

    def test_simple_japanese_sentence_segmented(self):
        """Should segment simple Japanese sentence with spaces."""
        extractor = DomainTermExtractor()

        # Mock tokenizer to return segmented tokens
        mock_tokenizer = MagicMock()
        mock_token1 = MagicMock()
        mock_token1.surface.return_value = "売上"
        mock_token2 = MagicMock()
        mock_token2.surface.return_value = "高"
        mock_token3 = MagicMock()
        mock_token3.surface.return_value = "は"
        mock_token4 = MagicMock()
        mock_token4.surface.return_value = "増加"
        mock_token5 = MagicMock()
        mock_token5.surface.return_value = "した"

        mock_tokenizer.tokenize.return_value = [
            mock_token1, mock_token2, mock_token3, mock_token4, mock_token5
        ]

        extractor._sudachi_tokenizer = mock_tokenizer
        extractor._sudachi_init_attempted = True

        result = extractor._segment_japanese("売上高は増加した")

        # Should return space-separated tokens
        assert " " in result
        assert len(result.split()) > 1

    def test_empty_input_returns_empty(self):
        """Should return empty string for empty input."""
        extractor = DomainTermExtractor()

        # Mock tokenizer
        mock_tokenizer = MagicMock()
        mock_tokenizer.tokenize.return_value = []
        extractor._sudachi_tokenizer = mock_tokenizer
        extractor._sudachi_init_attempted = True

        result = extractor._segment_japanese("")
        assert result == ""

    def test_segmentation_error_fallback(self):
        """Should return original text if segmentation fails."""
        extractor = DomainTermExtractor()

        # Mock tokenizer that raises exception
        mock_tokenizer = MagicMock()
        mock_tokenizer.tokenize.side_effect = Exception("Segmentation failed")
        extractor._sudachi_tokenizer = mock_tokenizer
        extractor._sudachi_init_attempted = True

        original_text = "日本語テスト"
        result = extractor._segment_japanese(original_text)

        # Should return original text as fallback
        assert result == original_text

    def test_tokenizer_unavailable_fallback(self):
        """Should return original text if tokenizer is None."""
        extractor = DomainTermExtractor()
        extractor._sudachi_tokenizer = None
        extractor._sudachi_init_attempted = True

        original_text = "日本語テスト"
        result = extractor._segment_japanese(original_text)

        # Should return original text as fallback
        assert result == original_text


class TestLazyTokenizerInitialization:
    """Tests for _get_sudachi_tokenizer() method."""

    def test_first_call_initializes_tokenizer(self):
        """Should initialize tokenizer on first call."""
        extractor = DomainTermExtractor()

        # Use actual sudachipy if available, otherwise skip
        try:
            import sudachipy
            result = extractor._get_sudachi_tokenizer()
            assert result is not None
            assert extractor._sudachi_init_attempted is True
        except ImportError:
            pytest.skip("sudachipy not installed")

    def test_second_call_returns_cached(self):
        """Should return cached tokenizer on second call."""
        extractor = DomainTermExtractor()

        # Pre-initialize
        mock_tokenizer = MagicMock()
        extractor._sudachi_tokenizer = mock_tokenizer
        extractor._sudachi_init_attempted = True

        result = extractor._get_sudachi_tokenizer()

        # Should return same instance
        assert result == mock_tokenizer

    def test_import_error_handling(self):
        """Should handle ImportError gracefully and mark as initialized."""
        # This test verifies the behavior when sudachipy is not installed
        # Since we can't easily mock the import without affecting other tests,
        # we verify the initialization flag behavior
        extractor = DomainTermExtractor()

        # Set up a scenario where tokenizer fails to initialize
        extractor._sudachi_init_attempted = False

        # Directly test that if we mark it as failed, subsequent calls return None
        extractor._sudachi_tokenizer = None
        extractor._sudachi_init_attempted = True

        result = extractor._get_sudachi_tokenizer()
        assert result is None

    def test_tokenizer_not_reinitialized_after_failure(self):
        """Should not retry initialization after first failure."""
        extractor = DomainTermExtractor()

        # Simulate a failed initialization state
        extractor._sudachi_tokenizer = None
        extractor._sudachi_init_attempted = True

        # Call multiple times
        result1 = extractor._get_sudachi_tokenizer()
        result2 = extractor._get_sudachi_tokenizer()

        # Should return None both times and not attempt reinitialization
        assert result1 is None
        assert result2 is None
        assert extractor._sudachi_init_attempted is True


class TestEnvironmentVariableControl:
    """Tests for environment variable control of segmentation."""

    def test_segmentation_enabled_with_env_true(self):
        """Should enable segmentation when ENABLE_CJK_SEGMENTATION=true."""
        # Use direct mocking instead of reload to avoid side effects
        with patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", True):
            from src.extraction_v2.domain_term_extractor import ENABLE_CJK_SEGMENTATION
            assert ENABLE_CJK_SEGMENTATION is True

    def test_segmentation_disabled_with_env_false(self):
        """Should disable segmentation when ENABLE_CJK_SEGMENTATION=false."""
        # Use direct mocking instead of reload to avoid side effects
        with patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", False):
            from src.extraction_v2.domain_term_extractor import ENABLE_CJK_SEGMENTATION
            assert ENABLE_CJK_SEGMENTATION is False

    def test_missing_env_var_defaults_to_true(self):
        """Should default to true when env var is missing."""
        # When env var is not set, os.getenv should return the default value
        # We test this by checking that the parsing logic works correctly
        result = os.getenv("NONEXISTENT_VAR_FOR_TEST", "true").lower() == "true"
        assert result is True

    def test_invalid_value_treated_as_false(self):
        """Should treat invalid value as false."""
        # Test environment variable parsing logic
        with patch.dict("os.environ", {"ENABLE_CJK_SEGMENTATION": "invalid"}):
            result = os.getenv("ENABLE_CJK_SEGMENTATION", "true").lower() == "true"
            assert result is False


class TestLLMExtractIntegration:
    """Tests for integration with _llm_extract() method."""

    @patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", True)
    def test_japanese_text_segmented_before_llm(self):
        """Should segment Japanese text before sending to LLM."""
        extractor = DomainTermExtractor()

        # Mock the client
        mock_client = MagicMock()
        mock_response = MagicMock()
        mock_response.choices[0].message.content = '["term1", "term2"]'
        mock_client.chat.completions.create.return_value = mock_response

        with patch.object(extractor, "_get_client", return_value=mock_client):
            with patch.object(extractor, "_segment_japanese", return_value="segmented text") as mock_segment:
                result = extractor._llm_extract("日本語テスト")

                # Should have called segmentation
                mock_segment.assert_called_once()

    @patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", False)
    def test_japanese_text_not_segmented_when_disabled(self):
        """Should not segment Japanese text when env var is false."""
        extractor = DomainTermExtractor()

        # Mock the client
        mock_client = MagicMock()
        mock_response = MagicMock()
        mock_response.choices[0].message.content = '["term1", "term2"]'
        mock_client.chat.completions.create.return_value = mock_response

        with patch.object(extractor, "_get_client", return_value=mock_client):
            with patch.object(extractor, "_segment_japanese") as mock_segment:
                result = extractor._llm_extract("日本語テスト")

                # Should NOT have called segmentation
                mock_segment.assert_not_called()

    @patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", True)
    def test_english_text_not_segmented(self):
        """Should not segment English text."""
        extractor = DomainTermExtractor()

        # Mock the client
        mock_client = MagicMock()
        mock_response = MagicMock()
        mock_response.choices[0].message.content = '["term1", "term2"]'
        mock_client.chat.completions.create.return_value = mock_response

        with patch.object(extractor, "_get_client", return_value=mock_client):
            with patch.object(extractor, "_segment_japanese") as mock_segment:
                result = extractor._llm_extract("English text")

                # Should NOT have called segmentation (no Japanese detected)
                mock_segment.assert_not_called()

    @patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", True)
    def test_segmentation_failure_fallback_to_original(self):
        """Should use original text if segmentation fails."""
        extractor = DomainTermExtractor()

        # Mock the client
        mock_client = MagicMock()
        mock_response = MagicMock()
        mock_response.choices[0].message.content = '["term1", "term2"]'
        mock_client.chat.completions.create.return_value = mock_response

        japanese_text = "日本語テスト"

        with patch.object(extractor, "_get_client", return_value=mock_client):
            with patch.object(extractor, "_segment_japanese", return_value=japanese_text):
                result = extractor._llm_extract(japanese_text)

                # Should complete successfully even if segmentation returns original
                assert result == ["term1", "term2"]

    @patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", True)
    def test_llm_receives_segmented_text(self):
        """Should verify LLM receives segmented text in prompt."""
        extractor = DomainTermExtractor()

        # Mock the client
        mock_client = MagicMock()
        mock_response = MagicMock()
        mock_response.choices[0].message.content = '["term1", "term2"]'
        mock_client.chat.completions.create.return_value = mock_response

        japanese_text = "日本語テスト"
        segmented_text = "日本 語 テスト"

        with patch.object(extractor, "_get_client", return_value=mock_client):
            with patch.object(extractor, "_segment_japanese", return_value=segmented_text):
                result = extractor._llm_extract(japanese_text)

                # Get the actual prompt sent to LLM
                call_args = mock_client.chat.completions.create.call_args
                messages = call_args.kwargs["messages"]
                user_message = messages[1]["content"]

                # Segmented text should be in the prompt
                assert segmented_text in user_message


class TestJapaneseIntegrationEndToEnd:
    """Integration tests demonstrating actual Japanese domain term extraction improvement."""

    @patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", True)
    def test_japanese_financial_terms_extracted_with_segmentation(self):
        """Should extract Japanese financial domain terms better with segmentation enabled."""
        extractor = DomainTermExtractor()

        # Japanese financial content with domain terms
        japanese_financial_text = "売上高は前年比で増加した。営業利益と純利益も改善している。"

        # Mock the client to return financial terms
        mock_client = MagicMock()
        mock_response = MagicMock()
        # Simulate LLM extracting terms from segmented text
        mock_response.choices[0].message.content = '["売上", "営業利益", "純利益", "前年比"]'
        mock_client.chat.completions.create.return_value = mock_response

        with patch.object(extractor, "_get_client", return_value=mock_client):
            result = extractor.extract_from_chunk(japanese_financial_text)

            # Should have extracted financial terms
            assert "domain_terms" in result
            assert len(result["domain_terms"]) > 0
            # Verify segmentation was called (LLM received segmented text)
            call_args = mock_client.chat.completions.create.call_args
            messages = call_args.kwargs["messages"]
            user_message = messages[1]["content"]
            # Segmented text should contain spaces
            assert " " in user_message

    @patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", False)
    def test_japanese_extraction_without_segmentation_baseline(self):
        """Should still extract some terms without segmentation (baseline comparison)."""
        extractor = DomainTermExtractor()

        # Same Japanese financial content
        japanese_financial_text = "売上高は前年比で増加した。営業利益と純利益も改善している。"

        # Mock the client
        mock_client = MagicMock()
        mock_response = MagicMock()
        # Without segmentation, LLM might extract fewer or different terms
        mock_response.choices[0].message.content = '["売上高", "営業利益"]'
        mock_client.chat.completions.create.return_value = mock_response

        with patch.object(extractor, "_get_client", return_value=mock_client):
            result = extractor.extract_from_chunk(japanese_financial_text)

            # Should still work but potentially with different results
            assert "domain_terms" in result
            assert isinstance(result["domain_terms"], list)
            # Verify segmentation was NOT called
            call_args = mock_client.chat.completions.create.call_args
            messages = call_args.kwargs["messages"]
            user_message = messages[1]["content"]
            # Original text without added spaces from segmentation
            assert japanese_financial_text in user_message


class TestBackwardCompatibilityNonJapanese:
    """Tests verifying backward compatibility for non-Japanese languages."""

    @patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", True)
    def test_english_financial_terms_unchanged_with_segmentation_enabled(self):
        """Should extract English terms correctly with segmentation enabled (no regression)."""
        extractor = DomainTermExtractor()

        english_text = "Revenue increased compared to last year. Operating profit and net income also improved."

        # Mock the client
        mock_client = MagicMock()
        mock_response = MagicMock()
        mock_response.choices[0].message.content = '["revenue", "operating profit", "net income"]'
        mock_client.chat.completions.create.return_value = mock_response

        with patch.object(extractor, "_get_client", return_value=mock_client):
            result = extractor.extract_from_chunk(english_text)

            # Should extract terms normally
            assert "domain_terms" in result
            terms = result["domain_terms"]
            assert len(terms) == 3
            assert "revenue" in terms
            # If LLM was called, verify segmentation was NOT applied to English text
            if mock_client.chat.completions.create.called:
                call_args = mock_client.chat.completions.create.call_args
                messages = call_args.kwargs["messages"]
                user_message = messages[1]["content"]
                # Original English text should be in prompt unchanged
                assert english_text in user_message

    @patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", True)
    def test_spanish_terms_unchanged_with_segmentation_enabled(self):
        """Should extract Spanish terms correctly (non-CJK language)."""
        extractor = DomainTermExtractor()

        spanish_text = "Los ingresos aumentaron. El beneficio operativo mejoró significativamente."

        # Mock the client
        mock_client = MagicMock()
        mock_response = MagicMock()
        mock_response.choices[0].message.content = '["ingresos", "beneficio operativo"]'
        mock_client.chat.completions.create.return_value = mock_response

        with patch.object(extractor, "_get_client", return_value=mock_client):
            result = extractor.extract_from_chunk(spanish_text)

            # Should extract terms normally
            assert "domain_terms" in result
            terms = result["domain_terms"]
            assert len(terms) == 2
            # Spanish should not be detected as Japanese
            if mock_client.chat.completions.create.called:
                call_args = mock_client.chat.completions.create.call_args
                messages = call_args.kwargs["messages"]
                user_message = messages[1]["content"]
                assert spanish_text in user_message

    @patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", True)
    def test_mixed_english_numbers_unchanged(self):
        """Should handle mixed English with numbers correctly."""
        extractor = DomainTermExtractor()

        mixed_text = "Q4 revenue reached $1.5M. EBITDA margin improved to 25%."

        # Mock the client
        mock_client = MagicMock()
        mock_response = MagicMock()
        mock_response.choices[0].message.content = '["revenue", "ebitda", "margin"]'
        mock_client.chat.completions.create.return_value = mock_response

        with patch.object(extractor, "_get_client", return_value=mock_client):
            result = extractor.extract_from_chunk(mixed_text)

            # Should work correctly
            assert "domain_terms" in result
            terms = result["domain_terms"]
            assert len(terms) > 0
            # No Japanese segmentation should occur if LLM was called
            if mock_client.chat.completions.create.called:
                call_args = mock_client.chat.completions.create.call_args
                messages = call_args.kwargs["messages"]
                user_message = messages[1]["content"]
                assert "$1.5M" in user_message  # Numbers preserved
                assert "25%" in user_message  # Percentages preserved


class TestConcurrentAccessPatterns:
    """Tests for concurrent access patterns in Celery/Airflow workers."""

    def test_concurrent_tokenizer_initialization(self):
        """Should safely handle concurrent initialization attempts."""
        extractor = DomainTermExtractor()
        errors = []
        init_count = []

        def init_tokenizer(thread_id):
            """Simulate worker thread initializing tokenizer."""
            try:
                tokenizer = extractor._get_sudachi_tokenizer()
                # Track initialization
                if tokenizer is not None:
                    init_count.append(thread_id)
            except Exception as e:
                errors.append((thread_id, str(e)))

        # Simulate 5 concurrent workers
        threads = []
        for i in range(5):
            t = threading.Thread(target=init_tokenizer, args=(i,))
            threads.append(t)
            t.start()

        # Wait for all threads
        for t in threads:
            t.join()

        # Should have no errors
        assert len(errors) == 0, f"Concurrent initialization errors: {errors}"

        # Flag should be set exactly once
        assert extractor._sudachi_init_attempted is True

    def test_concurrent_term_extraction(self):
        """Should safely handle concurrent term extraction calls."""
        extractor = DomainTermExtractor()
        results = []
        errors = []

        # Mock the client to avoid actual API calls
        mock_client = MagicMock()
        mock_response = MagicMock()
        mock_response.choices[0].message.content = '["test", "terms"]'
        mock_client.chat.completions.create.return_value = mock_response

        def extract_terms(thread_id, text):
            """Simulate worker thread extracting terms."""
            try:
                with patch.object(extractor, "_get_client", return_value=mock_client):
                    result = extractor.extract_from_chunk(text)
                    results.append((thread_id, result))
            except Exception as e:
                errors.append((thread_id, str(e)))

        # Test with multiple texts
        texts = [
            "Revenue increased in Q4",
            "Employee benefits improved",
            "API deployment successful",
            "Contract terms updated",
            "Inventory levels normal"
        ]

        # Simulate 5 concurrent workers
        threads = []
        for i, text in enumerate(texts):
            t = threading.Thread(target=extract_terms, args=(i, text))
            threads.append(t)
            t.start()

        # Wait for all threads
        for t in threads:
            t.join()

        # Should have no errors
        assert len(errors) == 0, f"Concurrent extraction errors: {errors}"

        # Should have 5 successful results
        assert len(results) == 5

        # All results should have expected structure
        for thread_id, result in results:
            assert "domain_terms" in result
            assert "domain_category" in result
            assert "extraction_method" in result

    @patch("src.extraction_v2.domain_term_extractor.ENABLE_CJK_SEGMENTATION", True)
    def test_concurrent_japanese_segmentation(self):
        """Should safely handle concurrent Japanese segmentation."""
        extractor = DomainTermExtractor()
        results = []
        errors = []

        # Japanese text for testing
        japanese_text = "売上高が増加しました。"

        def segment_text(thread_id):
            """Simulate worker thread segmenting Japanese text."""
            try:
                # Check if text is Japanese
                is_jp = extractor._is_japanese(japanese_text)
                results.append((thread_id, is_jp))
            except Exception as e:
                errors.append((thread_id, str(e)))

        # Simulate 10 concurrent workers
        threads = []
        for i in range(10):
            t = threading.Thread(target=segment_text, args=(i,))
            threads.append(t)
            t.start()

        # Wait for all threads
        for t in threads:
            t.join()

        # Should have no errors
        assert len(errors) == 0, f"Concurrent segmentation errors: {errors}"

        # All results should correctly identify Japanese text
        assert len(results) == 10
        for thread_id, is_jp in results:
            assert is_jp is True
