Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
331 changes: 331 additions & 0 deletions tests/test_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -2379,6 +2379,337 @@ def extract_tool_calls_streaming(
"total_tokens": 10,
}

@pytest.mark.anyio
async def test_stream_terminal_finish_reason_when_tool_parser_suppresses_eot(
self, monkeypatch
):
"""A tool_calls stream must still end with finish_reason when the engine's
finished output is swallowed by the tool parser's `continue`.

Reproduces the gemma4 production failure: the model emits the complete
canonical call (end marker included) in one delta, then a bare <turn|>
end-of-turn token as the terminal output. The tool parser finds nothing
new in that delta and suppresses it — without a guard, no chunk ever
carries finish_reason and strict OpenAI clients abort with
"Stream ended without finish_reason". Ref: #672.
"""
from vllm_mlx.engine.base import GenerationOutput
from vllm_mlx.server import (
ChatCompletionRequest,
Message,
stream_chat_completion,
)
import vllm_mlx.server as server

class FakeEngine:
model_name = "fake-engine"
tokenizer = None

async def stream_chat(self, messages, **kwargs):
chunks = [
GenerationOutput(
text="",
new_text=(
'<|tool_call>call:get_weather{<|"|>city<|"|>: '
'<|"|>Paris<|"|>}<tool_call|>'
),
finished=False,
),
GenerationOutput(
text="",
new_text="<turn|>",
finished=True,
finish_reason="stop",
prompt_tokens=11,
completion_tokens=5,
),
]
for chunk in chunks:
yield chunk

monkeypatch.setattr(server, "_model_name", "served-model")
monkeypatch.setattr(server, "_reasoning_parser", None)
monkeypatch.setattr(server, "_enable_auto_tool_choice", True)
monkeypatch.setattr(server, "_tool_call_parser", "gemma4")
monkeypatch.setattr(server, "_tool_parser_instance", None)

request = ChatCompletionRequest(
model="request-model",
messages=[Message(role="user", content="weather in Paris?")],
tools=[
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
},
}
],
stream=True,
)

chunks = [
chunk
async for chunk in stream_chat_completion(
FakeEngine(), request.messages, request
)
]

payloads = [
json.loads(chunk.removeprefix("data: ").strip())
for chunk in chunks
if chunk != "data: [DONE]\n\n"
]
tool_payloads = [
payload
for payload in payloads
if payload["choices"] and payload["choices"][0]["delta"].get("tool_calls")
]

assert len(tool_payloads) == 1
assert payloads[-1]["choices"][0]["finish_reason"] == "tool_calls"

@pytest.mark.anyio
async def test_stream_terminal_finish_reason_when_reasoning_path_suppresses_eot(
self, monkeypatch
):
"""Same terminal-swallow class through the reasoning-parser branch:
the finished output is consumed by the tool parser after a completed
call, so the guard must emit the missing finish_reason chunk. Ref: #672.
"""
from vllm_mlx.engine.base import GenerationOutput
from vllm_mlx.reasoning import DeltaMessage
from vllm_mlx.server import (
ChatCompletionRequest,
Message,
stream_chat_completion,
)
import vllm_mlx.server as server

class FakeEngine:
model_name = "fake-engine"

async def stream_chat(self, messages, **kwargs):
chunks = [
GenerationOutput(
text="",
new_text=(
'<|tool_call>call:search{<|"|>q<|"|>: <|"|>weather<|"|>}'
"<tool_call|>"
),
finished=False,
),
GenerationOutput(
text="",
new_text="<turn|>",
finished=True,
finish_reason="stop",
prompt_tokens=7,
completion_tokens=3,
),
]
for chunk in chunks:
yield chunk

class FakeReasoningParser:
def reset_state(self):
pass

def extract_reasoning_streaming(
self, previous_text, current_text, delta_text
):
return DeltaMessage(content=delta_text)

monkeypatch.setattr(server, "_model_name", "served-model")
monkeypatch.setattr(server, "_reasoning_parser", FakeReasoningParser())
monkeypatch.setattr(server, "_enable_auto_tool_choice", True)
monkeypatch.setattr(server, "_tool_call_parser", "gemma4")
monkeypatch.setattr(server, "_tool_parser_instance", None)

class FakeEngineWithTokenizer(FakeEngine):
tokenizer = None

request = ChatCompletionRequest(
model="request-model",
messages=[Message(role="user", content="hi")],
stream=True,
)

chunks = [
chunk
async for chunk in stream_chat_completion(
FakeEngineWithTokenizer(), request.messages, request
)
]

payloads = [
json.loads(chunk.removeprefix("data: ").strip())
for chunk in chunks
if chunk != "data: [DONE]\n\n"
]
tool_payloads = [
payload
for payload in payloads
if payload["choices"] and payload["choices"][0]["delta"].get("tool_calls")
]

