# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

from contextlib import AsyncExitStack, asynccontextmanager
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock

import pytest
import pytest_asyncio
from openai.types.responses import (
    ResponseFunctionToolCall,
    ResponseOutputItemDoneEvent,
    ResponseOutputMessage,
    ResponseOutputText,
    ResponseReasoningItem,
    ResponseReasoningTextDeltaEvent,
    ResponseReasoningTextDoneEvent,
    ResponseTextConfig,
    ResponseTextDeltaEvent,
)
from openai.types.responses import (
    ResponseUsage as OpenAIResponseUsage,
)
from openai.types.responses.response_format_text_json_schema_config import (
    ResponseFormatTextJSONSchemaConfig,
)
from openai.types.responses.tool import (
    CodeInterpreterContainerCodeInterpreterToolAuto,
    LocalShell,
    Mcp,
    Tool,
)
from openai_harmony import Message as OpenAIHarmonyMessage
from openai_harmony import Role

import vllm.envs as envs
from tests.entrypoints.openai.chat_completion.test_serving_chat import MockModelConfig
from vllm.entrypoints.generate.base.protocol import (
    DeltaFunctionCall,
    DeltaMessage,
    DeltaToolCall,
    RequestResponseMetadata,
)
from vllm.entrypoints.mcp.tool_server import ToolServer
from vllm.entrypoints.openai.responses.context import (
    ConversationContext,
    HarmonyContext,
    ParsableContext,
    SimpleContext,
)
from vllm.entrypoints.openai.responses.protocol import (
    ResponseCompletedEvent,
    ResponseCreatedEvent,
    ResponseRawMessageAndToken,
    ResponsesRequest,
    ResponsesResponse,
    serialize_message,
)
from vllm.entrypoints.openai.responses.serving import (
    OpenAIServingResponses,
    extract_tool_types,
)
from vllm.entrypoints.openai.responses.streaming_events import (
    StreamingState,
)
from vllm.entrypoints.serve.engine.protocol import ErrorResponse
from vllm.inputs import tokens_input
from vllm.logprobs import Logprob as SampleLogprob
from vllm.outputs import CompletionOutput, RequestOutput
from vllm.parser.harmony import ChunkResult, HarmonyParser, Segment
from vllm.renderers import TokenizeParams
from vllm.renderers.online_renderer import (
    OnlineRenderer,
    _extract_allowed_tools_from_mcp_requests,
)
from vllm.sampling_params import SamplingParams
from vllm.v1.metrics.stats import RequestStateStats

pytestmark = pytest.mark.skip_global_cleanup


def _new_online_renderer() -> OnlineRenderer:
    renderer = OnlineRenderer.__new__(OnlineRenderer)
    renderer.exclude_tools_when_tool_choice_none = False
    renderer.parser = None
    renderer.trust_request_chat_template = True
    return renderer


class MockConversationContext(ConversationContext):
    """Mock conversation context for testing."""

    def __init__(self):
        self.init_tool_sessions_called = False
        self.init_tool_sessions_args = None
        self.init_tool_sessions_kwargs = None

    def append_output(self, output) -> None:
        pass

    def append_tool_output(self, output) -> None:
        pass

    async def call_tool(self):
        return []

    def need_builtin_tool_call(self) -> bool:
        return False

    def render_for_completion(self):
        return []

    async def init_tool_sessions(self, tool_server, exit_stack, request_id, mcp_tools):
        self.init_tool_sessions_called = True
        self.init_tool_sessions_args = (tool_server, exit_stack, request_id, mcp_tools)

    async def cleanup_session(self) -> None:
        pass


def test_serialize_message_pydantic_model_returns_dict() -> None:
    msg = ResponseRawMessageAndToken(message="hello", tokens=[1, 2, 3])

    serialized = serialize_message(msg)

    assert isinstance(serialized, dict)
    assert serialized["type"] == "raw_message_tokens"
    assert serialized["message"] == "hello"


@pytest.fixture
def mock_context():
    """Create a mock conversation context."""
    return MockConversationContext()


@pytest.fixture
def mock_exit_stack():
    """Create a mock async exit stack."""
    return MagicMock(spec=AsyncExitStack)


def test_extract_tool_types(monkeypatch: pytest.MonkeyPatch) -> None:
    tools: list[Tool] = []
    assert extract_tool_types(tools) == set()

    tools.append(LocalShell(type="local_shell"))
    assert extract_tool_types(tools) == {"local_shell"}

    tools.append(CodeInterpreterContainerCodeInterpreterToolAuto(type="auto"))
    assert extract_tool_types(tools) == {"local_shell", "auto"}

    tools.extend(
        [
            Mcp(type="mcp", server_label="random", server_url=""),
            Mcp(type="mcp", server_label="container", server_url=""),
            Mcp(type="mcp", server_label="code_interpreter", server_url=""),
            Mcp(type="mcp", server_label="web_search_preview", server_url=""),
        ]
    )
    # When envs.VLLM_GPT_OSS_SYSTEM_TOOL_MCP_LABELS is not set,
    # mcp tool types are all ignored.
    assert extract_tool_types(tools) == {"local_shell", "auto"}

    # container is allowed, it would be extracted
    monkeypatch.setenv("VLLM_GPT_OSS_SYSTEM_TOOL_MCP_LABELS", "container")
    assert extract_tool_types(tools) == {"local_shell", "auto", "container"}

    # code_interpreter and web_search_preview are allowed,
    # they would be extracted
    monkeypatch.setenv(
        "VLLM_GPT_OSS_SYSTEM_TOOL_MCP_LABELS", "code_interpreter,web_search_preview"
    )
    assert extract_tool_types(tools) == {
        "local_shell",
        "auto",
        "code_interpreter",
        "web_search_preview",
    }


@pytest.mark.skip_global_cleanup
def test_response_created_event_uses_public_json_schema_alias() -> None:
    schema = {
        "type": "object",
        "properties": {
            "event_name": {"type": "string"},
            "date": {"type": "string"},
            "participants": {"type": "array", "items": {"type": "string"}},
        },
        "required": ["event_name", "date", "participants"],
        "additionalProperties": False,
    }
    text = ResponseTextConfig()
    text.format = ResponseFormatTextJSONSchemaConfig(
        type="json_schema",
        name="calendar_event",
        schema=schema,
        description="A calendar event.",
        strict=True,
    )
    request = ResponsesRequest(
        model="test-model",
        input="Alice and Bob are going to a science fair on Friday.",
        text=text,
    )
    sampling_params = request.to_sampling_params(default_max_tokens=64)
    initial_response = ResponsesResponse.from_request(
        request=request,
        sampling_params=sampling_params,
        model_name="test-model",
        created_time=0,
        output=[],
        status="in_progress",
        usage=None,
    ).model_dump(mode="json", by_alias=True)

    fmt = initial_response["text"]["format"]
    assert fmt["schema"] == schema
    assert "schema_" not in fmt

    event = ResponseCreatedEvent(
        type="response.created",
        sequence_number=0,
        response=initial_response,
    )
    assert event.response.text is not None
    assert event.response.text.format is not None
    assert event.response.text.format.model_dump(by_alias=True)["schema"] == schema


@pytest.mark.asyncio
async def test_online_renderer_renders_non_harmony_responses_with_explicit_history():
    renderer = _new_online_renderer()
    renderer.use_harmony = False
    renderer.chat_template = "server-template"
    renderer.chat_template_content_format = "string"
    renderer.default_chat_template_kwargs = {"server_default": "kept"}
    renderer.parser = None
    renderer.preprocess_chat = AsyncMock(
        return_value=(
            [],
            [tokens_input([11, 12, 13], cache_salt="request-salt")],
        )
    )
    previous_messages = [
        {"role": "system", "content": "old instructions"},
        {"role": "user", "content": "first"},
    ]
    previous_outputs = [
        ResponseOutputMessage(
            id="msg_1",
            content=[
                ResponseOutputText(
                    annotations=[],
                    text="prior answer",
                    type="output_text",
                )
            ],
            role="assistant",
            status="completed",
            type="message",
        )
    ]
    request = ResponsesRequest(
        input="next",
        instructions="new instructions",
        tools=[],
        cache_salt="request-salt",
        chat_template_kwargs={"request_default": "kept"},
    )

    result = await renderer.render_responses(
        request,
        previous_messages=previous_messages,
        previous_response_outputs=previous_outputs,
    )

    assert not isinstance(result, ErrorResponse)
    assert result.messages == [
        {"role": "system", "content": "new instructions"},
        {"role": "user", "content": "first"},
        {"role": "assistant", "content": "prior answer"},
        {"role": "user", "content": "next"},
    ]
    assert result.engine_input["prompt_token_ids"] == [11, 12, 13]
    assert result.engine_input["cache_salt"] == "request-salt"
    assert previous_messages[0]["content"] == "old instructions"
    assert previous_outputs[0].content[0].text == "prior answer"

    preprocess_kwargs = renderer.preprocess_chat.call_args.kwargs
    assert preprocess_kwargs["default_template"] == "server-template"
    assert preprocess_kwargs["default_template_content_format"] == "string"
    assert preprocess_kwargs["default_template_kwargs"]["server_default"] == "kept"
    assert preprocess_kwargs["default_template_kwargs"]["request_default"] == "kept"


_REUSED_IDS = [10, 20, 30]
_IMAGE_PART = {
    "type": "input_image",
    "image_url": "https://example.com/a.png",
    "detail": "auto",
}


async def _render_responses_with_reuse(request, previous_messages=None):
    online_renderer = OnlineRenderer(
        model_config=MockModelConfig(),
        renderer=MagicMock(),
        request_logger=None,
        chat_template=None,
        chat_template_content_format="auto",
    )
    render_chat = AsyncMock(return_value=([[]], [tokens_input([1, 2])]))
    online_renderer.renderer.render_chat_async = render_chat

    result = await online_renderer.render_responses(
        request, previous_messages=previous_messages
    )

    assert not isinstance(result, ErrorResponse)
    return render_chat, result.engine_input


@pytest.mark.asyncio
@pytest.mark.parametrize("ids", [_REUSED_IDS, [1.5]])
@pytest.mark.parametrize(
    "request_input",
    [
        [
            {
                "role": "user",
                "content": [{"type": "input_text", "text": "look"}, _IMAGE_PART],
            }
        ],
        [
            {
                "type": "function_call",
                "call_id": "call_1",
                "name": "shot",
                "arguments": "{}",
            },
            {
                "type": "function_call_output",
                "call_id": "call_1",
                "output": [_IMAGE_PART],
            },
        ],
    ],
    ids=["input_image", "tool_output_image"],
)
async def test_responses_kv_transfer_prompt_token_ids_ignored_with_non_text_input(
    request_input, ids
):
    request = ResponsesRequest(
        input=request_input,
        kv_transfer_params={"do_remote_prefill": True, "prompt_token_ids": ids},
    )

    render_chat, _ = await _render_responses_with_reuse(request)

    render_chat.assert_awaited_once()
    assert request.kv_transfer_params == {"do_remote_prefill": True}


@pytest.mark.asyncio
async def test_responses_kv_transfer_prompt_token_ids_ignored_with_media_in_history():
    request = ResponsesRequest(
        input="next", kv_transfer_params={"prompt_token_ids": _REUSED_IDS}
    )

    render_chat, _ = await _render_responses_with_reuse(
        request, previous_messages=[{"role": "user", "content": [_IMAGE_PART]}]
    )

    render_chat.assert_awaited_once()


@pytest.mark.asyncio
async def test_responses_kv_transfer_prompt_token_ids_kept_with_text_input():
    request = ResponsesRequest(
        input=[
            {"role": "user", "content": [{"type": "input_text", "text": "hi"}]},
            {"role": "assistant", "content": "ok"},
        ],
        kv_transfer_params={"prompt_token_ids": _REUSED_IDS},
    )

    _, engine_input = await _render_responses_with_reuse(request)

    assert engine_input["prompt_token_ids"] == _REUSED_IDS


@pytest.mark.asyncio
async def test_online_renderer_rejects_untrusted_responses_chat_template():
    renderer = _new_online_renderer()
    renderer.use_harmony = False
    renderer.chat_template = "server-template"
    renderer.chat_template_content_format = "string"
    renderer.default_chat_template_kwargs = {}
    renderer.exclude_tools_when_tool_choice_none = False
    renderer.trust_request_chat_template = False
    renderer.parser = None
    renderer.preprocess_chat = AsyncMock(
        return_value=([], [tokens_input([1])]),
    )
    request = ResponsesRequest(
        input="hello",
        chat_template_kwargs={"chat_template": "{{ messages }}"},
    )

    result = await renderer.render_responses(request)

    assert isinstance(result, ErrorResponse)
    assert result.error.code == 400
    assert "untrusted chat template" in result.error.message.lower()
    renderer.preprocess_chat.assert_not_awaited()


@pytest.mark.asyncio
async def test_online_renderer_rejects_harmony_history_for_non_harmony_model():
    renderer = _new_online_renderer()
    renderer.use_harmony = False
    request = ResponsesRequest(input="next")

    result = await renderer.render_responses(
        request,
        previous_messages=[
            OpenAIHarmonyMessage.from_role_and_content(Role.USER, "first")
        ],
    )

    assert isinstance(result, ErrorResponse)
    assert result.error.param == "previous_response_id"


@pytest.mark.asyncio
async def test_online_renderer_rejects_chat_history_for_harmony_model():
    renderer = _new_online_renderer()
    renderer.use_harmony = True
    request = ResponsesRequest(input="next")

    result = await renderer.render_responses(
        request,
        previous_messages=[{"role": "user", "content": "first"}],
    )

    assert isinstance(result, ErrorResponse)
    assert result.error.param == "previous_response_id"


@pytest.mark.asyncio
async def test_online_renderer_adjusts_harmony_tool_choice(
    monkeypatch: pytest.MonkeyPatch,
):
    class StubParser:
        def __init__(self, tokenizer, tools, *, model_config):
            pass

        def adjust_request(self, request):
            request.cache_salt = "adjusted"
            return request

    renderer = _new_online_renderer()
    renderer.use_harmony = True
    renderer.parser = StubParser
    renderer.model_config = SimpleNamespace(max_model_len=100)
    tokenizer = SimpleNamespace(truncation_side="left")
    renderer.renderer = SimpleNamespace(
        tokenizer=tokenizer,
        get_tokenizer=lambda: tokenizer,
    )
    monkeypatch.setattr(
        "vllm.renderers.online_renderer.render_for_completion",
        lambda messages: [1],
    )
    request = ResponsesRequest(
        input="next",
        tools=[
            {
                "type": "function",
                "name": "lookup",
                "parameters": {"type": "object", "properties": {}},
            }
        ],
        tool_choice={"type": "function", "name": "lookup"},
    )

    result = await renderer.render_responses(request)

    assert not isinstance(result, ErrorResponse)
    assert result.engine_input["cache_salt"] == "adjusted"


@pytest.mark.asyncio
async def test_online_renderer_applies_responses_token_budget_to_harmony_prompt(
    monkeypatch: pytest.MonkeyPatch,
):
    renderer = _new_online_renderer()
    renderer.use_harmony = True
    renderer.model_config = SimpleNamespace(max_model_len=5)
    renderer.renderer = SimpleNamespace(
        tokenizer=SimpleNamespace(truncation_side="left")
    )
    monkeypatch.setattr(
        "vllm.renderers.online_renderer.render_for_completion",
        lambda messages: [1, 2, 3, 4, 5, 6],
    )
    request = ResponsesRequest(
        input="hello",
        max_output_tokens=2,
        truncation="auto",
    )

    result = await renderer.render_responses(request)

    assert not isinstance(result, ErrorResponse)
    assert result.engine_input["prompt_token_ids"] == [4, 5, 6]


@pytest.mark.asyncio
async def test_online_renderer_renders_harmony_continuation_with_call_linkage(
    monkeypatch: pytest.MonkeyPatch,
):
    renderer = _new_online_renderer()
    renderer.use_harmony = True
    renderer.model_config = SimpleNamespace(max_model_len=100)
    renderer.renderer = SimpleNamespace(
        tokenizer=SimpleNamespace(truncation_side="left")
    )
    monkeypatch.setattr(
        "vllm.renderers.online_renderer.render_for_completion",
        lambda messages: [1],
    )
    previous_messages = [
        OpenAIHarmonyMessage.from_role_and_content(Role.USER, "first"),
    ]
    previous_outputs = [
        ResponseFunctionToolCall(
            arguments='{"city":"Paris"}',
            call_id="call_1",
            name="lookup_weather",
            type="function_call",
        )
    ]
    request = ResponsesRequest(
        input=[
            {
                "type": "function_call_output",
                "call_id": "call_1",
                "output": "sunny",
            }
        ],
        instructions="replacement instructions are ignored on continuations",
        reasoning={"effort": "high"},
        tools=[
            {
                "type": "function",
                "name": "lookup_weather",
                "parameters": {"type": "object", "properties": {}},
            }
        ],
        cache_salt="request-salt",
    )

    result = await renderer.render_responses(
        request,
        previous_messages=previous_messages,
        previous_response_outputs=previous_outputs,
    )

    assert not isinstance(result, ErrorResponse)
    assert result.messages[0] is previous_messages[0]
    tool_output = result.messages[-1]
    assert tool_output.author.role == "tool"
    assert tool_output.author.name == "functions.lookup_weather"
    assert tool_output.channel == "commentary"
    assert tool_output.recipient == "assistant"
    assert result.engine_input["cache_salt"] == "request-salt"
    assert "arrival_time" in result.engine_input
    assert previous_messages == [
        OpenAIHarmonyMessage.from_role_and_content(Role.USER, "first"),
    ]
    assert previous_outputs[0].call_id == "call_1"


@pytest.mark.asyncio
async def test_online_renderer_returns_typed_error_for_bad_harmony_call_reference():
    renderer = _new_online_renderer()
    renderer.use_harmony = True
    request = ResponsesRequest(
        input=[
            {
                "type": "function_call_output",
                "call_id": "missing",
                "output": "orphaned",
            }
        ],
        tools=[],
    )

    result = await renderer.render_responses(request)

    assert isinstance(result, ErrorResponse)
    assert result.error.type == "invalid_request_error"
    assert result.error.param == "input"
    assert "No call message found for missing" in result.error.message


@pytest.mark.asyncio
async def test_online_renderer_preserves_missing_harmony_builtin_tool_behavior(
    monkeypatch: pytest.MonkeyPatch,
):
    renderer = _new_online_renderer()
    renderer.use_harmony = True
    renderer.model_config = SimpleNamespace(max_model_len=100)
    renderer.renderer = SimpleNamespace(
        tokenizer=SimpleNamespace(truncation_side="left")
    )
    monkeypatch.setattr(
        "vllm.renderers.online_renderer.render_for_completion",
        lambda messages: [1],
    )
    request = ResponsesRequest(
        input="search",
        tools=[{"type": "web_search_preview"}],
    )

    result = await renderer.render_responses(request, tool_server=None)

    assert not isinstance(result, ErrorResponse)


@pytest.mark.parametrize(
    "engine_inputs",
    [[], [tokens_input([1]), tokens_input([2])]],
)
def test_responses_render_result_rejects_non_single_prompt(engine_inputs):
    renderer = _new_online_renderer()

    result = renderer._responses_render_result([], engine_inputs)

    assert isinstance(result, ErrorResponse)
    assert f"got {len(engine_inputs)}" in result.error.message


class TestInitializeToolSessions:
    """Test class for _initialize_tool_sessions method."""

    @pytest.fixture(params=[ParsableContext, HarmonyContext])
    def tool_context(self, request):
        context = request.param.__new__(request.param)
        context.available_tools = ["browser", "python"]
        context._tool_sessions = {}
        context.called_tools = set()
        return context

    @pytest.mark.asyncio
    @pytest.mark.parametrize(
        ("available_tools", "called_tools"),
        [
            ([], []),
            (["browser"], ["browser"]),
            (["browser", "python"], []),
            (["browser", "python"], ["browser"]),
            (["browser", "python"], ["python"]),
            (["browser", "python"], ["browser", "python"]),
        ],
        ids=["empty", "single", "unused", "first", "last", "both"],
    )
    async def test_mcp_sessions_cleaned_once_before_close(
        self, tool_context, available_tools, called_tools
    ):
        events = []
        sessions = {}

        @asynccontextmanager
        async def new_session(name, *_args):
            session = MagicMock()
            session.call_tool = AsyncMock(
                side_effect=lambda *_: events.append(("cleanup", name))
            )
            sessions[name] = session
            try:
                yield session
            finally:
                events.append(("close", name))

        tool_server = MagicMock(spec=ToolServer)
        tool_server.new_session.side_effect = new_session
        tool_context.available_tools = available_tools
        async with AsyncExitStack() as stack:
            await tool_context.init_tool_sessions(tool_server, stack, "req", {})
            # Reinitializing existing sessions must not duplicate cleanup either.
            await tool_context.init_tool_sessions(tool_server, stack, "req", {})
            tool_context.called_tools.update(called_tools)

        for name, session in sessions.items():
            if name in called_tools:
                session.call_tool.assert_awaited_once_with("cleanup_session", {})
                assert events.index(("cleanup", name)) < events.index(("close", name))
            else:
                session.call_tool.assert_not_awaited()
        assert [name for event, name in events if event == "close"] == list(
            reversed(available_tools)
        )

    @pytest.mark.asyncio
    async def test_mcp_partial_initialization_closes_opened_session(self, tool_context):
        closed = []
        session = MagicMock(call_tool=AsyncMock())

        @asynccontextmanager
        async def new_session(name, *_args):
            if name == "python":
                raise RuntimeError("session initialization failed")
            try:
                yield session
            finally:
                closed.append(name)

        tool_server = MagicMock(spec=ToolServer)
        tool_server.new_session.side_effect = new_session
        with pytest.raises(RuntimeError, match="session initialization failed"):
            async with AsyncExitStack() as stack:
                await tool_context.init_tool_sessions(tool_server, stack, "req", {})

        assert closed == ["browser"]
        session.call_tool.assert_not_awaited()

    @pytest_asyncio.fixture
    async def serving_responses_instance(self):
        """Create a real OpenAIServingResponses instance for testing."""
        # Create minimal mocks for required dependencies
        engine_client = MagicMock()

        model_config = MagicMock()
        model_config.max_model_len = 100
        model_config.hf_config.model_type = "test"
        model_config.get_diff_sampling_param.return_value = {}
        engine_client.model_config = model_config

        engine_client.input_processor = MagicMock()
        engine_client.renderer = MagicMock()

        models = MagicMock()

        tool_server = MagicMock(spec=ToolServer)

        # Create the actual instance
        instance = OpenAIServingResponses(
            engine_client=engine_client,
            models=models,
            online_renderer=MagicMock(),
            request_logger=None,
            chat_template=None,
            chat_template_content_format="auto",
            tool_server=tool_server,
        )

        return instance

    @pytest.mark.asyncio
    async def test_stateful_harmony_delegates_named_tool_choice_to_renderer(
        self,
        serving_responses_instance,
    ):
        serving_responses_instance.use_harmony = True
        serving_responses_instance.online_renderer.render_responses = AsyncMock(
            return_value=SimpleNamespace(
                messages=[],
                engine_input=tokens_input([1]),
            )
        )
        request = ResponsesRequest(
            input="hello",
            tool_choice={"type": "function", "name": "lookup"},
            tools=[
                {
                    "type": "function",
                    "name": "lookup",
                    "parameters": {"type": "object", "properties": {}},
                }
            ],
        )

        result = await serving_responses_instance._render_resolved_response_inputs(
            request,
            None,
        )

        assert result.engine_input["prompt_token_ids"] == [1]
        serving_responses_instance.online_renderer.render_responses.assert_awaited_once()

    @pytest.mark.asyncio
    async def test_render_next_turn_reuses_shared_responses_renderer(
        self, serving_responses_instance
    ):
        serving_responses_instance.online_renderer.render_responses = AsyncMock(
            return_value=SimpleNamespace(
                messages=[],
                engine_input=tokens_input([8, 9], cache_salt="request-salt"),
            )
        )
        serving_responses_instance.online_renderer.preprocess_chat = AsyncMock(
            return_value=([], [tokens_input([1])])
        )
        request = ResponsesRequest(
            input="first",
            instructions="initial instructions",
            cache_salt="request-salt",
        )
        turn_messages = [
            {"role": "user", "content": "first"},
            {"role": "assistant", "content": "tool call"},
            {"role": "tool", "content": "tool output"},
        ]

        engine_inputs = await serving_responses_instance._render_next_turn(
            request,
            turn_messages,
        )

        assert engine_inputs[0]["prompt_token_ids"] == [8, 9]
        render_request = (
            serving_responses_instance.online_renderer.render_responses.call_args.args[
                0
            ]
        )
        assert render_request.input == turn_messages
        assert render_request.instructions is None
        assert render_request.cache_salt == "request-salt"

    @pytest.mark.asyncio
    async def test_harmony_tool_followup_preserves_render_params(
        self, serving_responses_instance, monkeypatch: pytest.MonkeyPatch
    ):
        class ToolCallingHarmonyContext(HarmonyContext):
            def __init__(self):
                self._messages = []
                self.request = None
                self._needs_tool_call = True

            def append_output(self, output) -> None:
                pass

            def need_builtin_tool_call(self) -> bool:
                return self._needs_tool_call

            async def call_tool(self):
                self._needs_tool_call = False
                return []

            def append_tool_output(self, output) -> None:
                self._messages.extend(output)

        async def generate_output():
            yield MagicMock()

        context = ToolCallingHarmonyContext()
        serving_responses_instance.engine_client.generate.side_effect = (
            lambda *args, **kwargs: generate_output()
        )
        monkeypatch.setattr(
            "vllm.renderers.online_renderer.render_for_completion",
            lambda messages: [1, 2, 3, 4],
        )
        online_renderer = serving_responses_instance.online_renderer
        online_renderer.renderer = SimpleNamespace(
            tokenizer=SimpleNamespace(truncation_side="left")
        )
        serving_responses_instance.online_renderer.render_responses_harmony_messages = (
            OnlineRenderer.render_responses_harmony_messages.__get__(online_renderer)
        )
        serving_responses_instance._extract_prompt_len = MagicMock(return_value=2)

        async for _ in serving_responses_instance._generate_with_builtin_tools(
            request_id="req",
            engine_input=tokens_input([1], cache_salt="request-salt"),
            sampling_params=MagicMock(),
            context=context,
            tok_params=TokenizeParams(
                max_total_tokens=3,
                max_output_tokens=1,
                truncate_prompt_tokens=-1,
            ),
        ):
            pass

        assert serving_responses_instance.engine_client.generate.call_count == 2
        assert context.request_metrics_cover_all_generation_turns is False
        followup_engine_input = (
            serving_responses_instance.engine_client.generate.call_args_list[1].args[0]
        )
        assert followup_engine_input.get("cache_salt") == "request-salt"
        assert followup_engine_input["prompt_token_ids"] == [3, 4]

    @pytest.mark.asyncio
    async def test_harmony_tool_followup_respects_max_output_tokens(
        self, serving_responses_instance
    ):
        class ToolCallingHarmonyContext(HarmonyContext):
            def __init__(self, request):
                self._messages = []
                self.request = request
                self._needs_tool_call = True

            def append_output(self, output) -> None:
                pass

            def need_builtin_tool_call(self) -> bool:
                return self._needs_tool_call

            async def call_tool(self):
                self._needs_tool_call = False
                return []

            def append_tool_output(self, output) -> None:
                pass

        max_tokens_per_turn = []

        async def generate_output(engine_input, sampling_params, *args, **kwargs):
            max_tokens_per_turn.append(sampling_params.max_tokens)
            yield MagicMock()

        serving_responses_instance.engine_client.generate.side_effect = generate_output
        serving_responses_instance.online_renderer.render_responses_harmony_messages = (
            MagicMock(return_value=tokens_input([1, 2, 3]))
        )
        context = ToolCallingHarmonyContext(
            ResponsesRequest(input="test", max_output_tokens=10)
        )

        async for _ in serving_responses_instance._generate_with_builtin_tools(
            request_id="req",
            engine_input=tokens_input([1]),
            sampling_params=SamplingParams(max_tokens=10),
            context=context,
        ):
            pass

        assert max_tokens_per_turn == [10, 10]

    @pytest.mark.asyncio
    async def test_initialize_tool_sessions(
        self, serving_responses_instance, mock_context, mock_exit_stack
    ):
        """Test that method works correctly with only MCP tools."""
        request = ResponsesRequest(input="test input", tools=[])

        # Call the method
        await serving_responses_instance._initialize_tool_sessions(
            request, mock_context, mock_exit_stack
        )
        assert mock_context.init_tool_sessions_called is False

        # Create only MCP tools
        tools = [
            {"type": "web_search_preview"},
            {"type": "code_interpreter", "container": {"type": "auto"}},
        ]

        request = ResponsesRequest(input="test input", tools=tools)

        # Call the method
        await serving_responses_instance._initialize_tool_sessions(
            request, mock_context, mock_exit_stack
        )

        # Verify that init_tool_sessions was called
        assert mock_context.init_tool_sessions_called

    def test_validate_create_responses_input(
        self, serving_responses_instance, mock_context, mock_exit_stack
    ):
        request = ResponsesRequest(
            input="test input",
            previous_input_messages=[
                {
                    "role": "user",
                    "content": [
                        {
                            "type": "text",
                            "text": "What is my horoscope? I am an Aquarius.",
                        }
                    ],
                }
            ],
            previous_response_id="lol",
        )
        error = serving_responses_instance._validate_create_responses_input(request)
        assert error is not None
        assert error.error.type == "invalid_request_error"


class TestValidateGeneratorInput:
    """Test class for _validate_generator_input method."""

    @pytest_asyncio.fixture
    async def serving_responses_instance(self):
        """Create a real OpenAIServingResponses instance for testing."""
        # Create minimal mocks for required dependencies
        engine_client = MagicMock()

        model_config = MagicMock()
        model_config.max_model_len = 100
        model_config.hf_config.model_type = "test"
        model_config.get_diff_sampling_param.return_value = {}
        engine_client.model_config = model_config

        engine_client.input_processor = MagicMock()
        engine_client.renderer = MagicMock()

        models = MagicMock()

        # Create the actual instance
        instance = OpenAIServingResponses(
            engine_client=engine_client,
            models=models,
            online_renderer=MagicMock(),
            request_logger=None,
            chat_template=None,
            chat_template_content_format="auto",
        )

        return instance

    def test_validate_generator_input(self, serving_responses_instance):
        """Test _validate_generator_input with valid prompt length."""
        # Create an engine prompt with valid length (less than max_model_len)
        valid_prompt_token_ids = list(range(5))  # 5 tokens < 100 max_model_len
        engine_input = tokens_input(valid_prompt_token_ids)

        # Call the method
        result = serving_responses_instance._validate_generator_input(engine_input)

        # Should return None for valid input
        assert result is None

        # create an invalid engine prompt
        invalid_prompt_token_ids = list(range(200))  # 100 tokens >= 100 max_model_len
        engine_input = tokens_input(invalid_prompt_token_ids)

        # Call the method
        result = serving_responses_instance._validate_generator_input(engine_input)

        # Should return an ErrorResponse
        assert result is not None
        assert isinstance(result, ErrorResponse)


class _Qwen3FakeTokenizer:
    def __init__(self):
        self._vocab = {"<think>": 1, "</think>": 2, "reason": 3, "final": 4}

    def get_vocab(self):
        return self._vocab

    def decode(self, token_ids):
        id_to_token = {v: k for k, v in self._vocab.items()}
        return "".join(id_to_token.get(token_id, "x") for token_id in token_ids)


def _make_qwen3_serving() -> tuple[OpenAIServingResponses, _Qwen3FakeTokenizer]:
    engine_client = MagicMock()
    model_config = MagicMock()
    model_config.hf_config.model_type = "test"
    model_config.hf_text_config = MagicMock()
    model_config.get_diff_sampling_param.return_value = {}
    engine_client.model_config = model_config
    engine_client.input_processor = MagicMock()
    engine_client.renderer = MagicMock()

    tokenizer = _Qwen3FakeTokenizer()
    engine_client.renderer.get_tokenizer.return_value = tokenizer

    serving = OpenAIServingResponses(
        engine_client=engine_client,
        models=MagicMock(),
        online_renderer=MagicMock(),
        request_logger=None,
        chat_template=None,
        chat_template_content_format="auto",
        reasoning_parser="qwen3",
    )
    return serving, tokenizer


def _make_text_request_output(text: str, token_ids: list[int]) -> RequestOutput:
    return RequestOutput(
        request_id="req",
        prompt="hi",
        prompt_token_ids=[7, 8],
        prompt_logprobs=None,
        outputs=[
            CompletionOutput(
                index=0,
                text=text,
                token_ids=token_ids,
                cumulative_logprob=0.0,
                logprobs=None,
                finish_reason="stop",
                stop_reason=None,
            )
        ],
        finished=True,
        num_cached_tokens=0,
        num_cache_creation_tokens=3,
    )


async def _run_full_generator(serving, request, context, tokenizer):
    async def dummy_result_generator():
        yield None

    return await serving.responses_full_generator(
        request=request,
        sampling_params=SamplingParams(max_tokens=16),
        result_generator=dummy_result_generator(),
        context=context,
        model_name="test-model",
        tokenizer=tokenizer,
        request_metadata=RequestResponseMetadata(request_id="req"),
    )


@pytest.mark.asyncio
async def test_reasoning_tokens_counted_for_text_reasoning_model(monkeypatch):
    """Ensure reasoning_tokens usage is derived from thinking token spans."""
    # Force non-harmony, SimpleContext path
    monkeypatch.setattr(envs, "VLLM_USE_EXPERIMENTAL_PARSER_CONTEXT", False)

    serving, tokenizer = _make_qwen3_serving()
    request = ResponsesRequest(input="hi", tools=[], stream=False)
    response_parser = serving._make_response_parser(
        request,
        tokenizer,
        serving._effective_chat_template_kwargs(request),
    )

    # Build a SimpleContext with thinking tokens in the output.
    context = SimpleContext(response_parser=response_parser)
    # <think> 10 </think> 20 -> reasoning token count = 1
    context.append_output(
        _make_text_request_output("<think>reason</think>final", [1, 10, 2, 20])
    )

    response = await _run_full_generator(serving, request, context, tokenizer)

    assert response.usage.output_tokens_details.reasoning_tokens == 1
    assert response.usage.input_tokens_details.cache_write_tokens == 3

    OpenAIResponseUsage.model_validate(response.usage.model_dump(mode="json"))


class TestExtractAllowedToolsFromMcpRequests:
    """Test class for _extract_allowed_tools_from_mcp_requests function."""

    def test_extract_allowed_tools_basic_formats(self):
        """Test extraction with list format, object format, and None."""
        from openai.types.responses.tool import McpAllowedToolsMcpToolFilter

        tools = [
            # List format
            Mcp(
                type="mcp",
                server_label="server1",
                allowed_tools=["tool1", "tool2"],
            ),
            # Object format
            Mcp(
                type="mcp",
                server_label="server2",
                allowed_tools=McpAllowedToolsMcpToolFilter(
                    tool_names=["tool3", "tool4"]
                ),
            ),
            # None (no filter)
            Mcp(
                type="mcp",
                server_label="server3",
                allowed_tools=None,
            ),
        ]
        result = _extract_allowed_tools_from_mcp_requests(tools)
        assert result == {
            "server1": ["tool1", "tool2"],
            "server2": ["tool3", "tool4"],
            "server3": None,
        }

    def test_extract_allowed_tools_star_normalization(self):
        """Test that '*' wildcard is normalized to None (select all tools).

        This is the key test requested by reviewers to explicitly demonstrate
        that the "*" select-all scenario is handled correctly.
        """
        from openai.types.responses.tool import McpAllowedToolsMcpToolFilter

        tools = [
            # Star in list format
            Mcp(
                type="mcp",
                server_label="server1",
                allowed_tools=["*"],
            ),
            # Star mixed with other tools in list
            Mcp(
                type="mcp",
                server_label="server2",
                allowed_tools=["tool1", "*"],
            ),
            # Star in object format
            Mcp(
                type="mcp",
                server_label="server3",
                allowed_tools=McpAllowedToolsMcpToolFilter(tool_names=["*"]),
            ),
        ]
        result = _extract_allowed_tools_from_mcp_requests(tools)
        # All should be normalized to None (allows all tools)
        assert result == {
            "server1": None,
            "server2": None,
            "server3": None,
        }

    def test_extract_allowed_tools_filters_non_mcp(self):
        """Test that non-MCP tools are ignored during extraction."""
        tools = [
            Mcp(
                type="mcp",
                server_label="server1",
                allowed_tools=["tool1"],
            ),
            LocalShell(type="local_shell"),  # Non-MCP tool should be ignored
            Mcp(
                type="mcp",
                server_label="server2",
                allowed_tools=["tool2"],
            ),
        ]
        result = _extract_allowed_tools_from_mcp_requests(tools)
        # Non-MCP tools should be ignored
        assert result == {
            "server1": ["tool1"],
            "server2": ["tool2"],
        }


class TestHarmonyPreambleStreaming:
    """Tests for preamble (commentary with no recipient) streaming events."""

    @staticmethod
    def _make_segment(*, channel, recipient, delta="hello"):
        """Build a lightweight segment for Harmony streaming tests."""
        return Segment(channel=channel, recipient=recipient, delta=delta)

    @staticmethod
    def _make_previous_item(*, channel, recipient, text="preamble text"):
        """Build a lightweight mock previous_item (openai_harmony Message)."""
        content_part = MagicMock()
        content_part.text = text
        item = MagicMock()
        item.channel = channel
        item.recipient = recipient
        item.content = [content_part]
        return item

    def test_preamble_delta_emits_text_events(self) -> None:
        """Commentary + recipient=None should emit output_text.delta events."""
        from vllm.entrypoints.openai.responses.streaming_events import (
            emit_content_delta_events,
        )

        segment = self._make_segment(channel="commentary", recipient=None)
        state = StreamingState()

        events = emit_content_delta_events(segment, state)

        type_names = [e.type for e in events]
        assert "response.output_text.delta" in type_names
        assert "response.output_item.added" in type_names

    def test_preamble_delta_second_token_no_added(self) -> None:
        """Second preamble token should emit delta only, not added again."""
        from vllm.entrypoints.openai.responses.streaming_events import (
            emit_content_delta_events,
        )

        segment = self._make_segment(channel="commentary", recipient=None, delta="w")
        state = StreamingState()
        state.sent_output_item_added = True
        state.current_item_id = "msg_test"
        state.current_content_index = 0

        events = emit_content_delta_events(segment, state)

        type_names = [e.type for e in events]
        assert "response.output_text.delta" in type_names
        assert "response.output_item.added" not in type_names

    def test_commentary_with_function_recipient_not_preamble(self) -> None:
        """Commentary + recipient='functions.X' must NOT use preamble path."""
        from vllm.entrypoints.openai.responses.streaming_events import (
            emit_content_delta_events,
        )

        segment = self._make_segment(
            channel="commentary",
            recipient="functions.get_weather",
        )
        state = StreamingState()

        events = emit_content_delta_events(segment, state)

        type_names = [e.type for e in events]
        assert "response.output_text.delta" not in type_names

    def test_preamble_done_emits_text_done_events(self) -> None:
        """Completed preamble should emit text done + content_part done +
        output_item done, same shape as final channel."""
        from vllm.entrypoints.openai.responses.streaming_events import (
            emit_previous_item_done_events,
        )

        previous = self._make_previous_item(channel="commentary", recipient=None)
        state = StreamingState()
        state.sent_output_item_added = True
        state.current_item_id = "msg_test"
        state.current_output_index = 0
        state.current_content_index = 0

        events = emit_previous_item_done_events(previous, state)

        type_names = [e.type for e in events]
        assert "response.output_text.done" in type_names
        assert "response.content_part.done" in type_names
        assert "response.output_item.done" in type_names

    def test_commentary_with_recipient_no_preamble_done(self) -> None:
        """Commentary + recipient='functions.X' should route to function call
        done, not preamble done."""
        from vllm.entrypoints.openai.responses.streaming_events import (
            emit_previous_item_done_events,
        )

        previous = self._make_previous_item(
            channel="commentary", recipient="functions.get_weather"
        )
        state = StreamingState()
        state.is_first_function_call_delta = True
        state.current_item_id = "fc_test"
        state.current_call_id = "call_test"

        events = emit_previous_item_done_events(previous, state)

        type_names = [e.type for e in events]
        assert "response.output_text.done" not in type_names

    @pytest.mark.xfail(
        reason=(
            "TODO: Ensure added/in-progress events are emitted for zero-delta items."
            "So we can safely emit done events for zero-delta items."
        ),
        strict=True,
    )
    def test_zero_delta_items_should_preserve_streaming_lifecycle(
        self,
    ) -> None:
        """Zero-delta Harmony items should still produce a coherent lifecycle."""
        from vllm.entrypoints.openai.responses.streaming_events import (
            emit_previous_item_done_events,
        )

        cases: list[tuple[str, str | None, str]] = [
            ("commentary", None, "msg_stale"),
            ("analysis", None, "msg_stale"),
            ("commentary", "functions.get_weather", "fc_stale"),
            ("commentary", "python", "tool_stale"),
            ("commentary", "repo_browser.list", "mcp_stale"),
        ]

        for channel, recipient, current_item_id in cases:
            previous = self._make_previous_item(channel=channel, recipient=recipient)
            state = StreamingState()
            state.current_item_id = current_item_id
            state.current_call_id = "call_stale"
            state.current_content_index = 0

            events = emit_previous_item_done_events(
                previous, state, function_tool_names=None
            )

            type_names = [e.type for e in events]
            assert "response.output_item.added" in type_names
            assert "response.output_item.done" in type_names


_PER_REQUEST_STATS = RequestStateStats(
    queued_ts=1.0,
    scheduled_ts=1.5,
    first_token_ts=2.0,
    last_token_ts=3.0,
)


def _make_request_output(
    text,
    token_ids,
    metrics: RequestStateStats | None = None,
):
    completion = CompletionOutput(
        index=0,
        text=text,
        token_ids=token_ids,
        cumulative_logprob=0.0,
        logprobs=None,
        finish_reason=None,
        stop_reason=None,
    )
    return RequestOutput(
        request_id="req",
        prompt="hi",
        prompt_token_ids=[7, 8],
        prompt_logprobs=None,
        outputs=[completion],
        finished=False,
        num_cached_tokens=0,
        metrics=metrics,
    )


def _make_simple_context_with_output(
    text,
    token_ids,
    response_parser=None,
    metrics: RequestStateStats | None = None,
):
    """Create a SimpleContext with a RequestOutput containing the given text."""
    ctx = SimpleContext(response_parser=response_parser)
    req_output = _make_request_output(text, token_ids, metrics)
    ctx.append_output(req_output)
    ctx.request_metrics = metrics
    return ctx


def _make_serving_instance(
    *,
    reasoning_parser: str = "",
    tool_parser: str | None = None,
    enable_per_request_metrics: bool = False,
) -> OpenAIServingResponses:
    engine_client = MagicMock()
    model_config = MagicMock()
    model_config.max_model_len = 100
    model_config.hf_config.model_type = "test"
    model_config.hf_text_config = MagicMock()
    model_config.get_diff_sampling_param.return_value = {}
    engine_client.model_config = model_config
    engine_client.input_processor = MagicMock()
    engine_client.renderer = MagicMock()

    return OpenAIServingResponses(
        engine_client=engine_client,
        models=MagicMock(),
        online_renderer=MagicMock(),
        request_logger=None,
        chat_template=None,
        chat_template_content_format="auto",
        reasoning_parser=reasoning_parser,
        enable_auto_tools=tool_parser is not None,
        tool_parser=tool_parser,
        enable_per_request_metrics=enable_per_request_metrics,
    )


async def _empty_context_generator():
    if False:
        yield


async def _make_full_metrics_response(
    enable_per_request_metrics: bool,
    request_metrics_cover_all_generation_turns: bool = True,
):
    serving = _make_serving_instance(
        enable_per_request_metrics=enable_per_request_metrics
    )
    request = ResponsesRequest(input="hi", tools=[], stream=False, store=False)
    sampling_params = SamplingParams(max_tokens=16)
    context = SimpleContext()
    context.request_metrics_cover_all_generation_turns = (
        request_metrics_cover_all_generation_turns
    )

    async def generate(*args, **kwargs):
        yield _make_request_output("hello", [10, 20], _PER_REQUEST_STATS)

    serving.engine_client.generate.side_effect = generate
    result_generator = serving._generate_with_builtin_tools(
        request_id=request.request_id,
        engine_input=tokens_input([7, 8]),
        sampling_params=sampling_params,
        context=context,
    )
    response = await serving.responses_full_generator(
        request=request,
        sampling_params=sampling_params,
        result_generator=result_generator,
        context=context,
        model_name="test-model",
        tokenizer=MagicMock(),
        request_metadata=RequestResponseMetadata(request_id="req"),
    )
    assert isinstance(response, ResponsesResponse)
    return response


@pytest.mark.asyncio
async def test_responses_per_request_metrics_follow_server_flag():
    disabled_response = await _make_full_metrics_response(False)
    assert disabled_response.metrics is None
    assert "metrics" not in disabled_response.model_dump(mode="json")

    enabled_response = await _make_full_metrics_response(True)
    assert enabled_response.metrics is not None
    assert enabled_response.metrics.time_to_first_token_ms == pytest.approx(500.0)


@pytest.mark.asyncio
async def test_responses_streaming_metrics_only_on_completed_event():
    serving = _make_serving_instance(enable_per_request_metrics=True)
    request = ResponsesRequest(input="hi", tools=[], stream=True, store=False)
    context = _make_simple_context_with_output(
        "hello", [10, 20], metrics=_PER_REQUEST_STATS
    )

    events = [
        event
        async for event in serving.responses_stream_generator(
            request=request,
            sampling_params=SamplingParams(max_tokens=16),
            result_generator=_empty_context_generator(),
            context=context,
            model_name="test-model",
            tokenizer=MagicMock(),
            request_metadata=RequestResponseMetadata(request_id="req"),
        )
    ]

    for event in events[:-1]:
        assert "metrics" not in event.response.model_dump(mode="json")
    assert isinstance(events[-1], ResponseCompletedEvent)
    assert events[-1].response.metrics is not None
    assert events[-1].response.metrics.time_to_first_token_ms == pytest.approx(500.0)


@pytest.mark.asyncio
async def test_responses_metrics_suppressed_for_multiple_generation_turns():
    response = await _make_full_metrics_response(
        True, request_metrics_cover_all_generation_turns=False
    )

    assert response.metrics is None
    assert "metrics" not in response.model_dump(mode="json")


def _make_serving_instance_with_reasoning():
    """Create an OpenAIServingResponses with a mocked reasoning parser."""
    return _make_serving_instance(reasoning_parser="qwen3")


def _identity_increment(event):
    """Simple identity callable for _increment_sequence_number_and_return."""
    seq = getattr(_identity_increment, "_counter", 0)
    if hasattr(event, "sequence_number"):
        event.sequence_number = seq
    _identity_increment._counter = seq + 1  # type: ignore
    return event


def _mock_parser_with_reasoning(serving, delta_sequence: list[DeltaMessage]):
    """Set up serving.parser so that it returns a mock parser instance
    with a reasoning parser that returns the given delta_sequence.

    The mock has reasoning_parser set (truthy) but tool_parser as None,
    so the parser's parse_delta enters the reasoning-only branch.
    """
    call_count = 0

    def mock_parse_delta(**kwargs):
        nonlocal call_count
        if call_count >= len(delta_sequence):
            return None
        result = delta_sequence[call_count]
        call_count += 1
        return result

    mock_parser_instance = MagicMock()
    mock_parser_instance.reasoning_parser = MagicMock()  # truthy
    mock_parser_instance.tool_parser = None
    mock_parser_instance.parse_delta = mock_parse_delta
    mock_parser_instance.is_reasoning_end = MagicMock(return_value=False)
    serving.parser = MagicMock(return_value=mock_parser_instance)
    return mock_parser_instance


class TestStreamingReasoningToContentTransition:
    """Tests for _process_simple_streaming_events reasoning-to-content
    transition, specifically the fix for mixed deltas that carry both
    reasoning and content simultaneously."""

    @pytest.mark.asyncio
    async def test_mixed_delta_reasoning_and_content_emits_reasoning_delta(
        self, monkeypatch
    ):
        """When the reasoning parser produces a delta with both reasoning
        and content set (e.g. reasoning end and content start in the same
        chunk), the trailing reasoning text must be emitted as a
        ResponseReasoningTextDeltaEvent and included in the
        ResponseReasoningTextDoneEvent text."""
        monkeypatch.setattr(envs, "VLLM_USE_EXPERIMENTAL_PARSER_CONTEXT", False)
        serving = _make_serving_instance_with_reasoning()

        # Sequence of DeltaMessages the mock orchestrator will return
        delta_sequence = [
            DeltaMessage(reasoning="thinking..."),
            DeltaMessage(reasoning=" end", content="hello"),  # mixed delta
            DeltaMessage(content=" world"),
        ]
        response_parser = _mock_parser_with_reasoning(serving, delta_sequence)
        # Create contexts for each streaming chunk
        contexts = [
            _make_simple_context_with_output("chunk1", [10], response_parser),
            _make_simple_context_with_output("chunk2", [20], response_parser),
            _make_simple_context_with_output("chunk3", [30], response_parser),
        ]

        async def result_generator():
            for ctx in contexts:
                yield ctx

        request = ResponsesRequest(input="hi", tools=[], stream=True)
        sampling_params = SamplingParams(max_tokens=64)
        metadata = RequestResponseMetadata(request_id="req")
        _identity_increment._counter = 0  # type: ignore

        events = []
        async for event in serving._process_simple_streaming_events(
            request=request,
            sampling_params=sampling_params,
            result_generator=result_generator(),
            context=SimpleContext(response_parser=response_parser),
            model_name="test-model",
            tokenizer=MagicMock(),
            request_metadata=metadata,
            created_time=0,
            _increment_sequence_number_and_return=_identity_increment,
        ):
            events.append(event)

        # The first reasoning delta should be emitted
        reasoning_deltas = [
            e for e in events if isinstance(e, ResponseReasoningTextDeltaEvent)
        ]
        assert len(reasoning_deltas) == 2
        assert reasoning_deltas[0].delta == "thinking..."
        # The trailing reasoning from the mixed delta must also be emitted
        assert reasoning_deltas[1].delta == " end"

        # The done event must include both reasoning parts
        reasoning_done = [
            e for e in events if isinstance(e, ResponseReasoningTextDoneEvent)
        ]
        assert len(reasoning_done) == 1
        assert reasoning_done[0].text == "thinking... end"

        # Content deltas should be emitted for both the mixed delta's
        # content and the pure content delta
        text_deltas = [e for e in events if isinstance(e, ResponseTextDeltaEvent)]
        assert len(text_deltas) == 2
        assert text_deltas[0].delta == "hello"
        assert text_deltas[1].delta == " world"

    @pytest.mark.asyncio
    async def test_transition_without_mixed_delta_no_extra_reasoning_event(
        self, monkeypatch
    ):
        """When the transition from reasoning to content is clean (no mixed
        delta), no extra reasoning delta event should be emitted."""
        monkeypatch.setattr(envs, "VLLM_USE_EXPERIMENTAL_PARSER_CONTEXT", False)
        serving = _make_serving_instance_with_reasoning()

        delta_sequence = [
            DeltaMessage(reasoning="thinking"),
            DeltaMessage(content="answer"),
        ]
        response_parser = _mock_parser_with_reasoning(serving, delta_sequence)

        contexts = [
            _make_simple_context_with_output("chunk1", [10], response_parser),
            _make_simple_context_with_output("chunk2", [20], response_parser),
        ]

        async def result_generator():
            for ctx in contexts:
                yield ctx

        request = ResponsesRequest(input="hi", tools=[], stream=True)
        sampling_params = SamplingParams(max_tokens=64)
        metadata = RequestResponseMetadata(request_id="req")
        _identity_increment._counter = 0  # type: ignore

        events = []
        async for event in serving._process_simple_streaming_events(
            request=request,
            sampling_params=sampling_params,
            result_generator=result_generator(),
            context=SimpleContext(response_parser=response_parser),
            model_name="test-model",
            tokenizer=MagicMock(),
            request_metadata=metadata,
            created_time=0,
            _increment_sequence_number_and_return=_identity_increment,
        ):
            events.append(event)

        # Exactly one reasoning delta
        reasoning_deltas = [
            e for e in events if isinstance(e, ResponseReasoningTextDeltaEvent)
        ]
        assert len(reasoning_deltas) == 1
        assert reasoning_deltas[0].delta == "thinking"

        # Done event has just "thinking"
        reasoning_done = [
            e for e in events if isinstance(e, ResponseReasoningTextDoneEvent)
        ]
        assert len(reasoning_done) == 1
        assert reasoning_done[0].text == "thinking"

        # One content delta
        text_deltas = [e for e in events if isinstance(e, ResponseTextDeltaEvent)]
        assert len(text_deltas) == 1
        assert text_deltas[0].delta == "answer"

    @pytest.mark.asyncio
    async def test_reasoning_only_stream_no_content(self, monkeypatch):
        """When the stream has only reasoning deltas and no content, the
        reasoning done event should be emitted at finalization with the
        full accumulated text, and no text delta events should appear."""
        monkeypatch.setattr(envs, "VLLM_USE_EXPERIMENTAL_PARSER_CONTEXT", False)
        serving = _make_serving_instance_with_reasoning()

        delta_sequence = [
            DeltaMessage(reasoning="step 1"),
            DeltaMessage(reasoning=" step 2"),
        ]
        response_parser = _mock_parser_with_reasoning(serving, delta_sequence)

        contexts = [
            _make_simple_context_with_output("chunk1", [10], response_parser),
            _make_simple_context_with_output("chunk2", [20], response_parser),
        ]

        async def result_generator():
            for ctx in contexts:
                yield ctx

        request = ResponsesRequest(input="hi", tools=[], stream=True)
        sampling_params = SamplingParams(max_tokens=64)
        metadata = RequestResponseMetadata(request_id="req")
        _identity_increment._counter = 0  # type: ignore

        events = []
        async for event in serving._process_simple_streaming_events(
            request=request,
            sampling_params=sampling_params,
            result_generator=result_generator(),
            context=SimpleContext(response_parser=response_parser),
            model_name="test-model",
            tokenizer=MagicMock(),
            request_metadata=metadata,
            created_time=0,
            _increment_sequence_number_and_return=_identity_increment,
        ):
            events.append(event)

        # Two reasoning deltas
        reasoning_deltas = [
            e for e in events if isinstance(e, ResponseReasoningTextDeltaEvent)
        ]
        assert len(reasoning_deltas) == 2
        assert reasoning_deltas[0].delta == "step 1"
        assert reasoning_deltas[1].delta == " step 2"

        # Done event at finalization with accumulated text
        reasoning_done = [
            e for e in events if isinstance(e, ResponseReasoningTextDoneEvent)
        ]
        assert len(reasoning_done) == 1
        assert reasoning_done[0].text == "step 1 step 2"

        # No content text deltas
        text_deltas = [e for e in events if isinstance(e, ResponseTextDeltaEvent)]
        assert len(text_deltas) == 0

        # Final item should be a reasoning item
        item_done_events = [
            e for e in events if isinstance(e, ResponseOutputItemDoneEvent)
        ]
        assert len(item_done_events) == 1
        assert isinstance(item_done_events[0].item, ResponseReasoningItem)


class TestAutoToolStreaming:
    @staticmethod
    async def _collect_events(delta_sequence: list[DeltaMessage]):
        serving = _make_serving_instance_with_reasoning()
        response_parser = _mock_parser_with_reasoning(serving, delta_sequence)

        contexts = [
            _make_simple_context_with_output("chunk", [i], response_parser)
            for i in range(len(delta_sequence))
        ]

        async def result_generator():
            for ctx in contexts:
                yield ctx

        request = ResponsesRequest(
            input="hi",
            tools=[
                {
                    "type": "function",
                    "name": "get_weather",
                    "description": "Get weather.",
                    "parameters": {
                        "type": "object",
                        "properties": {"location": {"type": "string"}},
                        "required": ["location"],
                        "additionalProperties": False,
                    },
                }
            ],
            tool_choice="auto",
            stream=True,
        )
        sampling_params = SamplingParams(max_tokens=64)
        metadata = RequestResponseMetadata(request_id="req")
        _identity_increment._counter = 0  # type: ignore

        events = []
        async for event in serving._process_simple_streaming_events(
            request=request,
            sampling_params=sampling_params,
            result_generator=result_generator(),
            context=SimpleContext(response_parser=response_parser),
            model_name="test-model",
            tokenizer=MagicMock(),
            request_metadata=metadata,
            created_time=0,
            _increment_sequence_number_and_return=_identity_increment,
        ):
            events.append(event)
        return events

    @pytest.mark.skip_global_cleanup
    @pytest.mark.asyncio
    async def test_auto_multi_tool_streaming_opens_one_item_per_tool(self, monkeypatch):
        monkeypatch.setattr(envs, "VLLM_USE_EXPERIMENTAL_PARSER_CONTEXT", False)

        delta_sequence = [
            DeltaMessage(
                tool_calls=[
                    DeltaToolCall(
                        id="call_vienna",
                        type="function",
                        index=0,
                        function=DeltaFunctionCall(
                            name="get_weather",
                            arguments="",
                        ),
                    )
                ]
            ),
            DeltaMessage(
                tool_calls=[
                    DeltaToolCall(
                        index=0,
                        function=DeltaFunctionCall(
                            arguments='{"location":"Vienna"}',
                        ),
                    )
                ]
            ),
            DeltaMessage(
                tool_calls=[
                    DeltaToolCall(
                        id="call_berlin",
                        type="function",
                        index=1,
                        function=DeltaFunctionCall(
                            name="get_weather",
                            arguments='{"location":"Berlin"}',
                        ),
                    )
                ]
            ),
        ]
        events = await self._collect_events(delta_sequence)

        function_items = [
            event
            for event in events
            if event.type == "response.output_item.added"
            and getattr(event.item, "type", None) == "function_call"
        ]
        assert len(function_items) == 2
        assert [event.item.name for event in function_items] == [
            "get_weather",
            "get_weather",
        ]
        assert [event.output_index for event in function_items] == [0, 1]

        argument_deltas = [
            event.delta
            for event in events
            if event.type == "response.function_call_arguments.delta"
        ]
        assert argument_deltas == [
            '{"location":"Vienna"}',
            '{"location":"Berlin"}',
        ]

        argument_done = [
            event
            for event in events
            if event.type == "response.function_call_arguments.done"
        ]
        assert [event.arguments for event in argument_done] == [
            '{"location":"Vienna"}',
            '{"location":"Berlin"}',
        ]
        assert [event.output_index for event in argument_done] == [0, 1]

        function_done = [
            event
            for event in events
            if event.type == "response.output_item.done"
            and getattr(event.item, "type", None) == "function_call"
        ]
        assert [event.item.arguments for event in function_done] == [
            '{"location":"Vienna"}',
            '{"location":"Berlin"}',
        ]
        assert [event.output_index for event in function_done] == [0, 1]

    @pytest.mark.skip_global_cleanup
    @pytest.mark.asyncio
    async def test_auto_tool_choice_first_delta_tool_call_does_not_duplicate_item(
        self, monkeypatch
    ):
        monkeypatch.setattr(envs, "VLLM_USE_EXPERIMENTAL_PARSER_CONTEXT", False)

        delta_sequence = [
            DeltaMessage(
                tool_calls=[
                    DeltaToolCall(
                        id="call_test",
                        type="function",
                        index=0,
                        function=DeltaFunctionCall(
                            name="get_weather",
                            arguments="",
                        ),
                    )
                ]
            ),
            DeltaMessage(
                tool_calls=[
                    DeltaToolCall(
                        index=0,
                        function=DeltaFunctionCall(
                            arguments='{"location":"Berlin"}',
                        ),
                    )
                ]
            ),
        ]
        events = await self._collect_events(delta_sequence)

        function_items = [
            event
            for event in events
            if event.type == "response.output_item.added"
            and getattr(event.item, "type", None) == "function_call"
        ]
        assert len(function_items) == 1
        assert function_items[0].item.name == "get_weather"

        argument_deltas = [
            event.delta
            for event in events
            if event.type == "response.function_call_arguments.delta"
        ]
        assert "".join(argument_deltas) == '{"location":"Berlin"}'

    @pytest.mark.skip_global_cleanup
    @pytest.mark.asyncio
    async def test_compound_content_and_tool_name_args_same_delta(self, monkeypatch):
        monkeypatch.setattr(envs, "VLLM_USE_EXPERIMENTAL_PARSER_CONTEXT", False)

        tool_args = '{"location":"Berlin"}'

        delta_sequence = [
            DeltaMessage(
                content="Let me check.",
                tool_calls=[
                    DeltaToolCall(
                        id="call_weather",
                        type="function",
                        index=0,
                        function=DeltaFunctionCall(name="get_weather"),
                    ),
                    DeltaToolCall(
                        index=0,
                        function=DeltaFunctionCall(arguments=tool_args),
                    ),
                ],
            )
        ]

        events = await self._collect_events(delta_sequence)

        text_deltas = [
            event.delta
            for event in events
            if event.type == "response.output_text.delta"
        ]
        assert text_deltas == ["Let me check."]

        argument_deltas = [
            event.delta
            for event in events
            if event.type == "response.function_call_arguments.delta"
        ]
        assert argument_deltas == [tool_args]

        types = [event.type for event in events]
        assert types.index("response.output_text.delta") < types.index(
            "response.function_call_arguments.delta"
        )

        argument_done = [
            event
            for event in events
            if event.type == "response.function_call_arguments.done"
        ]
        assert [event.arguments for event in argument_done] == [tool_args]

        function_done = [
            event
            for event in events
            if event.type == "response.output_item.done"
            and getattr(event.item, "type", None) == "function_call"
        ]
        assert len(function_done) == 1
        assert function_done[0].item.name == "get_weather"
        assert function_done[0].item.arguments == tool_args


@pytest.mark.skip_global_cleanup
@pytest.mark.asyncio
async def test_stream_completed_response_reuses_streamed_items(monkeypatch):
    """response.completed must carry the streamed items (same ids,
    call_id, logprobs) rather than a reparse of the full output."""
    monkeypatch.setattr(envs, "VLLM_USE_EXPERIMENTAL_PARSER_CONTEXT", False)
    serving = _make_serving_instance_with_reasoning()
    call = DeltaToolCall(
        id="chatcmpl-tool-parser-id",
        index=0,
        function=DeltaFunctionCall(name="get_weather", arguments="{}"),
    )
    response_parser = _mock_parser_with_reasoning(
        serving,
        [
            DeltaMessage(reasoning="think"),
            DeltaMessage(content="Hi"),
            DeltaMessage(tool_calls=[call]),
        ],
    )
    response_parser.count_reasoning_tokens = MagicMock(return_value=1)
    context = SimpleContext(response_parser=response_parser)

    async def result_generator():
        for token_id, text in [(10, "think"), (20, "Hi"), (30, "call")]:
            output = _make_request_output(text, [token_id])
            output.outputs[0].logprobs = [
                {token_id: SampleLogprob(logprob=-0.5, decoded_token=text)}
            ]
            context.append_output(output)
            yield context

    events = [
        event
        async for event in serving.responses_stream_generator(
            request=ResponsesRequest(
                input="hi",
                include=["message.output_text.logprobs"],
                stream=True,
                store=False,
            ),
            sampling_params=SamplingParams(max_tokens=16),
            result_generator=result_generator(),
            context=context,
            model_name="test-model",
            tokenizer=MagicMock(),
            request_metadata=RequestResponseMetadata(request_id="req"),
        )
    ]

    streamed = [e.item for e in events if e.type == "response.output_item.done"]
    assert events[-1].response.output == streamed
    _, message, function_call = streamed
    assert [lp.token for lp in message.content[0].logprobs] == ["Hi"]
    assert function_call.call_id == "chatcmpl-tool-parser-id"


def _harmony_msg(channel: str, text: str, recipient: str | None = None):
    msg = OpenAIHarmonyMessage.from_role_and_content(Role.ASSISTANT, text)
    msg = msg.with_channel(channel)
    return msg.with_recipient(recipient) if recipient else msg


def _harmony_segments(msg, *, streamed: bool = True) -> list[Segment]:
    """Segments the parser yields for one message: a content delta (absent
    for a zero-delta message), then the completed message."""
    delta = [Segment(msg.channel, msg.recipient, msg.content[0].text)]
    completed = Segment(msg.channel, msg.recipient, "", completed_message=msg)
    return (delta if streamed else []) + [completed]


async def _harmony_stream_events(segment_chunks: list[list[Segment]]):
    serving = _make_serving_instance()
    serving.use_harmony = True
    parser = MagicMock(spec=HarmonyParser)
    parser.process_chunk.side_effect = [
        ChunkResult(segments=segments, reasoning_token_count=0)
        for segments in segment_chunks
    ]
    context = HarmonyContext([], [], frozenset({"get_weather"}), parser)

    async def result_generator():
        for token_id in range(len(segment_chunks)):
            context.append_output(_make_request_output("", [token_id]))
            yield context

    return [
        event
        async for event in serving.responses_stream_generator(
            request=ResponsesRequest(input="hi", stream=True, store=False),
            sampling_params=SamplingParams(max_tokens=16),
            result_generator=result_generator(),
            context=context,
            model_name="test-model",
            tokenizer=MagicMock(),
            request_metadata=RequestResponseMetadata(request_id="req"),
        )
    ]


@pytest.mark.asyncio
async def test_harmony_stream_completed_response_reuses_streamed_ids():
    """response.completed must carry the id and call_id each item was streamed
    with, not ones minted when the harmony messages are converted again."""
    messages = [
        _harmony_msg("analysis", "think"),
        _harmony_msg("commentary", '{"city": "Paris"}', "functions.get_weather"),
        _harmony_msg("final", "Sunny"),
    ]

    events = await _harmony_stream_events([_harmony_segments(msg) for msg in messages])

    streamed = [e.item for e in events if e.type == "response.output_item.done"]
    output = events[-1].response.output
    assert [item.type for item in output] == ["reasoning", "function_call", "message"]
    assert [item.id for item in output] == [item.id for item in streamed]
    assert output[1].call_id == streamed[1].call_id
    assert output[2].content[0].text == "Sunny"


@pytest.mark.asyncio
async def test_harmony_stream_unstreamed_item_keeps_own_id():
    """An item with no done event (zero-delta) must not take the streamed id
    of a later item of the same type."""
    unstreamed = _harmony_msg("analysis", "skipped")
    reasoning = _harmony_msg("analysis", "think")

    events = await _harmony_stream_events(
        [_harmony_segments(unstreamed, streamed=False), _harmony_segments(reasoning)]
    )

    (streamed,) = [e.item for e in events if e.type == "response.output_item.done"]
    first, second = events[-1].response.output
    assert first.content[0].text == "skipped"
    assert first.id != streamed.id
    assert second.id == streamed.id


@pytest.mark.asyncio
@pytest.mark.parametrize(("parallel_tool_calls", "num_calls"), [(False, 1), (None, 3)])
async def test_streaming_parallel_tool_calls(
    monkeypatch, parallel_tool_calls, num_calls
):
    """parallel_tool_calls=False keeps only the first call even when the grammar
    does not limit the model; None keeps every call."""
    monkeypatch.setattr(envs, "VLLM_USE_EXPERIMENTAL_PARSER_CONTEXT", False)
    serving = _make_serving_instance(tool_parser="hermes")
    request = ResponsesRequest(
        input="hi",
        tools=[{"type": "function", "name": "get_weather", "parameters": {}}],
        parallel_tool_calls=parallel_tool_calls,
        stream=True,
    )
    context = SimpleContext(
        response_parser=serving._make_response_parser(request, MagicMock(), {})
    )
    args = [f'{{"city": "{city}"}}' for city in ("Paris", "Berlin", "Tokyo")]

    async def result_generator():
        for i, arg in enumerate(args):
            text = f'<tool_call>\n{{"name": "get_weather", "arguments": {arg}}}'
            context.append_output(_make_request_output(f"{text}\n</tool_call>\n", [i]))
            yield context

    events = [
        event
        async for event in serving.responses_stream_generator(
            request=request,
            sampling_params=SamplingParams(max_tokens=64),
            result_generator=result_generator(),
            context=context,
            model_name="test-model",
            tokenizer=MagicMock(),
            request_metadata=RequestResponseMetadata(request_id="req"),
        )
    ]

    assert [
        event.output_index
        for event in events
        if event.type == "response.output_item.added"
    ] == list(range(num_calls))
    assert all(getattr(event, "output_index", 0) < num_calls for event in events)
    assert [
        event.item.arguments
        for event in events
        if event.type == "response.output_item.done"
    ] == args[:num_calls]
    assert [item.arguments for item in events[-1].response.output] == args[:num_calls]
