Merge pull request #41504 from BerriAI/litellm_bedrock_agent_runtime_strip_virtual_key

fix(proxy): stop forwarding LiteLLM credential headers on Bedrock agent-runtime passthrough
This commit is contained in:
Mateo Wang 2026-09-16 17:02:17 -07:00 • committed by GitHub
commit 6cdf398bea
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 153 additions and 2 deletions

View file

@ -1180,9 +1180,8 @@ async def bedrock_proxy_route(
endpoint_func: Final = create_pass_through_route(
endpoint=endpoint,
target=str(prepped.url),
custom_headers=prepped.headers,
custom_headers=_upstream_headers_for_bedrock_agent_runtime_route(request, user_api_key_dict, prepped.headers),
is_streaming_request=is_streaming_request,
_forward_headers=True,
) # dynamically construct pass-through endpoint based on incoming path
setattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, data)
# SigV4 signs an exact payload; pass-through must send prepped.body, not json.dumps
@ -2001,6 +2000,9 @@ _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS: Final = frozenset({"authorization", "x-a
_HEADERS_NEVER_FORWARDED_TO_ANTHROPIC: Final = frozenset({"content-length", "host", "accept-encoding"}) | (
SpecialHeaders.litellm_credential_header_names() - _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS
)
_HEADERS_NEVER_FORWARDED_TO_BEDROCK: Final = (
frozenset({"content-length", "host", "accept-encoding"}) | SpecialHeaders.litellm_credential_header_names()
)
_MAPPED_ROUTE_CALLER_KEY_HEADER: Final = "litellm_user_api_key"
@ -2099,6 +2101,17 @@ def _upstream_headers_for_anthropic_route(
return MappingProxyType({**caller_headers, **(proxy_auth_header or {})})
def _upstream_headers_for_bedrock_agent_runtime_route(
request: Request, user_api_key_dict: UserAPIKeyAuth, signed_headers: Mapping[str, object]
) -> Mapping[str, object]:
caller_headers: Final = _caller_headers_without_litellm_secrets(
request,
user_api_key_dict,
_HEADERS_NEVER_FORWARDED_TO_BEDROCK | frozenset(name.lower() for name in signed_headers),
)
return MappingProxyType({**caller_headers, **signed_headers})
async def _prepare_vertex_auth_headers(
request: Request,
vertex_credentials: VertexPassThroughCredentials | None,

View file

@ -1983,6 +1983,144 @@ class TestBedrockAgentRuntimePassthroughToggle:
create_route.assert_called_once()
class TestBedrockAgentRuntimePassthroughVirtualKeyLeak:
VKEY: Final = "sk-litellm-victim-key"
MASTER_KEY: Final = "sk-master-1234"
ENDPOINT: Final = "knowledgebases/KB1234567/retrieve"
AMBIENT_AWS_ENV: Final = (
"AWS_BEARER_TOKEN_BEDROCK",
"AWS_SESSION_TOKEN",
"AWS_SESSION_NAME",
"AWS_PROFILE_NAME",
"AWS_ROLE_NAME",
"AWS_WEB_IDENTITY_TOKEN",
"AWS_STS_ENDPOINT",
"AWS_EXTERNAL_ID",
)
async def _upstream_headers(self, monkeypatch, headers: list[tuple[bytes, bytes]]) -> dict:
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import HttpPassThroughEndpointHelpers
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", self.MASTER_KEY)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
for ambient in self.AMBIENT_AWS_ENV:
monkeypatch.delenv(ambient, raising=False)
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "ak")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "sk")
monkeypatch.setenv("AWS_REGION_NAME", "us-east-1")
caller: Final = UserAPIKeyAuth(api_key=self.VKEY)
async def receive():
return {"type": "http.request", "body": b'{"retrievalQuery": {"text": "hi"}}', "more_body": False}
request: Final = Request(
{
"type": "http",
"method": "POST",
"path": f"/bedrock/{self.ENDPOINT}",
"headers": headers,
"query_string": b"",
},
receive=receive,
)
captured: dict = {}
def fake_create_pass_through_route(**kwargs):
captured.update(kwargs)
return AsyncMock(return_value={"status": "success"})
module: Final = "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints"
with (
patch(f"{module}.create_request_copy", Mock()),
patch(f"{module}.create_pass_through_route", side_effect=fake_create_pass_through_route),
):
await bedrock_proxy_route(
endpoint=self.ENDPOINT,
request=request,
fastapi_response=Response(),
user_api_key_dict=caller,
)
return HttpPassThroughEndpointHelpers.forward_headers_from_request(
request_headers=dict(request.headers),
headers=dict(captured["custom_headers"] or {}),
forward_headers=captured.get("_forward_headers", False),
)
@staticmethod
def _blob(upstream: dict) -> str:
return " ".join(f"{name}:{value}" for name, value in upstream.items())
@staticmethod
def _names_matching(upstream: dict, lowercase_name: str) -> list[str]:
return [name for name in upstream if name.lower() == lowercase_name]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"header_name", ["x-api-key", "x-litellm-api-key", "api-key", "x-goog-api-key", "ocp-apim-subscription-key"]
)
async def test_virtual_key_in_a_credential_header_never_reaches_aws(self, monkeypatch, header_name: str):
upstream: Final = await self._upstream_headers(
monkeypatch,
[
(header_name.encode(), self.VKEY.encode()),
(b"content-type", b"application/json"),
(b"x-request-id", b"trace-1"),
],
)
assert self.VKEY not in self._blob(upstream)
assert self._names_matching(upstream, header_name) == []
assert upstream["x-request-id"] == "trace-1", "a benign caller header still reaches AWS"
assert upstream["Authorization"].startswith("AWS4-HMAC-SHA256")
assert self._names_matching(upstream, "content-type") == ["Content-Type"], "the signed header is the only one"
@pytest.mark.asyncio
async def test_credential_headers_are_dropped_by_name_even_when_they_carry_someone_elses_key(self, monkeypatch):
other_key: Final = "sk-other-tenant-key"
upstream: Final = await self._upstream_headers(
monkeypatch,
[
(b"x-api-key", other_key.encode()),
(b"x-litellm-api-key", other_key.encode()),
(b"x-request-id", b"trace-3"),
],
)
assert other_key not in self._blob(upstream)
assert self._names_matching(upstream, "x-api-key") == []
assert self._names_matching(upstream, "x-litellm-api-key") == []
assert upstream["x-request-id"] == "trace-3"
@pytest.mark.asyncio
async def test_virtual_key_in_authorization_bearer_is_replaced_by_the_sigv4_signature(self, monkeypatch):
upstream: Final = await self._upstream_headers(
monkeypatch,
[(b"authorization", f"Bearer {self.VKEY}".encode()), (b"content-type", b"application/json")],
)
assert self.VKEY not in self._blob(upstream)
assert self._names_matching(upstream, "authorization") == ["Authorization"]
assert upstream["Authorization"].startswith("AWS4-HMAC-SHA256")
@pytest.mark.asyncio
async def test_authenticated_secrets_in_any_other_header_never_reach_aws(self, monkeypatch):
upstream: Final = await self._upstream_headers(
monkeypatch,
[
(b"x-api-key", self.VKEY.encode()),
(b"x-forwarded-key", self.VKEY.encode()),
(b"x-operator-token", self.MASTER_KEY.encode()),
(b"x-request-id", b"trace-2"),
],
)
assert self.VKEY not in self._blob(upstream) and self.MASTER_KEY not in self._blob(upstream)
assert self._names_matching(upstream, "x-forwarded-key") == []
assert self._names_matching(upstream, "x-operator-token") == []
assert upstream["x-request-id"] == "trace-2"
class TestLLMPassthroughFactoryProxyRoute:
@pytest.mark.asyncio
async def test_llm_passthrough_factory_proxy_route_success(self):