[PD][OpenAI] Gate /v1/responses persistence behind --enable-response-store, default off (#39122)

Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
Shangming Cai
2026-09-13 21:20:47 +08:00
committed by GitHub
co-authored by Xinyuan Tong
parent f9fca05803
commit 14a131ad5b
8 changed files with 519 additions and 10 deletions
@@ -1063,6 +1063,16 @@ Please consult the documentation below and [server_args.py](https://github.com/s
</tbody> </tbody>
</table> </table>
## Responses storage
Responses storage is now disabled by default, including on standalone servers. Add `--enable-response-store` to retain responses for retrieval, `previous_response_id` chaining, and background requests. Storage is process-local and in-memory, has no TTL or size limit, and is lost on restart.
Foreground generation and streaming still accept the API's default `store: true` without retaining responses or message history when server storage is disabled. The response's `store` field echoes the client request; the server flag controls persistence. With storage enabled, `store: false` opts out for the current request while allowing it to read a stored predecessor.
Background requests, including streaming requests, require both `--enable-response-store` and `store: true`. Without the flag, chaining, background requests, retrieval, and cancellation return HTTP 400. Cancellation applies to detached background requests; active streams have no stored response until completion.
Combining `--enable-response-store` with `--disaggregation-mode=prefill` or `decode` fails at startup. PD clients must send explicit conversation history in foreground requests. Built-in web search and code interpreter calls are also unsupported under PD.
## API related ## API related
<table style={{width: "100%", borderCollapse: "collapse", tableLayout: "fixed"}}> <table style={{width: "100%", borderCollapse: "collapse", tableLayout: "fixed"}}>
<colgroup> <colgroup>
@@ -30,6 +30,12 @@ class Serving(msgspec.Struct):
"""Namespace ``serving``.""" """Namespace ``serving``."""
_NS_PATH = "serving" _NS_PATH = "serving"
enable_response_store: A[
bool,
"Enable in-memory Responses storage for retrieval, chaining, and background "
"requests. Disabled by default; unsupported with prefill-decode "
"disaggregation. Storage has no TTL or size limit.",
] = False
tokenizer_path: A[Optional[str], "The path of the tokenizer."] = None tokenizer_path: A[Optional[str], "The path of the tokenizer."] = None
tokenizer_mode: A[ tokenizer_mode: A[
str, str,
+2
View File
@@ -101,10 +101,12 @@ def run_resolution_pipeline(server_args: Any) -> None:
default_unset_prefill_decode_interval, default_unset_prefill_decode_interval,
validate_experimental_sgl_marlin, validate_experimental_sgl_marlin,
validate_prefill_decode_interval, validate_prefill_decode_interval,
validate_response_store,
validate_sampling_mask_max_tokens, validate_sampling_mask_max_tokens,
) )
run_hook(validate_prefill_decode_interval, server_args) run_hook(validate_prefill_decode_interval, server_args)
run_hook(validate_response_store, server_args)
run_hook(validate_sampling_mask_max_tokens, server_args) run_hook(validate_sampling_mask_max_tokens, server_args)
# Reject an explicitly enabled but incompatible hardware runtime before # Reject an explicitly enabled but incompatible hardware runtime before
@@ -49,6 +49,7 @@ _OVERRIDABLE_HOOKS: FrozenSet[str] = frozenset(
"handle_offload_compatibility", "handle_offload_compatibility",
"validate_prefill_decode_interval", "validate_prefill_decode_interval",
"default_unset_prefill_decode_interval", "default_unset_prefill_decode_interval",
"validate_response_store",
"validate_sampling_mask_max_tokens", "validate_sampling_mask_max_tokens",
"validate_prefill_cp_platform", "validate_prefill_cp_platform",
"handle_hardware_runtime_validation", "handle_hardware_runtime_validation",
@@ -25,6 +25,16 @@ from sglang.srt.utils.runai_utils import is_runai_obj_uri
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def validate_response_store(server_args: Any) -> None:
cfg = resolving_view(server_args)
if cfg.enable_response_store and cfg.disaggregation_mode != "null":
raise ValueError(
"--enable-response-store is not supported with "
"--disaggregation-mode=prefill or decode; response storage must "
"remain disabled in PD mode."
)
def check_server_args(server_args: Any): def check_server_args(server_args: Any):
from sglang.srt.arg_groups.lora_hook import check_lora_server_args from sglang.srt.arg_groups.lora_hook import check_lora_server_args
@@ -85,7 +85,7 @@ from sglang.srt.function_call.function_call_parser import FunctionCallParser
from sglang.srt.function_call.json_array_parser import JsonArrayParser from sglang.srt.function_call.json_array_parser import JsonArrayParser
from sglang.srt.managers.io_struct import GenerateReqInput from sglang.srt.managers.io_struct import GenerateReqInput
from sglang.srt.parser.reasoning_parser import ReasoningParser from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.runtime_context import get_serving from sglang.srt.runtime_context import get_disagg, get_serving
from sglang.srt.sampling.sampling_params import ( from sglang.srt.sampling.sampling_params import (
set_request_reasoning_end_token_ids, set_request_reasoning_end_token_ids,
) )
@@ -204,6 +204,8 @@ class OpenAIServingResponses(OpenAIServingChat):
self.msg_store: dict[str, Union[list[dict], list[OpenAIMessage]]] = {} self.msg_store: dict[str, Union[list[dict], list[OpenAIMessage]]] = {}
self.background_tasks: dict[str, asyncio.Task] = {} self.background_tasks: dict[str, asyncio.Task] = {}
self.enable_response_store = get_serving().enable_response_store
self.is_disaggregated = get_disagg().disaggregation_mode != "null"
@staticmethod @staticmethod
def _has_response_tool(request: ResponsesRequest, *tool_types: str) -> bool: def _has_response_tool(request: ResponsesRequest, *tool_types: str) -> bool:
@@ -245,6 +247,14 @@ class OpenAIServingResponses(OpenAIServingChat):
def _request_id_prefix(self) -> str: def _request_id_prefix(self) -> str:
return "resp_" return "resp_"
def _response_store_disabled_error(self, param: str) -> ORJSONResponse:
return self.create_error_response(
"Response store is disabled. Stateful Responses require "
"--enable-response-store on a standalone server; response storage "
"is unavailable in PD mode.",
param=param,
)
def _known_model_names(self) -> set[str]: def _known_model_names(self) -> set[str]:
"""Model ids a caller may address, mirroring what ``/v1/models`` lists.""" """Model ids a caller may address, mirroring what ``/v1/models`` lists."""
names = {self.tokenizer_manager.served_model_name} names = {self.tokenizer_manager.served_model_name}
@@ -281,6 +291,15 @@ class OpenAIServingResponses(OpenAIServingChat):
if not self.tokenizer_manager: if not self.tokenizer_manager:
return self.create_error_response("Model not loaded") return self.create_error_response("Model not loaded")
if not self.enable_response_store and request.previous_response_id is not None:
return self._response_store_disabled_error("previous_response_id")
if not self.enable_response_store and request.background:
return self._response_store_disabled_error("background")
if request.background and not request.store:
return self.create_error_response(
"background=true requires store=true.", param="store"
)
model_error = self._validate_model(request.model) model_error = self._validate_model(request.model)
if model_error is not None: if model_error is not None:
return model_error return model_error
@@ -337,6 +356,20 @@ class OpenAIServingResponses(OpenAIServingChat):
"https://dashboard.exa.ai/api-keys." "https://dashboard.exa.ai/api-keys."
) )
if (
self.use_harmony
and self.tool_server is not None
and self.is_disaggregated
and self._has_response_tool(
request, "web_search", "web_search_preview", "code_interpreter"
)
):
return self.create_error_response(
"built-in tools (web_search, code_interpreter) are not supported "
"with prefill-decode disaggregation",
param="tools",
)
# Handle the previous response ID # Handle the previous response ID
prev_response_id = request.previous_response_id prev_response_id = request.previous_response_id
if prev_response_id is not None: if prev_response_id is not None:
@@ -547,14 +580,15 @@ class OpenAIServingResponses(OpenAIServingChat):
(result_generator,) = generators (result_generator,) = generators
# Store the input messages # Store the input messages
if request.store: persist = self.enable_response_store and bool(request.store)
if persist:
self.msg_store[request.request_id] = ( self.msg_store[request.request_id] = (
messages[2:] messages[2:]
if self.use_harmony if self.use_harmony
else self._response_input_history(request) else self._response_input_history(request)
) )
if request.background and not request.stream: if request.background and not request.stream and persist:
created_time = int(time.time()) created_time = int(time.time())
response = ResponsesResponse.from_request( response = ResponsesResponse.from_request(
request, request,
@@ -832,7 +866,7 @@ class OpenAIServingResponses(OpenAIServingChat):
) )
response.error = self._error_from_finish_reason(finish_reason) response.error = self._error_from_finish_reason(finish_reason)
if request.store: if self.enable_response_store and request.store:
async with self.response_store_lock: async with self.response_store_lock:
stored_response = self.response_store.get(response.id) stored_response = self.response_store.get(response.id)
# If the response is already cancelled, don't update it # If the response is already cancelled, don't update it
@@ -1552,6 +1586,8 @@ class OpenAIServingResponses(OpenAIServingChat):
self, self,
response_id: str, response_id: str,
) -> Union[ResponsesResponse, ORJSONResponse]: ) -> Union[ResponsesResponse, ORJSONResponse]:
if not self.enable_response_store:
return self._response_store_disabled_error("response_id")
if not response_id.startswith("resp_"): if not response_id.startswith("resp_"):
return self._make_invalid_id_error(response_id) return self._make_invalid_id_error(response_id)
@@ -1566,6 +1602,8 @@ class OpenAIServingResponses(OpenAIServingChat):
self, self,
response_id: str, response_id: str,
) -> Union[ResponsesResponse, ORJSONResponse]: ) -> Union[ResponsesResponse, ORJSONResponse]:
if not self.enable_response_store:
return self._response_store_disabled_error("response_id")
if not response_id.startswith("resp_"): if not response_id.startswith("resp_"):
return self._make_invalid_id_error(response_id) return self._make_invalid_id_error(response_id)
@@ -2601,7 +2639,7 @@ class OpenAIServingResponses(OpenAIServingChat):
usage=usage, usage=usage,
) )
final_response.error = self._error_from_finish_reason(finish_reason) final_response.error = self._error_from_finish_reason(finish_reason)
if request.store: if self.enable_response_store and request.store:
async with self.response_store_lock: async with self.response_store_lock:
stored = self.response_store.get(final_response.id) stored = self.response_store.get(final_response.id)
if stored is None or stored.status != "cancelled": if stored is None or stored.status != "cancelled":
@@ -2646,6 +2684,12 @@ class OpenAIServingResponses(OpenAIServingChat):
# The model did not ask for a tool call, so we're done. # The model did not ask for a tool call, so we're done.
break break
if self.is_disaggregated:
raise ValueError(
"built-in tool calls are not supported with prefill-decode "
"disaggregation"
)
# Call the tool and update the context with the result. # Call the tool and update the context with the result.
tool_output = await context.call_tool() tool_output = await context.call_tool()
context.append_output(tool_output) context.append_output(tool_output)
@@ -1,20 +1,31 @@
import argparse
import asyncio import asyncio
import sys
import unittest import unittest
from unittest.mock import Mock, patch from contextlib import asynccontextmanager
from unittest.mock import AsyncMock, Mock, patch
import orjson
import pytest
from openai.types.responses import ( from openai.types.responses import (
ResponseOutputMessage, ResponseOutputMessage,
ResponseOutputText, ResponseOutputText,
ResponseReasoningItem, ResponseReasoningItem,
) )
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
from utils import StreamFixture, engine_chunk, make_serving from openai_harmony import Conversation, Message, Role, ToolNamespaceConfig
from utils import StreamFixture, engine_chunk, event_payloads, make_serving
from sglang.srt.entrypoints.context import SimpleContext from sglang.srt.entrypoints.context import (
HarmonyContext,
SimpleContext,
)
from sglang.srt.entrypoints.harmony_utils import get_encoding
from sglang.srt.entrypoints.openai.protocol import ( from sglang.srt.entrypoints.openai.protocol import (
MessageProcessingResult, MessageProcessingResult,
RequestResponseMetadata, RequestResponseMetadata,
ResponsesRequest, ResponsesRequest,
ResponsesResponse,
) )
from sglang.srt.entrypoints.openai.serving_responses import ( from sglang.srt.entrypoints.openai.serving_responses import (
OpenAIServingResponses, OpenAIServingResponses,
@@ -23,7 +34,7 @@ from sglang.srt.entrypoints.openai.serving_responses import (
) )
from sglang.srt.function_call.core_types import ToolCallItem from sglang.srt.function_call.core_types import ToolCallItem
from sglang.srt.parser.template_detection import ReasoningToggleConfig from sglang.srt.parser.template_detection import ReasoningToggleConfig
from sglang.srt.runtime_context import publish, reset_context from sglang.srt.runtime_context import get_serving, publish, reset_context
from sglang.srt.sampling.sampling_params import ( from sglang.srt.sampling.sampling_params import (
REQUEST_REASONING_END_TOKEN_IDS_KEY, REQUEST_REASONING_END_TOKEN_IDS_KEY,
) )
@@ -95,6 +106,9 @@ class InputMessageConstructionTestCase(CustomTestCase):
) )
def test_stored_tool_turn_matches_client_replay_without_old_instructions(self): def test_stored_tool_turn_matches_client_replay_without_old_instructions(self):
publish(
ServerArgs(model_path="dummy", enable_response_store=True), role="tokenizer"
)
for stream in (False, True): for stream in (False, True):
with self.subTest(stream=stream): with self.subTest(stream=stream):
serving = make_serving() serving = make_serving()
@@ -1263,6 +1277,9 @@ class EnginePassthroughTestCase(CustomTestCase):
class CancelIdempotencyTestCase(CustomTestCase): class CancelIdempotencyTestCase(CustomTestCase):
def test_cancelling_a_terminal_response_returns_it_not_an_error(self): def test_cancelling_a_terminal_response_returns_it_not_an_error(self):
publish(
ServerArgs(model_path="dummy", enable_response_store=True), role="tokenizer"
)
from sglang.srt.entrypoints.openai.protocol import ResponsesResponse from sglang.srt.entrypoints.openai.protocol import ResponsesResponse
for status in ("cancelled", "completed"): for status in ("cancelled", "completed"):
@@ -1302,5 +1319,414 @@ class StreamingLogprobsRejectionTestCase(CustomTestCase):
self.assertIn("streaming mode", body["error"]["message"]) self.assertIn("streaming mode", body["error"]["message"])
STORE_DISABLED_MESSAGE = (
"Response store is disabled. Stateful Responses require "
"--enable-response-store on a standalone server; response storage "
"is unavailable in PD mode."
)
STORE_PD_MESSAGE = (
"--enable-response-store is not supported with "
"--disaggregation-mode=prefill or decode; response storage must "
"remain disabled in PD mode."
)
@pytest.fixture(autouse=True)
def isolated_response_config():
reset_context()
yield
reset_context()
@pytest.fixture
def response_serving():
def build(enabled=False, mode="null", harmony=False):
reset_context()
publish(
ServerArgs(
model_path="dummy",
enable_response_store=enabled,
disaggregation_mode=mode,
),
role="tokenizer",
)
serving = make_serving()
serving.use_harmony = harmony
serving.default_chat_template_kwargs = {}
serving.template_manager.chat_template_name = None
serving.template_manager.jinja_template_content_format = "string"
serving.tokenizer_manager.tokenizer.apply_chat_template.return_value = [1, 2, 3]
serving.tokenizer_manager.abort_request = Mock()
serving.reasoning_parser = None
serving.tool_call_parser = None
async def generate(*args, **kwargs):
chunk = engine_chunk("ok", finish=True)
if harmony:
chunk["output_ids"] = get_encoding().render_conversation(
Conversation.from_messages(
[
Message.from_role_and_content(
Role.ASSISTANT, "ok"
).with_channel("final")
]
)
)
chunk["meta_info"]["completion_tokens"] = len(chunk["output_ids"])
yield chunk
serving.tokenizer_manager.generate_request = Mock(side_effect=generate)
return serving
return build
async def create_response_result(serving, request):
result = await serving.create_responses(request)
if request.stream:
payloads = event_payloads([event async for event in result])
assert payloads[0]["type"] == "response.created"
assert payloads[-1]["type"] == "response.completed"
return ResponsesResponse.model_validate(payloads[-1]["response"])
return result
def assert_response_error(response, param, message=STORE_DISABLED_MESSAGE, status=400):
assert response.status_code == status
assert orjson.loads(response.body) == {
"error": {
"message": message,
"type": "invalid_request_error",
"param": param,
"code": status,
}
}
@pytest.mark.parametrize(
"mode,enabled", [("null", False), ("null", True), ("decode", True)]
)
def test_response_store_configuration(mode, enabled):
parser = argparse.ArgumentParser()
ServerArgs.add_cli_args(parser)
argv = ["--model-path", "dummy", "--disaggregation-mode", mode]
if enabled:
argv.append("--enable-response-store")
args = ServerArgs.from_cli_args(parser.parse_args(argv))
if mode != "null":
with pytest.raises(ValueError) as error:
publish(args, role="tokenizer")
assert str(error.value) == STORE_PD_MESSAGE
else:
publish(args, role="tokenizer")
serving = make_serving()
assert get_serving().enable_response_store is enabled
assert serving.enable_response_store is enabled
assert not serving.is_disaggregated
@pytest.mark.parametrize("stream", [False, True])
@pytest.mark.parametrize("enabled", [False, True])
def test_response_store_persistence_and_continuation(response_serving, stream, enabled):
serving = response_serving(enabled=enabled, harmony=not stream)
async def run():
first = await create_response_result(
serving, ResponsesRequest(model="x", input="first", stream=stream)
)
assert first.status == "completed"
assert first.output[0].content[0].text == "ok"
assert first.store is True
assert not serving.background_tasks
if not enabled:
assert not serving.response_store and not serving.msg_store
return
assert serving.response_store[first.id].output == first.output
history = list(serving.msg_store[first.id])
assert len(history) == 2
second = await create_response_result(
serving,
ResponsesRequest(
model="x",
input="next",
previous_response_id=first.id,
store=False,
stream=stream,
),
)
if serving.use_harmony:
generated_request = (
serving.tokenizer_manager.generate_request.call_args.args[0]
)
prompt = get_encoding().decode(generated_request.input_ids)
else:
prompt = str(
serving.tokenizer_manager.tokenizer.apply_chat_template.call_args.args[
0
]
)
assert all(text in prompt for text in ("first", "ok", "next"))
assert second.status == "completed"
assert second.output[0].content[0].text == "ok"
assert second.store is False
assert set(serving.response_store) == {first.id}
assert set(serving.msg_store) == {first.id}
assert serving.msg_store[first.id] == history
asyncio.run(run())
@pytest.mark.parametrize(
"enabled,fields,param,message",
[
(
False,
{"previous_response_id": ""},
"previous_response_id",
STORE_DISABLED_MESSAGE,
),
(False, {}, "background", STORE_DISABLED_MESSAGE),
(True, {}, "store", "background=true requires store=true."),
],
ids=["empty-predecessor", "disabled-background", "background-without-store"],
)
def test_response_store_admission(response_serving, enabled, fields, param, message):
serving = response_serving(enabled=enabled, mode="null" if enabled else "prefill")
serving._make_request = AsyncMock(side_effect=AssertionError("prompt reached"))
serving.response_store = Mock()
serving.msg_store = Mock()
request = ResponsesRequest(
model="unknown", input="hi", background=True, stream=True, store=None, **fields
)
assert_response_error(
asyncio.run(serving.create_responses(request)), param, message
)
serving._make_request.assert_not_called()
serving.tokenizer_manager.generate_request.assert_not_called()
assert not serving.response_store.mock_calls
assert not serving.msg_store.mock_calls
assert not serving.background_tasks
@pytest.mark.parametrize("enabled", [False, True])
def test_response_store_read_endpoints(response_serving, enabled):
serving = response_serving(enabled=enabled)
async def run():
if not enabled:
serving.response_store = Mock()
serving.background_tasks = Mock()
for method in (serving.retrieve_responses, serving.cancel_responses):
assert_response_error(await method("bad"), "response_id")
assert not serving.response_store.mock_calls
assert not serving.background_tasks.mock_calls
else:
response = await create_response_result(
serving, ResponsesRequest(model="x", input="hi")
)
assert await serving.retrieve_responses(response.id) is response
for response_id, status in (("bad", 400), ("resp_missing", 404)):
for method in (serving.retrieve_responses, serving.cancel_responses):
assert (await method(response_id)).status_code == status
request = ResponsesRequest(
model="x", input="hi", previous_response_id=response_id
)
assert (await serving.create_responses(request)).status_code == status
serving.tokenizer_manager.abort_request.assert_not_called()
asyncio.run(run())
@pytest.mark.parametrize(
"outcome", ["success", "failure", "cancel_queued", "cancel_running"]
)
def test_background_task_lifecycle(response_serving, outcome):
serving = response_serving(enabled=True)
original_generate = serving.tokenizer_manager.generate_request
async def run():
entered, release = asyncio.Event(), asyncio.Event()
async def generate(*args, **kwargs):
entered.set()
await release.wait()
if outcome == "failure":
raise ValueError("generation failed")
async for chunk in original_generate(*args, **kwargs):
yield chunk
serving.tokenizer_manager.generate_request = generate
request = ResponsesRequest(model="x", input="hi", background=True)
queued = await serving.create_responses(request)
assert queued.status == "queued"
assert serving.response_store[queued.id] is queued
assert request.request_id in serving.msg_store
task = serving.background_tasks[queued.id]
if outcome != "cancel_queued":
await entered.wait()
assert (await serving.retrieve_responses(queued.id)).status == "in_progress"
if outcome.startswith("cancel"):
cancelled = await serving.cancel_responses(queued.id)
assert cancelled.status == "cancelled"
serving.tokenizer_manager.abort_request.assert_called_once_with(
rid=queued.id
)
assert task.done()
if outcome == "cancel_queued":
assert task.cancelled()
else:
release.set()
await task
assert serving.response_store[queued.id].status == (
"failed" if outcome == "failure" else "completed"
)
assert not serving.background_tasks
asyncio.run(run())
def test_active_stream_cancel_and_final_history(response_serving):
serving = response_serving(enabled=True, harmony=True)
async def run():
request = ResponsesRequest(model="x", input="hi", background=True, stream=True)
stream = await serving.create_responses(request)
assert "response.created" in await anext(stream)
assert not serving.background_tasks
assert (await serving.cancel_responses(request.request_id)).status_code == 404
serving.tokenizer_manager.abort_request.assert_not_called()
events = [event async for event in stream]
assert event_payloads(events)[-1]["type"] == "response.completed"
assert serving.response_store[request.request_id].status == "completed"
assert len(serving.msg_store[request.request_id]) == 2
asyncio.run(run())
@pytest.mark.parametrize("stream", [False, True])
def test_completion_preserves_cancelled_response(response_serving, stream):
serving = response_serving(enabled=True, harmony=not stream)
request = ResponsesRequest(model="x", input="hi", stream=stream)
cancelled = ResponsesResponse.from_request(
request,
{},
model_name="x",
created_time=0,
output=[],
status="cancelled",
usage=None,
)
serving.response_store[request.request_id] = cancelled
serving.msg_store[request.request_id] = ["original"]
if stream:
StreamFixture(serving, request).run([engine_chunk("ok", finish=True)])
else:
messages = serving._construct_input_messages_with_harmony(request, None)
context = HarmonyContext(messages, {})
async def generate():
async for chunk in serving.tokenizer_manager.generate_request():
context.append_output(chunk)
yield context
response = asyncio.run(
serving.responses_full_generator(
request,
{},
generate(),
context,
"x",
Mock(),
RequestResponseMetadata(request_id=request.request_id),
require_reasoning=False,
)
)
assert response.status == "completed"
assert serving.response_store[request.request_id] is cancelled
assert serving.msg_store[request.request_id] == ["original"]
def test_pd_builtin_tool_admission(response_serving):
serving = response_serving(mode="prefill", harmony=True)
serving.tool_server = Mock()
serving.supports_browsing = True
request = ResponsesRequest(model="x", input="hi", tools=[{"type": "web_search"}])
result = asyncio.run(serving.create_responses(request))
assert result.status_code == 400
assert orjson.loads(result.body)["error"]["param"] == "tools"
serving.tool_server.get_tool_session.assert_not_called()
serving.tokenizer_manager.generate_request.assert_not_called()
def test_standalone_builtin_tools_without_storage(response_serving):
serving = response_serving(harmony=True)
tool_session = Mock()
tool_session.call_tool = AsyncMock(
return_value=Mock(content=[Mock(text="Search result: 42")])
)
@asynccontextmanager
async def session(name):
yield tool_session
turns = iter(
[
Message.from_role_and_content(Role.ASSISTANT, '{"query":"answer"}')
.with_channel("commentary")
.with_recipient("browser.search"),
Message.from_role_and_content(
Role.ASSISTANT, "The answer is 42."
).with_channel("final"),
]
)
async def generate(*args, **kwargs):
chunk = engine_chunk("", finish=True)
chunk["output_ids"] = get_encoding().render_conversation(
Conversation.from_messages([next(turns)])
)
if serving.tokenizer_manager.generate_request.call_count == 1:
# The parser starts inside the prompt's open assistant header.
chunk["output_ids"] = chunk["output_ids"][2:]
chunk["meta_info"]["completion_tokens"] = len(chunk["output_ids"])
yield chunk
serving.tokenizer_manager.generate_request = Mock(side_effect=generate)
serving.tool_server = Mock()
serving.tool_server.get_tool_session = session
serving.tool_server.get_tool_description.return_value = ToolNamespaceConfig(
name="browser", description="Browser", tools=[]
)
serving.supports_browsing = True
request = ResponsesRequest(model="x", input="hi", tools=[{"type": "web_search"}])
response = asyncio.run(create_response_result(serving, request))
assert isinstance(response, ResponsesResponse), response.body
tool_session.call_tool.assert_awaited_once_with("search", {"query": "answer"})
assert response.status == "completed"
assert response.output[-1].content[0].text == "The answer is 42."
assert serving.tokenizer_manager.generate_request.call_count == 2
continuation = serving.tokenizer_manager.generate_request.call_args_list[1].args[0]
assert "Search result: 42" in get_encoding().decode(continuation.input_ids)
assert not serving.msg_store and not serving.response_store
def test_pd_tool_continuation_stops_before_side_effect(response_serving):
serving = response_serving(mode="decode")
context = Mock()
context.need_builtin_tool_call.return_value = True
context.call_tool = AsyncMock()
async def run():
async for _ in serving._generate_with_builtin_tools(
"resp_tool", "hi", Mock(), {}, context
):
pass
with pytest.raises(ValueError, match="disaggregation"):
asyncio.run(run())
context.call_tool.assert_not_awaited()
assert serving.tokenizer_manager.generate_request.call_count == 1
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() sys.exit(pytest.main([__file__]))
@@ -115,6 +115,11 @@ class NonHarmonyStreamTestCase(CustomTestCase):
) )
def test_truncated_and_aborted_streams_have_matching_terminal_events(self): def test_truncated_and_aborted_streams_have_matching_terminal_events(self):
reset_context()
self.addCleanup(reset_context)
publish(
ServerArgs(model_path="dummy", enable_response_store=True), role="tokenizer"
)
serving = make_serving() serving = make_serving()
for finish_reason, status in ( for finish_reason, status in (
({"type": "length"}, "incomplete"), ({"type": "length"}, "incomplete"),
@@ -418,6 +423,11 @@ class HarmonyStreamLifecycleTestCase(CustomTestCase):
self.assertEqual(added["call_id"], done["call_id"]) self.assertEqual(added["call_id"], done["call_id"])
def test_split_and_coalesced_messages_preserve_stream_items(self): def test_split_and_coalesced_messages_preserve_stream_items(self):
reset_context()
self.addCleanup(reset_context)
publish(
ServerArgs(model_path="dummy", enable_response_store=True), role="tokenizer"
)
from openai_harmony import Message, Role, StreamState from openai_harmony import Message, Role, StreamState
from sglang.srt.entrypoints.context import StreamingHarmonyContext from sglang.srt.entrypoints.context import StreamingHarmonyContext