[HiCache] return cached_tokens_details in sglext for streaming responses (#22055)

Signed-off-by: Vladislav Nosivskoy <vladnosiv@gmail.com>
This commit is contained in:
Vladislav Nosivskoy
2026-05-04 12:30:17 -07:00
committed by GitHub
parent 8ffd39e140
commit 60a1dacd89
5 changed files with 310 additions and 39 deletions
@@ -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:
@@ -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:
+22 -16
View File
@@ -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]]]]:
@@ -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.
@@ -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)