[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:
co-authored by
Xinyuan Tong
parent
f9fca05803
commit
14a131ad5b
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user