fix(bedrock/realtime): nova sonic tool calling (#33127)

* fix(bedrock/realtime): Nova Sonic tool calling

* increase coverage

* Apply suggestion from @greptile-apps[bot]

Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com>

* Update test_bedrock_realtime_transformation.py

* test(e2e): bound the post-/model/new servable wait at 40s

_await_model_servable used poll_timeout (120s), the spend/log read-back
budget. A stuck model reload therefore stalled every suite that creates a
deployment for two minutes before failing

Give create_model a fixed harness middle ground: model_servable_timeout=40s,
polled every 2s, with each /v1/models call capped at 5s and clamped to the
remaining deadline so one slow GET cannot overrun the wait. Happy path still
returns on the first listing. Not derived from proxy general_settings or env

Transport.get accepts an optional per-call timeout for that clamp. Unit tests
cover the deadline arithmetic and clamp without a live proxy

(cherry picked from commit c082a0e648)

* fix(e2e): wait one default DB reload interval of continuous listing

create_model returned after the first /v1/models hit that listed the model,
so chat could still land on a cold gateway worker (numWorkers>1 / peer pod)
and 400 Invalid model name. Require continuous listing for the product
default add_deployment interval (30s) after first sight so every worker has
synced from the DB; first listing still bounded at 40s

(cherry picked from commit 7d1ee2ff86)

* test(e2e): drop proxy_client model-servable unit tests

Keep the create_model DB-sync wait in the harness; the pure-function unit
file is not needed for this PR

(cherry picked from commit 89204651d1)

* fix(e2e): never skip the final deadline-clamped model-servable poll

When less than one full poll interval remained in the first-listing budget,
the pre-sleep check returned NotServable without another /v1/models call.
Sleep only min(interval, time left) so a model that becomes listable in the
last seconds of the timeout still gets a clamped final poll

(cherry picked from commit 8439195922)

* fix(e2e): reject first listing that returns after the 40s deadline

A poll may start with remaining budget and still return after started+timeout
if the transport overruns its clamp. Recheck the first-listing deadline after
the response so a late listing does not open the continuous DB-sync phase

(cherry picked from commit 7ff2bcbf14)

* test(e2e): poll MCP tools across multi-worker lag (#35047)

* fix(mcp): resolve call_tool by registry without requiring tool map

Multi-worker reloads put MCP servers in the registry from the DB but do
not re-run tools/list on every process. Gating call_tool on
tool_name_to_mcp_server_name_mapping made cold workers 500 with Tool not
found after another worker had already listed the tool. Treat a registry
match on server id/name/alias as enough; upstream rejects unknown tools

* test(e2e): poll MCP register, tools/list, and tools/call across multi-worker lag

Stage multi-worker gateways only load MCP servers and tool maps on the
process that handled the request. Poll until the server is listed, the
tool appears on tools/list, and tools/call is not a cold-worker 500 so
key-access and Datadog MCP e2e stop racing the LB

* Revert "fix(mcp): resolve call_tool by registry without requiring tool map"

This reverts commit 8b56e51e39.

* test(e2e): tighten MCP multi-worker lag classifier

Only retry tools/call on gateway shapes Tool <name> not found and
server_not_found, not any 500 that mentions tool/server not found, so
upstream failures are not retried until the poll deadline

* test(e2e): drop unit file for MCP lag classifier

The live await_call_tool polls already cover multi-worker lag; a separate
string-match unit module is not worth keeping

(cherry picked from commit c274cf321c)

---------

Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com>
Co-authored-by: yuneng-jiang <yuneng@berri.ai>
Co-authored-by: mubashir1osmani <mubashir.osmani777@gmail.com>
This commit is contained in:
Otavio Brito 2026-08-03 20:00:48 -03:00 • committed by GitHub
parent 9187b3d9ca
commit 1ca08f0bbd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 254 additions and 28 deletions

View file

@ -1029,11 +1029,13 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
if not current_output_item_id or not current_response_id:
return [], "", ""
# Parse the tool input
# Parse the tool input. Nova 2 Sonic sends arguments in `content`;
# fall back to `input` for backward compatibility.
tool_input = {}
if "input" in tool_use:
raw_input = tool_use["content"] if "content" in tool_use else tool_use.get("input")
if raw_input:
try:
tool_input = json.loads(tool_use["input"]) if isinstance(tool_use["input"], str) else tool_use["input"]
tool_input = json.loads(raw_input) if isinstance(raw_input, str) else raw_input
except json.JSONDecodeError:
tool_input = {}

View file

@ -75,12 +75,130 @@ from transport import HttpTransport, SplitTransport, Transport
RowsPredicate = Callable[[list[SpendLogRow]], bool]
# After /model/new, the control-plane writer reloads itself immediately, but every
# other gateway worker (and peer pod) only picks the model up on its add_deployment
# job. That job runs every proxy_config_reload_interval_seconds (product default 30).
# A single /v1/models hit can land on a hot worker while the next /chat hits a cold
# one ("Invalid model name"). Wait for first listing within MODEL_SERVABLE_TIMEOUT,
# then require continuous listing for MODEL_SERVABLE_DB_SYNC_SECONDS (the default
# reload interval) so every worker has had a chance to sync from the DB.
MODEL_SERVABLE_TIMEOUT = 40.0
MODEL_SERVABLE_DB_SYNC_SECONDS = 30.0
MODEL_SERVABLE_INTERVAL = 2.0
# Cap each /v1/models poll so one slow request cannot outlast the remaining budget.
MODEL_SERVABLE_REQUEST_TIMEOUT = 5.0
@dataclass(frozen=True, slots=True)
class Servable:
"""The data plane listed the model within the deadline."""
@dataclass(frozen=True, slots=True)
class NotServable:
"""The deadline passed without the data plane listing the model.
`last_result` is the final /v1/models read, so the caller can tell "the proxy
answered but omitted the model" (propagation) from "the read itself failed"
(network/auth) when reporting."""
last_result: Result[ModelsListResponse] | None
ServableOutcome = Servable | NotServable
def await_servable(
list_models: Callable[[float], Result[ModelsListResponse]],
*,
model_name: str,
timeout: float,
interval: float,
request_timeout: float,
db_sync_seconds: float,
now: Callable[[], float],
sleep: Callable[[float], None],
) -> ServableOutcome:
"""Poll until `model_name` is listed long enough for every worker to DB-sync.
First listing must happen within `timeout`. After that, the model must stay
listed continuously for `db_sync_seconds` (any miss resets the continuous
window). `db_sync_seconds=0` returns on the first listing. Each poll's request
timeout is clamped to the remaining budget. Sleeps only min(interval, time left)
so a final deadline-clamped poll is never skipped just because a full interval
does not fit. Clock and sleep are injected."""
started = now()
first_seen_at: float | None = None
last_result: Result[ModelsListResponse] | None = None
while True:
t = now()
phase_deadline = (
started + timeout if first_seen_at is None else first_seen_at + db_sync_seconds
)
remaining = phase_deadline - t
if remaining <= 0:
if (
last_result is not None
and first_seen_at is not None
and (db_sync_seconds <= 0 or t - first_seen_at >= db_sync_seconds)
):
return Servable()
return NotServable(last_result=last_result)
poll_timeout = min(request_timeout, remaining)
last_result = list_models(poll_timeout)
listed = isinstance(last_result, Success) and any(
entry.id == model_name for entry in last_result.data.data
)
t = now()
if not listed:
first_seen_at = None
elif first_seen_at is None:
if t > started + timeout:
return NotServable(last_result=last_result)
first_seen_at = t
if db_sync_seconds <= 0:
return Servable()
elif t - first_seen_at >= db_sync_seconds:
return Servable()
phase_deadline = (
started + timeout if first_seen_at is None else first_seen_at + db_sync_seconds
)
wait = min(interval, phase_deadline - now())
if wait > 0:
sleep(wait)
def servable_timeout_message(
*,
model_name: str,
timeout: float,
db_sync_seconds: float,
last_result: Result[ModelsListResponse] | None,
) -> str:
last_error = (
f"; last /v1/models poll did not succeed: {last_result}"
if last_result is not None and not isinstance(last_result, Success)
else ""
)
return (
f"model {model_name!r} was created but never became servable on the data "
f"plane within {timeout}s of first listing (plus {db_sync_seconds}s continuous "
f"DB sync) after /model/new (control/data-plane propagation or "
f"STORE_MODEL_IN_DB reload issue){last_error}"
)
@dataclass(frozen=True, slots=True)
class ProxyClient:
transport: Transport
poll_timeout: float = 120.0
poll_interval: float = 5.0
model_servable_timeout: float = MODEL_SERVABLE_TIMEOUT
model_servable_db_sync_seconds: float = MODEL_SERVABLE_DB_SYNC_SECONDS
model_servable_interval: float = MODEL_SERVABLE_INTERVAL
model_servable_request_timeout: float = MODEL_SERVABLE_REQUEST_TIMEOUT
# ---- keys / customers (satisfies lifecycle.ResourceClient) ----------
@ -167,7 +285,12 @@ class ProxyClient:
this returns can race the reload and 400 with "Invalid model name passed".
We therefore poll the data-plane /v1/models until the model appears before
handing back, so callers can invoke it immediately. In the monolithic case
it is already present on the first poll, so this adds one request."""
it is already present on the first poll, so this adds one request.
First listing must arrive within `model_servable_timeout` (not the longer
spend `poll_timeout`). The model must then stay listed for
`model_servable_db_sync_seconds` (product default DB reload interval) so every
gateway worker has run add_deployment before callers use the model."""
model_id = unwrap(
self.transport.post(
"/model/new",
@ -184,33 +307,38 @@ class ProxyClient:
return model_id
def _await_model_servable(self, model_name: str) -> None:
"""Block until the data plane lists `model_name`, or fail loudly if it does
not within poll_timeout (a real propagation/config problem, surfaced here
instead of as a downstream "Invalid model name passed")."""
deadline = time.monotonic() + self.poll_timeout
last_result: Result[ModelsListResponse] | None = None
while time.monotonic() < deadline:
last_result = self.transport.get(
"""Block until the data plane lists `model_name` long enough for DB sync.
Fails if first listing misses model_servable_timeout, or if continuous listing
for model_servable_db_sync_seconds never holds (multi-worker / peer reload)."""
outcome = await_servable(
lambda poll_timeout: self.transport.get(
"/v1/models",
headers=self.transport.master,
params=NoBody(),
response_type=ModelsListResponse,
)
if isinstance(last_result, Success) and any(
entry.id == model_name for entry in last_result.data.data
):
timeout=poll_timeout,
),
model_name=model_name,
timeout=self.model_servable_timeout,
interval=self.model_servable_interval,
request_timeout=self.model_servable_request_timeout,
db_sync_seconds=self.model_servable_db_sync_seconds,
now=time.monotonic,
sleep=time.sleep,
)
match outcome:
case Servable():
return
time.sleep(self.poll_interval)
last_error = (
f"; last /v1/models poll did not succeed: {last_result}"
if last_result is not None and not isinstance(last_result, Success)
else ""
)
raise AssertionError(
f"model {model_name!r} was created but never became servable on the data "
f"plane within {self.poll_timeout}s of /model/new (control/data-plane "
f"propagation or STORE_MODEL_IN_DB reload issue){last_error}"
)
case NotServable(last_result=last_result):
raise AssertionError(
servable_timeout_message(
model_name=model_name,
timeout=self.model_servable_timeout,
db_sync_seconds=self.model_servable_db_sync_seconds,
last_result=last_result,
)
)
def update_model(self, model_id: str, litellm_params: LiteLLMParamsBody) -> None:
"""Merge `litellm_params` over the deployment `model_id`'s stored params via

View file

@ -58,6 +58,7 @@ class Transport(Protocol):
headers: BaseModel,
params: BaseModel,
response_type: type[R],
timeout: float | None = None,
) -> Result[R]: ...
def delete[R: BaseModel](
@ -136,13 +137,16 @@ class HttpTransport:
headers: BaseModel,
params: BaseModel,
response_type: type[R],
timeout: float | None = None,
) -> Result[R]:
"""`timeout` overrides the transport-wide request_timeout for this call, for
pollers whose own deadline is shorter than it."""
return e2e_http.get(
self._url(path),
headers=headers,
params=params,
response_type=response_type,
timeout=self.request_timeout,
timeout=self.request_timeout if timeout is None else timeout,
)
def delete[R: BaseModel](
@ -336,9 +340,14 @@ class SplitTransport:
headers: BaseModel,
params: BaseModel,
response_type: type[R],
timeout: float | None = None,
) -> Result[R]:
return self._route(path).get(
path, headers=headers, params=params, response_type=response_type
path,
headers=headers,
params=params,
response_type=response_type,
timeout=timeout,
)
def delete[R: BaseModel](

View file

@ -571,6 +571,93 @@ class TestBedrockRealtimeResponseTransformation:
args = json.loads(function_call["arguments"])
assert args["location"] == "San Francisco"
def test_transform_tool_use_response_with_content_field(self):
"""Test toolUse response transformation with Nova 2 Sonic `content` field"""
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_123"
tool_use_message = {
"event": {
"toolUse": {
"toolUseId": "tool_call_123",
"toolName": "get_weather",
"content": json.dumps({"location": "San Francisco"}),
}
}
}
result = config.transform_realtime_response(
json.dumps(tool_use_message),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input={
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": "item_123",
"current_response_id": "resp_123",
"current_conversation_id": "conv_123",
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": "text",
},
)
# Check for function call event
assert len(result["response"]) == 1
function_call = result["response"][0]
assert function_call["type"] == "response.function_call_arguments.done"
assert function_call["call_id"] == "tool_call_123"
assert function_call["name"] == "get_weather"
# Verify arguments are properly formatted
args = json.loads(function_call["arguments"])
assert args["location"] == "San Francisco"
def test_transform_tool_use_event_directly(self):
"""Test transform_tool_use_event directly for guard clause and input parsing"""
config = BedrockRealtimeConfig()
# Guard clause: missing IDs returns empty result
events, tool_call_id, tool_name = config.transform_tool_use_event(
{"toolUse": {}}, None, "resp_123"
)
assert events == []
assert tool_call_id == ""
assert tool_name == ""
# JSON string content is parsed and converted to a function call event
events, tool_call_id, tool_name = config.transform_tool_use_event(
{
"toolUse": {
"toolUseId": "tool_call_123",
"toolName": "get_weather",
"content": json.dumps({"location": "San Francisco"}),
}
},
"item_123",
"resp_123",
)
assert len(events) == 1
assert events[0]["type"] == "response.function_call_arguments.done"
assert events[0]["call_id"] == "tool_call_123"
assert events[0]["name"] == "get_weather"
assert json.loads(events[0]["arguments"]) == {"location": "San Francisco"}
# Invalid JSON content falls back to empty arguments
events, _, _ = config.transform_tool_use_event(
{
"toolUse": {
"toolUseId": "tool_call_124",
"toolName": "get_weather",
"content": "not valid json",
}
},
"item_123",
"resp_123",
)
assert len(events) == 1
assert json.loads(events[0]["arguments"]) == {}
def test_transform_content_end_text(self):
"""Test contentEnd for text response"""
config = BedrockRealtimeConfig()