mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(bedrock/realtime): meter usage and log via RealTimeStreaming
Parse Nova Sonic usageEvent into turn-level response.done usage, reuse RealTimeStreaming store_message/store_input/log_messages for spend and budget accounting, and pass user key metadata through arealtime
This commit is contained in:
parent
7444897f13
commit
7d60bf6126
5 changed files with 337 additions and 41 deletions
|
|
@ -2,17 +2,20 @@
|
|||
This file contains the handler for AWS Bedrock Nova Sonic realtime API.
|
||||
|
||||
This uses aws_sdk_bedrock_runtime for bidirectional streaming with Nova Sonic.
|
||||
Spend / budget logging follows the same RealTimeStreaming path as OpenAI/Azure:
|
||||
store_message for backend events, store_input for client events, log_messages on close.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import json
|
||||
from typing import Any, Final
|
||||
from typing import Any, Final, cast
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._logging import _redact_string, verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
from ..common_utils import BedrockError
|
||||
|
|
@ -46,6 +49,8 @@ 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,
|
||||
litellm_metadata: dict | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
|
|
@ -120,6 +125,27 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
|
||||
transformation_config: Final = BedrockRealtimeConfig()
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=None,
|
||||
api_key=api_key or "",
|
||||
additional_args={
|
||||
"api_base": endpoint_uri,
|
||||
"complete_input_dict": {"model": model},
|
||||
},
|
||||
)
|
||||
|
||||
# RealTimeStreaming owns spend logging for other realtime providers. Bedrock cannot
|
||||
# use its WebSocket bidirectional_forward (AWS SDK stream instead), but store_message /
|
||||
# store_input / log_messages are the same path used by OpenAI and Azure.
|
||||
realtime_streaming: Final = RealTimeStreaming(
|
||||
websocket=websocket,
|
||||
backend_ws=cast(Any, object()),
|
||||
logging_obj=logging_obj,
|
||||
model=model,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data={"litellm_metadata": litellm_metadata or {}},
|
||||
)
|
||||
|
||||
try:
|
||||
# Initialize the bidirectional stream
|
||||
bedrock_stream: Final = await bedrock_client.invoke_model_with_bidirectional_stream(
|
||||
|
|
@ -128,7 +154,9 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Bidirectional stream established")
|
||||
|
||||
await websocket.send_text(json.dumps(transformation_config.session_created_event(model, logging_obj)))
|
||||
session_created: Final = transformation_config.session_created_event(model, logging_obj)
|
||||
realtime_streaming.store_message(session_created)
|
||||
await websocket.send_text(json.dumps(session_created))
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: sent session.created to client on connect")
|
||||
|
||||
# Track state for transformation
|
||||
|
|
@ -142,35 +170,38 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
"session_configuration_request": None,
|
||||
}
|
||||
|
||||
# Create tasks for bidirectional forwarding
|
||||
client_to_bedrock_task: Final = asyncio.create_task(
|
||||
self._forward_client_to_bedrock(
|
||||
websocket,
|
||||
bedrock_stream,
|
||||
transformation_config,
|
||||
model,
|
||||
session_state,
|
||||
logging_obj,
|
||||
try:
|
||||
client_to_bedrock_task: Final = asyncio.create_task(
|
||||
self._forward_client_to_bedrock(
|
||||
websocket,
|
||||
bedrock_stream,
|
||||
transformation_config,
|
||||
model,
|
||||
session_state,
|
||||
logging_obj,
|
||||
realtime_streaming,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
bedrock_to_client_task: Final = asyncio.create_task(
|
||||
self._forward_bedrock_to_client(
|
||||
bedrock_stream,
|
||||
websocket,
|
||||
transformation_config,
|
||||
model,
|
||||
logging_obj,
|
||||
session_state,
|
||||
bedrock_to_client_task: Final = asyncio.create_task(
|
||||
self._forward_bedrock_to_client(
|
||||
bedrock_stream,
|
||||
websocket,
|
||||
transformation_config,
|
||||
model,
|
||||
logging_obj,
|
||||
session_state,
|
||||
realtime_streaming,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
# Wait for both tasks to complete
|
||||
await asyncio.gather(
|
||||
client_to_bedrock_task,
|
||||
bedrock_to_client_task,
|
||||
return_exceptions=True,
|
||||
)
|
||||
await asyncio.gather(
|
||||
client_to_bedrock_task,
|
||||
bedrock_to_client_task,
|
||||
return_exceptions=True,
|
||||
)
|
||||
finally:
|
||||
await realtime_streaming.log_messages()
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error in BedrockRealtime.async_realtime: %s", e)
|
||||
|
|
@ -180,6 +211,24 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
pass
|
||||
raise
|
||||
|
||||
@staticmethod
|
||||
def _collect_tool_call_from_function_call_event(
|
||||
realtime_streaming: RealTimeStreaming,
|
||||
message: object,
|
||||
) -> None:
|
||||
if not isinstance(message, dict) or message.get("type") != "response.function_call_arguments.done":
|
||||
return
|
||||
realtime_streaming.tool_calls.append(
|
||||
{
|
||||
"id": message.get("call_id", ""),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": message.get("name", ""),
|
||||
"arguments": message.get("arguments", "{}"),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
async def _forward_client_to_bedrock(
|
||||
self,
|
||||
client_ws: Any,
|
||||
|
|
@ -188,6 +237,7 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
model: str,
|
||||
session_state: dict,
|
||||
logging_obj: LiteLLMLogging | None = None,
|
||||
realtime_streaming: RealTimeStreaming | None = None,
|
||||
):
|
||||
"""Forward messages from client WebSocket to Bedrock stream."""
|
||||
from aws_sdk_bedrock_runtime.models import (
|
||||
|
|
@ -208,6 +258,9 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
message = await client_ws.receive_text()
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Received from client: %s", message[:200])
|
||||
|
||||
if realtime_streaming is not None:
|
||||
realtime_streaming.store_input(message)
|
||||
|
||||
# Transform OpenAI format to Bedrock format
|
||||
transformed_messages = transformation_config.transform_realtime_request(
|
||||
message=message,
|
||||
|
|
@ -230,11 +283,12 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
parsed_client_message.get("session", {}).get("modalities")
|
||||
)
|
||||
if client_message_type == "session.update":
|
||||
await client_ws.send_text(
|
||||
json.dumps(
|
||||
transformation_config.session_updated_event(model, logging_obj, requested_modalities)
|
||||
)
|
||||
session_updated: Final = transformation_config.session_updated_event(
|
||||
model, logging_obj, requested_modalities
|
||||
)
|
||||
if realtime_streaming is not None:
|
||||
realtime_streaming.store_message(session_updated)
|
||||
await client_ws.send_text(json.dumps(session_updated))
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Client to Bedrock forwarding ended: %s", e, exc_info=True)
|
||||
|
|
@ -252,6 +306,7 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
model: str,
|
||||
logging_obj: LiteLLMLogging,
|
||||
session_state: dict,
|
||||
realtime_streaming: RealTimeStreaming | None = None,
|
||||
):
|
||||
"""Forward messages from Bedrock stream to client WebSocket."""
|
||||
try:
|
||||
|
|
@ -304,6 +359,9 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
# Send transformed messages to client
|
||||
openai_messages = transformed_response.get("response", [])
|
||||
for openai_message in openai_messages:
|
||||
if realtime_streaming is not None:
|
||||
realtime_streaming.store_message(openai_message)
|
||||
self._collect_tool_call_from_function_call_event(realtime_streaming, openai_message)
|
||||
message_json = json.dumps(openai_message)
|
||||
await client_ws.send_text(message_json)
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Sent to client: %s", message_json[:200])
|
||||
|
|
|
|||
|
|
@ -36,9 +36,6 @@ from litellm.types.realtime import (
|
|||
RealtimeResponseTransformInput,
|
||||
RealtimeResponseTypedDict,
|
||||
)
|
||||
from litellm.utils import get_empty_usage
|
||||
|
||||
|
||||
class BedrockContentEnd(BaseModel):
|
||||
stopReason: str | None = None
|
||||
|
||||
|
|
@ -61,6 +58,74 @@ def _parse_bedrock_tool_use_input(raw_input: object) -> object:
|
|||
return {}
|
||||
|
||||
|
||||
def _as_nonneg_int(value: object) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
return 0
|
||||
return max(0, int(value))
|
||||
|
||||
|
||||
def _empty_usage_snapshot() -> dict[str, int]:
|
||||
return {
|
||||
"input_speech": 0,
|
||||
"input_text": 0,
|
||||
"output_speech": 0,
|
||||
"output_text": 0,
|
||||
"total_input": 0,
|
||||
"total_output": 0,
|
||||
"total": 0,
|
||||
}
|
||||
|
||||
|
||||
def _usage_snapshot_from_event(usage_event: dict) -> dict[str, int]:
|
||||
details: Final = usage_event.get("details") if isinstance(usage_event.get("details"), dict) else {}
|
||||
total_block: Final = details.get("total") if isinstance(details.get("total"), dict) else {}
|
||||
input_block: Final = total_block.get("input") if isinstance(total_block.get("input"), dict) else {}
|
||||
output_block: Final = total_block.get("output") if isinstance(total_block.get("output"), dict) else {}
|
||||
input_speech: Final = _as_nonneg_int(input_block.get("speechTokens"))
|
||||
input_text: Final = _as_nonneg_int(input_block.get("textTokens"))
|
||||
output_speech: Final = _as_nonneg_int(output_block.get("speechTokens"))
|
||||
output_text: Final = _as_nonneg_int(output_block.get("textTokens"))
|
||||
total_input: Final = _as_nonneg_int(usage_event.get("totalInputTokens")) or (input_speech + input_text)
|
||||
total_output: Final = _as_nonneg_int(usage_event.get("totalOutputTokens")) or (output_speech + output_text)
|
||||
total: Final = _as_nonneg_int(usage_event.get("totalTokens")) or (total_input + total_output)
|
||||
return {
|
||||
"input_speech": input_speech,
|
||||
"input_text": input_text,
|
||||
"output_speech": output_speech,
|
||||
"output_text": output_text,
|
||||
"total_input": total_input,
|
||||
"total_output": total_output,
|
||||
"total": total,
|
||||
}
|
||||
|
||||
|
||||
def _usage_snapshot_delta(current: dict[str, int], previous: dict[str, int]) -> dict[str, int]:
|
||||
return {key: max(0, current.get(key, 0) - previous.get(key, 0)) for key in _empty_usage_snapshot()}
|
||||
|
||||
|
||||
def _openai_usage_from_snapshot(snapshot: dict[str, int]) -> dict[str, Any]:
|
||||
input_tokens: Final = snapshot.get("total_input", 0) or (
|
||||
snapshot.get("input_speech", 0) + snapshot.get("input_text", 0)
|
||||
)
|
||||
output_tokens: Final = snapshot.get("total_output", 0) or (
|
||||
snapshot.get("output_speech", 0) + snapshot.get("output_text", 0)
|
||||
)
|
||||
return {
|
||||
"total_tokens": snapshot.get("total", 0) or (input_tokens + output_tokens),
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": output_tokens,
|
||||
"input_token_details": {
|
||||
"text_tokens": snapshot.get("input_text", 0),
|
||||
"audio_tokens": snapshot.get("input_speech", 0),
|
||||
"cached_tokens": 0,
|
||||
},
|
||||
"output_token_details": {
|
||||
"text_tokens": snapshot.get("output_text", 0),
|
||||
"audio_tokens": snapshot.get("output_speech", 0),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class BedrockRealtimeConfig(BaseRealtimeConfig):
|
||||
"""Configuration for Bedrock Nova Sonic realtime transformations."""
|
||||
|
||||
|
|
@ -71,6 +136,8 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
|
|||
self.audio_content_name = str(uuid_lib.uuid4())
|
||||
self.prompt_started = False
|
||||
self.client_audio_streamed = False
|
||||
self._usage_totals = _empty_usage_snapshot()
|
||||
self._usage_at_last_response_done = _empty_usage_snapshot()
|
||||
|
||||
# Default configuration values
|
||||
# Inference configuration
|
||||
|
|
@ -98,6 +165,14 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
|
|||
# Text configuration
|
||||
self.text_media_type = "text/plain"
|
||||
|
||||
def record_usage_event(self, usage_event: dict) -> None:
|
||||
self._usage_totals = _usage_snapshot_from_event(usage_event)
|
||||
|
||||
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 validate_environment(self, headers: dict, model: str, api_key: str | None = None) -> dict:
|
||||
"""Validate environment - no special validation needed for Bedrock."""
|
||||
return headers
|
||||
|
|
@ -999,7 +1074,6 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
|
|||
if not current_response_id or not current_conversation_id:
|
||||
return [], None, None, None
|
||||
|
||||
usage_obj: Final = get_empty_usage()
|
||||
response_done: Final = OpenAIRealtimeDoneEvent(
|
||||
type="response.done",
|
||||
event_id=f"event_{uuid.uuid4()}",
|
||||
|
|
@ -1009,15 +1083,10 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
|
|||
status="completed",
|
||||
output=[],
|
||||
conversation_id=current_conversation_id,
|
||||
usage={
|
||||
"prompt_tokens": usage_obj.prompt_tokens,
|
||||
"completion_tokens": usage_obj.completion_tokens,
|
||||
"total_tokens": usage_obj.total_tokens,
|
||||
},
|
||||
usage=self.consume_usage_for_response_done(),
|
||||
),
|
||||
)
|
||||
|
||||
# Reset state for next response
|
||||
return [response_done], None, None, None
|
||||
|
||||
def transform_tool_use_event(
|
||||
|
|
@ -1187,6 +1256,11 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
|
|||
if "sessionStart" in event:
|
||||
session_configuration_request = json.dumps({"configured": True})
|
||||
|
||||
elif "usageEvent" in event:
|
||||
usage_event: Final = event["usageEvent"]
|
||||
if isinstance(usage_event, dict):
|
||||
self.record_usage_event(usage_event)
|
||||
|
||||
elif "contentStart" in event:
|
||||
(
|
||||
events,
|
||||
|
|
|
|||
|
|
@ -434,6 +434,8 @@ async def _arealtime(
|
|||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
|
||||
aws_external_id=aws_external_id,
|
||||
user_api_key_dict=kwargs.get("user_api_key_dict"),
|
||||
litellm_metadata=_build_litellm_metadata(kwargs),
|
||||
)
|
||||
elif _custom_llm_provider == "xai":
|
||||
api_base = (
|
||||
|
|
|
|||
|
|
@ -57,6 +57,33 @@ class FakeBedrockStream:
|
|||
class FakeLogging:
|
||||
def __init__(self, trace_id="trace-nova-sonic"):
|
||||
self.litellm_trace_id = trace_id
|
||||
self.dispatched_results = []
|
||||
self.model_call_details = {}
|
||||
self.pre_call_args = []
|
||||
|
||||
def pre_call(self, input=None, api_key="", model=None, additional_args=None):
|
||||
self.pre_call_args.append(
|
||||
{"input": input, "api_key": api_key, "model": model, "additional_args": additional_args or {}}
|
||||
)
|
||||
|
||||
async def dispatch_success_handlers(self, result=None, prefer_async_handlers=False, **kwargs):
|
||||
self.dispatched_results.append(result)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def drain_bedrock_realtime_logging_worker(monkeypatch):
|
||||
pending = []
|
||||
|
||||
def capture_enqueue(coro):
|
||||
pending.append(coro)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.litellm_core_utils.realtime_streaming.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
|
||||
capture_enqueue,
|
||||
)
|
||||
yield pending
|
||||
for coro in pending:
|
||||
coro.close()
|
||||
|
||||
|
||||
class DisconnectingClientWS:
|
||||
|
|
@ -313,6 +340,30 @@ class TestBedrockRealtimeSessionLifecycle:
|
|||
assert first_event["session"]["id"] == "trace-nova-sonic"
|
||||
assert first_event["session"]["model"] == "amazon.nova-sonic-v1:0"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_dispatches_logged_events_via_realtime_streaming(
|
||||
self, stub_aws_sdk_client, drain_bedrock_realtime_logging_worker
|
||||
):
|
||||
handler = BedrockRealtime()
|
||||
websocket = RealtimeClientWS()
|
||||
logging_obj = FakeLogging()
|
||||
|
||||
await handler.async_realtime(
|
||||
model="amazon.nova-sonic-v1:0",
|
||||
websocket=websocket,
|
||||
logging_obj=logging_obj,
|
||||
aws_region_name="us-east-1",
|
||||
aws_access_key_id="k",
|
||||
aws_secret_access_key="s",
|
||||
)
|
||||
|
||||
assert logging_obj.pre_call_args
|
||||
assert len(drain_bedrock_realtime_logging_worker) == 1
|
||||
await drain_bedrock_realtime_logging_worker.pop()
|
||||
assert logging_obj.dispatched_results
|
||||
dispatched = logging_obj.dispatched_results[0]
|
||||
assert any(event.get("type") == "session.created" for event in dispatched)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_update_is_acked_with_session_updated(self, stub_aws_models):
|
||||
handler = BedrockRealtime()
|
||||
|
|
|
|||
|
|
@ -1090,5 +1090,116 @@ class TestBedrockRealtimeSessionEvents:
|
|||
assert event["session"]["modalities"] == ["text", "audio"]
|
||||
|
||||
|
||||
class TestBedrockRealtimeUsageAccounting:
|
||||
def _usage_event(
|
||||
self,
|
||||
*,
|
||||
input_speech: int,
|
||||
input_text: int,
|
||||
output_speech: int,
|
||||
output_text: int,
|
||||
total_input: int | None = None,
|
||||
total_output: int | None = None,
|
||||
total: int | None = None,
|
||||
) -> dict:
|
||||
resolved_input = total_input if total_input is not None else input_speech + input_text
|
||||
resolved_output = total_output if total_output is not None else output_speech + output_text
|
||||
resolved_total = total if total is not None else resolved_input + resolved_output
|
||||
return {
|
||||
"event": {
|
||||
"usageEvent": {
|
||||
"completionId": "completion_1",
|
||||
"details": {
|
||||
"total": {
|
||||
"input": {"speechTokens": input_speech, "textTokens": input_text},
|
||||
"output": {"speechTokens": output_speech, "textTokens": output_text},
|
||||
}
|
||||
},
|
||||
"promptName": "prompt_1",
|
||||
"sessionId": "session_1",
|
||||
"totalInputTokens": resolved_input,
|
||||
"totalOutputTokens": resolved_output,
|
||||
"totalTokens": resolved_total,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
def test_usage_event_fills_response_done_turn_delta(self):
|
||||
config = BedrockRealtimeConfig()
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.litellm_trace_id = "trace_123"
|
||||
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",
|
||||
}
|
||||
|
||||
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,
|
||||
)
|
||||
first_done = 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,
|
||||
)
|
||||
done_events = [msg for msg in first_done["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 usage["input_token_details"]["audio_tokens"] == 10
|
||||
assert usage["input_token_details"]["text_tokens"] == 2
|
||||
assert usage["output_token_details"]["audio_tokens"] == 20
|
||||
assert usage["output_token_details"]["text_tokens"] == 3
|
||||
|
||||
config.transform_realtime_response(
|
||||
json.dumps(
|
||||
self._usage_event(
|
||||
input_speech=15,
|
||||
input_text=2,
|
||||
output_speech=30,
|
||||
output_text=3,
|
||||
total_input=17,
|
||||
total_output=33,
|
||||
total=50,
|
||||
)
|
||||
),
|
||||
"amazon.nova-sonic-v1:0",
|
||||
logging_obj,
|
||||
realtime_response_transform_input={
|
||||
**state,
|
||||
"current_output_item_id": "item_2",
|
||||
"current_response_id": "resp_2",
|
||||
},
|
||||
)
|
||||
second_done = 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,
|
||||
"current_output_item_id": "item_2",
|
||||
"current_response_id": "resp_2",
|
||||
},
|
||||
)
|
||||
second_usage = [msg for msg in second_done["response"] if msg["type"] == "response.done"][0]["response"][
|
||||
"usage"
|
||||
]
|
||||
assert second_usage["input_tokens"] == 5
|
||||
assert second_usage["output_tokens"] == 10
|
||||
assert second_usage["total_tokens"] == 15
|
||||
assert second_usage["input_token_details"]["audio_tokens"] == 5
|
||||
assert second_usage["output_token_details"]["audio_tokens"] == 10
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue