[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>
</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
<table style={{width: "100%", borderCollapse: "collapse", tableLayout: "fixed"}}>
<colgroup>
@@ -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,
+2
View File
@@ -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
@@ -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",
@@ -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
@@ -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)
@@ -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__]))
@@ -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