fix(bedrock): keep region_name and model id on Invoke Claude responses

This commit is contained in:
flex-hyunsook 2026-10-02 14:38:21 +09:00
parent 8d28e8d776
commit c2c5368ab8
4 changed files with 110 additions and 2 deletions

View file

@ -2684,7 +2684,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
model_response.model = completion_response["model"]
_hidden_params["provider_specific_fields"] = provider_specific_fields
model_response._hidden_params = _hidden_params
model_response._hidden_params = {**model_response._hidden_params, **_hidden_params}
return model_response
def get_prefix_prompt(self, messages: list[AllMessageValues]) -> str | None:

View file

@ -23,6 +23,7 @@ from litellm.llms.bedrock.common_utils import (
normalize_bedrock_opus_output_config_effort,
normalize_custom_field_on_tools,
normalize_tool_input_schema_types_for_bedrock_invoke,
strip_bedrock_routing_prefix,
strip_unsupported_bedrock_invoke_output_config_keys,
tools_without_eager_input_streaming,
)
@ -378,7 +379,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
api_key: str | None = None,
json_mode: bool | None = None,
) -> ModelResponse:
return AnthropicConfig.transform_response(
transformed: Final = AnthropicConfig.transform_response(
self,
model=model,
raw_response=raw_response,
@ -392,3 +393,5 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
api_key=api_key,
json_mode=json_mode,
)
transformed.model = strip_bedrock_routing_prefix(model)
return transformed

View file

@ -5095,6 +5095,39 @@ def test_transform_parsed_response_reverse_maps_tool_names():
assert _json.loads(tcs[0].function.arguments) == {"job_id": 123}
def test_transform_parsed_response_keeps_hidden_params_set_before_the_call():
"""main.py writes region_name / custom_llm_provider onto the response before the
provider call, for region-based pricing. Rebuilding _hidden_params dropped them (#44002)."""
from litellm.types.utils import ModelResponse
config = AnthropicConfig()
raw_response = MagicMock()
raw_response.headers = {}
raw_response.status_code = 200
completion_response = {
"id": "msg_1",
"model": "claude-opus-5",
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 20},
"content": [{"type": "text", "text": "hello there"}],
}
model_response = ModelResponse()
model_response._hidden_params["region_name"] = "us-gov-west-1"
model_response._hidden_params["custom_llm_provider"] = "bedrock"
out = config.transform_parsed_response(
completion_response=completion_response,
raw_response=raw_response,
model_response=model_response,
)
assert out._hidden_params["region_name"] == "us-gov-west-1"
assert out._hidden_params["custom_llm_provider"] == "bedrock"
assert out._hidden_params["original_response"] == completion_response["content"]
assert "additional_headers" in out._hidden_params
assert "provider_specific_fields" in out._hidden_params
def test_transform_parsed_response_does_not_rewrite_unmapped_names():
"""CRITICAL: a tool legitimately named `foo_bar` must NOT be rewritten
to `foo/bar` just because some other request had that pair. The reverse

View file

@ -1136,3 +1136,75 @@ def test_chat_flagged_model_replays_a_byte_identical_prefix_around_a_mid_convers
_assert_prefix_stable(requests)
assert [m["role"] for m in requests[1]["messages"]] == ["user", "assistant", "user", "system"]
assert [m["role"] for m in requests[2]["messages"]] == ["user", "assistant", "user", "system", "assistant", "user"]
_INVOKE_CLAUDE_BODY: Final = {
"id": "msg_1",
"type": "message",
"role": "assistant",
"model": "claude-opus-5",
"content": [{"type": "text", "text": "hello there"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 10, "output_tokens": 20},
}
_CONVERSE_BODY: Final = {
"output": {"message": {"role": "assistant", "content": [{"text": "hello there"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 10, "outputTokens": 20, "totalTokens": 30},
"metrics": {"latencyMs": 1},
}
def _complete_bedrock_claude(model: str, aws_region_name: str) -> litellm.ModelResponse:
from litellm.llms.custom_httpx.http_handler import HTTPHandler
def fake_post(self, url, *args, **kwargs):
body = _CONVERSE_BODY if url.endswith("/converse") else _INVOKE_CLAUDE_BODY
return httpx.Response(
200,
json=body,
headers={"content-type": "application/json"},
request=httpx.Request("POST", url),
)
with patch.object(HTTPHandler, "post", fake_post):
return litellm.completion(
model=model,
aws_region_name=aws_region_name,
messages=[{"role": "user", "content": "hi"}],
client=HTTPHandler(),
)
@pytest.mark.parametrize(
("aws_region_name", "bedrock_model", "cost_key"),
[
("us-gov-west-1", "anthropic.claude-opus-5", "bedrock/us-gov-west-1/anthropic.claude-opus-5"),
("us-west-2", "anthropic.claude-opus-5", "anthropic.claude-opus-5"),
("us-west-2", "us.anthropic.claude-opus-5", "us.anthropic.claude-opus-5"),
("us-west-2", "global.anthropic.claude-opus-5", "global.anthropic.claude-opus-5"),
("ap-northeast-2", "apac.anthropic.claude-opus-5", "apac.anthropic.claude-opus-5"),
],
)
def test_invoke_claude_prices_like_converse_for_the_same_region_and_model(
local_model_cost_map, monkeypatch, aws_region_name, bedrock_model, cost_key
):
"""Invoke Claude used to drop region_name and report the Anthropic body model, so a
GovCloud call was priced at the commercial rate while Converse got the GovCloud key
(#44002). Cross-region profiles (us., apac.) must keep their own rate too."""
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "AKIAFAKEFAKEFAKE")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "fake")
monkeypatch.delenv("AWS_PROFILE", raising=False)
monkeypatch.delenv("AWS_SESSION_TOKEN", raising=False)
prices: Final = litellm.model_cost[cost_key]
expected_cost: Final = 10 * prices["input_cost_per_token"] + 20 * prices["output_cost_per_token"]
invoke: Final = _complete_bedrock_claude(f"bedrock/invoke/{bedrock_model}", aws_region_name)
converse: Final = _complete_bedrock_claude(f"bedrock/converse/{bedrock_model}", aws_region_name)
for response in (invoke, converse):
assert response._hidden_params["region_name"] == aws_region_name
assert response._hidden_params["custom_llm_provider"] == "bedrock"
assert response.model == bedrock_model
assert response._hidden_params["response_cost"] == pytest.approx(expected_cost)