assert len(tool_payloads) == 1
assert payloads[-1]["choices"][0]["finish_reason"] == "tool_calls"

@pytest.mark.anyio
async def test_stream_terminal_finish_reason_when_engine_emits_none(
self, monkeypatch
):
"""Regression: SimpleEngine can emit finished=True with finish_reason=None
on natural stops before #681. The tracker must not falsely claim a reason
was emitted just because finished=True; the guard should still fire and
stamp the default "stop". Ref: #672.
"""
from vllm_mlx.engine.base import GenerationOutput
from vllm_mlx.server import (
ChatCompletionRequest,
Message,
stream_chat_completion,
)
import vllm_mlx.server as server

class FakeEngine:
model_name = "fake-engine"
tokenizer = None

async def stream_chat(self, messages, **kwargs):
yield GenerationOutput(
text="",
new_text="hello",
finished=False,
)
yield GenerationOutput(
text="",
new_text="",
finished=True,
finish_reason=None,
prompt_tokens=3,
completion_tokens=1,
)

monkeypatch.setattr(server, "_model_name", "served-model")
monkeypatch.setattr(server, "_reasoning_parser", None)
monkeypatch.setattr(server, "_enable_auto_tool_choice", False)
monkeypatch.setattr(server, "_tool_call_parser", None)
monkeypatch.setattr(server, "_tool_parser_instance", None)

request = ChatCompletionRequest(
model="request-model",
messages=[Message(role="user", content="hi")],
stream=True,
)

chunks = [
chunk
async for chunk in stream_chat_completion(
FakeEngine(), request.messages, request
)
]

payloads = [
json.loads(chunk.removeprefix("data: ").strip())
for chunk in chunks
if chunk != "data: [DONE]\n\n"
]

# The last payload before [DONE] must carry finish_reason="stop".
assert payloads[-1]["choices"][0]["finish_reason"] == "stop"
usage_payloads = [payload for payload in payloads if payload.get("usage")]
assert len(usage_payloads) == 1
assert usage_payloads[0] == payloads[-1]
assert usage_payloads[0]["usage"] == {
"prompt_tokens": 3,
"completion_tokens": 1,
"total_tokens": 4,
}

@pytest.mark.anyio
async def test_stream_terminal_when_reasoning_parser_swallows_finished_delta(
self, monkeypatch
):
"""A reasoning parser may consume the finished delta directly."""
from vllm_mlx.engine.base import GenerationOutput
from vllm_mlx.reasoning import DeltaMessage
from vllm_mlx.server import (
ChatCompletionRequest,
Message,
stream_chat_completion,
)
import vllm_mlx.server as server

class FakeEngine:
model_name = "fake-engine"
tokenizer = None

async def stream_chat(self, messages, **kwargs):
yield GenerationOutput(
text="",
new_text="visible",
finished=False,
)
yield GenerationOutput(
text="",
new_text="</think>",
finished=True,
finish_reason=None,
prompt_tokens=5,
completion_tokens=2,
)

class FakeReasoningParser:
def reset_state(self):
pass

def extract_reasoning_streaming(
self, previous_text, current_text, delta_text
):
if delta_text == "</think>":
return None
return DeltaMessage(content=delta_text)

monkeypatch.setattr(server, "_model_name", "served-model")
monkeypatch.setattr(server, "_reasoning_parser", FakeReasoningParser())
monkeypatch.setattr(server, "_enable_auto_tool_choice", False)
monkeypatch.setattr(server, "_tool_call_parser", None)
monkeypatch.setattr(server, "_tool_parser_instance", None)

request = ChatCompletionRequest(
model="request-model",
messages=[Message(role="user", content="hi")],
stream=True,
)
chunks = [
chunk
async for chunk in stream_chat_completion(
FakeEngine(), request.messages, request
)
]
payloads = [
json.loads(chunk.removeprefix("data: ").strip())
for chunk in chunks
if chunk != "data: [DONE]\n\n"
]

assert payloads[-1]["choices"][0]["finish_reason"] == "stop"
usage_payloads = [payload for payload in payloads if payload.get("usage")]
assert len(usage_payloads) == 1
assert usage_payloads[0] == payloads[-1]
assert usage_payloads[0]["usage"] == {
"prompt_tokens": 5,
"completion_tokens": 2,
"total_tokens": 7,
}

@pytest.mark.anyio
async def test_reasoning_stream_redirects_gemma4_tool_marker(self, monkeypatch):
"""Gemma 4 tool markup inside reasoning should reach the tool parser."""
Expand Down
Loading
Loading