"""Unit tests for symbol resolution in extracted content."""

from unittest.mock import MagicMock, Mock, patch
from uuid import uuid4

import pytest

from src.extraction.models import SymbolResolutionSummary
from src.extraction.symbol_resolver import (
    BATCH_SIZE,
    SYMBOL_CHARS,
    SymbolResolver,
)


@pytest.fixture
def resolver():
    """Create a SymbolResolver instance."""
    return SymbolResolver()


@pytest.fixture
def sample_symbol_lookup():
    """Create a sample symbol lookup dictionary."""
    return {
        "○": "Applicable/Yes",
        "×": "Not applicable/No",
        "△": "Pending review",
        "07": "No required fields",
        "E01": "Invalid input",
    }


class TestSymbolResolutionSummaryModel:
    """Tests for SymbolResolutionSummary dataclass."""

    def test_valid_summary_creation(self):
        """Test creating a valid resolution summary."""
        doc_id = uuid4()
        summary = SymbolResolutionSummary(
            document_id=doc_id,
            records_processed=100,
            symbols_resolved=50,
            unresolved_count=5,
            unresolved_symbols=["☆", "★"],
            warnings=[],
        )
        assert summary.document_id == doc_id
        assert summary.records_processed == 100
        assert summary.symbols_resolved == 50
        assert summary.unresolved_count == 5

    def test_summary_to_dict(self):
        """Test SymbolResolutionSummary.to_dict() serialization."""
        doc_id = uuid4()
        summary = SymbolResolutionSummary(
            document_id=doc_id,
            records_processed=10,
            symbols_resolved=5,
            unresolved_count=2,
            unresolved_symbols=["☆", "★"],
            warnings=["Test warning"],
        )
        result = summary.to_dict()

        assert result["document_id"] == str(doc_id)
        assert result["records_processed"] == 10
        assert result["symbols_resolved"] == 5
        assert result["unresolved_count"] == 2
        assert result["unresolved_symbols"] == ["☆", "★"]
        assert result["warnings"] == ["Test warning"]

    def test_summary_empty_lists(self):
        """Test summary with empty lists."""
        doc_id = uuid4()
        summary = SymbolResolutionSummary(
            document_id=doc_id,
            records_processed=0,
            symbols_resolved=0,
            unresolved_count=0,
            unresolved_symbols=[],
            warnings=[],
        )
        result = summary.to_dict()
        assert result["unresolved_symbols"] == []
        assert result["warnings"] == []


class TestLooksLikeSymbol:
    """Tests for the _looks_like_symbol heuristic."""

    def test_common_symbols_detected(self, resolver):
        """Test that common symbols are detected."""
        assert resolver._looks_like_symbol("○") is True
        assert resolver._looks_like_symbol("×") is True
        assert resolver._looks_like_symbol("△") is True
        assert resolver._looks_like_symbol("◎") is True
        assert resolver._looks_like_symbol("●") is True
        assert resolver._looks_like_symbol("□") is True
        assert resolver._looks_like_symbol("■") is True

    def test_short_alphanumeric_codes_detected(self, resolver):
        """Test that short alphanumeric codes are detected."""
        assert resolver._looks_like_symbol("07") is True
        assert resolver._looks_like_symbol("A1") is True
        assert resolver._looks_like_symbol("E01") is True
        assert resolver._looks_like_symbol("X") is True

    def test_long_alphanumeric_not_detected(self, resolver):
        """Test that long alphanumeric strings are not detected."""
        assert resolver._looks_like_symbol("ABCDE") is False
        assert resolver._looks_like_symbol("12345") is False

    def test_empty_string_not_detected(self, resolver):
        """Test that empty string is not detected."""
        assert resolver._looks_like_symbol("") is False

    def test_long_string_not_detected(self, resolver):
        """Test that long strings are not detected."""
        assert resolver._looks_like_symbol("This is a long string") is False
        assert resolver._looks_like_symbol("x" * 11) is False

    def test_normal_text_not_detected(self, resolver):
        """Test that normal text is not detected as symbol."""
        assert resolver._looks_like_symbol("Hello") is False
        assert resolver._looks_like_symbol("Product Name") is False


class TestResolveRecord:
    """Tests for _resolve_record method."""

    def test_resolve_single_symbol(self, resolver, sample_symbol_lookup):
        """Test resolving a single symbol in a record."""
        content = {"Item Code": "2024-4_0019", "Screen": "○"}
        resolved, count, unresolved = resolver._resolve_record(
            content, sample_symbol_lookup
        )

        assert resolved is not None
        assert count == 1
        assert resolved["Screen"] == "○"  # Original preserved
        assert resolved["Screen_resolved"] == "Applicable/Yes"
        assert resolved["Item Code"] == "2024-4_0019"  # Unchanged

    def test_resolve_multiple_symbols(self, resolver, sample_symbol_lookup):
        """Test resolving multiple symbols in a record."""
        content = {"Field1": "○", "Field2": "×", "Field3": "△"}
        resolved, count, unresolved = resolver._resolve_record(
            content, sample_symbol_lookup
        )

        assert resolved is not None
        assert count == 3
        assert resolved["Field1_resolved"] == "Applicable/Yes"
        assert resolved["Field2_resolved"] == "Not applicable/No"
        assert resolved["Field3_resolved"] == "Pending review"

    def test_no_symbols_to_resolve(self, resolver, sample_symbol_lookup):
        """Test when record has no symbols to resolve."""
        content = {"Name": "John Doe", "Email": "john@example.com"}
        resolved, count, unresolved = resolver._resolve_record(
            content, sample_symbol_lookup
        )

        assert resolved is None
        assert count == 0

    def test_unknown_symbol_tracking(self, resolver, sample_symbol_lookup):
        """Test that unknown symbols are tracked."""
        # Use ● which is in SYMBOL_CHARS but not in lookup
        content = {"Field1": "○", "Field2": "●"}  # ● is not in lookup
        resolved, count, unresolved = resolver._resolve_record(
            content, sample_symbol_lookup
        )

        assert count == 1  # Only ○ resolved
        assert "●" in unresolved

    def test_empty_content(self, resolver, sample_symbol_lookup):
        """Test with empty content dictionary."""
        content = {}
        resolved, count, unresolved = resolver._resolve_record(
            content, sample_symbol_lookup
        )

        assert resolved is None
        assert count == 0
        assert len(unresolved) == 0

    def test_none_values_skipped(self, resolver, sample_symbol_lookup):
        """Test that None values are skipped."""
        content = {"Field1": "○", "Field2": None, "Field3": "×"}
        resolved, count, unresolved = resolver._resolve_record(
            content, sample_symbol_lookup
        )

        assert count == 2
        assert "Field2_resolved" not in resolved

    def test_numeric_values_converted(self, resolver, sample_symbol_lookup):
        """Test that numeric values are converted to string."""
        content = {"Code": 7, "Status": "○"}  # 7 is numeric, not "07"
        resolved, count, unresolved = resolver._resolve_record(
            content, sample_symbol_lookup
        )

        assert count == 1  # Only ○ resolved, "7" != "07"

    def test_whitespace_stripped(self, resolver, sample_symbol_lookup):
        """Test that whitespace is stripped from values."""
        content = {"Field1": " ○ ", "Field2": "  ×  "}
        resolved, count, unresolved = resolver._resolve_record(
            content, sample_symbol_lookup
        )

        assert count == 2
        assert resolved["Field1_resolved"] == "Applicable/Yes"
        assert resolved["Field2_resolved"] == "Not applicable/No"

    def test_alphanumeric_code_resolution(self, resolver, sample_symbol_lookup):
        """Test resolving alphanumeric codes like '07'."""
        content = {"ErrorCode": "07", "Status": "E01"}
        resolved, count, unresolved = resolver._resolve_record(
            content, sample_symbol_lookup
        )

        assert count == 2
        assert resolved["ErrorCode_resolved"] == "No required fields"
        assert resolved["Status_resolved"] == "Invalid input"


class TestBuildSymbolLookup:
    """Tests for _build_symbol_lookup method."""

    def test_build_lookup_from_database(self, resolver):
        """Test building symbol lookup from database records."""
        doc_id = uuid4()

        # Create mock session and records
        mock_session = Mock()
        mock_record1 = Mock()
        mock_record1.symbol = "○"
        mock_record1.meaning = "Yes"
        mock_record1.context = None  # Auto-detected

        mock_record2 = Mock()
        mock_record2.symbol = "×"
        mock_record2.meaning = "No"
        mock_record2.context = None  # Auto-detected

        mock_query = Mock()
        mock_query.filter.return_value.order_by.return_value.all.return_value = [
            mock_record1,
            mock_record2,
        ]
        mock_session.query.return_value = mock_query

        lookup = resolver._build_symbol_lookup(doc_id, mock_session)

        assert len(lookup) == 2
        assert lookup["○"] == "Yes"
        assert lookup["×"] == "No"

    def test_build_lookup_empty_database(self, resolver):
        """Test building lookup when no symbols exist."""
        doc_id = uuid4()

        mock_session = Mock()
        mock_query = Mock()
        mock_query.filter.return_value.order_by.return_value.all.return_value = []
        mock_session.query.return_value = mock_query

        lookup = resolver._build_symbol_lookup(doc_id, mock_session)

        assert len(lookup) == 0

    def test_build_lookup_handles_duplicates(self, resolver):
        """Test that duplicate auto-detected symbols keep first occurrence."""
        doc_id = uuid4()

        mock_session = Mock()
        mock_record1 = Mock()
        mock_record1.symbol = "○"
        mock_record1.meaning = "First meaning"
        mock_record1.context = None  # Auto-detected

        mock_record2 = Mock()
        mock_record2.symbol = "○"  # Duplicate
        mock_record2.meaning = "Second meaning"
        mock_record2.context = None  # Auto-detected

        mock_query = Mock()
        mock_query.filter.return_value.order_by.return_value.all.return_value = [
            mock_record1,
            mock_record2,
        ]
        mock_session.query.return_value = mock_query

        lookup = resolver._build_symbol_lookup(doc_id, mock_session)

        assert len(lookup) == 1
        assert lookup["○"] == "First meaning"  # First auto-detected wins

    def test_build_lookup_custom_overrides_auto(self, resolver):
        """Test that custom symbols override auto-detected symbols."""
        doc_id = uuid4()

        mock_session = Mock()
        # Auto-detected symbol
        mock_record_auto = Mock()
        mock_record_auto.symbol = "○"
        mock_record_auto.meaning = "Auto meaning"
        mock_record_auto.context = None

        # Custom symbol (should come last due to ordering and override)
        mock_record_custom = Mock()
        mock_record_custom.symbol = "○"
        mock_record_custom.meaning = "Custom meaning"
        mock_record_custom.context = "api_upload"

        mock_query = Mock()
        # Order: auto-detected first, custom last (as per the order_by logic)
        mock_query.filter.return_value.order_by.return_value.all.return_value = [
            mock_record_auto,
            mock_record_custom,
        ]
        mock_session.query.return_value = mock_query

        lookup = resolver._build_symbol_lookup(doc_id, mock_session)

        assert len(lookup) == 1
        assert lookup["○"] == "Custom meaning"  # Custom overrides auto


class TestBatchUpdateResolvedContent:
    """Tests for _batch_update_resolved_content method."""

    def test_batch_update_single_batch(self, resolver):
        """Test updating records in a single batch."""
        mock_session = Mock()
        mock_query_obj = Mock()
        mock_session.query.return_value = mock_query_obj
        mock_query_obj.filter.return_value = mock_query_obj
        mock_query_obj.update.return_value = 1

        updates = [
            (uuid4(), {"Field1": "○", "Field1_resolved": "Yes"}),
            (uuid4(), {"Field2": "×", "Field2_resolved": "No"}),
        ]

        count = resolver._batch_update_resolved_content(updates, mock_session)

        assert count == 2
        mock_session.commit.assert_called_once()

    def test_batch_update_multiple_batches(self, resolver):
        """Test updating records across multiple batches."""
        mock_session = Mock()
        mock_query_obj = Mock()
        mock_session.query.return_value = mock_query_obj
        mock_query_obj.filter.return_value = mock_query_obj
        mock_query_obj.update.return_value = 1

        # Create more updates than BATCH_SIZE
        updates = [(uuid4(), {"Field": "value"}) for _ in range(BATCH_SIZE + 50)]

        count = resolver._batch_update_resolved_content(updates, mock_session)

        assert count == BATCH_SIZE + 50
        assert mock_session.commit.call_count == 2  # Two batches

    def test_batch_update_empty_list(self, resolver):
        """Test updating with empty list."""
        mock_session = Mock()

        count = resolver._batch_update_resolved_content([], mock_session)

        assert count == 0
        mock_session.commit.assert_not_called()

    def test_batch_update_rollback_on_error(self, resolver):
        """Test that rollback is called on error."""
        mock_session = Mock()
        mock_query_obj = Mock()
        mock_session.query.return_value = mock_query_obj
        mock_query_obj.filter.return_value = mock_query_obj
        mock_query_obj.update.side_effect = Exception("DB Error")

        updates = [(uuid4(), {"Field": "value"})]

        with pytest.raises(Exception, match="DB Error"):
            resolver._batch_update_resolved_content(updates, mock_session)

        mock_session.rollback.assert_called_once()


