mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(rubrik): sanitize proxy_server_request and harden tool_calls parsing
Address bugbot review concerns: - Sanitize proxy_server_request before forwarding to the Rubrik webhook. The previous code passed the entire inbound HTTP context (Authorization, Cookie, x-api-key, and the raw request body) through to a third-party endpoint, which exfiltrates proxy credentials and upstream secrets. The new _sanitize_proxy_server_request allowlists only url and method. (Cursor Bugbot HIGH severity #3192354895) - Treat a null choices[0].message.tool_calls as 'all blocked' rather than letting iteration raise and silently fall through the outer except in apply_guardrail (which would fail open). Iterate over a defensive fallback list instead of relying on the dict default. (Cursor Bugbot MEDIUM severity #3192349538) Co-authored-by: Cursor Bugbot <bugbot@cursor.com>
This commit is contained in:
parent
24bbf4988e
commit
acb37dc664
2 changed files with 113 additions and 2 deletions
|
|
@ -291,7 +291,23 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
return {
|
||||
"messages": call_details.get("messages"),
|
||||
"model": call_details.get("model"),
|
||||
"proxy_server_request": litellm_params.get("proxy_server_request"),
|
||||
"proxy_server_request": RubrikLogger._sanitize_proxy_server_request(
|
||||
litellm_params.get("proxy_server_request")
|
||||
),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _sanitize_proxy_server_request(proxy_server_request: Any) -> Any:
|
||||
"""Allowlist only routing fields (``url``, ``method``) when forwarding
|
||||
``proxy_server_request`` to the external Rubrik webhook, dropping
|
||||
inbound ``headers`` (Authorization, Cookie, x-api-key, ...) and the raw
|
||||
request ``body`` so proxy credentials are not exfiltrated."""
|
||||
if not isinstance(proxy_server_request, dict):
|
||||
return proxy_server_request
|
||||
return {
|
||||
key: proxy_server_request[key]
|
||||
for key in ("url", "method")
|
||||
if key in proxy_server_request
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -503,7 +519,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
raise Exception("Tool blocking service returned empty response")
|
||||
|
||||
message = choices[0].get("message", {})
|
||||
returned_tool_calls = message.get("tool_calls", [])
|
||||
returned_tool_calls = message.get("tool_calls") or []
|
||||
blocking_explanation = message.get("content", "")
|
||||
|
||||
allowed_ids = {tc["id"] for tc in returned_tool_calls if tc.get("id")}
|
||||
|
|
|
|||
|
|
@ -654,6 +654,51 @@ class TestApplyGuardrail:
|
|||
assert req["model"] == "gpt-4"
|
||||
assert req["messages"] == [{"role": "user", "content": "hi"}]
|
||||
|
||||
async def test_proxy_server_request_headers_stripped(self, handler):
|
||||
tc = make_tool_call_dict("call_1", "test_tool")
|
||||
inputs = make_inputs_with_tools([tc])
|
||||
|
||||
captured_payload: Dict[str, Any] = {}
|
||||
|
||||
async def mock_post(*_args, **kwargs):
|
||||
captured_payload.update(kwargs.get("json", {}))
|
||||
mock_resp = Mock()
|
||||
mock_resp.json.return_value = captured_payload.get("response", {})
|
||||
mock_resp.raise_for_status = Mock()
|
||||
return mock_resp
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post = mock_post
|
||||
handler.tool_blocking_client = mock_client
|
||||
|
||||
logging_obj = Mock()
|
||||
logging_obj.model_call_details = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"model": "gpt-4",
|
||||
"litellm_params": {
|
||||
"proxy_server_request": {
|
||||
"url": "/chat/completions",
|
||||
"method": "POST",
|
||||
"headers": {
|
||||
"authorization": "Bearer sk-litellm-secret",
|
||||
"cookie": "session=abc",
|
||||
"x-api-key": "leaked-key",
|
||||
},
|
||||
"body": {"api_key": "sk-upstream-secret"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
await handler.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={},
|
||||
input_type="response",
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
forwarded = captured_payload["request"]["proxy_server_request"]
|
||||
assert forwarded == {"url": "/chat/completions", "method": "POST"}
|
||||
|
||||
|
||||
# -- Anthropic format ----------------------------------------------------------
|
||||
|
||||
|
|
@ -818,6 +863,56 @@ class TestExtractBlockedTools:
|
|||
with pytest.raises(Exception, match="empty response"):
|
||||
RubrikLogger._extract_blocked_tools({"choices": []}, [])
|
||||
|
||||
def test_null_tool_calls_treated_as_all_blocked(self):
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Function
|
||||
|
||||
tc = ChatCompletionMessageToolCall(
|
||||
id="call_1", type="function", function=Function(name="fn", arguments="{}")
|
||||
)
|
||||
service_resp = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"tool_calls": None,
|
||||
"content": "blocked everything",
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
result = RubrikLogger._extract_blocked_tools(service_resp, [tc])
|
||||
assert result is not None
|
||||
assert result.allowed_tools == []
|
||||
assert "blocked everything" in result.explanation
|
||||
|
||||
|
||||
# -- Sanitize proxy server request -------------------------------------------
|
||||
|
||||
|
||||
class TestSanitizeProxyServerRequest:
|
||||
def test_drops_headers_and_body(self):
|
||||
proxy_request = {
|
||||
"url": "/chat/completions",
|
||||
"method": "POST",
|
||||
"headers": {
|
||||
"authorization": "Bearer sk-litellm-secret",
|
||||
"cookie": "session=abc",
|
||||
"content-type": "application/json",
|
||||
},
|
||||
"body": {"api_key": "sk-upstream-secret", "model": "gpt-4"},
|
||||
}
|
||||
result = RubrikLogger._sanitize_proxy_server_request(proxy_request)
|
||||
assert result == {"url": "/chat/completions", "method": "POST"}
|
||||
|
||||
def test_none_passthrough(self):
|
||||
assert RubrikLogger._sanitize_proxy_server_request(None) is None
|
||||
|
||||
def test_non_dict_passthrough(self):
|
||||
assert RubrikLogger._sanitize_proxy_server_request("not a dict") == "not a dict"
|
||||
|
||||
def test_partial_dict(self):
|
||||
result = RubrikLogger._sanitize_proxy_server_request({"url": "/v1/messages"})
|
||||
assert result == {"url": "/v1/messages"}
|
||||
|
||||
|
||||
# -- Resolve model -------------------------------------------------------------
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue