diff --git a/docs/docs/advanced_features/server_arguments.mdx b/docs/docs/advanced_features/server_arguments.mdx
index 5837c0f10..36f76a01f 100644
--- a/docs/docs/advanced_features/server_arguments.mdx
+++ b/docs/docs/advanced_features/server_arguments.mdx
@@ -1063,6 +1063,16 @@ Please consult the documentation below and [server_args.py](https://github.com/s
+## 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
diff --git a/python/sglang/srt/arg_groups/fields/serving.py b/python/sglang/srt/arg_groups/fields/serving.py
index d7a2ec0c7..a85a3220f 100644
--- a/python/sglang/srt/arg_groups/fields/serving.py
+++ b/python/sglang/srt/arg_groups/fields/serving.py
@@ -30,6 +30,12 @@ class Serving(msgspec.Struct):
"""Namespace ``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_mode: A[
str,
diff --git a/python/sglang/srt/arg_groups/pipeline.py b/python/sglang/srt/arg_groups/pipeline.py
index c063b5187..9b982c232 100644
--- a/python/sglang/srt/arg_groups/pipeline.py
+++ b/python/sglang/srt/arg_groups/pipeline.py
@@ -101,10 +101,12 @@ def run_resolution_pipeline(server_args: Any) -> None:
default_unset_prefill_decode_interval,
validate_experimental_sgl_marlin,
validate_prefill_decode_interval,
+ validate_response_store,
validate_sampling_mask_max_tokens,
)
run_hook(validate_prefill_decode_interval, server_args)
+ run_hook(validate_response_store, server_args)
run_hook(validate_sampling_mask_max_tokens, server_args)
# Reject an explicitly enabled but incompatible hardware runtime before
diff --git a/python/sglang/srt/arg_groups/resolution_hooks.py b/python/sglang/srt/arg_groups/resolution_hooks.py
index c8719e11f..674c4832f 100644
--- a/python/sglang/srt/arg_groups/resolution_hooks.py
+++ b/python/sglang/srt/arg_groups/resolution_hooks.py
@@ -49,6 +49,7 @@ _OVERRIDABLE_HOOKS: FrozenSet[str] = frozenset(
"handle_offload_compatibility",
"validate_prefill_decode_interval",
"default_unset_prefill_decode_interval",
+ "validate_response_store",
"validate_sampling_mask_max_tokens",
"validate_prefill_cp_platform",
"handle_hardware_runtime_validation",
diff --git a/python/sglang/srt/arg_groups/validation_hook.py b/python/sglang/srt/arg_groups/validation_hook.py
index 38fab7c6a..ec47d353a 100644
--- a/python/sglang/srt/arg_groups/validation_hook.py
+++ b/python/sglang/srt/arg_groups/validation_hook.py
@@ -25,6 +25,16 @@ from sglang.srt.utils.runai_utils import is_runai_obj_uri
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):
from sglang.srt.arg_groups.lora_hook import check_lora_server_args
diff --git a/python/sglang/srt/entrypoints/openai/serving_responses.py b/python/sglang/srt/entrypoints/openai/serving_responses.py
index d82033393..9fdcbe393 100644
--- a/python/sglang/srt/entrypoints/openai/serving_responses.py
+++ b/python/sglang/srt/entrypoints/openai/serving_responses.py
@@ -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.managers.io_struct import GenerateReqInput
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 (
set_request_reasoning_end_token_ids,
)
@@ -204,6 +204,8 @@ class OpenAIServingResponses(OpenAIServingChat):
self.msg_store: dict[str, Union[list[dict], list[OpenAIMessage]]] = {}
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
def _has_response_tool(request: ResponsesRequest, *tool_types: str) -> bool:
@@ -245,6 +247,14 @@ class OpenAIServingResponses(OpenAIServingChat):
def _request_id_prefix(self) -> str:
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]:
"""Model ids a caller may address, mirroring what ``/v1/models`` lists."""
names = {self.tokenizer_manager.served_model_name}
@@ -281,6 +291,15 @@ class OpenAIServingResponses(OpenAIServingChat):
if not self.tokenizer_manager:
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)
if model_error is not None:
return model_error
@@ -337,6 +356,20 @@ class OpenAIServingResponses(OpenAIServingChat):
"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
prev_response_id = request.previous_response_id
if prev_response_id is not None:
@@ -547,14 +580,15 @@ class OpenAIServingResponses(OpenAIServingChat):
(result_generator,) = generators
# Store the input messages
- if request.store:
+ persist = self.enable_response_store and bool(request.store)
+ if persist:
self.msg_store[request.request_id] = (
messages[2:]
if self.use_harmony
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())
response = ResponsesResponse.from_request(
request,
@@ -832,7 +866,7 @@ class OpenAIServingResponses(OpenAIServingChat):
)
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:
stored_response = self.response_store.get(response.id)
# If the response is already cancelled, don't update it
@@ -1552,6 +1586,8 @@ class OpenAIServingResponses(OpenAIServingChat):
self,
response_id: str,
) -> Union[ResponsesResponse, ORJSONResponse]:
+ if not self.enable_response_store:
+ return self._response_store_disabled_error("response_id")
if not response_id.startswith("resp_"):
return self._make_invalid_id_error(response_id)
@@ -1566,6 +1602,8 @@ class OpenAIServingResponses(OpenAIServingChat):
self,
response_id: str,
) -> Union[ResponsesResponse, ORJSONResponse]:
+ if not self.enable_response_store:
+ return self._response_store_disabled_error("response_id")
if not response_id.startswith("resp_"):
return self._make_invalid_id_error(response_id)
@@ -2601,7 +2639,7 @@ class OpenAIServingResponses(OpenAIServingChat):
usage=usage,
)
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:
stored = self.response_store.get(final_response.id)
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.
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.
tool_output = await context.call_tool()
context.append_output(tool_output)
diff --git a/test/registered/unit/entrypoints/openai/test_serving_responses.py b/test/registered/unit/entrypoints/openai/test_serving_responses.py
index e60f3193a..9b7e425bd 100644
--- a/test/registered/unit/entrypoints/openai/test_serving_responses.py
+++ b/test/registered/unit/entrypoints/openai/test_serving_responses.py
@@ -1,20 +1,31 @@
+import argparse
import asyncio
+import sys
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 (
ResponseOutputMessage,
ResponseOutputText,
ResponseReasoningItem,
)
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 (
MessageProcessingResult,
RequestResponseMetadata,
ResponsesRequest,
+ ResponsesResponse,
)
from sglang.srt.entrypoints.openai.serving_responses import (
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.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 (
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):
+ publish(
+ ServerArgs(model_path="dummy", enable_response_store=True), role="tokenizer"
+ )
for stream in (False, True):
with self.subTest(stream=stream):
serving = make_serving()
@@ -1263,6 +1277,9 @@ class EnginePassthroughTestCase(CustomTestCase):
class CancelIdempotencyTestCase(CustomTestCase):
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
for status in ("cancelled", "completed"):
@@ -1302,5 +1319,414 @@ class StreamingLogprobsRejectionTestCase(CustomTestCase):
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__":
- unittest.main()
+ sys.exit(pytest.main([__file__]))
diff --git a/test/registered/unit/entrypoints/openai/test_serving_responses_stream.py b/test/registered/unit/entrypoints/openai/test_serving_responses_stream.py
index 3fdbccad6..c9d6974cc 100644
--- a/test/registered/unit/entrypoints/openai/test_serving_responses_stream.py
+++ b/test/registered/unit/entrypoints/openai/test_serving_responses_stream.py
@@ -115,6 +115,11 @@ class NonHarmonyStreamTestCase(CustomTestCase):
)
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()
for finish_reason, status in (
({"type": "length"}, "incomplete"),
@@ -418,6 +423,11 @@ class HarmonyStreamLifecycleTestCase(CustomTestCase):
self.assertEqual(added["call_id"], done["call_id"])
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 sglang.srt.entrypoints.context import StreamingHarmonyContext