class TestResolveDocument:
    """Tests for resolve_document main method."""

    def test_resolve_document_no_symbols(self, resolver):
        """Test when document has no symbol dictionaries."""
        doc_id = uuid4()
        mock_session = Mock()

        # No symbols in database
        mock_query = Mock()
        mock_query.filter.return_value.order_by.return_value.all.return_value = []
        mock_session.query.return_value = mock_query

        summary = resolver.resolve_document(doc_id, mock_session)

        assert summary.records_processed == 0
        assert summary.symbols_resolved == 0
        assert summary.unresolved_count == 0

    def test_resolve_document_no_records(self, resolver):
        """Test when document has symbols but no extracted records."""
        doc_id = uuid4()
        mock_session = Mock()

        # Mock symbol lookup
        mock_symbol = Mock()
        mock_symbol.symbol = "○"
        mock_symbol.meaning = "Yes"
        mock_symbol.context = None

        # First query returns symbols (with order_by), second returns empty records
        mock_query = Mock()
        call_count = [0]

        def filter_side_effect(*args, **kwargs):
            mock_filter = Mock()
            call_count[0] += 1
            if call_count[0] == 1:
                # Symbol query uses order_by
                mock_filter.order_by.return_value.all.return_value = [mock_symbol]
            else:
                # Record query doesn't use order_by
                mock_filter.all.return_value = []
            return mock_filter

        mock_query.filter.side_effect = filter_side_effect
        mock_session.query.return_value = mock_query

        summary = resolver.resolve_document(doc_id, mock_session)

        assert summary.records_processed == 0
        assert "No extracted records found" in summary.warnings[0]

    def test_resolve_document_full_flow(self, resolver):
        """Test full resolution flow with symbols and records."""
        doc_id = uuid4()
        record_id = uuid4()
        mock_session = Mock()

        # Mock symbol
        mock_symbol = Mock()
        mock_symbol.symbol = "○"
        mock_symbol.meaning = "Yes"
        mock_symbol.context = None

        # Mock record
        mock_record = Mock()
        mock_record.id = record_id
        mock_record.content = {"Status": "○", "Name": "Test"}

        # Setup query mock
        mock_query = Mock()
        call_count = [0]

        def filter_side_effect(*args, **kwargs):
            mock_filter = Mock()
            call_count[0] += 1
            if call_count[0] == 1:
                # Symbol query uses order_by
                mock_filter.order_by.return_value.all.return_value = [mock_symbol]
            elif call_count[0] == 2:
                # Record query
                mock_filter.all.return_value = [mock_record]
            else:
                # Update query
                mock_filter.update.return_value = 1
                return mock_filter
            return mock_filter

        mock_query.filter.side_effect = filter_side_effect
        mock_session.query.return_value = mock_query

        summary = resolver.resolve_document(doc_id, mock_session)

        assert summary.document_id == doc_id
        assert summary.records_processed == 1
        assert summary.symbols_resolved == 1

    def test_resolve_document_with_unknown_symbols(self, resolver):
        """Test resolution with unknown symbols in records."""
        doc_id = uuid4()
        mock_session = Mock()

        # Mock symbol
        mock_symbol = Mock()
        mock_symbol.symbol = "○"
        mock_symbol.meaning = "Yes"
        mock_symbol.context = None

        # Mock record with unknown symbol (use ● which is in SYMBOL_CHARS)
        mock_record = Mock()
        mock_record.id = uuid4()
        mock_record.content = {"Status": "○", "Other": "●"}  # ● unknown but in SYMBOL_CHARS

        # Setup query mock
        mock_query = Mock()
        call_count = [0]

        def filter_side_effect(*args, **kwargs):
            mock_filter = Mock()
            call_count[0] += 1
            if call_count[0] == 1:
                # Symbol query uses order_by
                mock_filter.order_by.return_value.all.return_value = [mock_symbol]
            elif call_count[0] == 2:
                # Record query
                mock_filter.all.return_value = [mock_record]
            else:
                # Update query
                mock_filter.update.return_value = 1
                return mock_filter
            return mock_filter

        mock_query.filter.side_effect = filter_side_effect
        mock_session.query.return_value = mock_query

        summary = resolver.resolve_document(doc_id, mock_session)

        assert summary.symbols_resolved == 1
        assert summary.unresolved_count == 1
        assert "●" in summary.unresolved_symbols


class TestUnicodeSymbols:
    """Tests for Unicode symbol handling."""

    def test_japanese_symbols(self, resolver):
        """Test handling of Japanese symbols."""
        lookup = {
            "○": "Yes/Applicable",
            "×": "No/Not applicable",
            "△": "Pending",
            "◎": "Priority/Excellent",
        }
        content = {"Status1": "○", "Status2": "◎"}

        resolved, count, unresolved = resolver._resolve_record(content, lookup)

        assert count == 2
        assert resolved["Status1_resolved"] == "Yes/Applicable"
        assert resolved["Status2_resolved"] == "Priority/Excellent"

    def test_mixed_unicode_and_ascii(self, resolver):
        """Test mixing Unicode symbols and ASCII codes."""
        lookup = {
            "○": "Yes",
            "07": "Code Seven",
        }
        content = {"Check": "○", "Code": "07"}

        resolved, count, unresolved = resolver._resolve_record(content, lookup)

        assert count == 2
        assert resolved["Check_resolved"] == "Yes"
        assert resolved["Code_resolved"] == "Code Seven"

    def test_checkmark_symbols(self, resolver):
        """Test checkmark and cross symbols."""
        lookup = {
            "✓": "Checked",
            "✗": "Unchecked",
        }
        content = {"Task1": "✓", "Task2": "✗"}

        resolved, count, unresolved = resolver._resolve_record(content, lookup)

        assert count == 2
        assert resolved["Task1_resolved"] == "Checked"
        assert resolved["Task2_resolved"] == "Unchecked"
