"""Tests for the interrupt system.

Run with: python -m pytest tests/test_interrupt.py -v
"""

import threading
import time
import pytest


# ---------------------------------------------------------------------------
# Unit tests: shared interrupt module
# ---------------------------------------------------------------------------

class TestInterruptModule:
    """Tests for tools/interrupt.py"""

    def test_set_and_check(self):
        from tools.interrupt import set_interrupt, is_interrupted
        set_interrupt(False)
        assert not is_interrupted()

        set_interrupt(True)
        assert is_interrupted()

        set_interrupt(False)
        assert not is_interrupted()

    def test_is_thread_interrupted_checks_target_tid_not_caller(self):
        from tools.interrupt import (
            set_interrupt, is_interrupted, is_thread_interrupted, _interrupted_threads, _lock,
        )
        with _lock:
            _interrupted_threads.clear()
        other_tid = threading.get_ident() + 1
        set_interrupt(True, thread_id=other_tid)
        assert not is_interrupted()
        assert is_thread_interrupted(other_tid)
        assert is_thread_interrupted(None) is False
        set_interrupt(False, thread_id=other_tid)
        assert not is_thread_interrupted(other_tid)


    def test_clear_current_thread_interrupt_leaves_other_threads(self):
        """clear_current_thread_interrupt only touches the calling thread."""
        from tools.interrupt import (
            set_interrupt, is_interrupted, clear_current_thread_interrupt,
            _interrupted_threads, _lock,
        )
        with _lock:
            _interrupted_threads.clear()
        other_tid = threading.get_ident() + 1  # an ident that isn't us
        set_interrupt(True, thread_id=other_tid)
        set_interrupt(True)  # current thread
        assert is_interrupted()

        clear_current_thread_interrupt()

        assert not is_interrupted()  # ours cleared
        with _lock:
            assert other_tid in _interrupted_threads  # other thread untouched
            _interrupted_threads.discard(other_tid)

    def test_run_if_not_interrupted_skips_callback_when_already_interrupted(self):
        from tools.interrupt import run_if_not_interrupted, set_interrupt

        callbacks = []
        set_interrupt(True)
        try:
            assert run_if_not_interrupted(lambda: callbacks.append("claimed")) is False
        finally:
            set_interrupt(False)

        assert callbacks == []

    @pytest.mark.parametrize("callback_should_fail", [False, True])
    def test_run_if_not_interrupted_orders_callback_before_concurrent_interrupt(
        self, callback_should_fail, monkeypatch
    ):
        import tools.interrupt as interrupt

        class CallbackFailure(Exception):
            pass

        original_lock = interrupt._lock
        attempting_interrupt_lock = threading.Event()
        interrupt_published = threading.Event()
        publisher_lock_contention = []
        callback_observations = []
        setters = []
        setter_tids = []

        class ObservedLock:
            def __enter__(self):
                if (
                    threading.current_thread() in setters
                    and not attempting_interrupt_lock.is_set()
                ):
                    acquired = original_lock.acquire(blocking=False)
                    publisher_lock_contention.append(not acquired)
                    attempting_interrupt_lock.set()
                    if acquired:
                        return self
                original_lock.acquire()
                return self

            def __exit__(self, exc_type, exc_value, traceback):
                original_lock.release()

        interrupt.set_interrupt(False)
        with original_lock:
            baseline = (
                set(interrupt._interrupted_threads),
                dict(interrupt._interrupt_reasons),
            )
        monkeypatch.setattr(interrupt, "_lock", ObservedLock())

        def publish_interrupt():
            setter_tids.append(threading.get_ident())
            try:
                interrupt.set_interrupt(True)
                interrupt_published.set()
            finally:
                interrupt.set_interrupt(False)

        def callback():
            setter = threading.Thread(target=publish_interrupt)
            setters.append(setter)
            setter.start()
            assert attempting_interrupt_lock.wait(5)
            assert publisher_lock_contention == [True]
            callback_observations.append(interrupt_published.is_set())
            if callback_should_fail:
                raise CallbackFailure

        try:
            if callback_should_fail:
                with pytest.raises(CallbackFailure):
                    interrupt.run_if_not_interrupted(callback)
            else:
                assert interrupt.run_if_not_interrupted(callback) is True
            assert interrupt_published.wait(5)
        finally:
            for setter in setters:
                if setter.ident is not None:
                    setter.join(timeout=5)
            interrupt.set_interrupt(False)

        assert setters
        assert all(not setter.is_alive() for setter in setters)
        assert setter_tids
        assert callback_observations == [False]
        assert interrupt_published.is_set()
        with original_lock:
            final_state = (
                set(interrupt._interrupted_threads),
                dict(interrupt._interrupt_reasons),
            )
        assert final_state == baseline
        assert all(setter_tid not in final_state[0] for setter_tid in setter_tids)
        assert all(setter_tid not in final_state[1] for setter_tid in setter_tids)


# ---------------------------------------------------------------------------
# Unit tests: pre-tool interrupt check
# ---------------------------------------------------------------------------

class TestPreToolCheck:
    """Verify that _execute_tool_calls skips all tools when interrupted."""

    def test_all_tools_skipped_when_interrupted(self):
        """Mock an interrupted agent and verify no tools execute."""
        from unittest.mock import MagicMock

        # Build a fake assistant_message with 3 tool calls
        tc1 = MagicMock()
        tc1.id = "tc_1"
        tc1.function.name = "terminal"
        tc1.function.arguments = '{"command": "rm -rf /"}'

        tc2 = MagicMock()
        tc2.id = "tc_2"
        tc2.function.name = "terminal"
        tc2.function.arguments = '{"command": "echo hello"}'

        tc3 = MagicMock()
        tc3.id = "tc_3"
        tc3.function.name = "web_search"
        tc3.function.arguments = '{"query": "test"}'

        assistant_msg = MagicMock()
        assistant_msg.tool_calls = [tc1, tc2, tc3]

        messages = []

        # Create a minimal mock agent with _interrupt_requested = True
        agent = MagicMock()
        agent._interrupt_requested = True
        agent.log_prefix = ""
        agent._persist_session = MagicMock()
        # PR #72425: execute_tool_calls_* read _incremental_persistence_failed
        # via getattr at loop top. A bare MagicMock auto-creates a truthy value
        # for any attribute access, which would short-circuit the interrupt
        # skip path before any cancelled-tool messages are appended.
        agent._incremental_persistence_failed = False

        # Import and call the method
        import types
        from run_agent import AIAgent
        # Bind the real methods to our mock so dispatch works correctly
        agent._execute_tool_calls_sequential = types.MethodType(AIAgent._execute_tool_calls_sequential, agent)
        agent._execute_tool_calls_concurrent = types.MethodType(AIAgent._execute_tool_calls_concurrent, agent)
        AIAgent._execute_tool_calls(agent, assistant_msg, messages, "default")

        # All 3 should be skipped
        assert len(messages) == 3
        for msg in messages:
            assert msg["role"] == "tool"
            assert "cancelled" in msg["content"].lower() or "interrupted" in msg["content"].lower()

        # No actual tool handlers should have been called
        # (handle_function_call should NOT have been invoked)


# ---------------------------------------------------------------------------
# Unit tests: message combining
# ---------------------------------------------------------------------------



# ---------------------------------------------------------------------------
# Integration tests (require local terminal)
# ---------------------------------------------------------------------------

