diff --git a/python/sglang/srt/entrypoints/openai/serving_base.py b/python/sglang/srt/entrypoints/openai/serving_base.py index 51e1ddfe3..6e01d2fd0 100644 --- a/python/sglang/srt/entrypoints/openai/serving_base.py +++ b/python/sglang/srt/entrypoints/openai/serving_base.py @@ -282,3 +282,8 @@ class OpenAIServingBase(ABC): if label in self.allowed_custom_labels } return custom_labels + + def extract_routing_key(self, raw_request): + if raw_request is None: + return None + return raw_request.headers.get("x-smg-routing-key") diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index dbdb4a6b6..73dbc6d94 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -246,6 +246,7 @@ class OpenAIServingChat(OpenAIServingBase): extra_key=self._compute_extra_key(request), require_reasoning=self._get_reasoning_from_request(request), priority=request.priority, + routing_key=self.extract_routing_key(raw_request), custom_labels=custom_labels, custom_logit_processor=request.custom_logit_processor, image_max_dynamic_patch=img_max_dynamic_patch, diff --git a/python/sglang/srt/entrypoints/openai/serving_completions.py b/python/sglang/srt/entrypoints/openai/serving_completions.py index 9fbfb8841..8229de122 100644 --- a/python/sglang/srt/entrypoints/openai/serving_completions.py +++ b/python/sglang/srt/entrypoints/openai/serving_completions.py @@ -121,6 +121,7 @@ class OpenAIServingCompletion(OpenAIServingBase): rid=request.rid, extra_key=self._compute_extra_key(request), priority=request.priority, + routing_key=self.extract_routing_key(raw_request), custom_labels=custom_labels, custom_logit_processor=request.custom_logit_processor, ) diff --git a/python/sglang/srt/entrypoints/openai/serving_embedding.py b/python/sglang/srt/entrypoints/openai/serving_embedding.py index 223d6c3e6..a42325ae5 100644 --- a/python/sglang/srt/entrypoints/openai/serving_embedding.py +++ b/python/sglang/srt/entrypoints/openai/serving_embedding.py @@ -130,6 +130,7 @@ class OpenAIServingEmbedding(OpenAIServingBase): **prompt_kwargs, rid=request.rid, priority=request.priority, + routing_key=self.extract_routing_key(raw_request), dimensions=request.dimensions, ) diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 2ecd8542f..e7c60bc74 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -243,6 +243,9 @@ class GenerateReqInput(BaseReq, APIServingTimingMixin): # Extra key for classifying the request (e.g. cache_salt) extra_key: Optional[Union[List[str], str]] = None + # Routing key for routing-key schedule policy + routing_key: Optional[str] = None + # Whether to disallow logging for this request (e.g. due to ZDR) no_logs: bool = False @@ -740,6 +743,9 @@ class TokenizedGenerateReqInput(BaseReq): # Extra key for classifying the request (e.g. cache_salt) extra_key: Optional[str] = None + # Routing key for routing-key schedule policy + routing_key: Optional[str] = None + # Whether to disallow logging for this request (e.g. due to ZDR) no_logs: bool = False @@ -802,6 +808,8 @@ class EmbeddingReqInput(BaseReq, APIServingTimingMixin): is_cross_encoder_request: bool = False # Priority for the request priority: Optional[int] = None + # Routing key for routing-key schedule policy + routing_key: Optional[str] = None # For background responses (OpenAI responses API) background: bool = False diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index ab4fffb44..4cfa135e1 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -521,6 +521,7 @@ class Req: priority: Optional[int] = None, metrics_collector: Optional[SchedulerMetricsCollector] = None, extra_key: Optional[str] = None, + routing_key: Optional[str] = None, dimensions: Optional[int] = None, http_worker_ipc: Optional[str] = None, ): @@ -580,6 +581,7 @@ class Req: self.extra_key = extra_key self.lora_id = lora_id + self.routing_key = routing_key # Memory pool info self.req_pool_idx: Optional[int] = None diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 67a3633da..37ad13227 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -925,6 +925,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi data_parallel_rank=obj.data_parallel_rank, priority=obj.priority, extra_key=obj.extra_key, + routing_key=obj.routing_key, need_wait_for_image=obj.need_wait_for_image, num_items_assigned=obj.num_items_assigned, )