From 9fc353c16a8dec77192b20dc234dae0aecb31071 Mon Sep 17 00:00:00 2001 From: Shobhit Verma Date: Wed, 29 Apr 2026 15:17:01 -0700 Subject: [PATCH] Fix disaggregated cached token usage Signed-off-by: Shobhit Verma --- tensorrt_llm/serve/openai_disagg_service.py | 139 +++++++++++++++++- .../test_openai_disagg_service.py | 81 +++++++++- 2 files changed, 214 insertions(+), 6 deletions(-) diff --git a/tensorrt_llm/serve/openai_disagg_service.py b/tensorrt_llm/serve/openai_disagg_service.py index 7ab886580608..b1fc646a9529 100644 --- a/tensorrt_llm/serve/openai_disagg_service.py +++ b/tensorrt_llm/serve/openai_disagg_service.py @@ -13,7 +13,9 @@ # limitations under the License. import asyncio +import json import os +from collections.abc import AsyncIterator from typing import Any, Callable, Dict, Optional from tensorrt_llm.llmapi.disagg_utils import ( @@ -34,6 +36,7 @@ CompletionRequest, DisaggregatedParams, DisaggScheduleStyle, + PromptTokensDetails, UCompletionRequest, UCompletionResponse, ) @@ -162,7 +165,10 @@ async def _send_disagg_request_ctx_first( gen_server, _ = await self._gen_router.get_next_server( gen_req, exclude_server=ctx_server ) - return await self._gen_client.send_request(gen_req, server=gen_server, hooks=hooks) + gen_response = await self._gen_client.send_request( + gen_req, server=gen_server, hooks=hooks + ) + return self._rewrite_disagg_usage(gen_response, ctx_response) else: if request.stream: # ctx client will never return a generator when streaming is requested @@ -170,6 +176,128 @@ async def _send_disagg_request_ctx_first( return done_generator() return ctx_response + def _ctx_usage_for_client( + self, ctx_response: Optional[UCompletionResponse] + ) -> tuple[Optional[int], int]: + if ctx_response is None or ctx_response.usage is None: + return None, 0 + + prompt_tokens = ctx_response.usage.prompt_tokens + cached_tokens = 0 + prompt_tokens_details = ctx_response.usage.prompt_tokens_details + if prompt_tokens_details is not None: + cached_tokens = prompt_tokens_details.cached_tokens + return prompt_tokens, cached_tokens + + def _rewrite_usage_payload_from_ctx( + self, + usage: dict[str, Any], + ctx_response: Optional[UCompletionResponse], + ) -> None: + prompt_tokens, cached_tokens = self._ctx_usage_for_client(ctx_response) + if prompt_tokens is None: + return + + usage["prompt_tokens"] = prompt_tokens + usage["total_tokens"] = prompt_tokens + (usage.get("completion_tokens") or 0) + prompt_tokens_details = usage.get("prompt_tokens_details") + if not isinstance(prompt_tokens_details, dict): + prompt_tokens_details = {} + usage["prompt_tokens_details"] = prompt_tokens_details + prompt_tokens_details["cached_tokens"] = cached_tokens + + def _rewrite_usage_response_from_ctx( + self, + response: UCompletionResponse, + ctx_response: Optional[UCompletionResponse], + ) -> UCompletionResponse: + prompt_tokens, cached_tokens = self._ctx_usage_for_client(ctx_response) + if prompt_tokens is None or response.usage is None: + return response + + response.usage.prompt_tokens = prompt_tokens + response.usage.total_tokens = prompt_tokens + (response.usage.completion_tokens or 0) + response.usage.prompt_tokens_details = PromptTokensDetails(cached_tokens=cached_tokens) + return response + + @staticmethod + def _sse_separator_index(data: bytes) -> tuple[int, bytes] | None: + indexes = [(data.find(sep), sep) for sep in (b"\n\n", b"\r\n\r\n")] + indexes = [(idx, sep) for idx, sep in indexes if idx >= 0] + if not indexes: + return None + return min(indexes, key=lambda item: item[0]) + + def _rewrite_usage_sse_event_from_ctx( + self, + event: bytes, + ctx_response: Optional[UCompletionResponse], + ) -> bytes: + separator_match = self._sse_separator_index(event) + separator = separator_match[1] if separator_match else b"" + event_body = event[: -len(separator)] if separator else event + data_lines = [ + line.removeprefix(b"data:").strip() + for line in event_body.splitlines() + if line.startswith(b"data:") + ] + if len(data_lines) != 1 or data_lines[0] == b"[DONE]": + return event + + try: + payload = json.loads(data_lines[0]) + except json.JSONDecodeError: + return event + + usage = payload.get("usage") if isinstance(payload, dict) else None + if not isinstance(usage, dict): + return event + + self._rewrite_usage_payload_from_ctx(usage, ctx_response) + return b"data: " + json.dumps(payload, separators=(",", ":")).encode("utf-8") + separator + + async def _rewrite_streaming_usage_from_ctx( + self, + response: AsyncIterator[Any], + ctx_response: Optional[UCompletionResponse], + ) -> AsyncIterator[Any]: + pending = b"" + pending_is_str = False + async for chunk in response: + is_str = isinstance(chunk, str) + chunk_bytes = chunk.encode("utf-8") if is_str else chunk + if not isinstance(chunk_bytes, bytes): + yield chunk + continue + + if not pending: + pending_is_str = is_str + pending += chunk_bytes + while separator_match := self._sse_separator_index(pending): + separator_index, separator = separator_match + event_end = separator_index + len(separator) + event = pending[:event_end] + pending = pending[event_end:] + event = self._rewrite_usage_sse_event_from_ctx(event, ctx_response) + yield event.decode("utf-8") if pending_is_str else event + if not pending: + pending_is_str = False + + if pending: + event = self._rewrite_usage_sse_event_from_ctx(pending, ctx_response) + yield event.decode("utf-8") if pending_is_str else event + + def _rewrite_disagg_usage( + self, + response: UCompletionResponseOrGenerator, + ctx_response: Optional[UCompletionResponse], + ) -> UCompletionResponseOrGenerator: + if ctx_response is None: + return response + if hasattr(response, "__aiter__"): + return self._rewrite_streaming_usage_from_ctx(response, ctx_response) + return self._rewrite_usage_response_from_ctx(response, ctx_response) + def _need_gen(self, response: UCompletionResponse) -> bool: if response and response.choices[0].finish_reason not in ["length", "not_finished"]: del response.choices[0].disaggregated_params @@ -449,7 +577,9 @@ async def _consume_gen(): # Now send ctx request — gen server has received its request try: - await self._ctx_client.send_request(ctx_req, server=ctx_server, hooks=hooks) + ctx_response = await self._ctx_client.send_request( + ctx_req, server=ctx_server, hooks=hooks + ) except Exception: consume_task.cancel() try: @@ -475,7 +605,7 @@ async def _yield_from_queue(): except asyncio.CancelledError: pass - return _yield_from_queue() + return self._rewrite_disagg_usage(_yield_from_queue(), ctx_response) else: # Non-streaming or no ctx needed: both HTTP POSTs fire eagerly # through generator consumption, so asyncio.gather works fine. @@ -492,4 +622,5 @@ async def _yield_from_queue(): ) ) responses = await asyncio.gather(*tasks) - return responses[-1] + ctx_response = responses[0] if need_ctx else None + return self._rewrite_disagg_usage(responses[-1], ctx_response) diff --git a/tests/unittest/disaggregated/test_openai_disagg_service.py b/tests/unittest/disaggregated/test_openai_disagg_service.py index 8e5c8e51a57e..59f2c8022b07 100644 --- a/tests/unittest/disaggregated/test_openai_disagg_service.py +++ b/tests/unittest/disaggregated/test_openai_disagg_service.py @@ -1,4 +1,5 @@ import asyncio +import json from unittest.mock import AsyncMock import pytest @@ -13,6 +14,7 @@ CompletionResponseChoice, DisaggregatedParams, DisaggScheduleStyle, + PromptTokensDetails, UsageInfo, _deserialize_first_gen_log_probs, _deserialize_first_gen_logits, @@ -41,12 +43,20 @@ def _make_completion_response( disagg_request_id: int = 42, prompt_token_ids=None, context_only=True, + prompt_tokens=1, + completion_tokens=1, + cached_tokens=0, ) -> CompletionResponse: if prompt_token_ids is None: prompt_token_ids = [1, 2, 3] return CompletionResponse( model="test-model", - usage=UsageInfo(prompt_tokens=1, completion_tokens=1), + usage=UsageInfo( + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=prompt_tokens + completion_tokens, + prompt_tokens_details=PromptTokensDetails(cached_tokens=cached_tokens), + ), prompt_token_ids=prompt_token_ids, choices=[ CompletionResponseChoice( @@ -110,6 +120,9 @@ async def _delayed_ctx_response(*_args, **_kwargs): finish_reason="length", disagg_request_id=request.disaggregated_params.disagg_request_id, context_only=True, + prompt_tokens=101, + completion_tokens=0, + cached_tokens=7, ) async def _delayed_gen_response(*_args, **_kwargs): @@ -123,6 +136,9 @@ async def _delayed_gen_response(*_args, **_kwargs): finish_reason="stop", disagg_request_id=request.disaggregated_params.disagg_request_id, context_only=False, + prompt_tokens=101, + completion_tokens=13, + cached_tokens=101, ) service._ctx_client.send_request = AsyncMock(side_effect=_delayed_ctx_response) @@ -149,7 +165,10 @@ async def _delayed_gen_response(*_args, **_kwargs): assert chunks == stream_chunks else: assert result.model == "test-model" - assert result.usage.prompt_tokens == 1 + assert result.usage.prompt_tokens == 101 + assert result.usage.completion_tokens == 13 + assert result.usage.total_tokens == 114 + assert result.usage.prompt_tokens_details.cached_tokens == 7 assert len(result.choices) == 1 assert result.choices[0].text == resp_text assert result.choices[0].finish_reason == "stop" @@ -159,6 +178,64 @@ async def _delayed_gen_response(*_args, **_kwargs): ) +@pytest.mark.asyncio +@pytest.mark.parametrize("schedule_style", ["context_first", "generation_first"]) +async def test_send_disagg_request_rewrites_streaming_usage(schedule_style): + service = _make_service(schedule_style) + service._ctx_client = AsyncMock() + service._gen_client = AsyncMock() + service._ctx_router.get_next_server = AsyncMock(return_value=("ctx:9000", {"server_info": {}})) + service._gen_router.get_next_server = AsyncMock(return_value=("gen:9001", {"server_info": {}})) + + async def _ctx_response(request, *_args, **_kwargs): + return _make_completion_response( + "", + finish_reason="length", + disagg_request_id=request.disaggregated_params.disagg_request_id, + context_only=True, + prompt_tokens=128, + completion_tokens=0, + cached_tokens=9, + ) + + async def _gen_response(*_args, **_kwargs): + usage_chunk = { + "choices": [], + "model": "test-model", + "usage": { + "prompt_tokens": 128, + "completion_tokens": 5, + "total_tokens": 133, + "prompt_tokens_details": { + "cached_tokens": 128, + }, + }, + } + return _mock_streaming_response( + [ + ( + b'data: {"choices":[{"delta":{"content":"hello"},"index":0}],' + b'"model":"test-model"}\n\n' + ), + f"data: {json.dumps(usage_chunk)}\n\n".encode(), + b"data: [DONE]\n\n", + ] + ) + + service._ctx_client.send_request = AsyncMock(side_effect=_ctx_response) + service._gen_client.send_request = AsyncMock(side_effect=_gen_response) + + request = CompletionRequest(model="test-model", prompt="hello", stream=True) + result = await service._send_disagg_request(request) + chunks = [chunk async for chunk in result] + + usage = json.loads(chunks[1].decode().removeprefix("data: "))["usage"] + assert usage["prompt_tokens"] == 128 + assert usage["completion_tokens"] == 5 + assert usage["total_tokens"] == 133 + assert usage["prompt_tokens_details"]["cached_tokens"] == 9 + + class TestVerifyCtxResponseDiagnostics: """Test enriched error messages in _verify_ctx_response (TRTLLM-11123)."""