mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-10 14:27:00 -04:00
fix(bedrock): route ARNs via converse, named AWS profiles, and au. re… (#1456)
## 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>
This commit is contained in:
parent
c75ebdee6d
commit
7d87aa2f1c
8 changed files with 352 additions and 18 deletions
147
docs/claude-code-bedrock-headroom.md
Normal file
147
docs/claude-code-bedrock-headroom.md
Normal file
|
|
@ -0,0 +1,147 @@
|
|||
# Claude Code + AWS Bedrock, with Headroom compression
|
||||
|
||||
*Validated end-to-end on 2026-06-26 (Claude Code 2.1, Headroom 0.27.0, ap-southeast-2).*
|
||||
|
||||
This is the **working, tested** way to run **Claude Code** against **Claude models on
|
||||
AWS Bedrock** with **Headroom compressing the context** in the middle.
|
||||
|
||||
## TL;DR
|
||||
|
||||
Run Claude Code in **normal Anthropic mode** (NOT Bedrock mode) pointed at a local
|
||||
Headroom proxy, and let **Headroom** be the thing that talks to Bedrock:
|
||||
|
||||
```
|
||||
Claude Code ──ANTHROPIC_BASE_URL──▶ Headroom proxy ──LiteLLM (bedrock)──▶ AWS Bedrock
|
||||
(normal mode) (plain http) (compresses) (your AWS creds) (Claude)
|
||||
```
|
||||
|
||||
One non-obvious requirement makes the difference between "works" and "silently bypasses
|
||||
the proxy":
|
||||
|
||||
1. **`CLAUDE_CODE_USE_BEDROCK=0`** — Without this, Claude Code sees the
|
||||
`CLAUDE_CODE_USE_BEDROCK=1` flag and calls Bedrock directly via the AWS SDK,
|
||||
completely bypassing `ANTHROPIC_BASE_URL` and the proxy.
|
||||
|
||||
## Why not "just set CLAUDE_CODE_USE_BEDROCK=1"?
|
||||
|
||||
That approach **does not work** with Headroom. When `CLAUDE_CODE_USE_BEDROCK=1` is set,
|
||||
Claude Code calls Bedrock directly using the AWS SDK — `ANTHROPIC_BASE_URL` is ignored
|
||||
entirely and the proxy never receives a byte. Use the Anthropic-mode path below.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- **AWS credentials** configured for your environment (env vars, `~/.aws/credentials`,
|
||||
instance profile, or SSO via `aws sso login`). Confirm direct access works before
|
||||
involving Headroom:
|
||||
```bash
|
||||
aws bedrock-runtime invoke-model \
|
||||
--region us-east-1 \
|
||||
--model-id anthropic.claude-3-haiku-20240307-v1:0 \
|
||||
--body '{"anthropic_version":"bedrock-2023-05-31","max_tokens":20,"messages":[{"role":"user","content":"hi"}]}' \
|
||||
/tmp/out.json
|
||||
```
|
||||
- **boto3** in the proxy's Python environment (for dynamic inference profile discovery):
|
||||
```bash
|
||||
pip install boto3
|
||||
```
|
||||
- **IAM permissions** for the models you intend to use — at minimum
|
||||
`bedrock:InvokeModel` and `bedrock:InvokeModelWithResponseStream`. For application
|
||||
inference profiles, scope to the specific profile ARN:
|
||||
```json
|
||||
{
|
||||
"Effect": "Allow",
|
||||
"Action": ["bedrock:InvokeModel", "bedrock:InvokeModelWithResponseStream"],
|
||||
"Resource": ["arn:aws:bedrock:<region>:<account>:application-inference-profile/<id>"]
|
||||
}
|
||||
```
|
||||
|
||||
## Terminal 1 — start the Headroom proxy (Bedrock backend)
|
||||
|
||||
```bash
|
||||
headroom proxy --port 8787 \
|
||||
--backend bedrock \
|
||||
--region us-east-1
|
||||
```
|
||||
|
||||
With a named AWS SSO profile:
|
||||
|
||||
```bash
|
||||
headroom proxy --port 8787 \
|
||||
--backend bedrock \
|
||||
--region us-east-1 \
|
||||
--bedrock-profile my-sso-profile
|
||||
```
|
||||
|
||||
On startup the proxy calls `list_inference_profiles` to build a model map. Confirm it
|
||||
is routing correctly by checking the LiteLLM log lines — you should see:
|
||||
|
||||
```
|
||||
LiteLLM completion() model= converse/arn:aws:... provider = bedrock
|
||||
```
|
||||
|
||||
## Terminal 2 — run Claude Code (normal Anthropic mode) against the proxy
|
||||
|
||||
```bash
|
||||
export CLAUDE_CODE_USE_BEDROCK=0 # REQUIRED — prevents Claude Code bypassing the proxy
|
||||
export ANTHROPIC_BASE_URL=http://127.0.0.1:8787
|
||||
export ANTHROPIC_API_KEY=headroom # Claude Code needs *a* key to start; value is ignored
|
||||
export ANTHROPIC_MODEL=claude-opus-4-6
|
||||
export ANTHROPIC_DEFAULT_SONNET_MODEL=claude-sonnet-4-6
|
||||
export ANTHROPIC_DEFAULT_OPUS_MODEL=claude-opus-4-6
|
||||
export ANTHROPIC_DEFAULT_HAIKU_MODEL=claude-haiku-4-5-20251001
|
||||
|
||||
claude
|
||||
```
|
||||
|
||||
Or via `~/.claude/settings.json`:
|
||||
|
||||
```json
|
||||
{
|
||||
"env": {
|
||||
"CLAUDE_CODE_USE_BEDROCK": "0",
|
||||
"ANTHROPIC_BASE_URL": "http://127.0.0.1:8787",
|
||||
"ANTHROPIC_API_KEY": "headroom",
|
||||
"ANTHROPIC_MODEL": "claude-opus-4-6",
|
||||
"ANTHROPIC_DEFAULT_SONNET_MODEL": "claude-sonnet-4-6",
|
||||
"ANTHROPIC_DEFAULT_OPUS_MODEL": "claude-opus-4-6",
|
||||
"ANTHROPIC_DEFAULT_HAIKU_MODEL": "claude-haiku-4-5-20251001"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Claude Code now talks plain Anthropic `/v1/messages` to Headroom; Headroom compresses
|
||||
and forwards to Bedrock via LiteLLM, then translates the answer back.
|
||||
|
||||
## Application inference profiles (account-specific ARNs)
|
||||
|
||||
If your IAM policy only permits **application inference profiles** (account-specific
|
||||
ARNs) rather than system-defined cross-region profiles, pass the ARN directly as the
|
||||
model value in `ANTHROPIC_DEFAULT_*_MODEL`. The proxy detects `arn:aws:` prefixed model
|
||||
IDs and routes them via `bedrock/converse/<arn>` automatically — no extra configuration
|
||||
required.
|
||||
|
||||
## Region prefix notes
|
||||
|
||||
| AWS region | Cross-region inference prefix |
|
||||
|---|---|
|
||||
| `us-*` | `us.` |
|
||||
| `eu-*` | `eu.` |
|
||||
| `ap-*` (except `ap-southeast-2`) | `apac.` |
|
||||
| `ap-southeast-2` (Sydney) | `au.` |
|
||||
|
||||
The proxy uses the correct prefix automatically when constructing fallback model IDs.
|
||||
|
||||
## Verify compression is happening
|
||||
|
||||
- Dashboard: <http://localhost:8787/dashboard> — "tokens saved" climbs as you work.
|
||||
- `curl -s localhost:8787/stats` → `tokens.saved` and `request_logs[].transforms_applied`.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
| Symptom | Cause | Fix |
|
||||
|---|---|---|
|
||||
| Proxy receives no requests | Claude Code is in Bedrock mode, bypassing proxy | Set `CLAUDE_CODE_USE_BEDROCK=0` |
|
||||
| `400 The provided model identifier is invalid` | Bedrock rejected the model name format | Use standard cross-region profile names (`claude-sonnet-4-6`) or a valid application inference profile ARN |
|
||||
| `403 AccessDeniedException` on system-defined profiles | IAM policy only permits application profiles | Use `--bedrock-profile` with an authorized profile and pass application inference profile ARNs as model values |
|
||||
| `400 … Try calling via converse route` | Old proxy version routing ARNs to invoke path | Upgrade to headroom ≥ 0.27.1 |
|
||||
| Model map empty at startup | boto3 not installed or wrong AWS profile | `pip install boto3`; check `--bedrock-profile` / `AWS_PROFILE` |
|
||||
|
|
@ -71,8 +71,10 @@ _bedrock_profiles_cache: dict[str, dict[str, str]] = {} # region -> model_map
|
|||
|
||||
# Region prefix used in cross-region Bedrock inference profile IDs.
|
||||
# EU regions use "eu.", AP regions use "apac.", US (and everything else) use "us.".
|
||||
# ap-southeast-2 (Sydney/Australia) uses "au." — distinct from the rest of APAC.
|
||||
_BEDROCK_REGION_PREFIXES: dict[str, str] = {
|
||||
"eu": "eu",
|
||||
"ap-southeast-2": "au",
|
||||
"ap": "apac",
|
||||
}
|
||||
|
||||
|
|
@ -137,7 +139,9 @@ def _build_bedrock_fallback_map(region: str) -> dict[str, str]:
|
|||
return {name: f"bedrock/{prefix}.{model_id}" for name, model_id in _CLAUDE_MODELS}
|
||||
|
||||
|
||||
def _fetch_bedrock_inference_profiles(region: str | None) -> dict[str, str]:
|
||||
def _fetch_bedrock_inference_profiles(
|
||||
region: str | None, profile_name: str | None = None
|
||||
) -> dict[str, str]:
|
||||
"""Fetch available Bedrock inference profiles from AWS API.
|
||||
|
||||
Uses boto3 list_inference_profiles() to get all available profiles
|
||||
|
|
@ -149,15 +153,21 @@ def _fetch_bedrock_inference_profiles(region: str | None) -> dict[str, str]:
|
|||
|
||||
Args:
|
||||
region: AWS region (e.g., "us-east-1", "eu-central-1")
|
||||
profile_name: AWS named profile (e.g., "my-sso-profile"). When set,
|
||||
a boto3.Session is created with this profile name so
|
||||
the correct SSO or credential file is used. Falls back
|
||||
to ambient credentials (AWS_PROFILE env var, instance
|
||||
metadata, etc.) when not provided.
|
||||
|
||||
Returns:
|
||||
Model map: anthropic_model_name -> bedrock inference profile ID
|
||||
"""
|
||||
region = region or "us-east-1"
|
||||
|
||||
# Check cache first
|
||||
if region in _bedrock_profiles_cache:
|
||||
return _bedrock_profiles_cache[region]
|
||||
# Cache key includes profile_name so different profiles don't collide
|
||||
cache_key = f"{region}:{profile_name or ''}"
|
||||
if cache_key in _bedrock_profiles_cache:
|
||||
return _bedrock_profiles_cache[cache_key]
|
||||
|
||||
model_map: dict[str, str] = {}
|
||||
|
||||
|
|
@ -169,11 +179,12 @@ def _fetch_bedrock_inference_profiles(region: str | None) -> dict[str, str]:
|
|||
"Install boto3 for dynamic model discovery: pip install boto3"
|
||||
)
|
||||
model_map = _build_bedrock_fallback_map(region)
|
||||
_bedrock_profiles_cache[region] = model_map
|
||||
_bedrock_profiles_cache[cache_key] = model_map
|
||||
return model_map
|
||||
|
||||
try:
|
||||
bedrock_client = boto3.client("bedrock", region_name=region)
|
||||
session = boto3.Session(profile_name=profile_name) if profile_name else boto3.Session()
|
||||
bedrock_client = session.client("bedrock", region_name=region)
|
||||
response = bedrock_client.list_inference_profiles(typeEquals="SYSTEM_DEFINED")
|
||||
|
||||
for profile in response.get("inferenceProfileSummaries", []):
|
||||
|
|
@ -211,7 +222,7 @@ def _fetch_bedrock_inference_profiles(region: str | None) -> dict[str, str]:
|
|||
model_map = _build_bedrock_fallback_map(region)
|
||||
|
||||
# Cache the result
|
||||
_bedrock_profiles_cache[region] = model_map
|
||||
_bedrock_profiles_cache[cache_key] = model_map
|
||||
return model_map
|
||||
|
||||
|
||||
|
|
@ -222,18 +233,23 @@ def _normalize_bedrock_profile_id(profile_id: str) -> str | None:
|
|||
profile_id: e.g., "us.anthropic.claude-sonnet-4-20250514-v1:0"
|
||||
or "anthropic.claude-sonnet-4-20250514-v1:0"
|
||||
or "claude-sonnet-4-20250514"
|
||||
or "arn:aws:bedrock:...:application-inference-profile/..."
|
||||
|
||||
Returns:
|
||||
Normalized name like "claude-sonnet-4-20250514", or None if not parseable
|
||||
"""
|
||||
import re
|
||||
|
||||
# ARNs are opaque identifiers — cannot be normalized to a standard model name
|
||||
if profile_id.startswith("arn:aws:"):
|
||||
return None
|
||||
|
||||
# Strip "bedrock/" prefix if present
|
||||
if profile_id.startswith("bedrock/"):
|
||||
profile_id = profile_id[8:]
|
||||
|
||||
# Strip region prefix (us., eu., apac.)
|
||||
for prefix in ["us.", "eu.", "apac."]:
|
||||
# Strip region prefix (us., eu., apac., au.)
|
||||
for prefix in ["us.", "eu.", "apac.", "au."]:
|
||||
if profile_id.startswith(prefix):
|
||||
profile_id = profile_id[len(prefix) :]
|
||||
break
|
||||
|
|
@ -402,6 +418,7 @@ class LiteLLMBackend(Backend):
|
|||
self,
|
||||
provider: str = "bedrock",
|
||||
region: str | None = None,
|
||||
profile_name: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""Initialize LiteLLM backend.
|
||||
|
|
@ -409,6 +426,9 @@ class LiteLLMBackend(Backend):
|
|||
Args:
|
||||
provider: LiteLLM provider prefix (bedrock, vertex_ai, openrouter, etc.)
|
||||
region: Cloud region (provider-specific)
|
||||
profile_name: AWS named profile for credential resolution (bedrock only).
|
||||
When set, boto3 uses this profile (e.g. an SSO profile) instead
|
||||
of the ambient credentials. Ignored for non-bedrock providers.
|
||||
**kwargs: Additional provider-specific config
|
||||
"""
|
||||
if not LITELLM_AVAILABLE:
|
||||
|
|
@ -418,6 +438,7 @@ class LiteLLMBackend(Backend):
|
|||
|
||||
self.provider = provider
|
||||
self.region = region
|
||||
self.profile_name = profile_name
|
||||
self.kwargs = kwargs
|
||||
|
||||
# Get provider config from registry
|
||||
|
|
@ -438,7 +459,7 @@ class LiteLLMBackend(Backend):
|
|||
"botocore, which is not installed. Install the bedrock extra: "
|
||||
"pip install 'headroom-ai[bedrock]' (or pip install botocore)."
|
||||
)
|
||||
self._model_map = _fetch_bedrock_inference_profiles(region)
|
||||
self._model_map = _fetch_bedrock_inference_profiles(region, profile_name=profile_name)
|
||||
litellm.set_verbose = False # Reduce noise
|
||||
else:
|
||||
self._model_map = self._config.model_map
|
||||
|
|
@ -457,6 +478,7 @@ class LiteLLMBackend(Backend):
|
|||
- "anthropic.claude-sonnet-4-20250514-v1:0" (Bedrock without region)
|
||||
- "us.anthropic.claude-sonnet-4-20250514-v1:0" (Bedrock with region)
|
||||
- "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0" (LiteLLM format)
|
||||
- "arn:aws:bedrock:...:application-inference-profile/..." (application inference profile)
|
||||
"""
|
||||
# Check direct mapping first
|
||||
if anthropic_model in self._model_map:
|
||||
|
|
@ -464,6 +486,11 @@ class LiteLLMBackend(Backend):
|
|||
|
||||
# For Bedrock, try to normalize various input formats
|
||||
if self.provider == "bedrock":
|
||||
# Application inference profile ARNs must use the converse route —
|
||||
# the invoke route rejects ARNs with HTTP 400.
|
||||
if anthropic_model.startswith("arn:aws:"):
|
||||
return f"bedrock/converse/{anthropic_model}"
|
||||
|
||||
normalized = _normalize_bedrock_profile_id(anthropic_model)
|
||||
if normalized and normalized in self._model_map:
|
||||
return self._model_map[normalized]
|
||||
|
|
@ -696,6 +723,9 @@ class LiteLLMBackend(Backend):
|
|||
elif self.provider in ("vertex_ai", "vertex_ai_beta"):
|
||||
kwargs["vertex_location"] = self.region
|
||||
|
||||
if self.provider == "bedrock" and self.profile_name:
|
||||
kwargs["aws_profile_name"] = self.profile_name
|
||||
|
||||
# Forward API key from request headers if present.
|
||||
# Skip for Bedrock/Vertex: they use env-based auth (AWS SigV4 / Google ADC).
|
||||
# Forwarding x-api-key (e.g. sk-ant-dummy) would override their credentials.
|
||||
|
|
@ -800,6 +830,9 @@ class LiteLLMBackend(Backend):
|
|||
elif self.provider in ("vertex_ai", "vertex_ai_beta"):
|
||||
kwargs["vertex_location"] = self.region
|
||||
|
||||
if self.provider == "bedrock" and self.profile_name:
|
||||
kwargs["aws_profile_name"] = self.profile_name
|
||||
|
||||
# Forward API key from request headers if present.
|
||||
# Skip for Bedrock/Vertex: they use env-based auth (AWS SigV4 / Google ADC).
|
||||
# Forwarding x-api-key (e.g. sk-ant-dummy) would override their credentials.
|
||||
|
|
@ -1024,6 +1057,9 @@ class LiteLLMBackend(Backend):
|
|||
elif self.provider in ("vertex_ai", "vertex_ai_beta"):
|
||||
kwargs["vertex_location"] = self.region
|
||||
|
||||
if self.provider == "bedrock" and self.profile_name:
|
||||
kwargs["aws_profile_name"] = self.profile_name
|
||||
|
||||
# Forward API key from request headers if present.
|
||||
# Skip for Bedrock/Vertex: they use env-based auth (AWS SigV4 / Google ADC).
|
||||
# Forwarding x-api-key (e.g. sk-ant-dummy) would override their credentials.
|
||||
|
|
@ -1199,6 +1235,9 @@ class LiteLLMBackend(Backend):
|
|||
elif self.provider in ("vertex_ai", "vertex_ai_beta"):
|
||||
kwargs["vertex_location"] = self.region
|
||||
|
||||
if self.provider == "bedrock" and self.profile_name:
|
||||
kwargs["aws_profile_name"] = self.profile_name
|
||||
|
||||
# Forward API key from request headers if present.
|
||||
# Skip for Bedrock/Vertex: they use env-based auth (AWS SigV4 / Google ADC).
|
||||
# Forwarding x-api-key (e.g. sk-ant-dummy) would override their credentials.
|
||||
|
|
|
|||
|
|
@ -148,6 +148,7 @@ def create_proxy_backend(
|
|||
backend: str,
|
||||
anyllm_provider: str,
|
||||
bedrock_region: str | None,
|
||||
bedrock_profile: str | None = None,
|
||||
logger: logging.Logger,
|
||||
openai_api_url: str | None = None,
|
||||
anyllm_backend_cls: Any | None = None,
|
||||
|
|
@ -181,7 +182,10 @@ def create_proxy_backend(
|
|||
provider = "vertex_ai"
|
||||
try:
|
||||
backend_cls = litellm_backend_cls or _load_litellm_backend()
|
||||
instance = cast("Backend", backend_cls(provider=provider, region=bedrock_region))
|
||||
instance = cast(
|
||||
"Backend",
|
||||
backend_cls(provider=provider, region=bedrock_region, profile_name=bedrock_profile),
|
||||
)
|
||||
logger.info("LiteLLM backend enabled (provider=%s, region=%s)", provider, bedrock_region)
|
||||
return instance
|
||||
except ImportError as exc:
|
||||
|
|
|
|||
|
|
@ -945,6 +945,7 @@ class HeadroomProxy(
|
|||
backend=config.backend,
|
||||
anyllm_provider=config.anyllm_provider,
|
||||
bedrock_region=config.bedrock_region,
|
||||
bedrock_profile=config.bedrock_profile,
|
||||
logger=logger,
|
||||
openai_api_url=config.openai_api_url,
|
||||
anyllm_backend_cls=AnyLLMBackend,
|
||||
|
|
|
|||
|
|
@ -150,7 +150,9 @@ class TestFetchBedrockInferenceProfiles:
|
|||
mock_client.list_inference_profiles.side_effect = Exception(
|
||||
"AccessDeniedException: not authorized"
|
||||
)
|
||||
mock_boto3.client.return_value = mock_client
|
||||
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
|
||||
|
|
@ -188,7 +190,9 @@ class TestFetchBedrockInferenceProfiles:
|
|||
{"inferenceProfileId": "eu.meta.llama-3-70b-v1:0"}, # non-Anthropic, should skip
|
||||
]
|
||||
}
|
||||
mock_boto3.client.return_value = mock_client
|
||||
mock_session = MagicMock()
|
||||
mock_session.client.return_value = mock_client
|
||||
mock_boto3.Session.return_value = mock_session
|
||||
|
||||
import builtins
|
||||
|
||||
|
|
@ -212,13 +216,25 @@ class TestFetchBedrockInferenceProfiles:
|
|||
)
|
||||
|
||||
def test_caching_prevents_repeated_api_calls(self):
|
||||
"""Second call for same region should return cached result."""
|
||||
"""Second call for same region+profile should return cached result."""
|
||||
_bedrock_profiles_cache.clear()
|
||||
_bedrock_profiles_cache["us-east-1"] = {"test": "bedrock/test-model"}
|
||||
_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
|
||||
|
|
@ -310,6 +326,27 @@ class TestBedrockModelMapping:
|
|||
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)
|
||||
|
|
@ -352,3 +389,107 @@ class TestNormalizeBedrockProfileId:
|
|||
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
|
||||
|
|
|
|||
|
|
@ -119,7 +119,7 @@ def test_create_proxy_backend_handles_missing_litellm_backend(caplog) -> None:
|
|||
anyllm_provider="ignored",
|
||||
bedrock_region="us-east-1",
|
||||
logger=logger,
|
||||
litellm_backend_cls=lambda provider, region: (_ for _ in ()).throw(
|
||||
litellm_backend_cls=lambda provider, region, profile_name=None: (_ for _ in ()).throw(
|
||||
ImportError("missing")
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -113,7 +113,7 @@ def test_create_proxy_backend_uses_injected_backend_types() -> None:
|
|||
anyllm_provider="ignored",
|
||||
bedrock_region="us-east-1",
|
||||
logger=logger,
|
||||
litellm_backend_cls=lambda provider, region: {
|
||||
litellm_backend_cls=lambda provider, region, profile_name=None: {
|
||||
"kind": "litellm",
|
||||
"provider": provider,
|
||||
"region": region,
|
||||
|
|
|
|||
|
|
@ -78,7 +78,9 @@ def _capture_provider(backend: str) -> dict[str, Any]:
|
|||
captured: dict[str, Any] = {}
|
||||
|
||||
class FakeLiteLLM:
|
||||
def __init__(self, provider: str, region: str | None = None) -> None:
|
||||
def __init__(
|
||||
self, provider: str, region: str | None = None, profile_name: str | None = None
|
||||
) -> None:
|
||||
captured["provider"] = provider
|
||||
captured["region"] = region
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue