mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge a86ca4152b into b781d157d7
This commit is contained in:
commit
191b0a4a21
2 changed files with 51 additions and 2 deletions
|
|
@ -45,6 +45,7 @@ from litellm.proxy.common_utils.sse_keepalive import (
|
|||
wrap_sse_stream_with_keepalive_pings,
|
||||
)
|
||||
from litellm.proxy.utils import ProxyLogging, get_custom_url
|
||||
from litellm.types.agents import AGENT_CALLER_TEAM_ID_HEADER, AGENT_CALLER_USER_ID_HEADER
|
||||
from litellm.types.utils import all_litellm_params
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -172,6 +173,21 @@ async def _resolve_backend_auth_header(
|
|||
return await resolve_a2a_hop_auth_header(litellm_params, custom_llm_provider)
|
||||
|
||||
|
||||
# Headers the proxy mints itself on the backend call. Forged client copies of
|
||||
# these must not be forwarded (agent-id is minted later in asend_message, where
|
||||
# agent_extra_headers would otherwise overwrite it), but any other x-litellm-*
|
||||
# header (e.g. an admin-allowlisted x-litellm-api-key from extra_headers) is
|
||||
# valid passthrough.
|
||||
_MINTED_A2A_IDENTITY_HEADERS: Final = frozenset(
|
||||
{
|
||||
AGENT_CALLER_USER_ID_HEADER.lower(),
|
||||
AGENT_CALLER_TEAM_ID_HEADER.lower(),
|
||||
"x-litellm-trace-id",
|
||||
"x-litellm-agent-id",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _forwarding_headers(
|
||||
caller_identity: Mapping[str, str],
|
||||
request_data: Mapping[str, object],
|
||||
|
|
@ -179,11 +195,13 @@ def _forwarding_headers(
|
|||
backend_auth_header: Mapping[str, str] | None,
|
||||
) -> dict[str, str] | None:
|
||||
backend_auth: Final = tuple(backend_auth_header.items()) if backend_auth_header else ()
|
||||
minted_names: Final = frozenset(name.lower() for name, _ in backend_auth)
|
||||
minted_names: Final = frozenset(
|
||||
name.lower() for name, _ in backend_auth
|
||||
) | _MINTED_A2A_IDENTITY_HEADERS
|
||||
passthrough: Final = tuple(
|
||||
(name, value)
|
||||
for name, value in (agent_extra_headers.items() if agent_extra_headers else ())
|
||||
if not name.lower().startswith("x-litellm-") and name.lower() not in minted_names
|
||||
if name.lower() not in minted_names
|
||||
)
|
||||
trace_id: Final = request_data.get("litellm_trace_id")
|
||||
trace: Final = (("X-LiteLLM-Trace-Id", str(trace_id)),) if trace_id else ()
|
||||
|
|
|
|||
|
|
@ -645,6 +645,37 @@ async def test_message_methods_caller_identity_headers_cannot_be_spoofed(method:
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
|
||||
async def test_message_methods_forward_allowlisted_x_litellm_api_key(method: str):
|
||||
"""Issue #43450: an x-litellm-* header the admin allowlisted in extra_headers must be
|
||||
forwarded to the backend agent, while forged copies of the proxy-minted identity
|
||||
headers are still stripped."""
|
||||
agent = _make_agent_mock()
|
||||
agent.extra_headers = ["x-litellm-api-key"]
|
||||
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
|
||||
mock_request.headers = {
|
||||
"x-litellm-api-key": "sk-allowlisted",
|
||||
"x-a2a-test-agent-x-litellm-team-id": "attacker-team",
|
||||
"x-a2a-test-agent-x-litellm-agent-id": "attacker-agent",
|
||||
}
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="real-user", team_id="real-team")
|
||||
|
||||
captured = await _invoke_message_method(method, mock_request, user_api_key_dict, agent=agent)
|
||||
|
||||
forwarded_headers = captured.agent_extra_headers or {}
|
||||
assert forwarded_headers.get("x-litellm-api-key") == "sk-allowlisted", (
|
||||
"admin-allowlisted x-litellm-api-key must reach the backend agent"
|
||||
)
|
||||
assert forwarded_headers.get("X-LiteLLM-User-Id") == "real-user"
|
||||
assert forwarded_headers.get("X-LiteLLM-Team-Id") == "real-team", (
|
||||
"proxy-minted team id must not be overridden by a forged client header"
|
||||
)
|
||||
assert "x-litellm-agent-id" not in {k.lower() for k in forwarded_headers}, (
|
||||
"forged agent id must be stripped before asend_message mints the real one"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
|
||||
async def test_message_methods_forward_key_bound_identity_not_pre_call_rewrite(method: str):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue