mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
## Description
Fix three related gaps in Bedrock support that prevented headroom from
working with Claude Code when `CLAUDE_CODE_USE_BEDROCK=0` and
`ANTHROPIC_BASE_URL` is pointed at the proxy:
1. **ARN passthrough used the wrong LiteLLM route** — application
inference profile ARNs (e.g.
`arn:aws:bedrock:ap-southeast-2:<account>:application-inference-profile/<id>`)
were forwarded as `bedrock/<arn>`, which LiteLLM rejects with HTTP 400
"Try calling via converse route". Fixed to `bedrock/converse/<arn>`.
2. **Named AWS profile not forwarded to completion calls** —
`--bedrock-profile` was wired through the CLI → config →
`LiteLLMBackend.__init__` and used to fetch the model map at startup,
but never stored on `self`. All four `acompletion()` call sites
(`send_message`, `stream_message`, `send_openai_message`,
`stream_openai_message`) passed only `aws_region_name` — the
actual Bedrock calls used ambient credentials regardless of the flag.
Fixed by storing `self.profile_name` and passing `aws_profile_name=` to
every `acompletion()` call.
3. **`ap-southeast-2` used the wrong region prefix** — Australia should
use `au.` for cross-region inference profile IDs, not `apac.`. Added
`ap-southeast-2 → "au"` to `_BEDROCK_REGION_PREFIXES` and `"au."` to the
strip list in `_normalize_bedrock_profile_id`.
Closes #
## Type of Change
- [x] Bug fix (non-breaking change that fixes an issue)
## Changes Made
- `backends/litellm.py`: route `arn:aws:` model IDs via
`bedrock/converse/<arn>` in `map_model_id`
- `backends/litellm.py`: store `profile_name` as `self.profile_name` in
`LiteLLMBackend.__init__`; pass `aws_profile_name=` to `acompletion()`
in all four call sites; use
`boto3.Session(profile_name=...)` for startup discovery; cache key is
`region:profile_name` to prevent cross-profile collisions
- `backends/litellm.py`: add `ap-southeast-2 → "au"` to
`_BEDROCK_REGION_PREFIXES`; add `"au."` to prefix strip list in
`_normalize_bedrock_profile_id`
- `providers/registry.py`: pass `profile_name=bedrock_profile` to
`LiteLLMBackend`
- `proxy/server.py`: pass `config.bedrock_profile` to
`create_proxy_backend`
- `docs/claude-code-bedrock-headroom.md`: remove false claim that ARNs
in `ANTHROPIC_DEFAULT_*_MODEL` bypass the proxy; fix troubleshooting
table
- `tests/test_bedrock_region.py`: update `test_arn_passthrough` to
expect `bedrock/converse/<arn>`; update cache key format; add
`test_profile_cache_isolation`,
`test_ap_southeast_2_uses_au_prefix`, and
`TestBedrockProfileForwardedToCompletion` (3 async tests asserting
`aws_profile_name` appears in `acompletion()` kwargs for named profiles
and is
absent for the no-profile case)
- `tests/test_provider_registry*.py`,
`test_vertex_claude_compression.py`: update `litellm_backend_cls` stubs
to accept `profile_name=None`
## Testing
- [x] Unit tests pass (`pytest`)
- [x] New tests added for new functionality
- [x] Manual testing performed
### Test Output
```text
$ pytest tests/test_bedrock_region.py tests/test_provider_registry.py tests/test_provider_registry_extended.py \
-k "not test_fallback_when_boto3_import_fails and not test_fallback_when_api_call_fails and not test_successful_fetch" -q
collected 51 items / 3 deselected / 48 selected
tests/test_bedrock_region.py ...........................
tests/test_provider_registry.py ...........
tests/test_provider_registry_extended.py .......
48 passed, 3 deselected in 2.00s
```
Note: 3 deselected tests use patch("builtins.__import__") which hangs
under Python 3.13 — pre-existing issue unrelated to these changes.
## Real Behavior Proof
- Environment: macOS, Python 3.13, Claude Code with
`CLAUDE_CODE_USE_BEDROCK=0`, `ANTHROPIC_BASE_URL=http://127.0.0.1:8787`,
AWS ap-southeast-2, application inference profile ARNs in
`ANTHROPIC_DEFAULT_*_MODEL`
- Exact command / steps: `headroom proxy --port 8787 --backend bedrock
--region ap-southeast-2 --bedrock-profile "my-sso-profile"`
- Observed result: Requests routed correctly to
`bedrock/converse/arn:aws:bedrock:ap-southeast-2:...:application-inference-profile/<id>`
as confirmed in LiteLLM logs
- Not tested: EU/APAC region ARN passthrough (logic is identical);
non-SSO credential flows
```text
15:29:44 - LiteLLM:INFO: utils.py:4090 -
LiteLLM completion() model= converse/arn:aws:bedrock:ap-southeast-2:<account>:application-inference-profile/<id>; provider = bedrock
2026-06-26 15:29:44,322 - LiteLLM - INFO -
LiteLLM completion() model= converse/arn:aws:bedrock:ap-southeast-2:<account>:application-inference-profile/<id>; provider = bedrock
15:31:09 - LiteLLM:INFO: utils.py:4090 -
LiteLLM completion() model= converse/arn:aws:bedrock:ap-southeast-2:<account>:application-inference-profile/<id>; provider = bedrock
2026-06-26 15:31:09,928 - LiteLLM - INFO -
LiteLLM completion() model= converse/arn:aws:bedrock:ap-southeast-2:<account>:application-inference-profile/<id>; provider = bedrock
15:34:26 - LiteLLM:INFO: utils.py:4090 -
LiteLLM completion() model= converse/arn:aws:bedrock:ap-southeast-2:<account>:application-inference-profile/<id>; provider = bedrock
2026-06-26 15:34:26,811 - LiteLLM - INFO -
LiteLLM completion() model= converse/arn:aws:bedrock:ap-southeast-2:<account>:application-inference-profile/<id>; provider = bedrock
```
## Review Readiness
- [x] I have performed a self-review
- [x] This PR is ready for human review
## Checklist
- [x] My code follows the project's style guidelines
- [x] I have performed a self-review of my code
- [x] I have commented my code, particularly in hard-to-understand areas
- [x] My changes generate no new warnings
- [x] I have added tests that prove my fix is effective or that my
feature works
- [x] New and existing unit tests pass locally with my changes
## Additional Notes
The 3 skipped tests (`test_fallback_when_boto3_import_fails`,
`test_fallback_when_api_call_fails`, `test_successful_fetch`) pre-exist
in the repo and use `patch("builtins.__import__")` which hangs under
Python 3.13. Not affected by these changes.
---------
Co-authored-by: Matt Haitana <mhaitana@costar.com>
Co-authored-by: JerrettDavis <mxjerrett@gmail.com>
495 lines
20 KiB
Python
495 lines
20 KiB
Python
"""Tests for Bedrock region support and fallback model mapping.
|
|
|
|
Ensures that EU, AP, and US regions all produce valid Bedrock model IDs,
|
|
and that the proxy degrades gracefully when boto3 is unavailable or the
|
|
AWS API call fails.
|
|
"""
|
|
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
from tests._dotenv import importorskip_no_env_leak
|
|
|
|
importorskip_no_env_leak("litellm")
|
|
|
|
from headroom.backends.litellm import ( # noqa: E402 (must follow importorskip)
|
|
LiteLLMBackend,
|
|
_bedrock_profiles_cache,
|
|
_bedrock_region_prefix,
|
|
_build_bedrock_fallback_map,
|
|
_fetch_bedrock_inference_profiles,
|
|
_normalize_bedrock_profile_id,
|
|
)
|
|
|
|
# =============================================================================
|
|
# Region Prefix Mapping
|
|
# =============================================================================
|
|
|
|
|
|
class TestBedrockRegionPrefix:
|
|
"""Test AWS region -> inference profile prefix mapping."""
|
|
|
|
def test_us_regions(self):
|
|
assert _bedrock_region_prefix("us-east-1") == "us"
|
|
assert _bedrock_region_prefix("us-west-2") == "us"
|
|
|
|
def test_eu_regions(self):
|
|
assert _bedrock_region_prefix("eu-central-1") == "eu"
|
|
assert _bedrock_region_prefix("eu-west-1") == "eu"
|
|
assert _bedrock_region_prefix("eu-west-3") == "eu"
|
|
assert _bedrock_region_prefix("eu-north-1") == "eu"
|
|
|
|
def test_ap_regions(self):
|
|
assert _bedrock_region_prefix("ap-southeast-1") == "apac"
|
|
assert _bedrock_region_prefix("ap-northeast-1") == "apac"
|
|
|
|
def test_unknown_region_defaults_to_us(self):
|
|
assert _bedrock_region_prefix("me-south-1") == "us"
|
|
assert _bedrock_region_prefix("sa-east-1") == "us"
|
|
|
|
|
|
# =============================================================================
|
|
# Static Fallback Model Map
|
|
# =============================================================================
|
|
|
|
|
|
class TestBuildBedrockFallbackMap:
|
|
"""Test static fallback model map construction."""
|
|
|
|
def test_us_region_uses_us_prefix(self):
|
|
model_map = _build_bedrock_fallback_map("us-east-1")
|
|
assert "claude-sonnet-4-20250514" in model_map
|
|
assert model_map["claude-sonnet-4-20250514"] == (
|
|
"bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"
|
|
)
|
|
|
|
def test_eu_region_uses_eu_prefix(self):
|
|
model_map = _build_bedrock_fallback_map("eu-central-1")
|
|
assert "claude-sonnet-4-20250514" in model_map
|
|
assert model_map["claude-sonnet-4-20250514"] == (
|
|
"bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0"
|
|
)
|
|
|
|
def test_ap_region_uses_apac_prefix(self):
|
|
model_map = _build_bedrock_fallback_map("ap-southeast-1")
|
|
assert "claude-sonnet-4-20250514" in model_map
|
|
assert model_map["claude-sonnet-4-20250514"] == (
|
|
"bedrock/apac.anthropic.claude-sonnet-4-20250514-v1:0"
|
|
)
|
|
|
|
def test_all_models_present(self):
|
|
model_map = _build_bedrock_fallback_map("us-east-1")
|
|
expected_models = [
|
|
"claude-opus-4-6",
|
|
"claude-sonnet-4-6",
|
|
"claude-sonnet-4-20250514",
|
|
"claude-opus-4-20250514",
|
|
"claude-3-7-sonnet-20250219",
|
|
"claude-3-5-sonnet-20241022",
|
|
"claude-3-5-haiku-20241022",
|
|
"claude-3-opus-20240229",
|
|
"claude-3-haiku-20240307",
|
|
"claude-haiku-4-5-20251001",
|
|
]
|
|
for model in expected_models:
|
|
assert model in model_map, f"Missing model: {model}"
|
|
|
|
def test_all_values_are_valid_bedrock_format(self):
|
|
"""Every value must start with 'bedrock/' and contain 'anthropic.'."""
|
|
for region in ("us-east-1", "eu-west-1", "ap-northeast-1"):
|
|
model_map = _build_bedrock_fallback_map(region)
|
|
for name, litellm_id in model_map.items():
|
|
assert litellm_id.startswith("bedrock/"), (
|
|
f"{name}: expected bedrock/ prefix, got {litellm_id}"
|
|
)
|
|
assert "anthropic." in litellm_id, (
|
|
f"{name}: expected anthropic. in id, got {litellm_id}"
|
|
)
|
|
|
|
|
|
# =============================================================================
|
|
# Fetch with Graceful Fallback
|
|
# =============================================================================
|
|
|
|
|
|
class TestFetchBedrockInferenceProfiles:
|
|
"""Test dynamic fetch with fallback on failure."""
|
|
|
|
def setup_method(self):
|
|
"""Clear the cache before each test."""
|
|
_bedrock_profiles_cache.clear()
|
|
|
|
def test_fallback_when_boto3_import_fails(self):
|
|
"""Should return static map when boto3 is not installed."""
|
|
with patch.dict("sys.modules", {"boto3": None}):
|
|
# Force reimport failure
|
|
|
|
# Temporarily break boto3 import inside the function
|
|
original_import = (
|
|
__builtins__.__import__ if hasattr(__builtins__, "__import__") else __import__
|
|
)
|
|
|
|
def mock_import(name, *args, **kwargs):
|
|
if name == "boto3":
|
|
raise ImportError("No module named 'boto3'")
|
|
return original_import(name, *args, **kwargs)
|
|
|
|
with patch("builtins.__import__", side_effect=mock_import):
|
|
_bedrock_profiles_cache.clear()
|
|
result = _fetch_bedrock_inference_profiles("eu-central-1")
|
|
|
|
assert len(result) > 0
|
|
# Should use EU prefix
|
|
assert result["claude-sonnet-4-20250514"] == (
|
|
"bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0"
|
|
)
|
|
|
|
def test_fallback_when_api_call_fails(self):
|
|
"""Should return static map when list_inference_profiles raises."""
|
|
mock_boto3 = MagicMock()
|
|
mock_client = MagicMock()
|
|
mock_client.list_inference_profiles.side_effect = Exception(
|
|
"AccessDeniedException: not authorized"
|
|
)
|
|
mock_session = MagicMock()
|
|
mock_session.client.return_value = mock_client
|
|
mock_boto3.Session.return_value = mock_session
|
|
|
|
with patch("headroom.backends.litellm.boto3", mock_boto3, create=True):
|
|
# Patch the import inside the function
|
|
_fetch_bedrock_inference_profiles.__code__ # noqa: B018
|
|
_bedrock_profiles_cache.clear()
|
|
|
|
# We need to actually test the function, so let's just use the
|
|
# mock_boto3 and make sure the function catches the exception
|
|
import builtins
|
|
|
|
real_import = builtins.__import__
|
|
|
|
def patched_import(name, *args, **kwargs):
|
|
if name == "boto3":
|
|
return mock_boto3
|
|
return real_import(name, *args, **kwargs)
|
|
|
|
with patch("builtins.__import__", side_effect=patched_import):
|
|
_bedrock_profiles_cache.clear()
|
|
result = _fetch_bedrock_inference_profiles("eu-west-1")
|
|
|
|
assert len(result) > 0
|
|
# Should use EU prefix
|
|
for litellm_id in result.values():
|
|
assert "eu.anthropic." in litellm_id
|
|
|
|
def test_successful_fetch_uses_api_results(self):
|
|
"""When API works, should use dynamic results (not fallback)."""
|
|
mock_boto3 = MagicMock()
|
|
mock_client = MagicMock()
|
|
mock_client.list_inference_profiles.return_value = {
|
|
"inferenceProfileSummaries": [
|
|
{"inferenceProfileId": "eu.anthropic.claude-sonnet-4-20250514-v1:0"},
|
|
{"inferenceProfileId": "eu.anthropic.claude-3-5-sonnet-20241022-v2:0"},
|
|
{"inferenceProfileId": "eu.meta.llama-3-70b-v1:0"}, # non-Anthropic, should skip
|
|
]
|
|
}
|
|
mock_session = MagicMock()
|
|
mock_session.client.return_value = mock_client
|
|
mock_boto3.Session.return_value = mock_session
|
|
|
|
import builtins
|
|
|
|
real_import = builtins.__import__
|
|
|
|
def patched_import(name, *args, **kwargs):
|
|
if name == "boto3":
|
|
return mock_boto3
|
|
return real_import(name, *args, **kwargs)
|
|
|
|
with patch("builtins.__import__", side_effect=patched_import):
|
|
_bedrock_profiles_cache.clear()
|
|
result = _fetch_bedrock_inference_profiles("eu-central-1")
|
|
|
|
assert len(result) == 2
|
|
assert result["claude-sonnet-4-20250514"] == (
|
|
"bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0"
|
|
)
|
|
assert result["claude-3-5-sonnet-20241022"] == (
|
|
"bedrock/eu.anthropic.claude-3-5-sonnet-20241022-v2:0"
|
|
)
|
|
|
|
def test_caching_prevents_repeated_api_calls(self):
|
|
"""Second call for same region+profile should return cached result."""
|
|
_bedrock_profiles_cache.clear()
|
|
_bedrock_profiles_cache["us-east-1:"] = {"test": "bedrock/test-model"}
|
|
|
|
result = _fetch_bedrock_inference_profiles("us-east-1")
|
|
assert result == {"test": "bedrock/test-model"}
|
|
|
|
def test_profile_cache_isolation(self):
|
|
"""Different profiles for the same region must not share a cache entry."""
|
|
_bedrock_profiles_cache.clear()
|
|
_bedrock_profiles_cache["us-east-1:profileA"] = {"model": "bedrock/profile-a-model"}
|
|
_bedrock_profiles_cache["us-east-1:profileB"] = {"model": "bedrock/profile-b-model"}
|
|
|
|
result_a = _fetch_bedrock_inference_profiles("us-east-1", profile_name="profileA")
|
|
result_b = _fetch_bedrock_inference_profiles("us-east-1", profile_name="profileB")
|
|
assert result_a["model"] == "bedrock/profile-a-model"
|
|
assert result_b["model"] == "bedrock/profile-b-model"
|
|
assert result_a != result_b
|
|
|
|
|
|
# =============================================================================
|
|
# LiteLLMBackend.map_model_id with EU Regions
|
|
# =============================================================================
|
|
|
|
|
|
class TestBedrockModelMapping:
|
|
"""Test model ID mapping for different regions."""
|
|
|
|
def setup_method(self):
|
|
_bedrock_profiles_cache.clear()
|
|
|
|
def test_eu_region_maps_correctly(self):
|
|
"""EU region should produce eu.anthropic.* model IDs."""
|
|
with patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value={
|
|
"claude-sonnet-4-20250514": "bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0",
|
|
},
|
|
):
|
|
backend = LiteLLMBackend(provider="bedrock", region="eu-central-1")
|
|
result = backend.map_model_id("claude-sonnet-4-20250514")
|
|
assert result == "bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0"
|
|
|
|
def test_us_region_maps_correctly(self):
|
|
"""US region should produce us.anthropic.* model IDs."""
|
|
with patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value={
|
|
"claude-sonnet-4-20250514": "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0",
|
|
},
|
|
):
|
|
backend = LiteLLMBackend(provider="bedrock", region="us-west-2")
|
|
result = backend.map_model_id("claude-sonnet-4-20250514")
|
|
assert result == "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"
|
|
|
|
def test_fallback_for_unknown_model_in_eu(self):
|
|
"""Unknown models in EU should get eu.anthropic.* fallback, not bare 'bedrock/claude-...'."""
|
|
with patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value={},
|
|
):
|
|
backend = LiteLLMBackend(provider="bedrock", region="eu-west-1")
|
|
result = backend.map_model_id("claude-sonnet-4-20250514")
|
|
assert result == "bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0"
|
|
|
|
def test_fallback_for_unknown_model_in_ap(self):
|
|
"""Unknown models in AP should get apac.anthropic.* fallback."""
|
|
with patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value={},
|
|
):
|
|
backend = LiteLLMBackend(provider="bedrock", region="ap-southeast-1")
|
|
result = backend.map_model_id("claude-3-5-haiku-20241022")
|
|
assert result == "bedrock/apac.anthropic.claude-3-5-haiku-20241022-v1:0"
|
|
|
|
def test_bedrock_format_passthrough(self):
|
|
"""Already-formatted Bedrock IDs should pass through unchanged."""
|
|
with patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value={},
|
|
):
|
|
backend = LiteLLMBackend(provider="bedrock", region="eu-central-1")
|
|
model = "bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0"
|
|
result = backend.map_model_id(model)
|
|
assert result == model
|
|
|
|
def test_anthropic_dot_format_normalized(self):
|
|
"""Raw Bedrock IDs like 'anthropic.claude-...-v1:0' should normalize and map."""
|
|
with patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value={
|
|
"claude-sonnet-4-20250514": "bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0",
|
|
},
|
|
):
|
|
backend = LiteLLMBackend(provider="bedrock", region="eu-central-1")
|
|
result = backend.map_model_id("anthropic.claude-sonnet-4-20250514-v1:0")
|
|
assert result == "bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0"
|
|
|
|
def test_region_prefixed_format_normalized(self):
|
|
"""'eu.anthropic.claude-...-v1:0' should normalize and map."""
|
|
with patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value={
|
|
"claude-sonnet-4-20250514": "bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0",
|
|
},
|
|
):
|
|
backend = LiteLLMBackend(provider="bedrock", region="eu-central-1")
|
|
result = backend.map_model_id("eu.anthropic.claude-sonnet-4-20250514-v1:0")
|
|
assert result == "bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0"
|
|
|
|
def test_arn_passthrough(self):
|
|
"""Application inference profile ARNs must use the converse route."""
|
|
with patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value={},
|
|
):
|
|
backend = LiteLLMBackend(provider="bedrock", region="ap-southeast-2")
|
|
arn = "arn:aws:bedrock:ap-southeast-2:123456789012:application-inference-profile/abc123"
|
|
result = backend.map_model_id(arn)
|
|
assert result == f"bedrock/converse/{arn}"
|
|
|
|
def test_ap_southeast_2_uses_au_prefix(self):
|
|
"""ap-southeast-2 (Sydney/Australia) should use 'au.' prefix, not 'apac.'."""
|
|
with patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value={},
|
|
):
|
|
backend = LiteLLMBackend(provider="bedrock", region="ap-southeast-2")
|
|
result = backend.map_model_id("claude-sonnet-4-5-20250929")
|
|
assert result == "bedrock/au.anthropic.claude-sonnet-4-5-20250929-v1:0"
|
|
|
|
|
|
# =============================================================================
|
|
# Normalize Bedrock Profile ID (edge cases)
|
|
# =============================================================================
|
|
|
|
|
|
class TestNormalizeBedrockProfileId:
|
|
"""Test normalization of various Bedrock profile ID formats."""
|
|
|
|
def test_eu_prefixed(self):
|
|
assert _normalize_bedrock_profile_id("eu.anthropic.claude-sonnet-4-20250514-v1:0") == (
|
|
"claude-sonnet-4-20250514"
|
|
)
|
|
|
|
def test_apac_prefixed(self):
|
|
assert _normalize_bedrock_profile_id("apac.anthropic.claude-3-5-sonnet-20241022-v2:0") == (
|
|
"claude-3-5-sonnet-20241022"
|
|
)
|
|
|
|
def test_us_prefixed(self):
|
|
assert _normalize_bedrock_profile_id("us.anthropic.claude-opus-4-20250514-v1:0") == (
|
|
"claude-opus-4-20250514"
|
|
)
|
|
|
|
def test_no_region_prefix(self):
|
|
assert _normalize_bedrock_profile_id("anthropic.claude-3-haiku-20240307-v1:0") == (
|
|
"claude-3-haiku-20240307"
|
|
)
|
|
|
|
def test_with_bedrock_slash_prefix(self):
|
|
assert (
|
|
_normalize_bedrock_profile_id("bedrock/eu.anthropic.claude-sonnet-4-20250514-v1:0")
|
|
== "claude-sonnet-4-20250514"
|
|
)
|
|
|
|
def test_non_claude_returns_none(self):
|
|
assert _normalize_bedrock_profile_id("eu.meta.llama-3-70b-v1:0") is None
|
|
|
|
def test_already_normalized(self):
|
|
assert _normalize_bedrock_profile_id("claude-sonnet-4-20250514") == (
|
|
"claude-sonnet-4-20250514"
|
|
)
|
|
|
|
|
|
# =============================================================================
|
|
# Named profile forwarded to acompletion kwargs
|
|
# =============================================================================
|
|
|
|
_MODEL_MAP_US = {"claude-sonnet-4-20250514": "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"}
|
|
_BODY = {
|
|
"model": "claude-sonnet-4-20250514",
|
|
"messages": [{"role": "user", "content": "hi"}],
|
|
"max_tokens": 10,
|
|
}
|
|
|
|
|
|
def _make_fake_completion_resp():
|
|
mock_resp = MagicMock()
|
|
mock_resp.choices = [MagicMock()]
|
|
mock_resp.choices[0].message.content = "hello"
|
|
mock_resp.choices[0].message.tool_calls = None
|
|
mock_resp.choices[0].finish_reason = "stop"
|
|
mock_resp.usage.prompt_tokens = 10
|
|
mock_resp.usage.completion_tokens = 5
|
|
return mock_resp
|
|
|
|
|
|
class TestBedrockProfileForwardedToCompletion:
|
|
"""Regression: --bedrock-profile must be passed to acompletion(), not just to
|
|
_fetch_bedrock_inference_profiles() at startup. Without self.profile_name the
|
|
actual Bedrock call still uses ambient/default credentials even when the user
|
|
explicitly supplied a named SSO profile."""
|
|
|
|
def setup_method(self):
|
|
_bedrock_profiles_cache.clear()
|
|
|
|
async def test_send_message_passes_aws_profile_name(self):
|
|
"""send_message() must include aws_profile_name in the acompletion() kwargs."""
|
|
captured_kwargs: dict = {}
|
|
|
|
async def fake_acompletion(**kwargs):
|
|
captured_kwargs.update(kwargs)
|
|
return _make_fake_completion_resp()
|
|
|
|
with (
|
|
patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value=_MODEL_MAP_US,
|
|
),
|
|
patch("headroom.backends.litellm.acompletion", side_effect=fake_acompletion),
|
|
):
|
|
backend = LiteLLMBackend(
|
|
provider="bedrock", region="us-east-1", profile_name="my-sso-profile"
|
|
)
|
|
await backend.send_message(body=_BODY, headers={})
|
|
|
|
assert captured_kwargs.get("aws_profile_name") == "my-sso-profile"
|
|
|
|
async def test_stream_message_passes_aws_profile_name(self):
|
|
"""stream_message() must include aws_profile_name in the acompletion() kwargs."""
|
|
captured_kwargs: dict = {}
|
|
|
|
async def fake_acompletion(**kwargs):
|
|
captured_kwargs.update(kwargs)
|
|
|
|
async def _empty():
|
|
return
|
|
yield # pragma: no cover — makes this an async generator
|
|
|
|
return _empty()
|
|
|
|
with (
|
|
patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value=_MODEL_MAP_US,
|
|
),
|
|
patch("headroom.backends.litellm.acompletion", side_effect=fake_acompletion),
|
|
):
|
|
backend = LiteLLMBackend(
|
|
provider="bedrock", region="us-east-1", profile_name="my-sso-profile"
|
|
)
|
|
async for _ in backend.stream_message(body=_BODY, headers={}):
|
|
pass
|
|
|
|
assert captured_kwargs.get("aws_profile_name") == "my-sso-profile"
|
|
|
|
async def test_no_profile_does_not_set_aws_profile_name(self):
|
|
"""When no profile is configured, aws_profile_name must not appear in kwargs
|
|
(LiteLLM falls back to ambient credentials correctly)."""
|
|
captured_kwargs: dict = {}
|
|
|
|
async def fake_acompletion(**kwargs):
|
|
captured_kwargs.update(kwargs)
|
|
return _make_fake_completion_resp()
|
|
|
|
with (
|
|
patch(
|
|
"headroom.backends.litellm._fetch_bedrock_inference_profiles",
|
|
return_value=_MODEL_MAP_US,
|
|
),
|
|
patch("headroom.backends.litellm.acompletion", side_effect=fake_acompletion),
|
|
):
|
|
backend = LiteLLMBackend(provider="bedrock", region="us-east-1")
|
|
await backend.send_message(body=_BODY, headers={})
|
|
|
|
assert "aws_profile_name" not in captured_kwargs
|