Skip to content

Commit f011d8c

Browse files
authored
[Anthropic] Fix missing cache_read_input_tokens in streaming responses (sgl-project#29703)
1 parent 860244d commit f011d8c

2 files changed

Lines changed: 72 additions & 0 deletions

File tree

python/sglang/srt/entrypoints/openai/serving_chat.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@ class ThinkingMode(str, Enum):
3838
FunctionResponse,
3939
LogProbs,
4040
MessageProcessingResult,
41+
PromptTokensDetails,
4142
ResponseParserProtocol,
4243
SglExt,
4344
ToolCall,
@@ -336,6 +337,15 @@ def _get_parsed_response_fields(
336337
"""Post-process reasoning and tool_calls before building response."""
337338
return reasoning_text, tool_calls
338339

340+
def _continuous_usage_cached_details(
341+
self, content: Dict[str, Any]
342+
) -> Optional[PromptTokensDetails]:
343+
if not self.tokenizer_manager.server_args.enable_cache_report:
344+
return None
345+
return UsageProcessor._details_if_cached(
346+
content["meta_info"].get("cached_tokens", 0)
347+
)
348+
339349
async def _generate_stream_content(
340350
self,
341351
content: Dict[str, Any],
@@ -377,6 +387,7 @@ async def _generate_stream_content(
377387
prompt_tokens=prompt_tokens.get(index, 0),
378388
reasoning_tokens=reasoning_tokens.get(index, 0),
379389
completion_tokens=completion_tokens.get(index, 0),
390+
cached_tokens=self._continuous_usage_cached_details(content),
380391
).model_dump()
381392

382393
yield build_sse_content(
@@ -422,6 +433,7 @@ async def _generate_stream_content(
422433
prompt_tokens=prompt_tokens.get(index, 0),
423434
reasoning_tokens=reasoning_tokens.get(index, 0),
424435
completion_tokens=completion_tokens.get(index, 0),
436+
cached_tokens=self._continuous_usage_cached_details(content),
425437
).model_dump()
426438

427439
yield build_sse_content(
@@ -449,6 +461,7 @@ async def _generate_stream_content(
449461
prompt_tokens=prompt_tokens.get(index, 0),
450462
reasoning_tokens=reasoning_tokens.get(index, 0),
451463
completion_tokens=completion_tokens.get(index, 0),
464+
cached_tokens=self._continuous_usage_cached_details(content),
452465
).model_dump()
453466

454467
yield build_sse_content(
@@ -1898,6 +1911,7 @@ async def _process_tool_call_stream(
18981911
prompt_tokens=prompt_tokens,
18991912
completion_tokens=completion_tokens,
19001913
reasoning_tokens=reasoning_tokens,
1914+
cached_tokens=self._continuous_usage_cached_details(content),
19011915
)
19021916

19031917
yield f"data: {chunk.model_dump_json()}\n\n"
@@ -1950,6 +1964,7 @@ async def _process_tool_call_stream(
19501964
prompt_tokens=prompt_tokens,
19511965
completion_tokens=completion_tokens,
19521966
reasoning_tokens=reasoning_tokens,
1967+
cached_tokens=self._continuous_usage_cached_details(content),
19531968
)
19541969

19551970
yield f"data: {chunk.model_dump_json()}\n\n"

test/registered/unit/entrypoints/openai/test_serving_chat.py

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1800,6 +1800,63 @@ async def run_stream():
18001800
},
18011801
)
18021802

1803+
def _collect_continuous_usage(self, cached_tokens):
1804+
content = {
1805+
"text": "Hello",
1806+
"meta_info": {
1807+
"id": "chatcmpl-cont-usage",
1808+
"prompt_tokens": 10,
1809+
"completion_tokens": 2,
1810+
"cached_tokens": cached_tokens,
1811+
"finish_reason": {"type": "stop", "matched": None},
1812+
"output_token_logprobs": None,
1813+
"output_top_logprobs": None,
1814+
},
1815+
"index": 0,
1816+
}
1817+
req = ChatCompletionRequest(
1818+
model="x",
1819+
messages=[{"role": "user", "content": "Hi?"}],
1820+
stream=True,
1821+
)
1822+
1823+
async def _collect():
1824+
chunks = []
1825+
async for chunk in self.chat._generate_stream_content(
1826+
content=content,
1827+
index=0,
1828+
request=req,
1829+
stream_offsets={},
1830+
reasoning_parser_dict={},
1831+
parser_dict={},
1832+
has_tool_calls={},
1833+
choice_logprobs=None,
1834+
finish_reason_type="stop",
1835+
continuous_usage_stats=True,
1836+
prompt_tokens={0: 10},
1837+
reasoning_tokens={0: 0},
1838+
completion_tokens={0: 2},
1839+
):
1840+
chunks.append(chunk)
1841+
return chunks
1842+
1843+
chunks = get_or_create_event_loop().run_until_complete(_collect())
1844+
return [c["usage"] for c in self._parse_chunks(chunks) if c.get("usage")]
1845+
1846+
def test_continuous_usage_reports_cached_tokens(self):
1847+
"""continuous_usage_stats chunks include cached tokens when cache reporting is on."""
1848+
self.tm.server_args.enable_cache_report = True
1849+
usages = self._collect_continuous_usage(cached_tokens=6)
1850+
self.assertTrue(usages, "continuous_usage_stats attached no usage")
1851+
self.assertEqual(usages[0]["prompt_tokens_details"]["cached_tokens"], 6)
1852+
1853+
def test_continuous_usage_omits_cached_tokens_when_report_disabled(self):
1854+
"""With cache reporting off, continuous_usage_stats must not leak cached tokens."""
1855+
self.tm.server_args.enable_cache_report = False
1856+
usages = self._collect_continuous_usage(cached_tokens=6)
1857+
self.assertTrue(usages, "continuous_usage_stats attached no usage")
1858+
self.assertIsNone(usages[0].get("prompt_tokens_details"))
1859+
18031860
# ------------- incremental streaming output tests -------------
18041861
def test_incremental_streaming_output_delta(self):
18051862
"""Test that streaming with incremental_streaming_output produces correct deltas.

0 commit comments

Comments
 (0)