[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,
|
should_include_usage,
|
||||||
to_openai_style_logprobs,
|
to_openai_style_logprobs,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.entrypoints.request_headers import apply_header_overrides
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.function_call.core_types import ToolCallItem
|
from sglang.srt.function_call.core_types import ToolCallItem
|
||||||
from sglang.srt.function_call.function_call_parser import FunctionCallParser
|
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
|
return_prompt_token_ids=request.return_prompt_token_ids
|
||||||
or request.return_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
|
return adapted_request, request
|
||||||
|
|
||||||
|
|||||||
@@ -32,6 +32,7 @@ from sglang.srt.entrypoints.openai.serving_chat import (
|
|||||||
OpenAIServingChat,
|
OpenAIServingChat,
|
||||||
normalize_tool_content,
|
normalize_tool_content,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.function_call.kimik3_format import TOOLS_CLOSE
|
from sglang.srt.function_call.kimik3_format import TOOLS_CLOSE
|
||||||
from sglang.srt.managers.io_struct import GenerateReqInput
|
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||||
from sglang.srt.parser.template_detection import ReasoningToggleConfig
|
from sglang.srt.parser.template_detection import ReasoningToggleConfig
|
||||||
@@ -100,6 +101,7 @@ class _MockTokenizerManager:
|
|||||||
"prompt_tokens": 10,
|
"prompt_tokens": 10,
|
||||||
"completion_tokens": 5,
|
"completion_tokens": 5,
|
||||||
"cached_tokens": 0,
|
"cached_tokens": 0,
|
||||||
|
"weight_version": "test-version",
|
||||||
"finish_reason": {"type": "stop", "matched": None},
|
"finish_reason": {"type": "stop", "matched": None},
|
||||||
"output_token_logprobs": [(0.1, 1, "Test"), (0.2, 2, "response")],
|
"output_token_logprobs": [(0.1, 1, "Test"), (0.2, 2, "response")],
|
||||||
"output_top_logprobs": None,
|
"output_top_logprobs": None,
|
||||||
@@ -109,6 +111,7 @@ class _MockTokenizerManager:
|
|||||||
|
|
||||||
self.generate_request = Mock(return_value=_mock_generate())
|
self.generate_request = Mock(return_value=_mock_generate())
|
||||||
self.create_abort_task = Mock()
|
self.create_abort_task = Mock()
|
||||||
|
self.request_logger = Mock(log_requests=False, log_requests_level=0)
|
||||||
|
|
||||||
def config_value(self, name: str):
|
def config_value(self, name: str):
|
||||||
"""The manager's overlay accessor: no control-plane update recorded."""
|
"""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(adapted.session_id, "session-1")
|
||||||
self.assertEqual(processed, self.basic_req)
|
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):
|
def test_convert_to_internal_request_rejects_stream_token_ids(self):
|
||||||
for field in ("return_prompt_token_ids", "return_token_ids"):
|
for field in ("return_prompt_token_ids", "return_token_ids"):
|
||||||
req = ChatCompletionRequest(
|
req = ChatCompletionRequest(
|
||||||
|
|||||||
Reference in New Issue
Block a user