class TestSIGKILLEscalation:
    """Test that SIGTERM-resistant processes get SIGKILL'd."""

    @pytest.mark.skipif(
        not __import__("shutil").which("bash"),
        reason="Requires bash"
    )
    @pytest.mark.platforms("posix")
    def test_sigterm_trap_killed_within_2s(self, tmp_path):
        """A process that traps SIGTERM should be SIGKILL'd after 1s grace."""
        from tools.interrupt import set_interrupt
        from tools.environments.local import LocalEnvironment

        set_interrupt(False)
        env = LocalEnvironment(cwd=str(tmp_path), timeout=30)

        # Start execution in a thread, interrupt after 0.5s
        result_holder = {"value": None}

        def _run():
            result_holder["value"] = env.execute(
                "trap '' TERM; sleep 60",
                timeout=30,
            )

        t = threading.Thread(target=_run)
        t.start()

        time.sleep(0.5)
        set_interrupt(True, thread_id=t.ident)

        t.join(timeout=5)
        set_interrupt(False, thread_id=t.ident)

        assert result_holder["value"] is not None
        assert result_holder["value"]["returncode"] == 130
        assert "interrupted" in result_holder["value"]["output"].lower()


# ---------------------------------------------------------------------------
# Regression: _run_tool cleanup on BaseException (issue #35309)
# ---------------------------------------------------------------------------

class TestRunToolCleanupOnBaseException:
    """Verify that _run_tool cleans up _interrupted_threads even when
    _invoke_tool raises a BaseException (e.g. CancelledError).

    Regression test for #35309: without the finally block, a BaseException
    bypasses ``except Exception``, leaking the worker tid into
    _interrupted_threads.  ThreadPoolExecutor recycles tids, so the next
    tool scheduled on the same thread is instantly "interrupted".
    """

    def test_cleanup_on_base_exception(self):
        from unittest.mock import MagicMock
        import types
        from tools.interrupt import set_interrupt, _interrupted_threads, _lock

        # Clear global state
        with _lock:
            _interrupted_threads.clear()

        # Build a minimal mock agent with the attributes _run_tool needs
        agent = MagicMock()
        agent._interrupt_requested = False
        agent._tool_worker_threads = set()
        agent._tool_worker_threads_lock = threading.Lock()

        # _set_interrupt delegates to the real module
        def _mock_set_interrupt(active, tid=None):
            set_interrupt(active, tid)
        agent._set_interrupt = _mock_set_interrupt

        # _invoke_tool raises BaseException (simulating CancelledError)
        agent._invoke_tool = MagicMock(side_effect=BaseException("simulated CancelledError"))

        # Bind the real concurrent method so we get _run_tool
        from run_agent import AIAgent
        agent._execute_tool_calls_concurrent = types.MethodType(
            AIAgent._execute_tool_calls_concurrent, agent
        )

        # Build a single tool call
        tc = MagicMock()
        tc.id = "tc_base_exc"
        tc.function.name = "dummy_tool"
        tc.function.arguments = "{}"

        assistant_msg = MagicMock()
        assistant_msg.tool_calls = [tc]

        # _execute_tool_calls_concurrent will submit _run_tool to a
        # ThreadPoolExecutor.  The BaseException propagates out of the
        # worker, but the finally block should still clean up.
        try:
            agent._execute_tool_calls_concurrent(assistant_msg, [], "default")
        except Exception:
            pass  # ThreadPoolExecutor may re-raise

        # After the worker finishes (even with BaseException), the worker
        # tid should have been removed from _interrupted_threads and
        # _tool_worker_threads.
        assert len(agent._tool_worker_threads) == 0, (
            f"_tool_worker_threads not cleaned up: {agent._tool_worker_threads}"
        )

        # Verify no stale tid is left in the global interrupt set.  The
        # worker thread is recycled by ThreadPoolExecutor, so a leaked tid
        # would poison the next task on that thread.  We cleared the set at
        # the start and never set any interrupt ourselves, so a leak from
        # _run_tool is the only way an entry could land here.
        with _lock:
            leaked = set(_interrupted_threads)
        assert leaked == set(), f"leaked tids in _interrupted_threads: {leaked}"
