diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index b2b63b1fd..a67e2034c 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1161,6 +1161,7 @@ class EmbeddingReqInput: lora_id=self.lora_id[i] if self.lora_id is not None else None, positional_embed_overrides=self._get_positional_embed_overrides_item(i), http_worker_ipc=self.http_worker_ipc, + priority=self.priority, return_pooled_hidden_states=self.return_pooled_hidden_states, return_prompt_token_ids=self.return_prompt_token_ids, multi_item_delimiter_indices=( @@ -1188,6 +1189,7 @@ class EmbeddingReqInput: lora_id=self.lora_id[i] if self.lora_id is not None else None, positional_embed_overrides=self._get_positional_embed_overrides_item(i), http_worker_ipc=self.http_worker_ipc, + priority=self.priority, dimensions=self.dimensions, return_pooled_hidden_states=self.return_pooled_hidden_states, return_prompt_token_ids=self.return_prompt_token_ids, diff --git a/test/registered/unit/managers/test_io_struct.py b/test/registered/unit/managers/test_io_struct.py index 835c846c5..50d232822 100644 --- a/test/registered/unit/managers/test_io_struct.py +++ b/test/registered/unit/managers/test_io_struct.py @@ -1,7 +1,7 @@ import copy import unittest -from sglang.srt.managers.io_struct import GenerateReqInput +from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput from sglang.test.ci.ci_register import ( register_amd_ci, register_cpu_ci, @@ -704,5 +704,25 @@ class TestGenerateReqInputNormalization(CustomTestCase): req.normalize_batch_and_arguments() +class TestEmbeddingReqInputGetItem(CustomTestCase): + """Test EmbeddingReqInput.__getitem__.""" + + def test_priority_is_preserved(self): + """Priority must survive the batch split, in both __getitem__ branches.""" + req = EmbeddingReqInput(text=["Hello", "World"], priority=7) + req.normalize_batch_and_arguments() + self.assertEqual([req[0].priority, req[1].priority], [7, 7]) + + cross_encoder_req = EmbeddingReqInput( + text=[["query 1", "doc 1"], ["query 2", "doc 2"]], + is_cross_encoder_request=True, + priority=3, + ) + cross_encoder_req.normalize_batch_and_arguments() + self.assertEqual( + [cross_encoder_req[0].priority, cross_encoder_req[1].priority], [3, 3] + ) + + if __name__ == "__main__": unittest.main()