fix(bedrock/realtime): flush pending usage and type UserAPIKeyAuth

Bill late usageEvent after response ids clear, TOOL END_TURN response.done,
session-close drain, and completionEnd mint. Type user_api_key_dict as
UserAPIKeyAuth on Bedrock realtime and RealTimeStreaming.
This commit is contained in:
mubashir1osmani 2026-08-08 12:59:49 -07:00
parent a098330b51
commit e4e5f778bc
4 changed files with 183 additions and 15 deletions

View file

@ -20,6 +20,8 @@ from .litellm_logging import Logging as LiteLLMLogging
if TYPE_CHECKING:
from websockets.asyncio.client import ClientConnection
from litellm.proxy._types import UserAPIKeyAuth
CLIENT_CONNECTION_CLASS = ClientConnection
else:
CLIENT_CONNECTION_CLASS = Any
@ -48,7 +50,7 @@ class RealTimeStreaming:
logging_obj: LiteLLMLogging,
provider_config: BaseRealtimeConfig | None = None,
model: str = "",
user_api_key_dict: Any | None = None,
user_api_key_dict: "UserAPIKeyAuth | None" = None,
request_data: dict | None = None,
backend_uses_beta_protocol: bool | None = None,
force_transcription_model: str | None = None,
@ -83,7 +85,7 @@ class RealTimeStreaming:
self.current_item_chunks: list[OpenAIRealtimeOutputItemDone] | None = None
self.current_delta_type: ALL_DELTA_TYPES | None = None
self.session_configuration_request: str | None = None
self.user_api_key_dict = user_api_key_dict
self.user_api_key_dict: "UserAPIKeyAuth | None" = user_api_key_dict
self.request_data: dict = request_data or {}
# Violation counter for end_session_after_n_fails support
self._violation_count: int = 0

View file

@ -9,7 +9,7 @@ store_message for backend events, store_input for client events, log_messages on
import asyncio
import contextlib
import json
from typing import Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, cast
from pydantic import TypeAdapter
@ -21,6 +21,9 @@ from ..base_aws_llm import BaseAWSLLM
from ..common_utils import BedrockError
from .transformation import BedrockRealtimeConfig
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
_CLIENT_MODALITIES_ADAPTER: Final[TypeAdapter["list[str] | None"]] = TypeAdapter(list[str] | None)
@ -49,7 +52,7 @@ class BedrockRealtime(BaseAWSLLM):
aws_sts_endpoint: str | None = None,
aws_bedrock_runtime_endpoint: str | None = None,
aws_external_id: str | None = None,
user_api_key_dict: Any | None = None,
user_api_key_dict: "UserAPIKeyAuth | None" = None,
litellm_metadata: dict | None = None,
**kwargs,
):
@ -201,6 +204,11 @@ class BedrockRealtime(BaseAWSLLM):
return_exceptions=True,
)
finally:
for pending_usage_event in transformation_config.flush_pending_usage_as_response_done(
session_state.get("current_response_id"),
session_state.get("current_conversation_id"),
):
realtime_streaming.store_message(pending_usage_event)
await realtime_streaming.log_messages()
except Exception as e:

View file

@ -36,6 +36,8 @@ from litellm.types.realtime import (
RealtimeResponseTransformInput,
RealtimeResponseTypedDict,
)
class BedrockContentEnd(BaseModel):
stopReason: str | None = None
@ -168,11 +170,29 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
def record_usage_event(self, usage_event: dict) -> None:
self._usage_totals = _usage_snapshot_from_event(usage_event)
def has_unbilled_usage(self) -> bool:
delta: Final = _usage_snapshot_delta(self._usage_totals, self._usage_at_last_response_done)
return any(value > 0 for value in delta.values())
def consume_usage_for_response_done(self) -> dict[str, Any]:
delta: Final = _usage_snapshot_delta(self._usage_totals, self._usage_at_last_response_done)
self._usage_at_last_response_done = dict(self._usage_totals)
return _openai_usage_from_snapshot(delta)
def flush_pending_usage_as_response_done(
self,
current_response_id: str | None = None,
current_conversation_id: str | None = None,
) -> list[OpenAIRealtimeEvents]:
if not self.has_unbilled_usage():
return []
events, _, _, _ = self._response_done_events(
current_response_id,
current_conversation_id,
mint_ids_if_missing=True,
)
return events
def validate_environment(self, headers: dict, model: str, api_key: str | None = None) -> dict:
"""Validate environment - no special validation needed for Bedrock."""
return headers
@ -830,9 +850,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
output_index=0,
event_id=f"event_{uuid.uuid4()}",
item_id=current_output_item_id,
part=(
{"type": "text", "text": ""} if next_delta_type == "text" else {"type": "audio", "transcript": ""}
),
part=({"type": "text", "text": ""} if next_delta_type == "text" else {"type": "audio", "transcript": ""}),
response_id=current_response_id,
)
returned_messages.append(content_part_added)
@ -1051,19 +1069,27 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
Tuple of (events, reset_output_item_id, reset_response_id, reset_delta_type)
"""
verbose_logger.debug("Handling promptEnd")
return self._response_done_events(current_response_id, current_conversation_id)
return self._response_done_events(
current_response_id,
current_conversation_id,
mint_ids_if_missing=self.has_unbilled_usage(),
)
def _response_done_events(
self,
current_response_id: str | None,
current_conversation_id: str | None,
*,
mint_ids_if_missing: bool = False,
) -> tuple[
list[OpenAIRealtimeEvents],
str | None,
str | None,
ALL_DELTA_TYPES | None,
]:
if not current_response_id or not current_conversation_id:
response_id: Final = current_response_id or (f"resp_{uuid.uuid4()}" if mint_ids_if_missing else None)
conversation_id: Final = current_conversation_id or (f"conv_{uuid.uuid4()}" if mint_ids_if_missing else None)
if not response_id or not conversation_id:
return [], None, None, None
response_done: Final = OpenAIRealtimeDoneEvent(
@ -1071,10 +1097,10 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
event_id=f"event_{uuid.uuid4()}",
response=OpenAIRealtimeResponseDoneObject(
object="realtime.response",
id=current_response_id,
id=response_id,
status="completed",
output=[],
conversation_id=current_conversation_id,
conversation_id=conversation_id,
usage=self.consume_usage_for_response_done(),
),
)
@ -1252,6 +1278,19 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
usage_event: Final = event["usageEvent"]
if isinstance(usage_event, dict):
self.record_usage_event(usage_event)
if current_response_id is None and self.has_unbilled_usage():
(
done_events,
current_output_item_id,
current_response_id,
current_delta_type,
) = self._response_done_events(
None,
current_conversation_id,
mint_ids_if_missing=True,
)
returned_messages.extend(done_events)
current_delta_chunks = None
elif "contentStart" in event:
(
@ -1297,9 +1336,8 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
current_delta_chunks = None
current_delta_type = None
current_output_item_id = None
if content_end.get("type") == "TOOL":
current_response_id = None
if BedrockContentEnd.model_validate(content_end).stopReason == "END_TURN":
is_end_turn: Final = BedrockContentEnd.model_validate(content_end).stopReason == "END_TURN"
if is_end_turn:
(
done_events,
current_output_item_id,
@ -1308,6 +1346,8 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
) = self._response_done_events(current_response_id, current_conversation_id)
returned_messages.extend(done_events)
current_delta_chunks = None
elif content_end.get("type") == "TOOL":
current_response_id = None
elif "toolUse" in event:
(

View file

@ -15,7 +15,6 @@ from litellm.llms.bedrock.realtime.transformation import (
BedrockRealtimeConfig,
)
from litellm.llms.bedrock.realtime.trigger_audio import ready_trigger_pcm
from litellm.types.llms.openai import OpenAIRealtimeEventTypes
class TestBedrockRealtimeConfig:
@ -1356,6 +1355,125 @@ class TestBedrockRealtimeUsageAccounting:
assert second_usage["input_token_details"]["audio_tokens"] == 5
assert second_usage["output_token_details"]["audio_tokens"] == 10
def test_late_usage_event_after_response_id_cleared_emits_response_done(self):
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_late_usage"
state = {
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": "item_1",
"current_response_id": "resp_1",
"current_conversation_id": "conv_1",
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": "audio",
}
end_turn = config.transform_realtime_response(
json.dumps({"event": {"contentEnd": {"stopReason": "END_TURN", "type": "AUDIO"}}}),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
assert any(msg["type"] == "response.done" for msg in end_turn["response"])
assert end_turn["current_response_id"] is None
state["current_response_id"] = end_turn["current_response_id"]
state["current_output_item_id"] = end_turn["current_output_item_id"]
state["current_conversation_id"] = end_turn["current_conversation_id"]
late_usage = config.transform_realtime_response(
json.dumps(self._usage_event(input_speech=10, input_text=2, output_speech=20, output_text=3)),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
done_events = [msg for msg in late_usage["response"] if msg["type"] == "response.done"]
assert len(done_events) == 1
usage = done_events[0]["response"]["usage"]
assert usage["input_tokens"] == 12
assert usage["output_tokens"] == 23
assert usage["total_tokens"] == 35
assert not config.has_unbilled_usage()
def test_tool_end_turn_emits_response_done_before_clearing_ids(self):
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_tool_end_turn"
state = {
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": "item_tool",
"current_response_id": "resp_tool",
"current_conversation_id": "conv_tool",
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": None,
}
config.transform_realtime_response(
json.dumps(self._usage_event(input_speech=4, input_text=1, output_speech=0, output_text=2)),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
tool_end = config.transform_realtime_response(
json.dumps({"event": {"contentEnd": {"stopReason": "END_TURN", "type": "TOOL"}}}),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
done_events = [msg for msg in tool_end["response"] if msg["type"] == "response.done"]
assert len(done_events) == 1
assert done_events[0]["response"]["id"] == "resp_tool"
assert done_events[0]["response"]["usage"]["input_tokens"] == 5
assert done_events[0]["response"]["usage"]["output_tokens"] == 2
assert tool_end["current_response_id"] is None
assert not config.has_unbilled_usage()
def test_flush_pending_usage_on_session_close(self):
config = BedrockRealtimeConfig()
config.record_usage_event(
self._usage_event(input_speech=8, input_text=1, output_speech=16, output_text=2)["event"]["usageEvent"]
)
assert config.has_unbilled_usage()
flushed = config.flush_pending_usage_as_response_done(None, None)
assert len(flushed) == 1
assert flushed[0]["type"] == "response.done"
usage = flushed[0]["response"]["usage"]
assert usage["input_tokens"] == 9
assert usage["output_tokens"] == 18
assert usage["total_tokens"] == 27
assert not config.has_unbilled_usage()
assert config.flush_pending_usage_as_response_done(None, None) == []
def test_completion_end_with_unbilled_usage_mints_response_done(self):
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_completion_end"
config.record_usage_event(
self._usage_event(input_speech=3, input_text=0, output_speech=6, output_text=0)["event"]["usageEvent"]
)
result = config.transform_realtime_response(
json.dumps({"event": {"completionEnd": {}}}),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input={
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": None,
"current_response_id": None,
"current_conversation_id": None,
"current_delta_chunks": None,
"current_item_chunks": None,
"current_delta_type": None,
},
)
done_events = [msg for msg in result["response"] if msg["type"] == "response.done"]
assert len(done_events) == 1
assert done_events[0]["response"]["usage"]["input_tokens"] == 3
assert done_events[0]["response"]["usage"]["output_tokens"] == 6
assert not config.has_unbilled_usage()
if __name__ == "__main__":
pytest.main([__file__, "-v"])