mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
commit
6cdf398bea
2 changed files with 153 additions and 2 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue