"""Unit tests for ProviderGroup aggregate."""

from unittest.mock import MagicMock

import pytest

from mcp_hangar.domain.model.mcp_server_group import GroupCircuitOpened, GroupCreated, GroupMemberAdded, McpServerGroup
from mcp_hangar.domain.value_objects import GroupState, LoadBalancerStrategy, ProviderState


def create_mock_provider(mcp_server_id: str, state: ProviderState = ProviderState.READY):
    """Create a mock provider for testing."""
    mock = MagicMock()
    mock.id = mcp_server_id
    mock.mcp_server_id = mcp_server_id
    mock.state = state
    mock.state_snapshot = state
    mock.ensure_ready = MagicMock()
    mock.shutdown = MagicMock()
    mock.tools = []
    mock.get_tool_names = MagicMock(return_value=[])
    return mock


class TestProviderGroupCreation:
    """Tests for ProviderGroup initialization."""

    def test_creates_with_defaults(self):
        """Should create group with default configuration."""
        group = McpServerGroup(group_id="test-group")

        assert group.id == "test-group"
        assert group.state == GroupState.INACTIVE
        assert group.strategy == LoadBalancerStrategy.ROUND_ROBIN
        assert group.healthy_count == 0
        assert group.total_count == 0
        assert group.is_available is False

    def test_creates_with_custom_strategy(self):
        """Should create group with specified strategy."""
        group = McpServerGroup(
            group_id="weighted-group",
            strategy=LoadBalancerStrategy.WEIGHTED_ROUND_ROBIN,
        )

        assert group.strategy == LoadBalancerStrategy.WEIGHTED_ROUND_ROBIN

    def test_creates_with_custom_min_healthy(self):
        """Should create group with specified min_healthy."""
        group = McpServerGroup(
            group_id="high-availability",
            min_healthy=3,
        )

        assert group._min_healthy == 3

    def test_emits_group_created_event(self):
        """Should emit GroupCreated event on creation."""
        group = McpServerGroup(group_id="event-test")
        events = group.collect_events()

        assert len(events) == 1
        assert isinstance(events[0], GroupCreated)
        assert events[0].group_id == "event-test"


class TestMemberManagement:
    """Tests for adding/removing group members."""

    def test_add_member(self):
        """Should add member to group."""
        group = McpServerGroup(group_id="test-group", auto_start=False)
        provider = create_mock_provider("provider-1")

        group.add_member(provider, weight=2, priority=1)

        assert group.total_count == 1
        member = group.get_member("provider-1")
        assert member is not None
        assert member.weight == 2
        assert member.priority == 1

    def test_add_member_emits_event(self):
        """Should emit GroupMemberAdded event."""
        group = McpServerGroup(group_id="test-group", auto_start=False)
        group.collect_events()  # Clear creation event
        provider = create_mock_provider("provider-1")

        group.add_member(provider)
        events = group.collect_events()

        assert len(events) >= 1
        add_event = [e for e in events if isinstance(e, GroupMemberAdded)][0]
        assert add_event.member_id == "provider-1"

    def test_add_duplicate_member_raises(self):
        """Should raise error when adding duplicate member."""
        group = McpServerGroup(group_id="test-group", auto_start=False)
        provider = create_mock_provider("provider-1")
        group.add_member(provider)

        with pytest.raises(ValueError, match="already in group"):
            group.add_member(provider)

    def test_remove_member(self):
        """Should remove member from group."""
        group = McpServerGroup(group_id="test-group", auto_start=False)
        provider = create_mock_provider("provider-1")
        group.add_member(provider)

        result = group.remove_member("provider-1")

        assert result is True
        assert group.total_count == 0
        assert group.get_member("provider-1") is None

    def test_remove_nonexistent_member_returns_false(self):
        """Should return False when removing non-existent member."""
        group = McpServerGroup(group_id="test-group")

        result = group.remove_member("nonexistent")

        assert result is False

    def test_auto_start_adds_to_rotation(self):
        """With auto_start=True, ready members are added to rotation."""
        group = McpServerGroup(group_id="test-group", auto_start=True)
        provider = create_mock_provider("provider-1", state=ProviderState.READY)

        group.add_member(provider)

        # Member should be in rotation if provider is READY
        member = group.get_member("provider-1")
        # Note: actual in_rotation depends on ensure_ready success
        assert member is not None


class TestLoadBalancing:
    """Tests for load balancing functionality."""

    def test_select_member_returns_provider(self):
        """Should return a provider from available members."""
        group = McpServerGroup(group_id="test-group", auto_start=False)
        provider = create_mock_provider("provider-1", ProviderState.READY)
        group.add_member(provider)

        # Manually put in rotation and update state for test
        member = group.get_member("provider-1")
        member.in_rotation = True
        group._update_state()  # Update state to HEALTHY

        selected = group.select_member()

        assert selected is provider

    def test_select_member_returns_none_when_empty(self):
        """Should return None when no members available."""
        group = McpServerGroup(group_id="test-group")

        selected = group.select_member()

        assert selected is None

    def test_select_member_serves_healthy_backup_when_primary_evicted_and_circuit_open(self):
        """A healthy backup must stay selectable after the group circuit
        breaker opens from the primary's eviction (LIVE-P1-01).

        Regression: the group CB used to veto selection *before* checking
        member health, so an evicted primary that opened the group CB took the
        whole group offline even though a healthy backup remained in rotation.
        """
        group = McpServerGroup(
            group_id="test-group",
            auto_start=False,
            unhealthy_threshold=2,
            circuit_failure_threshold=2,
        )
        primary = create_mock_provider("primary")
        backup = create_mock_provider("backup")
        group.add_member(primary)
        group.add_member(backup)
        group.get_member("primary").in_rotation = True
        group.get_member("backup").in_rotation = True

        # Two failures on the primary: evict it AND open the group CB.
        group.report_failure("primary")
        group.report_failure("primary")

        assert group.circuit_open is True
        assert group.get_member("primary").in_rotation is False
        assert group.get_member("backup").in_rotation is True

        # The healthy backup must still be selected despite the open group CB.
        selected = group.select_member()

        assert selected is backup

    def test_select_member_returns_none_when_circuit_open_and_no_member_in_rotation(self):
        """The open circuit breaker still (correctly) blocks selection when no
        member remains in rotation -- the group is genuinely down."""
        group = McpServerGroup(
            group_id="test-group",
            auto_start=False,
            unhealthy_threshold=1,
            circuit_failure_threshold=1,
        )
        provider = create_mock_provider("provider-1")
        group.add_member(provider)
        group.get_member("provider-1").in_rotation = True

        # One failure evicts the only member AND opens the group CB.
        group.report_failure("provider-1")
        assert group.circuit_open is True
        assert group.get_member("provider-1").in_rotation is False

        selected = group.select_member()

        assert selected is None


class TestHealthReporting:
    """Tests for health reporting and rotation management."""

    def test_report_success_resets_failures(self):
        """report_success should reset consecutive failures."""
        group = McpServerGroup(group_id="test-group", auto_start=False)
        provider = create_mock_provider("provider-1", ProviderState.READY)
        group.add_member(provider)
        member = group.get_member("provider-1")
        member.consecutive_failures = 5

        group.report_success("provider-1")

        assert member.consecutive_failures == 0

    def test_report_failure_increments_counter(self):
        """report_failure should increment failure counter."""
        group = McpServerGroup(group_id="test-group", auto_start=False)
        provider = create_mock_provider("provider-1", ProviderState.READY)
        group.add_member(provider)
        member = group.get_member("provider-1")
        member.in_rotation = True

        group.report_failure("provider-1")
        group.report_failure("provider-1")

        assert member.consecutive_failures == 2

    def test_report_failure_removes_from_rotation_at_threshold(self):
        """Should remove from rotation when unhealthy threshold reached."""
        group = McpServerGroup(
            group_id="test-group",
            auto_start=False,
            unhealthy_threshold=2,
        )
        provider = create_mock_provider("provider-1", ProviderState.READY)
        group.add_member(provider)
        member = group.get_member("provider-1")
        member.in_rotation = True

        group.report_failure("provider-1")
        assert member.in_rotation is True  # Still in rotation

        group.report_failure("provider-1")
        assert member.in_rotation is False  # Removed after threshold


class TestCircuitBreaker:
    """Tests for circuit breaker functionality."""

    def test_circuit_opens_at_failure_threshold(self):
        """Circuit should open when failure threshold reached."""
        group = McpServerGroup(
            group_id="test-group",
            auto_start=False,
            circuit_failure_threshold=3,
        )
        provider = create_mock_provider("provider-1", ProviderState.READY)
        group.add_member(provider)
        member = group.get_member("provider-1")
        member.in_rotation = True

        # Trigger failures
        for _ in range(3):
            group.report_failure("provider-1")

        assert group.circuit_open is True

    def test_circuit_emits_event_when_opened(self):
        """Should emit GroupCircuitOpened event."""
        group = McpServerGroup(
            group_id="test-group",
            auto_start=False,
            circuit_failure_threshold=1,
        )
        provider = create_mock_provider("provider-1", ProviderState.READY)
        group.add_member(provider)
        member = group.get_member("provider-1")
        member.in_rotation = True
        group.collect_events()  # Clear previous events

        group.report_failure("provider-1")
        events = group.collect_events()

        circuit_events = [e for e in events if isinstance(e, GroupCircuitOpened)]
        assert len(circuit_events) == 1


class TestStateManagement:
    """Tests for group state transitions."""

    def test_state_is_inactive_with_no_healthy(self):
        """State should be INACTIVE when no healthy members."""
        group = McpServerGroup(group_id="test-group", auto_start=False)
        provider = create_mock_provider("provider-1", ProviderState.COLD)
        group.add_member(provider)

        assert group.state == GroupState.INACTIVE

    def test_state_is_partial_below_min_healthy(self):
        """State should be PARTIAL when healthy < min_healthy."""
        group = McpServerGroup(
            group_id="test-group",
            auto_start=False,
            min_healthy=2,
        )
        provider = create_mock_provider("provider-1", ProviderState.READY)
        group.add_member(provider)
        member = group.get_member("provider-1")
        member.in_rotation = True

        # Force state update
        group._update_state()

        assert group.state == GroupState.PARTIAL

    def test_state_is_healthy_at_min_healthy(self):
        """State should be HEALTHY when healthy >= min_healthy."""
        group = McpServerGroup(
            group_id="test-group",
            auto_start=False,
            min_healthy=1,
        )
        provider = create_mock_provider("provider-1", ProviderState.READY)
        group.add_member(provider)
        member = group.get_member("provider-1")
        member.in_rotation = True

        # Force state update
        group._update_state()

        assert group.state == GroupState.HEALTHY


class TestRebalance:
    """Tests for rebalance functionality."""

    def test_rebalance_adds_ready_members_to_rotation(self):
        """Rebalance should add READY members to rotation."""
        group = McpServerGroup(group_id="test-group", auto_start=False)
        provider = create_mock_provider("provider-1", ProviderState.READY)
        group.add_member(provider)
        member = group.get_member("provider-1")
        member.in_rotation = False  # Start out of rotation

        group.rebalance()

        assert member.in_rotation is True

    def test_rebalance_removes_non_ready_members(self):
        """Rebalance should remove non-READY members from rotation."""
        group = McpServerGroup(group_id="test-group", auto_start=False)
        provider = create_mock_provider("provider-1", ProviderState.COLD)
        group.add_member(provider)
        member = group.get_member("provider-1")
        member.in_rotation = True  # Somehow in rotation but not ready

        group.rebalance()

        assert member.in_rotation is False

    def test_rebalance_resets_circuit_breaker(self):
        """Rebalance should reset circuit breaker."""
        group = McpServerGroup(
            group_id="test-group",
            auto_start=False,
            circuit_failure_threshold=1,
        )
        # Open circuit breaker via report_failure
        provider = create_mock_provider("provider-1", ProviderState.READY)
        group.add_member(provider)
        member = group.get_member("provider-1")
        member.in_rotation = True
        group.report_failure("provider-1")  # This opens circuit

        assert group.circuit_open is True

        group.rebalance()

        assert group.circuit_open is False


class TestSerialization:
    """Tests for to_status_dict serialization."""

    def test_to_status_dict_includes_all_fields(self):
        """to_status_dict should include all relevant fields."""
        group = McpServerGroup(
            group_id="test-group",
            strategy=LoadBalancerStrategy.WEIGHTED_ROUND_ROBIN,
            min_healthy=2,
            description="Test group",
        )
        provider = create_mock_provider("provider-1", ProviderState.READY)
        group.add_member(provider, weight=3, priority=1)

        status = group.to_status_dict()

        assert status["group_id"] == "test-group"
        assert status["strategy"] == "weighted_round_robin"
        assert status["min_healthy"] == 2
        assert status["description"] == "Test group"
        assert len(status["members"]) == 1
        assert status["members"][0]["id"] == "provider-1"
        assert status["members"][0]["weight"] == 3


# --- Lock Hierarchy Tests (CONC-01) ---


class TestProviderGroupLockHierarchy:
    """Verify group lock is NOT held when calling provider operations.

    Lock hierarchy: Provider._lock (level 10) < ProviderGroup._lock (level 11).
    Group lock must be released before calling ensure_ready() or shutdown()
    to avoid level-11-holds-level-10 deadlock.
    """

    def test_add_member_does_not_hold_group_lock_during_ensure_ready(self):
        """add_member() with auto_start must release group lock before ensure_ready()."""
        from mcp_hangar.lock_hierarchy import LockLevel, get_current_thread_locks

        group = McpServerGroup(group_id="lock-test", auto_start=True)
        lock_held_during_ensure_ready = []

        def check_lock_not_held():
            # Check thread-local held locks for PROVIDER_GROUP level
            held = get_current_thread_locks()
            group_held = any(level == LockLevel.PROVIDER_GROUP for level, _ in held)
            lock_held_during_ensure_ready.append(group_held)

        provider = create_mock_provider("p1", ProviderState.COLD)
        provider.ensure_ready = MagicMock(side_effect=check_lock_not_held)
        # After ensure_ready, provider.state should be READY
        provider.state = ProviderState.READY

        group.add_member(provider, weight=1)

        assert len(lock_held_during_ensure_ready) == 1, "ensure_ready() should have been called once"
        assert lock_held_during_ensure_ready[0] is False, (
            "Group lock must NOT be held when calling ensure_ready(). "
            "This violates lock hierarchy: level 11 holding level 10."
        )

    def test_start_all_does_not_hold_group_lock_during_ensure_ready(self):
        """start_all() must release group lock before ensure_ready()."""
        from mcp_hangar.lock_hierarchy import LockLevel, get_current_thread_locks

        group = McpServerGroup(group_id="lock-test", auto_start=False)
        lock_held_during_ensure_ready = []

        def check_lock_not_held():
            held = get_current_thread_locks()
            group_held = any(level == LockLevel.PROVIDER_GROUP for level, _ in held)
            lock_held_during_ensure_ready.append(group_held)

        provider = create_mock_provider("p1", ProviderState.COLD)
        provider.ensure_ready = MagicMock(side_effect=check_lock_not_held)
        provider.state = ProviderState.READY

        group.add_member(provider)
        group.start_all()

        assert len(lock_held_during_ensure_ready) == 1, "ensure_ready() should have been called once"
        assert lock_held_during_ensure_ready[0] is False, (
            "Group lock must NOT be held during ensure_ready() in start_all()"
        )

    def test_stop_all_does_not_hold_group_lock_during_shutdown(self):
        """stop_all() must release group lock before shutdown()."""
        from mcp_hangar.lock_hierarchy import LockLevel, get_current_thread_locks

        group = McpServerGroup(group_id="lock-test", auto_start=False)
        lock_held_during_shutdown = []

        def check_lock_not_held():
            held = get_current_thread_locks()
            group_held = any(level == LockLevel.PROVIDER_GROUP for level, _ in held)
            lock_held_during_shutdown.append(group_held)

        provider = create_mock_provider("p1", ProviderState.READY)
        provider.shutdown = MagicMock(side_effect=check_lock_not_held)

        group.add_member(provider)
        group.stop_all()

        assert len(lock_held_during_shutdown) >= 1, "shutdown() should have been called"
        assert lock_held_during_shutdown[0] is False, "Group lock must NOT be held during shutdown() in stop_all()"

    def test_member_removed_between_phases_handled_gracefully(self):
        """Member removed between lock release and re-acquire is handled."""
        group = McpServerGroup(group_id="lock-test", auto_start=True)

        def ensure_ready_and_remove():
            # Simulate: another thread removes the member during ensure_ready
            group.remove_member("p1")

        provider = create_mock_provider("p1", ProviderState.COLD)
        provider.ensure_ready = MagicMock(side_effect=ensure_ready_and_remove)
        provider.state = ProviderState.READY

        # Should not raise KeyError
        group.add_member(provider, weight=1)

        # Member was removed during Phase 2, so it should not be in the group
        assert group.get_member("p1") is None

    def test_concurrent_add_and_start_all_no_deadlock(self):
        """Concurrent add_member() + start_all() must not deadlock."""
        import threading

        group = McpServerGroup(group_id="lock-test", auto_start=True)
        # The barrier holds each thread inside its lock long enough to overlap.
        # When only one side reaches it the wait runs to this timeout, which the
        # assertions below already tolerate (`BrokenBarrierError` is filtered
        # out) -- so it is the floor on how long the test takes, not a deadline
        # anything depends on. Five seconds bought nothing over one.
        barrier = threading.Barrier(2, timeout=1)
        errors: list[Exception] = []

        def add_mcp_server():
            try:
                provider = create_mock_provider("p-add", ProviderState.COLD)
                provider.state = ProviderState.READY

                def slow_start():
                    barrier.wait()

                provider.ensure_ready = MagicMock(side_effect=slow_start)
                group.add_member(provider)
            except Exception as e:  # noqa: BLE001 -- thread error collector
                errors.append(e)

        def start_providers():
            try:
                from mcp_hangar.domain.model.mcp_server_group import GroupMember

                provider2 = create_mock_provider("p-start", ProviderState.COLD)
                provider2.state = ProviderState.READY

                def slow_start():
                    barrier.wait()

                provider2.ensure_ready = MagicMock(side_effect=slow_start)
                with group._lock:
                    member = GroupMember(
                        mcp_server=provider2,
                        weight=1,
                        priority=1,
                    )
                    group._members["p-start"] = member

                group.start_all()
            except Exception as e:  # noqa: BLE001 -- thread error collector
                errors.append(e)

        t1 = threading.Thread(target=add_mcp_server)
        t2 = threading.Thread(target=start_providers)
        t1.start()
        t2.start()
        t1.join(timeout=10)
        t2.join(timeout=10)

        assert not t1.is_alive(), "add_member thread deadlocked"
        assert not t2.is_alive(), "start_all thread deadlocked"
        deadlock_errors = [e for e in errors if not isinstance(e, threading.BrokenBarrierError)]
        assert not deadlock_errors, f"Unexpected errors: {deadlock_errors}"
