"""Tests for Omni (STS) mode tool call routing."""

import json
import asyncio
import pytest
from unittest.mock import MagicMock, AsyncMock, patch
from conftest import MockSession

from qwen_pipeline_server import PipelineSession


@pytest.fixture
def omni_session(sample_extensions):
    """Create a mock session configured for Omni mode testing."""
    session = MockSession(
        extensions_directory=sample_extensions,
        omni_mode=True,
        active_agent='a',
        transfer_agents_enabled=False,
    )
    session.agents = {
        'a': {
            'name': 'Receptionist',
            'history': [],
            'instructions_raw_map': {'en': 'Test'},
            'instructions_decorations': '',
        }
    }
    # Mock async methods
    session._omni_wait_response_done = AsyncMock()
    session._emit = AsyncMock()
    session._omni_muted_audio_buffer = []
    session._omni_pre_emit_buffer = []
    session._omni_pre_emit_gen = 0
    session._omni_pre_emit_timer = None
    session._omni_pre_emit_buffering = False
    session._omni_audio_buffer_delay = 0.3
    session._omni_conv = MagicMock()
    session._omni_conv.send_raw = MagicMock()
    session._omni_suppress_flush = False

    # Mock loop.run_in_executor to call the function synchronously
    mock_loop = MagicMock()
    mock_loop.run_in_executor = AsyncMock(
        side_effect=lambda ex, fn, *a: fn(*a))
    session.loop = mock_loop

    # Bind _handle_search_extension from PipelineSession
    session._handle_search_extension = lambda args: PipelineSession._handle_search_extension(session, args)

    # Bind _execute_tool_endpoint (used by dial_extension path before handle_dial)
    session._execute_tool_endpoint = MagicMock(return_value=None)

    return session


class TestOmniSearchExtensionRouting:
    """Tests for search_extension routing in Omni mode."""

    @pytest.mark.asyncio
    async def test_search_emits_tool_call_event(self, omni_session):
        """search_extension should emit tool_call event to browser."""
        args = {'name': 'Zack'}
        arguments = json.dumps(args)

        await PipelineSession._handle_omni_tool_call_inner(
            omni_session, 'search_extension', args, arguments, 'call_123')

        # Check _emit was called with tool_call type
        emit_calls = omni_session._emit.call_args_list
        tool_call_emitted = any(
            c[0][0].get('type') == 'tool_call' and
            c[0][0].get('name') == 'search_extension'
            for c in emit_calls
        )
        assert tool_call_emitted

    @pytest.mark.asyncio
    async def test_search_submits_result_to_omni(self, omni_session):
        """search_extension result should be sent back to Omni as function_call_output."""
        args = {'name': 'Zack'}
        arguments = json.dumps(args)

        await PipelineSession._handle_omni_tool_call_inner(
            omni_session, 'search_extension', args, arguments, 'call_123')

        # Check send_raw was called with conversation.item.create
        send_calls = omni_session._omni_conv.send_raw.call_args_list
        found_result = False
        for call in send_calls:
            msg = json.loads(call[0][0])
            if msg.get('type') == 'conversation.item.create':
                item = msg.get('item', {})
                if item.get('type') == 'function_call_output':
                    assert item['call_id'] == 'call_123'
                    assert 'Zack Jong' in item['output']
                    found_result = True
        assert found_result

    @pytest.mark.asyncio
    async def test_search_triggers_response_create(self, omni_session):
        """After search_extension, a response.create should be triggered."""
        args = {'name': 'Zack'}
        arguments = json.dumps(args)

        await PipelineSession._handle_omni_tool_call_inner(
            omni_session, 'search_extension', args, arguments, 'call_123')

        # Check response.create was sent
        send_calls = omni_session._omni_conv.send_raw.call_args_list
        response_created = any(
            json.loads(c[0][0]).get('type') == 'response.create'
            for c in send_calls
        )
        assert response_created

    @pytest.mark.asyncio
    async def test_search_not_handled_by_dial(self, omni_session):
        """search_extension should NOT be routed to handle_dial."""
        args = {'name': 'Zack'}
        arguments = json.dumps(args)

        with patch('qwen_pipeline_server.handle_dial', new_callable=AsyncMock) as mock_dial:
            await PipelineSession._handle_omni_tool_call_inner(
                omni_session, 'search_extension', args, arguments, 'call_123')
            mock_dial.assert_not_called()


class TestOmniDialExtensionRouting:
    """Tests for dial_extension routing in Omni mode."""

    @pytest.mark.asyncio
    async def test_dial_calls_handle_dial(self, omni_session):
        """dial_extension should be routed to handle_dial handler."""
        args = {'extension': '8703', 'name': 'Zack Jong'}
        arguments = json.dumps(args)

        with patch('qwen_pipeline_server.handle_dial', new_callable=AsyncMock) as mock_dial:
            await PipelineSession._handle_omni_tool_call_inner(
                omni_session, 'dial_extension', args, arguments, 'call_456')
            mock_dial.assert_called_once()


class TestOmniEndCallRouting:
    """Tests for end_call routing in Omni mode."""

    @pytest.mark.asyncio
    async def test_end_call_routed(self, omni_session):
        """end_call should be routed to handle_end_call."""
        args = {'reason': 'Caller said goodbye'}
        arguments = json.dumps(args)

        with patch('qwen_pipeline_server.handle_end_call', new_callable=AsyncMock) as mock_end:
            await PipelineSession._handle_omni_tool_call_inner(
                omni_session, 'end_call', args, arguments, 'call_789')
            mock_end.assert_called_once()
