From 14a131ad5b43bbb9dd7c5e91968e21e0c6432e4c Mon Sep 17 00:00:00 2001 From: Shangming Cai Date: Sun, 13 Sep 2026 21:20:47 +0800 Subject: [PATCH] [PD][OpenAI] Gate /v1/responses persistence behind --enable-response-store, default off (#39122) Co-authored-by: Xinyuan Tong --- .../advanced_features/server_arguments.mdx | 10 + .../sglang/srt/arg_groups/fields/serving.py | 6 + python/sglang/srt/arg_groups/pipeline.py | 2 + .../sglang/srt/arg_groups/resolution_hooks.py | 1 + .../sglang/srt/arg_groups/validation_hook.py | 10 + .../entrypoints/openai/serving_responses.py | 54 ++- .../openai/test_serving_responses.py | 436 +++++++++++++++++- .../openai/test_serving_responses_stream.py | 10 + 8 files changed, 519 insertions(+), 10 deletions(-) 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