mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
a098330b51
commit
e4e5f778bc
4 changed files with 183 additions and 15 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
(
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue