From 60a1dacd898e087c36d873b88fef123bc3853cca Mon Sep 17 00:00:00 2001 From: Vladislav Nosivskoy Date: Mon, 4 May 2026 22:30:17 +0300 Subject: [PATCH] [HiCache] return cached_tokens_details in sglext for streaming responses (#22055) Signed-off-by: Vladislav Nosivskoy --- .../srt/entrypoints/openai/serving_chat.py | 39 ++++-- .../entrypoints/openai/serving_completions.py | 41 ++++-- python/sglang/srt/entrypoints/openai/utils.py | 38 +++--- .../entrypoints/openai/test_serving_chat.py | 120 ++++++++++++++++++ .../openai/test_serving_completions.py | 111 ++++++++++++++++ 5 files changed, 310 insertions(+), 39 deletions(-) diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index 72ef2780a..9b698ed34 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -39,6 +39,7 @@ from sglang.srt.entrypoints.openai.protocol import ( from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor from sglang.srt.entrypoints.openai.utils import ( + cached_tokens_details_from_dict, process_cached_tokens_details_from_ret, process_hidden_states_from_ret, process_routed_experts_from_ret, @@ -766,6 +767,7 @@ class OpenAIServingChat(OpenAIServingBase): cached_tokens = {} hidden_states = {} routed_experts = {} + cached_tokens_details = {} stream_started = False try: @@ -789,6 +791,9 @@ class OpenAIServingChat(OpenAIServingBase): cached_tokens[index] = content["meta_info"].get("cached_tokens", 0) hidden_states[index] = content["meta_info"].get("hidden_states", None) routed_experts[index] = content["meta_info"].get("routed_experts", None) + cached_tokens_details[index] = content["meta_info"].get( + "cached_tokens_details", None + ) # Handle logprobs choice_logprobs = None @@ -963,20 +968,32 @@ class OpenAIServingChat(OpenAIServingBase): ) yield f"data: {hidden_states_chunk.model_dump_json()}\n\n" + sglext_routed = None if request.return_routed_experts and routed_experts: - # Get first non-None routed_experts value - first_routed_experts = next( + sglext_routed = next( (v for v in routed_experts.values() if v is not None), None ) - if first_routed_experts is not None: - routed_experts_chunk = ChatCompletionStreamResponse( - id=content["meta_info"]["id"], - created=int(time.time()), - choices=[], # sglext is at response level - model=request.model, - sglext=SglExt(routed_experts=first_routed_experts), - ) - yield f"data: {routed_experts_chunk.model_dump_json()}\n\n" + + sglext_details = None + if request.return_cached_tokens_details and cached_tokens_details: + first_details = next( + (v for v in cached_tokens_details.values() if v is not None), None + ) + if first_details is not None: + sglext_details = cached_tokens_details_from_dict(first_details) + + if sglext_routed is not None or sglext_details is not None: + sglext_chunk = ChatCompletionStreamResponse( + id=content["meta_info"]["id"], + created=int(time.time()), + choices=[], # sglext is at response level + model=request.model, + sglext=SglExt( + routed_experts=sglext_routed, + cached_tokens_details=sglext_details, + ), + ) + yield f"data: {sglext_chunk.model_dump_json()}\n\n" # Additional usage chunk if include_usage: diff --git a/python/sglang/srt/entrypoints/openai/serving_completions.py b/python/sglang/srt/entrypoints/openai/serving_completions.py index 2b80469b6..8598620c4 100644 --- a/python/sglang/srt/entrypoints/openai/serving_completions.py +++ b/python/sglang/srt/entrypoints/openai/serving_completions.py @@ -20,6 +20,7 @@ from sglang.srt.entrypoints.openai.protocol import ( from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor from sglang.srt.entrypoints.openai.utils import ( + cached_tokens_details_from_dict, process_cached_tokens_details_from_ret, process_hidden_states_from_ret, process_routed_experts_from_ret, @@ -224,6 +225,7 @@ class OpenAIServingCompletion(OpenAIServingBase): cached_tokens = {} hidden_states = {} routed_experts = {} + cached_tokens_details = {} stream_started = False try: @@ -248,6 +250,9 @@ class OpenAIServingCompletion(OpenAIServingBase): cached_tokens[index] = content["meta_info"].get("cached_tokens", 0) hidden_states[index] = content["meta_info"].get("hidden_states", None) routed_experts[index] = content["meta_info"].get("routed_experts", None) + cached_tokens_details[index] = content["meta_info"].get( + "cached_tokens_details", None + ) is_first_chunk = index not in stream_offsets offset = stream_offsets.get(index, 0) @@ -379,21 +384,33 @@ class OpenAIServingCompletion(OpenAIServingBase): ) yield f"data: {hidden_states_chunk.model_dump_json()}\n\n" + sglext_routed = None if request.return_routed_experts and routed_experts: - # Get first non-None routed_experts value - first_routed_experts = next( + sglext_routed = next( (v for v in routed_experts.values() if v is not None), None ) - if first_routed_experts is not None: - routed_experts_chunk = CompletionStreamResponse( - id=content["meta_info"]["id"], - created=created, - object="text_completion", - choices=[], # sglext is at response level - model=request.model, - sglext=SglExt(routed_experts=first_routed_experts), - ) - yield f"data: {routed_experts_chunk.model_dump_json()}\n\n" + + sglext_details = None + if request.return_cached_tokens_details and cached_tokens_details: + first_details = next( + (v for v in cached_tokens_details.values() if v is not None), None + ) + if first_details is not None: + sglext_details = cached_tokens_details_from_dict(first_details) + + if sglext_routed is not None or sglext_details is not None: + sglext_chunk = CompletionStreamResponse( + id=content["meta_info"]["id"], + created=created, + object="text_completion", + choices=[], # sglext is at response level + model=request.model, + sglext=SglExt( + routed_experts=sglext_routed, + cached_tokens_details=sglext_details, + ), + ) + yield f"data: {sglext_chunk.model_dump_json()}\n\n" # Handle final usage chunk if include_usage: diff --git a/python/sglang/srt/entrypoints/openai/utils.py b/python/sglang/srt/entrypoints/openai/utils.py index 0756994aa..7586f62f6 100644 --- a/python/sglang/srt/entrypoints/openai/utils.py +++ b/python/sglang/srt/entrypoints/openai/utils.py @@ -106,22 +106,10 @@ def process_routed_experts_from_ret( return ret_item["meta_info"].get("routed_experts", None) -def process_cached_tokens_details_from_ret( - ret_item: Dict[str, Any], - request: Union[ - ChatCompletionRequest, - CompletionRequest, - ], -) -> Optional[CachedTokensDetails]: - """Process cached tokens details from a ret item in non-streaming response.""" - if not getattr(request, "return_cached_tokens_details", False): - return None - - details = ret_item["meta_info"].get("cached_tokens_details", None) - if details is None: - return None - - # Check if L3 storage fields are present +def cached_tokens_details_from_dict( + details: Dict[str, Any], +) -> CachedTokensDetails: + """Convert a raw cached_tokens_details dict to a CachedTokensDetails object.""" if "storage" in details: return CachedTokensDetails( device=details.get("device", 0), @@ -136,6 +124,24 @@ def process_cached_tokens_details_from_ret( ) +def process_cached_tokens_details_from_ret( + ret_item: Dict[str, Any], + request: Union[ + ChatCompletionRequest, + CompletionRequest, + ], +) -> Optional[CachedTokensDetails]: + """Process cached tokens details from a ret item in non-streaming response.""" + if not request.return_cached_tokens_details: + return None + + details = ret_item["meta_info"].get("cached_tokens_details", None) + if details is None: + return None + + return cached_tokens_details_from_dict(details) + + def convert_embeds_to_tensors( embeds: Optional[Union[List[Optional[List[List[float]]]], List[List[float]]]], ) -> Optional[List[Optional[List[torch.Tensor]]]]: diff --git a/test/registered/unit/entrypoints/openai/test_serving_chat.py b/test/registered/unit/entrypoints/openai/test_serving_chat.py index 19f30cf08..6fc9bf0a0 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_chat.py +++ b/test/registered/unit/entrypoints/openai/test_serving_chat.py @@ -775,6 +775,126 @@ class ServingChatTestCase(unittest.TestCase): self.assertEqual(len(chunks), 2) self.assertIn("error", chunks[0]) + def test_non_streaming_cached_tokens_details_emits_sglext(self): + """Test that non-streaming chat responses emit cached token details in sglext.""" + + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi?"}], + max_tokens=100, + return_cached_tokens_details=True, + ) + ret = [ + { + "text": "Cached response", + "meta_info": { + "id": "chatcmpl-cache-test", + "prompt_tokens": 10, + "completion_tokens": 2, + "cached_tokens": 6, + "cached_tokens_details": { + "device": 4, + "host": 1, + "storage": 1, + "storage_backend": "file", + }, + "finish_reason": {"type": "stop", "matched": None}, + "weight_version": "default", + }, + } + ] + + response = self.chat._build_chat_response(req, ret, 1234567890) + + self.assertIsNotNone(response.sglext) + self.assertEqual( + response.sglext.cached_tokens_details.model_dump(exclude_none=True), + { + "device": 4, + "host": 1, + "storage": 1, + "storage_backend": "file", + }, + ) + + def test_streaming_cached_tokens_details_emits_sglext(self): + """Test that streaming chat responses emit cached token details in sglext.""" + + async def _mock_generate_with_cached_tokens_details(): + yield { + "text": "Cached response", + "meta_info": { + "id": "chatcmpl-cache-test", + "prompt_tokens": 10, + "completion_tokens": 2, + "cached_tokens": 6, + "cached_tokens_details": { + "device": 4, + "host": 1, + "storage": 1, + "storage_backend": "file", + }, + "finish_reason": {"type": "stop", "matched": None}, + "output_token_logprobs": None, + "output_top_logprobs": None, + }, + "index": 0, + } + + self.tm.generate_request.return_value = ( + _mock_generate_with_cached_tokens_details() + ) + + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi?"}], + max_tokens=100, + stream=True, + return_cached_tokens_details=True, + ) + + with patch( + "sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv" + ) as conv_mock: + conv_ins = Mock() + conv_ins.get_prompt.return_value = "Test prompt" + conv_mock.return_value = conv_ins + + adapted_request, _ = self.chat._convert_to_internal_request( + req, self.fastapi_request + ) + + async def run_stream(): + chunks = [] + async for chunk in self.chat._generate_chat_stream( + adapted_request, req, self.fastapi_request + ): + chunks.append(chunk) + return chunks + + loop = get_or_create_event_loop() + chunks = loop.run_until_complete(run_stream()) + + sglext_chunks = [] + for chunk in chunks: + if not chunk.startswith("data: ") or chunk.strip() == "data: [DONE]": + continue + data = json.loads(chunk[len("data: ") :]) + if "sglext" in data: + sglext_chunks.append(data) + + self.assertEqual(len(sglext_chunks), 1) + self.assertEqual(sglext_chunks[0]["choices"], []) + self.assertEqual( + sglext_chunks[0]["sglext"]["cached_tokens_details"], + { + "device": 4, + "host": 1, + "storage": 1, + "storage_backend": "file", + }, + ) + # ------------- incremental streaming output tests ------------- def test_incremental_streaming_output_delta(self): """Test that streaming with incremental_streaming_output produces correct deltas. diff --git a/test/registered/unit/entrypoints/openai/test_serving_completions.py b/test/registered/unit/entrypoints/openai/test_serving_completions.py index f090cb4da..8f79b14a0 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_completions.py +++ b/test/registered/unit/entrypoints/openai/test_serving_completions.py @@ -256,6 +256,117 @@ class ServingCompletionTestCase(unittest.TestCase): self.assertGreaterEqual(len(chunks), 2) self.assertIn("error", chunks[0]) + def test_non_streaming_cached_tokens_details_emits_sglext(self): + """Test that non-streaming completion responses emit cached token details in sglext.""" + + req = CompletionRequest( + model="x", + prompt="Hello world", + max_tokens=100, + return_cached_tokens_details=True, + ) + ret = [ + { + "text": "Cached response", + "meta_info": { + "id": "cmpl-cache-test", + "prompt_tokens": 10, + "completion_tokens": 2, + "cached_tokens": 6, + "cached_tokens_details": { + "device": 4, + "host": 1, + "storage": 1, + "storage_backend": "file", + }, + "finish_reason": {"type": "stop", "matched": None}, + "weight_version": "default", + }, + } + ] + + response = self.sc._build_completion_response(req, ret, 1234567890) + + self.assertIsNotNone(response.sglext) + self.assertEqual( + response.sglext.cached_tokens_details.model_dump(exclude_none=True), + { + "device": 4, + "host": 1, + "storage": 1, + "storage_backend": "file", + }, + ) + + def test_streaming_cached_tokens_details_emits_sglext(self): + """Test that streaming completion responses emit cached token details in sglext.""" + + async def _mock_generate_with_cached_tokens_details(*args, **kwargs): + yield { + "text": "Cached response", + "meta_info": { + "id": "cmpl-cache-test", + "prompt_tokens": 10, + "completion_tokens": 2, + "cached_tokens": 6, + "cached_tokens_details": { + "device": 4, + "host": 1, + "storage": 1, + "storage_backend": "file", + }, + "finish_reason": {"type": "stop", "matched": None}, + "output_token_logprobs": None, + "output_top_logprobs": None, + }, + "index": 0, + } + + self.sc.tokenizer_manager.generate_request = ( + _mock_generate_with_cached_tokens_details + ) + + req = CompletionRequest( + model="x", + prompt="Hello world", + max_tokens=100, + stream=True, + return_cached_tokens_details=True, + ) + + adapted_request, _ = self.sc._convert_to_internal_request(req) + + async def run_stream(): + chunks = [] + async for chunk in self.sc._generate_completion_stream( + adapted_request, req, self.fastapi_request + ): + chunks.append(chunk) + return chunks + + loop = get_or_create_event_loop() + chunks = loop.run_until_complete(run_stream()) + + sglext_chunks = [] + for chunk in chunks: + if not chunk.startswith("data: ") or chunk.strip() == "data: [DONE]": + continue + data = json.loads(chunk[len("data: ") :]) + if "sglext" in data: + sglext_chunks.append(data) + + self.assertEqual(len(sglext_chunks), 1) + self.assertEqual(sglext_chunks[0]["choices"], []) + self.assertEqual( + sglext_chunks[0]["sglext"]["cached_tokens_details"], + { + "device": 4, + "host": 1, + "storage": 1, + "storage_backend": "file", + }, + ) + if __name__ == "__main__": unittest.main(verbosity=2)