[Frontend] Apply request header overrides to chat completions (#35001)
Co-authored-by: Ye (Charlotte) Qi <ye.charlotte.qi@gmail.com>
This commit is contained in:
@@ -61,6 +61,7 @@ from sglang.srt.entrypoints.openai.utils import (
|
||||
should_include_usage,
|
||||
to_openai_style_logprobs,
|
||||
)
|
||||
from sglang.srt.entrypoints.request_headers import apply_header_overrides
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.function_call.core_types import ToolCallItem
|
||||
from sglang.srt.function_call.function_call_parser import FunctionCallParser
|
||||
@@ -1037,6 +1038,11 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
return_prompt_token_ids=request.return_prompt_token_ids
|
||||
or request.return_token_ids,
|
||||
)
|
||||
if (
|
||||
raw_request is not None
|
||||
and envs.SGLANG_ENABLE_REQUEST_HEADER_OVERRIDES.get()
|
||||
):
|
||||
apply_header_overrides(adapted_request, raw_request.headers)
|
||||
|
||||
return adapted_request, request
|
||||
|
||||
|
||||
@@ -32,6 +32,7 @@ from sglang.srt.entrypoints.openai.serving_chat import (
|
||||
OpenAIServingChat,
|
||||
normalize_tool_content,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.function_call.kimik3_format import TOOLS_CLOSE
|
||||
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||
from sglang.srt.parser.template_detection import ReasoningToggleConfig
|
||||
@@ -100,6 +101,7 @@ class _MockTokenizerManager:
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"cached_tokens": 0,
|
||||
"weight_version": "test-version",
|
||||
"finish_reason": {"type": "stop", "matched": None},
|
||||
"output_token_logprobs": [(0.1, 1, "Test"), (0.2, 2, "response")],
|
||||
"output_top_logprobs": None,
|
||||
@@ -109,6 +111,7 @@ class _MockTokenizerManager:
|
||||
|
||||
self.generate_request = Mock(return_value=_mock_generate())
|
||||
self.create_abort_task = Mock()
|
||||
self.request_logger = Mock(log_requests=False, log_requests_level=0)
|
||||
|
||||
def config_value(self, name: str):
|
||||
"""The manager's overlay accessor: no control-plane update recorded."""
|
||||
@@ -287,6 +290,53 @@ class ServingChatTestCase(unittest.TestCase):
|
||||
self.assertEqual(adapted.session_id, "session-1")
|
||||
self.assertEqual(processed, self.basic_req)
|
||||
|
||||
def test_chat_applies_pd_header_overrides(self):
|
||||
request = ChatCompletionRequest(
|
||||
model="x",
|
||||
messages=[{"role": "user", "content": "Hi?"}],
|
||||
rid="body-rid",
|
||||
routed_dp_rank=3,
|
||||
disagg_prefill_dp_rank=4,
|
||||
priority=5,
|
||||
)
|
||||
self.fastapi_request.headers = {
|
||||
"x-override-rid": "header-rid",
|
||||
"x-override-bootstrap-host": "header-host",
|
||||
"x-override-bootstrap-port": "8998",
|
||||
"x-override-bootstrap-room": "456",
|
||||
"x-override-conversation-id": "conversation-1",
|
||||
"x-override-routed-dp-rank": "6",
|
||||
"x-override-disagg-prefill-dp-rank": "7",
|
||||
"x-override-priority": "8",
|
||||
}
|
||||
body = request.model_dump()
|
||||
|
||||
processed_messages = MessageProcessingResult(
|
||||
"Test prompt", [1, 2, 3], None, None, [], [], None
|
||||
)
|
||||
with (
|
||||
envs.SGLANG_ENABLE_REQUEST_HEADER_OVERRIDES.override(True),
|
||||
patch.object(
|
||||
self.chat, "_process_messages", return_value=processed_messages
|
||||
),
|
||||
):
|
||||
response = get_or_create_event_loop().run_until_complete(
|
||||
self.chat.handle_request(request, self.fastapi_request)
|
||||
)
|
||||
|
||||
self.assertEqual(response.choices[0].message.content, "Test response")
|
||||
adapted_request = self.tm.generate_request.call_args.args[0]
|
||||
self.assertEqual(adapted_request.bootstrap_room, 456)
|
||||
self.assertEqual(adapted_request.bootstrap_host, "header-host")
|
||||
self.assertEqual(adapted_request.bootstrap_port, 8998)
|
||||
self.assertEqual(adapted_request.rid, "header-rid")
|
||||
self.assertEqual(adapted_request.conversation_id, "conversation-1")
|
||||
self.assertEqual(adapted_request.routed_dp_rank, 6)
|
||||
self.assertEqual(adapted_request.disagg_prefill_dp_rank, 7)
|
||||
self.assertEqual(adapted_request.priority, 8)
|
||||
self.assertEqual(request.model_dump(), body)
|
||||
self.assertFalse(hasattr(request, "conversation_id"))
|
||||
|
||||
def test_convert_to_internal_request_rejects_stream_token_ids(self):
|
||||
for field in ("return_prompt_token_ids", "return_token_ids"):
|
||||
req = ChatCompletionRequest(
|
||||
|
||||
Reference in New Issue
Block a user