mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(hosted_vllm): keep reasoning_content on replayed assistant messages (#43599)
* fix(hosted_vllm): keep reasoning_content on assistant messages in _transform_messages vLLM accepts reasoning_content (200 on the wire) and qwen/deepseek/glm chat templates consume it, so popping it made reasoning models lose earlier reasoning across tool loops. thinking_blocks is still removed for vLLM compatibility. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(hosted_vllm): forward replayed reasoning_content only when it is a string * test(integration): cover hosted_vllm reasoning_content replay across endpoints * test(integration): require the surviving worker to serve its held requests in the sigkill chaos cell --------- Co-authored-by: mateo <mateo@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Co-authored-by: yassin <yassin@berri.ai>
This commit is contained in:
parent
405ed414cb
commit
657bb777fa
5 changed files with 1220 additions and 7 deletions
|
|
@ -161,13 +161,14 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
|
|||
"""
|
||||
Support translating:
|
||||
- video files from file_id or file_data to video_url
|
||||
- thinking_blocks and reasoning_content on assistant messages are removed,
|
||||
and content lists are converted to strings for vLLM compatibility
|
||||
- thinking_blocks and non-string reasoning_content on assistant messages
|
||||
are removed, and content lists are converted to strings for vLLM compatibility
|
||||
"""
|
||||
for message in messages:
|
||||
if message["role"] == "assistant":
|
||||
message.pop("thinking_blocks", None)
|
||||
message.pop("reasoning_content", None)
|
||||
if not isinstance(message.get("reasoning_content"), str):
|
||||
message.pop("reasoning_content", None)
|
||||
existing_content = message.get("content")
|
||||
if isinstance(existing_content, list):
|
||||
text_parts = []
|
||||
|
|
|
|||
|
|
@ -0,0 +1,223 @@
|
|||
import json
|
||||
import uuid
|
||||
from typing import Final
|
||||
|
||||
import anthropic
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_BACKEND: Final = "glm-reasoning"
|
||||
_API_KEY: Final = "synthetic-hosted-vllm-key"
|
||||
_TOOL_USE_ID: Final = "toolu_weather_1"
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_MESSAGES: Final = TypeAdapter(list[dict[str, JsonValue]])
|
||||
_TOOLS: Final[list[dict[str, JsonValue]]] = [
|
||||
{
|
||||
"name": "get_weather",
|
||||
"description": "Get the current weather for a city",
|
||||
"input_schema": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]},
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def _completion(identity: str) -> bytes:
|
||||
return json.dumps(
|
||||
{
|
||||
"id": identity,
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": _BACKEND,
|
||||
"choices": [
|
||||
{"index": 0, "message": {"role": "assistant", "content": "It is raining."}, "finish_reason": "stop"}
|
||||
],
|
||||
"usage": {"prompt_tokens": 30, "completion_tokens": 5, "total_tokens": 35},
|
||||
}
|
||||
).encode()
|
||||
|
||||
|
||||
def _streamed_completion(identity: str) -> Reply:
|
||||
chunk: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": _BACKEND}
|
||||
frames: Final = (
|
||||
{**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "It is raining."}}]},
|
||||
{
|
||||
**chunk,
|
||||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 30, "completion_tokens": 5, "total_tokens": 35},
|
||||
},
|
||||
)
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=(*(b"data: " + json.dumps(frame).encode() + b"\n\n" for frame in frames), b"data: [DONE]\n\n"),
|
||||
)
|
||||
|
||||
|
||||
def _tool_loop(thinking: str, marker: str) -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
{"role": "user", "content": f"What is the weather in Paris? {marker}"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "thinking", "thinking": thinking, "signature": "opaque-signature"},
|
||||
{"type": "text", "text": "Let me check."},
|
||||
{"type": "tool_use", "id": _TOOL_USE_ID, "name": "get_weather", "input": {"city": "Paris"}},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "tool_result", "tool_use_id": _TOOL_USE_ID, "content": "light rain, 14C"}],
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _expected_upstream(thinking: str, marker: str) -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
{"role": "user", "content": f"What is the weather in Paris? {marker}"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Let me check.",
|
||||
"reasoning_content": thinking,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": _TOOL_USE_ID,
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": json.dumps({"city": "Paris"})},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": _TOOL_USE_ID, "content": "light rain, 14C"},
|
||||
]
|
||||
|
||||
|
||||
def _only_body(wire: Wire) -> dict[str, JsonValue]:
|
||||
received: Final = wire.drain()
|
||||
assert [(request.method, request.target) for request in received] == [("POST", "/v1/chat/completions")]
|
||||
return _JSON_OBJECT.validate_json(received[0].body)
|
||||
|
||||
|
||||
def _sent_messages(body: dict[str, JsonValue]) -> list[dict[str, JsonValue]]:
|
||||
return _MESSAGES.validate_python(body["messages"])
|
||||
|
||||
|
||||
def _spend_status(identity: str) -> JsonValue:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT status FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (identity,)),
|
||||
lambda found: len(found) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
return rows[0]["status"]
|
||||
|
||||
|
||||
def _post_messages(gateway: Gateway, model: str, messages: list[dict[str, JsonValue]]) -> dict[str, JsonValue]:
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/messages",
|
||||
{"model": model, "max_tokens": 256, "messages": messages, "cache": {"no-cache": True}},
|
||||
headers={"anthropic-version": "2023-06-01"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
return _JSON_OBJECT.validate_json(response.content)
|
||||
|
||||
|
||||
def test_anthropic_sdk_thinking_block_reaches_hosted_vllm_as_reasoning_content(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
identity: Final = f"chatcmpl-messages-{marker}"
|
||||
thinking: Final = f"The user wants Paris weather, codeword mango{marker[:4]}."
|
||||
with wire_server(lambda _: Reply(body=_completion(identity))) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
client: Final = anthropic.Anthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0)
|
||||
message: Final = client.messages.create(
|
||||
model=model,
|
||||
max_tokens=256,
|
||||
tools=_TOOLS, # pyright: ignore[reportArgumentType] # plain JSON tool definitions
|
||||
messages=_tool_loop(thinking, marker), # pyright: ignore[reportArgumentType] # plain JSON content blocks
|
||||
)
|
||||
assert message.id == identity
|
||||
assert [(block.type, getattr(block, "text", None)) for block in message.content] == [("text", "It is raining.")]
|
||||
body: Final = _only_body(wire)
|
||||
assert _sent_messages(body) == _expected_upstream(thinking, marker)
|
||||
assert "thinking_blocks" not in json.dumps(body) and "opaque-signature" not in json.dumps(body), body
|
||||
assert _spend_status(identity) == "success"
|
||||
|
||||
|
||||
async def test_async_anthropic_sdk_stream_forwards_thinking_to_hosted_vllm(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
identity: Final = f"chatcmpl-messages-stream-{marker}"
|
||||
thinking: Final = f"Streaming thought {marker}."
|
||||
with wire_server(lambda _: _streamed_completion(identity)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
client: Final = anthropic.AsyncAnthropic(
|
||||
base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0
|
||||
)
|
||||
stream: Final = await client.messages.create(
|
||||
model=model,
|
||||
max_tokens=256,
|
||||
tools=_TOOLS, # pyright: ignore[reportArgumentType] # plain JSON tool definitions
|
||||
messages=_tool_loop(thinking, marker), # pyright: ignore[reportArgumentType] # plain JSON content blocks
|
||||
stream=True,
|
||||
)
|
||||
events: Final = [event async for event in stream]
|
||||
assert events[0].type == "message_start" and events[-1].type == "message_stop"
|
||||
message_id: Final = events[0].message.id
|
||||
assert "".join(
|
||||
event.delta.text
|
||||
for event in events
|
||||
if event.type == "content_block_delta" and event.delta.type == "text_delta"
|
||||
) == ("It is raining.")
|
||||
body: Final = _only_body(wire)
|
||||
assert body["stream"] is True
|
||||
assert _sent_messages(body) == _expected_upstream(thinking, marker)
|
||||
assert _spend_status(message_id) == "success"
|
||||
|
||||
|
||||
def test_redacted_thinking_alone_sends_no_reasoning_content(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with (
|
||||
wire_server(lambda _: Reply(body=_completion(f"chatcmpl-{marker}"))) as wire,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
_post_messages(
|
||||
gateway,
|
||||
model,
|
||||
[
|
||||
{"role": "user", "content": f"Hello {marker}"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "redacted_thinking", "data": "opaque-redacted"},
|
||||
{"type": "text", "text": "Hi."},
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "Again"},
|
||||
],
|
||||
)
|
||||
assert _sent_messages(_only_body(wire)) == [
|
||||
{"role": "user", "content": f"Hello {marker}"},
|
||||
{"role": "assistant", "content": "Hi."},
|
||||
{"role": "user", "content": "Again"},
|
||||
]
|
||||
|
||||
|
||||
def test_assistant_turn_without_thinking_sends_no_reasoning_content(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with (
|
||||
wire_server(lambda _: Reply(body=_completion(f"chatcmpl-{marker}"))) as wire,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
_post_messages(
|
||||
gateway,
|
||||
model,
|
||||
[
|
||||
{"role": "user", "content": f"Hello {marker}"},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "Hi."}]},
|
||||
{"role": "user", "content": "Again"},
|
||||
],
|
||||
)
|
||||
assert _sent_messages(_only_body(wire)) == [
|
||||
{"role": "user", "content": f"Hello {marker}"},
|
||||
{"role": "assistant", "content": "Hi."},
|
||||
{"role": "user", "content": "Again"},
|
||||
]
|
||||
|
|
@ -0,0 +1,393 @@
|
|||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import signal
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
import psutil
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import owned_proxy_process
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_BACKEND: Final = "qwen3-reasoning-chaos"
|
||||
_API_KEY: Final = "synthetic-hosted-vllm-key"
|
||||
_CONFIG_MODEL: Final = "hosted-vllm-reasoning-chaos"
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_MESSAGES: Final = TypeAdapter(list[dict[str, JsonValue]])
|
||||
_MARKER: Final = re.compile(r"marker-([0-9a-f]{32})")
|
||||
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
|
||||
_MODEL_LIST: Final = json.dumps(
|
||||
{"object": "list", "data": [{"id": _BACKEND, "object": "model", "owned_by": "vllm"}]}
|
||||
).encode()
|
||||
|
||||
Endpoint = Literal["chat", "messages", "responses"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Call:
|
||||
endpoint: Endpoint
|
||||
stream: bool
|
||||
marker: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Served:
|
||||
call: _Call
|
||||
status: int
|
||||
text: str
|
||||
|
||||
|
||||
def _thought(marker: str) -> str:
|
||||
return f"private thought for {marker}"
|
||||
|
||||
|
||||
def _answer(marker: str) -> str:
|
||||
return f"answer marker-{marker}"
|
||||
|
||||
|
||||
def _path(endpoint: Endpoint) -> str:
|
||||
match endpoint:
|
||||
case "chat":
|
||||
return "/v1/chat/completions"
|
||||
case "messages":
|
||||
return "/v1/messages"
|
||||
case "responses":
|
||||
return "/v1/responses"
|
||||
|
||||
|
||||
def _body(model: str, call: _Call) -> dict[str, JsonValue]:
|
||||
question: Final = f"Question marker-{call.marker}"
|
||||
common: Final[dict[str, JsonValue]] = {"model": model, "stream": call.stream, "num_retries": 0}
|
||||
match call.endpoint:
|
||||
case "chat":
|
||||
return {
|
||||
**common,
|
||||
"messages": [
|
||||
{"role": "user", "content": question},
|
||||
{"role": "assistant", "content": "Working on it.", "reasoning_content": _thought(call.marker)},
|
||||
{"role": "user", "content": "Go on."},
|
||||
],
|
||||
}
|
||||
case "messages":
|
||||
return {
|
||||
**common,
|
||||
"max_tokens": 64,
|
||||
"messages": [
|
||||
{"role": "user", "content": question},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "thinking", "thinking": _thought(call.marker), "signature": "sig"},
|
||||
{"type": "text", "text": "Working on it."},
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "Go on."},
|
||||
],
|
||||
}
|
||||
case "responses":
|
||||
return {
|
||||
**common,
|
||||
"input": [
|
||||
{"role": "user", "content": question},
|
||||
{
|
||||
"id": f"rs_{call.marker}",
|
||||
"type": "reasoning",
|
||||
"summary": [{"type": "summary_text", "text": _thought(call.marker)}],
|
||||
},
|
||||
{"role": "user", "content": "Go on."},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _chat_reply(marker: str, stream: bool, abort_after: int | None = None, pause: float = 0) -> Reply:
|
||||
usage: Final = {"prompt_tokens": 30, "completion_tokens": 5, "total_tokens": 35}
|
||||
if not stream:
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": f"chatcmpl-{marker}",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": _BACKEND,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": _answer(marker)},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": usage,
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
chunk: Final = {"id": f"chatcmpl-{marker}", "object": "chat.completion.chunk", "created": 1, "model": _BACKEND}
|
||||
frames: Final = (
|
||||
{**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "answer "}}]},
|
||||
{**chunk, "choices": [{"index": 0, "delta": {"content": f"marker-{marker}"}}]},
|
||||
{**chunk, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], "usage": usage},
|
||||
)
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=(*(b"data: " + json.dumps(frame).encode() + b"\n\n" for frame in frames), b"data: [DONE]\n\n"),
|
||||
abort_after=abort_after,
|
||||
pause_between_chunks=pause,
|
||||
)
|
||||
|
||||
|
||||
def _responses_reply(marker: str, stream: bool) -> Reply:
|
||||
identity: Final = f"resp_upstream_{marker}"
|
||||
response: Final = {
|
||||
"id": identity,
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": _BACKEND,
|
||||
"output": [
|
||||
{
|
||||
"id": f"msg_{marker}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": _answer(marker), "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 30, "output_tokens": 5, "total_tokens": 35},
|
||||
}
|
||||
if not stream:
|
||||
return Reply(body=json.dumps(response).encode())
|
||||
events: Final = (
|
||||
{
|
||||
"type": "response.created",
|
||||
"sequence_number": 0,
|
||||
"response": {**response, "status": "in_progress", "output": []},
|
||||
},
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"sequence_number": 1,
|
||||
"item_id": f"msg_{marker}",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": _answer(marker),
|
||||
},
|
||||
{"type": "response.completed", "sequence_number": 2, "response": response},
|
||||
)
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events),
|
||||
)
|
||||
|
||||
|
||||
def _marker_of(request: Request) -> str:
|
||||
found: Final = _MARKER.search(request.body.decode())
|
||||
assert found is not None, request.body
|
||||
return found.group(1)
|
||||
|
||||
|
||||
def _echo(request: Request) -> Reply:
|
||||
marker: Final = _marker_of(request)
|
||||
stream: Final = _JSON_OBJECT.validate_json(request.body).get("stream") is True
|
||||
if request.target == "/v1/responses":
|
||||
return _responses_reply(marker, stream)
|
||||
return _chat_reply(marker, stream)
|
||||
|
||||
|
||||
def _forwarded_reasoning(request: Request) -> tuple[str, JsonValue]:
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
if request.target == "/v1/responses":
|
||||
reasoning_item: Final = _MESSAGES.validate_python(body["input"])[1]
|
||||
return _marker_of(request), _MESSAGES.validate_python(reasoning_item["summary"])[0]["text"]
|
||||
assert request.target == "/v1/chat/completions", request.target
|
||||
return _marker_of(request), _MESSAGES.validate_python(body["messages"])[1].get("reasoning_content")
|
||||
|
||||
|
||||
def _assert_no_bleed(received: tuple[Request, ...], markers: frozenset[str]) -> None:
|
||||
forwarded: Final = [_forwarded_reasoning(request) for request in received]
|
||||
assert sorted(marker for marker, _ in forwarded) == sorted(markers)
|
||||
assert all(reasoning == _thought(marker) for marker, reasoning in forwarded), forwarded
|
||||
|
||||
|
||||
def _spend_statuses(model: str, expected: int) -> list[JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
|
||||
lambda found: len(found) >= expected,
|
||||
seconds=60,
|
||||
)
|
||||
assert len({row["request_id"] for row in rows}) == len(rows), rows
|
||||
return [row["status"] for row in rows]
|
||||
|
||||
|
||||
async def _send(client: httpx.AsyncClient, key: str, model: str, call: _Call) -> _Served:
|
||||
async with client.stream(
|
||||
"POST",
|
||||
_path(call.endpoint),
|
||||
json=_body(model, call),
|
||||
headers={"Authorization": f"Bearer {key}", "anthropic-version": "2023-06-01"},
|
||||
) as response:
|
||||
raw: Final = await response.aread()
|
||||
return _Served(call=call, status=response.status_code, text=raw.decode())
|
||||
|
||||
|
||||
async def _burst(
|
||||
base_url: str, key: str, model: str, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False
|
||||
) -> tuple[_Served, ...]:
|
||||
async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client:
|
||||
results: Final = await asyncio.gather(
|
||||
*(_send(client, key, model, call) for call in calls), return_exceptions=tolerate_transport_errors
|
||||
)
|
||||
for result in results:
|
||||
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
|
||||
return tuple(result for result in results if isinstance(result, _Served))
|
||||
|
||||
|
||||
def _calls(count: int, endpoints: tuple[Endpoint, ...], stream: Callable[[int], bool]) -> tuple[_Call, ...]:
|
||||
return tuple(
|
||||
_Call(endpoint=endpoints[index % len(endpoints)], stream=stream(index), marker=uuid.uuid4().hex)
|
||||
for index in range(count)
|
||||
)
|
||||
|
||||
|
||||
def _assert_answered_with_its_own_marker(served: _Served) -> None:
|
||||
assert served.status == 200, served.text
|
||||
assert set(_MARKER.findall(served.text)) == {served.call.marker}, served.text
|
||||
|
||||
|
||||
async def test_concurrent_replays_across_endpoints_keep_each_reasoning_with_its_request(gateway: Gateway) -> None:
|
||||
calls: Final = _calls(30, ("chat", "messages", "responses"), lambda index: index % 2 == 0)
|
||||
with wire_server(_echo) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls)
|
||||
assert len(served) == 30
|
||||
for item in served:
|
||||
_assert_answered_with_its_own_marker(item)
|
||||
_assert_no_bleed(wire.drain(), frozenset(call.marker for call in calls))
|
||||
assert _spend_statuses(model, 30) == ["success"] * 30
|
||||
|
||||
|
||||
async def test_upstream_stream_aborts_reach_callers_and_later_replays_still_forward_reasoning(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
calls: Final = _calls(12, ("chat",), lambda _: True)
|
||||
aborted: Final = frozenset(call.marker for index, call in enumerate(calls) if index % 3 == 0)
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
marker: Final = _marker_of(request)
|
||||
return _chat_reply(marker, stream=True, abort_after=0 if marker in aborted else None)
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls)
|
||||
assert len(served) == 12
|
||||
for item in served:
|
||||
if item.call.marker in aborted:
|
||||
assert item.status == 500, item.text
|
||||
assert "APIConnectionError" in item.text and "marker-" not in item.text, item.text
|
||||
else:
|
||||
_assert_answered_with_its_own_marker(item)
|
||||
assert item.text.rstrip().endswith("data: [DONE]"), item.text
|
||||
recovery: Final = _Call(endpoint="chat", stream=True, marker=uuid.uuid4().hex)
|
||||
(recovered,) = await _burst(str(gateway.client.base_url), gateway.key, model, (recovery,))
|
||||
_assert_answered_with_its_own_marker(recovered)
|
||||
_assert_no_bleed(wire.drain(), frozenset(call.marker for call in (*calls, recovery)))
|
||||
|
||||
|
||||
async def test_slow_upstream_streams_are_forwarded_once_with_their_own_reasoning(gateway: Gateway) -> None:
|
||||
calls: Final = _calls(10, ("chat",), lambda _: True)
|
||||
with (
|
||||
wire_server(lambda request: _chat_reply(_marker_of(request), stream=True, pause=0.3)) as wire,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls)
|
||||
assert len(served) == 10
|
||||
for item in served:
|
||||
_assert_answered_with_its_own_marker(item)
|
||||
assert item.text.rstrip().endswith("data: [DONE]"), item.text
|
||||
_assert_no_bleed(wire.drain(), frozenset(call.marker for call in calls))
|
||||
assert _spend_statuses(model, 10) == ["success"] * 10
|
||||
|
||||
|
||||
def _chaos_config(wire: Wire, tmp_path: Path) -> Path:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["model_list"] = [
|
||||
{
|
||||
"model_name": _CONFIG_MODEL,
|
||||
"litellm_params": {"model": f"hosted_vllm/{_BACKEND}", "api_base": wire.url + "/v1", "api_key": _API_KEY},
|
||||
}
|
||||
]
|
||||
path: Final = tmp_path / "hosted-vllm-reasoning-chaos.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
def _open_upstream_connections(pid: int, upstream: str) -> int:
|
||||
port: Final = urlsplit(upstream).port
|
||||
return sum(
|
||||
1
|
||||
for connection in psutil.Process(pid).net_connections(kind="tcp")
|
||||
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(180)
|
||||
async def test_worker_sigkill_mid_burst_leaves_the_sibling_forwarding_reasoning(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
calls: Final = _calls(20, ("chat",), lambda _: False)
|
||||
release: Final = threading.Event()
|
||||
held_markers: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
|
||||
def held(request: Request) -> Reply:
|
||||
if (request.method, request.target) == ("GET", "/v1/models"):
|
||||
return Reply(body=_MODEL_LIST)
|
||||
held_markers.put(_marker_of(request))
|
||||
assert release.wait(timeout=60), "The burst was never released"
|
||||
return _echo(request)
|
||||
|
||||
with wire_server(held) as wire:
|
||||
path: Final = _chaos_config(wire, tmp_path)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
|
||||
candidate: Final = owned.gateway
|
||||
workers: Final = eventually(
|
||||
lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())),
|
||||
lambda pids: len(pids) == 2,
|
||||
seconds=30,
|
||||
)
|
||||
burst: Final = asyncio.create_task(
|
||||
_burst(
|
||||
str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, calls, tolerate_transport_errors=True
|
||||
)
|
||||
)
|
||||
await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == 20, 60)
|
||||
held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, wire.url) for pid in workers})
|
||||
assert sum(held_by.values()) == 20, held_by
|
||||
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
|
||||
victim: Final = psutil.Process(victim_pid)
|
||||
victim.suspend()
|
||||
victim.send_signal(signal.SIGKILL)
|
||||
release.set()
|
||||
served: Final = await burst
|
||||
assert held_by[survivor_pid] >= 10, held_by
|
||||
assert len(served) == held_by[survivor_pid], (held_by, len(served))
|
||||
for item in served:
|
||||
_assert_answered_with_its_own_marker(item)
|
||||
follow_up: Final = _Call(endpoint="chat", stream=False, marker=uuid.uuid4().hex)
|
||||
(answered,) = await _burst(str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, (follow_up,))
|
||||
_assert_answered_with_its_own_marker(answered)
|
||||
received: Final = wire.drain()
|
||||
chats: Final = tuple(request for request in received if request.method == "POST")
|
||||
probes: Final = [(request.method, request.target) for request in received if request.method != "POST"]
|
||||
assert set(probes) <= {("GET", "/v1/models")}, probes
|
||||
_assert_no_bleed(chats, frozenset(call.marker for call in (*calls, follow_up)))
|
||||
|
|
@ -0,0 +1,565 @@
|
|||
import json
|
||||
import uuid
|
||||
from collections.abc import Sequence
|
||||
from typing import Final
|
||||
|
||||
import openai
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_BACKEND: Final = "qwen3-reasoning"
|
||||
_FALLBACK_BACKEND: Final = "qwen3-reasoning-fallback"
|
||||
_API_KEY: Final = "synthetic-hosted-vllm-key"
|
||||
_REASONING: Final = "I compared the two invoices and the totals differ by 42."
|
||||
_ANSWER_REASONING: Final = "The user wants the difference, which is 42."
|
||||
_TOOL_CALL_ID: Final = "call_reasoning_wire_1"
|
||||
_NO_CACHE: Final[dict[str, JsonValue]] = {"no-cache": True}
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_MESSAGES: Final = TypeAdapter(list[dict[str, JsonValue]])
|
||||
|
||||
|
||||
def _completion(identity: str, content: str) -> bytes:
|
||||
return json.dumps(
|
||||
{
|
||||
"id": identity,
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": _BACKEND,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": content, "reasoning_content": _ANSWER_REASONING},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 30, "completion_tokens": 5, "total_tokens": 35},
|
||||
}
|
||||
).encode()
|
||||
|
||||
|
||||
def _streamed_completion(identity: str, content: str) -> Reply:
|
||||
chunk: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": _BACKEND}
|
||||
frames: Final = (
|
||||
{**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "reasoning_content": _ANSWER_REASONING}}]},
|
||||
{**chunk, "choices": [{"index": 0, "delta": {"content": content}}]},
|
||||
{
|
||||
**chunk,
|
||||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 30, "completion_tokens": 5, "total_tokens": 35},
|
||||
},
|
||||
)
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=(*(b"data: " + json.dumps(frame).encode() + b"\n\n" for frame in frames), b"data: [DONE]\n\n"),
|
||||
)
|
||||
|
||||
|
||||
def _replayed_conversation(reasoning: JsonValue, marker: str) -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
{"role": "user", "content": f"Compare these invoices {marker}."},
|
||||
{"role": "assistant", "content": "Checking the totals.", "reasoning_content": reasoning},
|
||||
{"role": "user", "content": "What is the difference?"},
|
||||
]
|
||||
|
||||
|
||||
def _sent_messages(request: Request) -> list[dict[str, JsonValue]]:
|
||||
return _MESSAGES.validate_python(_JSON_OBJECT.validate_json(request.body)["messages"])
|
||||
|
||||
|
||||
def _only_request(wire: Wire) -> Request:
|
||||
received: Final = wire.drain()
|
||||
assert [(request.method, request.target) for request in received] == [("POST", "/v1/chat/completions")]
|
||||
return received[0]
|
||||
|
||||
|
||||
def _spend_row(identity: str) -> dict[str, JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT model_group, status, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
|
||||
(identity,),
|
||||
),
|
||||
lambda found: len(found) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
return rows[0]
|
||||
|
||||
|
||||
def _model_spend_statuses(model: str) -> list[JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
|
||||
lambda found: len(found) >= 1,
|
||||
seconds=70,
|
||||
)
|
||||
return [row["status"] for row in rows]
|
||||
|
||||
|
||||
def _openai_client(gateway: Gateway) -> openai.OpenAI:
|
||||
return openai.OpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0)
|
||||
|
||||
|
||||
def _async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI:
|
||||
return openai.AsyncOpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0)
|
||||
|
||||
|
||||
def _post_chat(gateway: Gateway, model: str, messages: Sequence[dict[str, JsonValue]]) -> dict[str, JsonValue]:
|
||||
response: Final = gateway.request(
|
||||
"POST", "/v1/chat/completions", {"model": model, "messages": list(messages), "cache": _NO_CACHE}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
return _JSON_OBJECT.validate_json(response.content)
|
||||
|
||||
|
||||
def test_hosted_vllm_assistant_reasoning_content_reaches_the_wire(gateway: Gateway) -> None:
|
||||
identity: Final = f"hosted-vllm-reasoning-{uuid.uuid4().hex}"
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == "/v1/chat/completions"
|
||||
assert request.headers["authorization"] == f"Bearer {_API_KEY}"
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
assert body["model"] == _BACKEND
|
||||
assert body["messages"] == [
|
||||
{"role": "user", "content": "Compare these invoices."},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Checking the totals.",
|
||||
"reasoning_content": _REASONING,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": _TOOL_CALL_ID,
|
||||
"type": "function",
|
||||
"function": {"name": "lookup_invoice", "arguments": json.dumps({"id": "inv-7"})},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": _TOOL_CALL_ID, "content": "invoice total is 1042"},
|
||||
], body["messages"]
|
||||
return Reply(body=_completion(identity, "The totals differ by 42."))
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [
|
||||
{"role": "user", "content": "Compare these invoices."},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Checking the totals.",
|
||||
"reasoning_content": _REASONING,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": _TOOL_CALL_ID,
|
||||
"type": "function",
|
||||
"function": {"name": "lookup_invoice", "arguments": json.dumps({"id": "inv-7"})},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": _TOOL_CALL_ID, "content": "invoice total is 1042"},
|
||||
],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
payload: Final = _JSON_OBJECT.validate_json(response.content)
|
||||
assert payload["id"] == identity
|
||||
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/v1/chat/completions")]
|
||||
|
||||
|
||||
def test_openai_sdk_replayed_reasoning_reaches_hosted_vllm_and_is_billed_once(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
identity: Final = f"chatcmpl-sdk-{marker}"
|
||||
with wire_server(lambda _: Reply(body=_completion(identity, "They differ by 42."))) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
completion: Final = _openai_client(gateway).chat.completions.create(
|
||||
model=model,
|
||||
messages=_replayed_conversation(_REASONING, marker), # pyright: ignore[reportArgumentType] # reasoning_content is a provider extension the SDK types omit
|
||||
)
|
||||
assert completion.id == identity
|
||||
assert completion.choices[0].message.content == "They differ by 42."
|
||||
assert (completion.choices[0].message.model_extra or {})["reasoning_content"] == _ANSWER_REASONING
|
||||
assert _sent_messages(_only_request(wire)) == _replayed_conversation(_REASONING, marker)
|
||||
assert _spend_row(identity) == {
|
||||
"model_group": model,
|
||||
"status": "success",
|
||||
"prompt_tokens": 30,
|
||||
"completion_tokens": 5,
|
||||
}
|
||||
|
||||
|
||||
async def test_async_openai_sdk_stream_forwards_replayed_reasoning_to_hosted_vllm(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
identity: Final = f"chatcmpl-stream-{marker}"
|
||||
with wire_server(lambda _: _streamed_completion(identity, "They differ by 42.")) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
stream: Final = await _async_openai_client(gateway).chat.completions.create(
|
||||
model=model,
|
||||
messages=_replayed_conversation(_REASONING, marker), # pyright: ignore[reportArgumentType] # reasoning_content is a provider extension the SDK types omit
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
)
|
||||
chunks: Final = [chunk async for chunk in stream]
|
||||
assert {chunk.id for chunk in chunks} == {identity}
|
||||
assert "".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices) == (
|
||||
"They differ by 42."
|
||||
)
|
||||
sent: Final = _only_request(wire)
|
||||
assert _JSON_OBJECT.validate_json(sent.body)["stream"] is True
|
||||
assert _sent_messages(sent) == _replayed_conversation(_REASONING, marker)
|
||||
assert _spend_row(identity)["status"] == "success"
|
||||
|
||||
|
||||
def test_each_replayed_turn_keeps_its_own_reasoning_in_order(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
conversation: Final[list[dict[str, JsonValue]]] = [
|
||||
{"role": "user", "content": f"Plan the migration {marker}."},
|
||||
{"role": "assistant", "content": "Step one.", "reasoning_content": f"first thought {marker}"},
|
||||
{"role": "user", "content": "Continue."},
|
||||
{"role": "assistant", "content": "Step two.", "reasoning_content": f"second thought {marker}"},
|
||||
{"role": "user", "content": "Summarize."},
|
||||
]
|
||||
with wire_server(lambda _: Reply(body=_completion(f"chatcmpl-{marker}", "Done."))) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
assert _post_chat(gateway, model, conversation)["id"] == f"chatcmpl-{marker}"
|
||||
assert _sent_messages(_only_request(wire)) == conversation
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("reasoning", "forwarded"),
|
||||
[
|
||||
pytest.param("", "", id="empty-string-forwarded"),
|
||||
pytest.param("x" * 5120, "x" * 5120, id="5kb-string-forwarded-intact"),
|
||||
pytest.param(None, None, id="null-dropped"),
|
||||
pytest.param(42, None, id="int-dropped"),
|
||||
pytest.param(["step one", "step two"], None, id="list-dropped"),
|
||||
pytest.param({"text": "step one"}, None, id="object-dropped"),
|
||||
],
|
||||
)
|
||||
def test_only_string_reasoning_content_is_forwarded_to_hosted_vllm(
|
||||
gateway: Gateway, reasoning: JsonValue, forwarded: str | None
|
||||
) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(lambda _: Reply(body=_completion(f"chatcmpl-{marker}", "Done."))) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
assert _post_chat(gateway, model, _replayed_conversation(reasoning, marker))["id"] == f"chatcmpl-{marker}"
|
||||
sent_assistant: Final = _sent_messages(_only_request(wire))[1]
|
||||
expected_assistant: Final[dict[str, JsonValue]] = {"role": "assistant", "content": "Checking the totals."}
|
||||
assert sent_assistant == (
|
||||
expected_assistant if forwarded is None else {**expected_assistant, "reasoning_content": forwarded}
|
||||
)
|
||||
|
||||
|
||||
def test_assistant_turn_without_reasoning_gets_no_reasoning_key(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
conversation: Final[list[dict[str, JsonValue]]] = [
|
||||
{"role": "user", "content": f"Hello {marker}"},
|
||||
{"role": "assistant", "content": "Hi there."},
|
||||
{"role": "user", "content": "Again"},
|
||||
]
|
||||
with wire_server(lambda _: Reply(body=_completion(f"chatcmpl-{marker}", "Hello again."))) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
_post_chat(gateway, model, conversation)
|
||||
assert _sent_messages(_only_request(wire)) == conversation
|
||||
|
||||
|
||||
def test_same_reasoning_on_two_turns_is_forwarded_on_both(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
reasoning: Final = f"repeated thought {marker}"
|
||||
conversation: Final[list[dict[str, JsonValue]]] = [
|
||||
{"role": "user", "content": "One"},
|
||||
{"role": "assistant", "content": "First.", "reasoning_content": reasoning},
|
||||
{"role": "user", "content": "Two"},
|
||||
{"role": "assistant", "content": "Second.", "reasoning_content": reasoning},
|
||||
{"role": "user", "content": "Three"},
|
||||
]
|
||||
with wire_server(lambda _: Reply(body=_completion(f"chatcmpl-{marker}", "Third."))) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
_post_chat(gateway, model, conversation)
|
||||
assert _sent_messages(_only_request(wire)) == conversation
|
||||
|
||||
|
||||
def test_thinking_blocks_are_stripped_while_reasoning_content_is_kept(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(lambda _: Reply(body=_completion(f"chatcmpl-{marker}", "Done."))) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
_post_chat(
|
||||
gateway,
|
||||
model,
|
||||
[
|
||||
{"role": "user", "content": f"Hello {marker}"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Hi.",
|
||||
"reasoning_content": _REASONING,
|
||||
"thinking_blocks": [{"type": "thinking", "thinking": _REASONING, "signature": "sig"}],
|
||||
},
|
||||
{"role": "user", "content": "Again"},
|
||||
],
|
||||
)
|
||||
assert _sent_messages(_only_request(wire))[1] == {
|
||||
"role": "assistant",
|
||||
"content": "Hi.",
|
||||
"reasoning_content": _REASONING,
|
||||
}
|
||||
|
||||
|
||||
def test_list_content_is_flattened_while_reasoning_content_is_kept(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(lambda _: Reply(body=_completion(f"chatcmpl-{marker}", "Done."))) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
_post_chat(
|
||||
gateway,
|
||||
model,
|
||||
[
|
||||
{"role": "user", "content": f"Hello {marker}"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "Part one."}, {"type": "text", "text": "Part two."}],
|
||||
"reasoning_content": _REASONING,
|
||||
},
|
||||
{"role": "user", "content": "Again"},
|
||||
],
|
||||
)
|
||||
assert _sent_messages(_only_request(wire))[1] == {
|
||||
"role": "assistant",
|
||||
"content": "Part one.\nPart two.",
|
||||
"reasoning_content": _REASONING,
|
||||
}
|
||||
|
||||
|
||||
def test_unauthenticated_replay_is_rejected_before_hosted_vllm(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(lambda _: Reply(body=_completion(f"chatcmpl-{marker}", "Done."))) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": _replayed_conversation(_REASONING, marker)},
|
||||
key=f"sk-not-a-key-{marker}",
|
||||
)
|
||||
assert response.status_code == 401, response.text
|
||||
assert wire.drain() == ()
|
||||
|
||||
|
||||
def test_hosted_vllm_auth_error_reaches_the_caller_after_one_attempt_with_reasoning(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
error_message: Final = f"invalid api key for deployment {marker}"
|
||||
reply: Final = Reply(
|
||||
status=401,
|
||||
body=json.dumps({"error": {"message": error_message, "type": "authentication_error"}}).encode(),
|
||||
)
|
||||
with wire_server(lambda _: reply) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": _replayed_conversation(_REASONING, marker), "cache": _NO_CACHE},
|
||||
)
|
||||
assert response.status_code == 401, response.text
|
||||
assert error_message in response.text, response.text
|
||||
assert _sent_messages(_only_request(wire)) == _replayed_conversation(_REASONING, marker)
|
||||
|
||||
|
||||
def test_fallback_attempt_replays_reasoning_to_the_second_deployment(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
if _JSON_OBJECT.validate_json(request.body)["model"] == _BACKEND:
|
||||
return Reply(status=500, body=b'{"error": {"message": "primary deployment is down"}}')
|
||||
return Reply(body=_completion(f"chatcmpl-fallback-{marker}", "Recovered."))
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
primary: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
fallback: Final = scenario.model(
|
||||
model=f"hosted_vllm/{_FALLBACK_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY
|
||||
)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": primary,
|
||||
"messages": _replayed_conversation(_REASONING, marker),
|
||||
"fallbacks": [fallback],
|
||||
"num_retries": 0,
|
||||
"cache": _NO_CACHE,
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert _JSON_OBJECT.validate_json(response.content)["id"] == f"chatcmpl-fallback-{marker}"
|
||||
attempts: Final = wire.drain()
|
||||
assert [_JSON_OBJECT.validate_json(attempt.body)["model"] for attempt in attempts] == [
|
||||
_BACKEND,
|
||||
_FALLBACK_BACKEND,
|
||||
]
|
||||
assert [_sent_messages(attempt) for attempt in attempts] == [
|
||||
_replayed_conversation(_REASONING, marker),
|
||||
_replayed_conversation(_REASONING, marker),
|
||||
]
|
||||
|
||||
|
||||
def test_identical_uncached_replays_are_each_forwarded_and_billed_once(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
identities: Final = iter((f"chatcmpl-first-{marker}", f"chatcmpl-second-{marker}"))
|
||||
with wire_server(lambda _: Reply(body=_completion(next(identities), "Done."))) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
first: Final = _post_chat(gateway, model, _replayed_conversation(_REASONING, marker))
|
||||
second: Final = _post_chat(gateway, model, _replayed_conversation(_REASONING, marker))
|
||||
assert (first["id"], second["id"]) == (f"chatcmpl-first-{marker}", f"chatcmpl-second-{marker}")
|
||||
assert [_sent_messages(request) for request in wire.drain()] == [
|
||||
_replayed_conversation(_REASONING, marker),
|
||||
_replayed_conversation(_REASONING, marker),
|
||||
]
|
||||
assert _spend_row(f"chatcmpl-first-{marker}")["status"] == "success"
|
||||
assert _spend_row(f"chatcmpl-second-{marker}")["status"] == "success"
|
||||
|
||||
|
||||
def test_cached_replay_hits_only_for_the_same_reasoning(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
identities: Final = iter((f"chatcmpl-cached-{marker}", f"chatcmpl-other-{marker}"))
|
||||
with wire_server(lambda _: Reply(body=_completion(next(identities), "Done."))) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
|
||||
def ask(reasoning: str) -> dict[str, JsonValue]:
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": _replayed_conversation(reasoning, marker)},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
return _JSON_OBJECT.validate_json(response.content)
|
||||
|
||||
assert ask(_REASONING)["id"] == f"chatcmpl-cached-{marker}"
|
||||
assert ask(_REASONING)["id"] == f"chatcmpl-cached-{marker}"
|
||||
assert ask(f"a different thought {marker}")["id"] == f"chatcmpl-other-{marker}"
|
||||
assert [_sent_messages(request)[1].get("reasoning_content") for request in wire.drain()] == [
|
||||
_REASONING,
|
||||
f"a different thought {marker}",
|
||||
]
|
||||
|
||||
|
||||
def _responses_input(marker: str) -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
{"role": "user", "content": f"Compare these invoices {marker}."},
|
||||
{
|
||||
"id": f"rs_{marker}",
|
||||
"type": "reasoning",
|
||||
"summary": [{"type": "summary_text", "text": _REASONING}],
|
||||
},
|
||||
{
|
||||
"id": f"msg_prior_{marker}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": "Checking the totals.", "annotations": []}],
|
||||
},
|
||||
{"role": "user", "content": "What is the difference?"},
|
||||
]
|
||||
|
||||
|
||||
def _responses_reply(identity: str, stream: bool) -> Reply:
|
||||
response: Final = {
|
||||
"id": identity,
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": _BACKEND,
|
||||
"output": [
|
||||
{
|
||||
"id": "msg_" + identity,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": "They differ by 42.", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 30, "output_tokens": 5, "total_tokens": 35},
|
||||
}
|
||||
if not stream:
|
||||
return Reply(body=json.dumps(response).encode())
|
||||
events: Final = (
|
||||
{
|
||||
"type": "response.created",
|
||||
"sequence_number": 0,
|
||||
"response": {**response, "status": "in_progress", "output": []},
|
||||
},
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"sequence_number": 1,
|
||||
"item_id": "msg_" + identity,
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": "They differ by 42.",
|
||||
},
|
||||
{"type": "response.completed", "sequence_number": 2, "response": response},
|
||||
)
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events),
|
||||
)
|
||||
|
||||
|
||||
def _only_responses_body(wire: Wire) -> dict[str, JsonValue]:
|
||||
received: Final = wire.drain()
|
||||
assert [(request.method, request.target) for request in received] == [("POST", "/v1/responses")]
|
||||
return _JSON_OBJECT.validate_json(received[0].body)
|
||||
|
||||
|
||||
def test_openai_sdk_responses_replay_reaches_hosted_vllm_with_its_reasoning_item(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(lambda _: _responses_reply(f"resp_upstream_{marker}", stream=False)) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
response: Final = _openai_client(gateway).responses.create(
|
||||
model=model,
|
||||
input=_responses_input(marker), # pyright: ignore[reportArgumentType] # plain JSON input items
|
||||
)
|
||||
assert response.output_text == "They differ by 42."
|
||||
assert _only_responses_body(wire)["input"] == _responses_input(marker)
|
||||
assert _spend_row(response.id) == {
|
||||
"model_group": model,
|
||||
"status": "success",
|
||||
"prompt_tokens": 30,
|
||||
"completion_tokens": 5,
|
||||
}
|
||||
|
||||
|
||||
async def test_async_openai_sdk_responses_stream_reaches_hosted_vllm_with_its_reasoning_item(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(lambda _: _responses_reply(f"resp_upstream_{marker}", stream=True)) as wire:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY)
|
||||
stream: Final = await _async_openai_client(gateway).responses.create(
|
||||
model=model,
|
||||
input=_responses_input(marker), # pyright: ignore[reportArgumentType] # plain JSON input items
|
||||
stream=True,
|
||||
)
|
||||
events: Final = [event async for event in stream]
|
||||
assert [event.type for event in events] == [
|
||||
"response.created",
|
||||
"response.output_text.delta",
|
||||
"response.completed",
|
||||
]
|
||||
completed: Final = events[-1]
|
||||
assert completed.type == "response.completed"
|
||||
body: Final = _only_responses_body(wire)
|
||||
assert body["stream"] is True
|
||||
assert body["input"] == _responses_input(marker)
|
||||
assert completed.response.output_text == "They differ by 42."
|
||||
assert _model_spend_statuses(model) == ["success"]
|
||||
|
|
@ -1,6 +1,8 @@
|
|||
import json
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
from litellm.constants import (
|
||||
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
|
||||
|
|
@ -112,10 +114,10 @@ def test_hosted_vllm_supports_thinking():
|
|||
assert optional_params["reasoning_effort"] == "low"
|
||||
|
||||
|
||||
def test_hosted_vllm_thinking_blocks_prepended_to_assistant_content():
|
||||
def test_hosted_vllm_reasoning_content_kept_and_thinking_blocks_removed():
|
||||
"""
|
||||
Test that thinking_blocks on assistant messages are removed and content
|
||||
stays a string for vLLM compatibility.
|
||||
Test that reasoning_content on assistant messages is forwarded to vLLM
|
||||
while thinking_blocks are removed and content stays a string.
|
||||
"""
|
||||
config = HostedVLLMChatConfig()
|
||||
messages = [
|
||||
|
|
@ -152,7 +154,36 @@ def test_hosted_vllm_thinking_blocks_prepended_to_assistant_content():
|
|||
assert isinstance(assistant_msg["content"], str)
|
||||
assert assistant_msg["content"] == "Here is my answer."
|
||||
assert "thinking_blocks" not in assistant_msg
|
||||
assert "reasoning_content" not in assistant_msg
|
||||
assert assistant_msg["reasoning_content"] == "Let me reason about this..."
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("reasoning_content", "expected"),
|
||||
[
|
||||
("step one, then step two", "step one, then step two"),
|
||||
("", ""),
|
||||
(None, "absent"),
|
||||
(42, "absent"),
|
||||
(["step one", "step two"], "absent"),
|
||||
({"text": "step one"}, "absent"),
|
||||
],
|
||||
)
|
||||
def test_hosted_vllm_forwards_only_string_reasoning_content(reasoning_content, expected):
|
||||
config = HostedVLLMChatConfig()
|
||||
transformed = config.transform_request(
|
||||
model="hosted_vllm/qwen3",
|
||||
messages=[
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "Hi", "reasoning_content": reasoning_content},
|
||||
{"role": "user", "content": "Again"},
|
||||
],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assistant_msg = transformed["messages"][1]
|
||||
assert assistant_msg.get("reasoning_content", "absent") == expected
|
||||
assert assistant_msg["content"] == "Hi"
|
||||
|
||||
|
||||
def test_hosted_vllm_thinking_blocks_with_list_content():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue