diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index 09f7ff06c..559934702 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -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 diff --git a/test/registered/unit/entrypoints/openai/test_serving_chat.py b/test/registered/unit/entrypoints/openai/test_serving_chat.py index 3140e9faf..6391a65f6 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_chat.py +++ b/test/registered/unit/entrypoints/openai/test_serving_chat.py @@ -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(