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

import pytest
from openai.types.responses import ResponseFunctionWebSearch, ResponseStreamEvent
from openai_harmony import Message, Role
from pydantic import TypeAdapter

from vllm.entrypoints.generate.base.protocol import (
    DeltaFunctionCall,
    DeltaMessage,
    DeltaToolCall,
)
from vllm.entrypoints.openai.responses.streaming_events import (
    SimpleStreamingEventProcessor,
    StreamingState,
    _StateType,
    emit_browser_tool_events,
    emit_reasoning_delta_events,
    emit_reasoning_done_events,
    split_delta,
)


def test_browser_find_uses_responses_action_type():
    """Both streaming output items must use the Responses find action type."""
    message = (
        Message.from_role_and_content(
            Role.ASSISTANT, '{"pattern": "vLLM", "cursor": 42}'
        )
        .with_channel("analysis")
        .with_recipient("browser.find")
    )

    events = emit_browser_tool_events(message, StreamingState())

    added, done = events[0], events[-1]
    assert added.type == "response.output_item.added"
    assert done.type == "response.output_item.done"
    for event in (added, done):
        assert isinstance(event.item, ResponseFunctionWebSearch)
        assert event.item.action.type == "find_in_page"
        assert event.item.action.pattern == "vLLM"
        assert event.item.action.url == "cursor:42"


@pytest.mark.parametrize("harmony", [False, True], ids=["simple", "harmony"])
def test_reasoning_content_parts_use_sdk_events(harmony):
    text = "Let me think."
    if harmony:
        state = StreamingState()
        events = emit_reasoning_delta_events(text, state)
        events.extend(emit_reasoning_done_events(text, state))
    else:
        processor = SimpleStreamingEventProcessor()
        events = _run_through_processor(processor, DeltaMessage(reasoning=text))
        events.extend(processor.close_current())

    assert [event.type for event in events] == [
        "response.output_item.added",
        "response.content_part.added",
        "response.reasoning_text.delta",
        "response.reasoning_text.done",
        "response.content_part.done",
        "response.output_item.done",
    ]
    added, part_added, delta, text_done, part_done, done = events
    assert part_added.part.type == part_done.part.type == "reasoning_text"
    assert part_added.part.text == ""
    assert delta.delta == text_done.text == part_done.part.text == text
    assert done.item.content[0].text == text
    assert done.item.id == added.item.id
    for event in events[1:-1]:
        assert event.item_id == added.item.id
        assert event.content_index == part_added.content_index
    validator = TypeAdapter(ResponseStreamEvent)
    for event in events:
        assert event.output_index == added.output_index
        validator.validate_python(event.model_dump())


def _make_tool_call(
    index: int, name: str | None = None, arguments: str | None = None
) -> DeltaToolCall:
    fn = DeltaFunctionCall(name=name, arguments=arguments)
    return DeltaToolCall(index=index, function=fn)


class TestSplitDelta:
    def test_all_three_fields(self):
        tc = _make_tool_call(0, name="f")
        delta = DeltaMessage(reasoning="r", content="c", tool_calls=[tc])
        result = split_delta(delta)

        assert len(result) == 3
        assert result[0].reasoning == "r" and result[0].content is None
        assert result[1].content == "c" and result[1].reasoning is None
        assert len(result[2].tool_calls) == 1 and result[2].content is None

    def test_tool_calls_grouped_by_index(self):
        tc0 = _make_tool_call(0, name="f1")
        tc1 = _make_tool_call(1, name="f2")
        tc0b = _make_tool_call(0, arguments='{"a":1}')

        # Different indices → split
        result = split_delta(DeltaMessage(tool_calls=[tc0, tc1]))
        assert len(result) == 2
        assert result[0].tool_calls == [tc0]
        assert result[1].tool_calls == [tc1]

        # Same index → stays together
        delta = DeltaMessage(tool_calls=[tc0, tc0b])
        result = split_delta(delta)
        assert len(result) == 1
        assert result[0] is delta


def _run_through_processor(
    processor: SimpleStreamingEventProcessor,
    delta_message: DeltaMessage,
) -> list:
    """Simulate the streaming loop from serving.py for a single delta."""
    events = []
    for dm in split_delta(delta_message):
        target_state, tool_call = processor.resolve_target_state(dm)
        if target_state == _StateType.NONE:
            continue
        if processor.needs_transition(target_state, tool_call):
            events.extend(processor.close_current())
            events.extend(processor.open(target_state, tool_call))
        events.extend(processor.emit_delta(dm, None))
    return events


class TestProcessorCompoundDeltas:
    def test_all_three_states(self):
        tc = _make_tool_call(0, name="f", arguments="{}")
        delta = DeltaMessage(reasoning="r", content="c", tool_calls=[tc])

        processor = SimpleStreamingEventProcessor()
        events = _run_through_processor(processor, delta)

        types = [e.type for e in events]
        r_idx = types.index("response.reasoning_text.delta")
        c_idx = types.index("response.output_text.delta")
        fc_idx = types.index("response.function_call_arguments.delta")
        assert r_idx < c_idx < fc_idx

    def test_parallel_tool_calls(self):
        tc0 = _make_tool_call(0, name="f1", arguments='{"a":1}')
        tc1 = _make_tool_call(1, name="f2", arguments='{"b":2}')
        delta = DeltaMessage(tool_calls=[tc0, tc1])

        processor = SimpleStreamingEventProcessor()
        events = _run_through_processor(processor, delta)

        added = [e for e in events if e.type == "response.output_item.added"]
        deltas = [
            e for e in events if e.type == "response.function_call_arguments.delta"
        ]
        assert len(added) == 2
        assert len(deltas) == 2

    def test_split_name_and_args_same_index(self):
        """Regression: parsers like KimiK2 emit name and args as separate
        DeltaToolCalls at the same index within one DeltaMessage."""
        tc_name = _make_tool_call(0, name="get_weather")
        tc_args = _make_tool_call(0, arguments='{"city":"SF"}')
        delta = DeltaMessage(tool_calls=[tc_name, tc_args])

        processor = SimpleStreamingEventProcessor()
        events = _run_through_processor(processor, delta)

        deltas = [
            e for e in events if e.type == "response.function_call_arguments.delta"
        ]
        assert len(deltas) == 1
        assert deltas[0].delta == '{"city":"SF"}'

    def test_reasoning_to_content_transition(self):
        """Regression: the old special case in emit_delta handled this;
        now split_delta handles it generically."""
        processor = SimpleStreamingEventProcessor()
        _run_through_processor(processor, DeltaMessage(reasoning="think"))
        assert processor.state.current_state == _StateType.REASONING

        events = _run_through_processor(
            processor, DeltaMessage(reasoning="more", content="answer")
        )
        types = [e.type for e in events]
        assert "response.reasoning_text.delta" in types
        assert "response.output_text.delta" in types
