From 62f903c0a2fd578179a3bb8d0850debe6ccda0a4 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 28 Feb 2026 17:36:58 -0800 Subject: [PATCH] fix(guardrails): address review feedback - SSRF, silent failure, redundant query - Validate api_base URL scheme (http/https only) and hostname in register_guardrail to prevent SSRF via team submissions - Return warning field in approve response when in-memory initialization fails so admins know the guardrail won't work until next sync cycle - Eliminate redundant DB query in list_guardrail_submissions by fetching all team guardrails once and deriving both filtered list and summary counts from the single result set Co-Authored-By: Claude Opus 4.6 --- .../proxy/guardrails/guardrail_endpoints.py | 50 +++-- .../guardrails/test_guardrail_endpoints.py | 208 +++++++++++++++++- 2 files changed, 235 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 3e890fc83ec..57b4529eed3 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -7,6 +7,7 @@ import inspect import json from datetime import datetime, timezone from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast +from urllib.parse import urlparse from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel @@ -616,11 +617,23 @@ async def register_guardrail( status_code=400, detail=f"Only guardrails with litellm_params.guardrail={GENERIC_GUARDRAIL_API!r} are accepted for registration", ) - if not params.get("api_base"): + api_base = params.get("api_base") + if not api_base: raise HTTPException( status_code=400, detail="litellm_params.api_base is required for generic_guardrail_api", ) + parsed = urlparse(api_base) + if parsed.scheme not in ("http", "https"): + raise HTTPException( + status_code=400, + detail="litellm_params.api_base must use http or https scheme", + ) + if not parsed.hostname: + raise HTTPException( + status_code=400, + detail="litellm_params.api_base must contain a valid hostname", + ) mode = params.get("mode") if mode is None: raise HTTPException( @@ -742,22 +755,13 @@ async def list_guardrail_submissions( raise HTTPException(status_code=500, detail="Prisma client not initialized") try: - # Only team guardrails (team_id is not null) - where: Dict[str, Any] = {"team_id": {"not": None}} - if status: - where["status"] = status - if team_id: - where["team_id"] = team_id - - rows = await prisma_client.db.litellm_guardrailstable.find_many( - where=where, + # Single query: fetch all team guardrails (team_id is not null) + all_team_rows = await prisma_client.db.litellm_guardrailstable.find_many( + where={"team_id": {"not": None}}, order={"created_at": "desc"}, ) - # Summary counts (team guardrails only) - all_team_rows = await prisma_client.db.litellm_guardrailstable.find_many( - where={"team_id": {"not": None}} - ) + # Derive summary counts from the full result set total = len(all_team_rows) pending_review = sum( 1 for r in all_team_rows if (r.status or "active") == "pending_review" @@ -769,6 +773,12 @@ async def list_guardrail_submissions( 1 for r in all_team_rows if (r.status or "active") == "rejected" ) + # Apply filters to get the submissions list + rows = all_team_rows + if status: + rows = [r for r in rows if r.status == status] + if team_id: + rows = [r for r in rows if r.team_id == team_id] if search: search_lower = search.lower() rows = [ @@ -785,9 +795,8 @@ async def list_guardrail_submissions( and search_lower in r.guardrail_info.lower() ) ] - items = [] - for r in rows: - items.append(_row_to_submission_item(r)) + + items = [_row_to_submission_item(r) for r in rows] return ListGuardrailSubmissionsResponse( submissions=items, summary=GuardrailSubmissionSummary( @@ -902,6 +911,13 @@ async def approve_guardrail_submission( guardrail_id, init_err, ) + return { + "guardrail_id": guardrail_id, + "status": "active", + "message": "Guardrail approved", + "warning": f"Guardrail was marked active but failed to initialize in memory: {init_err}. " + "It will be picked up on the next sync cycle.", + } return { "guardrail_id": guardrail_id, diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 9ec00994b89..c7147f6a785 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -1159,7 +1159,7 @@ async def test_register_guardrail_rejects_non_generic_api(mocker): guardrail_name="other-guard", litellm_params={"guardrail": "bedrock", "mode": "pre_call", "api_base": "https://x.com"}, ) - user = UserAPIKeyAuth(user_id="u1", user_email="a@b.com") + user = UserAPIKeyAuth(user_id="u1", user_email="a@b.com", team_id="team-1") with pytest.raises(HTTPException) as exc_info: await register_guardrail(req, user) @@ -1259,9 +1259,9 @@ async def test_list_guardrail_submissions_returns_only_team_guardrails(mocker): @pytest.mark.asyncio async def test_list_guardrail_submissions_team_id_filter(mocker): - """List submissions with team_id returns only that team's guardrails.""" + """List submissions with team_id filter returns only that team's guardrails.""" mock_prisma = mocker.Mock() - row = mocker.Mock( + row_abc = mocker.Mock( guardrail_id="team-1", guardrail_name="team-guard", status="active", @@ -1273,7 +1273,19 @@ async def test_list_guardrail_submissions_team_id_filter(mocker): created_at=datetime.now(), updated_at=datetime.now(), ) - find_many = AsyncMock(side_effect=[[row], [row]]) + row_other = mocker.Mock( + guardrail_id="team-2", + guardrail_name="other-guard", + status="active", + team_id="team-xyz", + litellm_params={}, + guardrail_info={}, + submitted_at=None, + reviewed_at=None, + created_at=datetime.now(), + updated_at=datetime.now(), + ) + find_many = AsyncMock(return_value=[row_abc, row_other]) mock_prisma.db.litellm_guardrailstable.find_many = find_many mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) @@ -1283,8 +1295,9 @@ async def test_list_guardrail_submissions_team_id_filter(mocker): ) assert len(result.submissions) == 1 + assert result.submissions[0].guardrail_id == "team-1" assert result.submissions[0].team_guardrail is True - assert find_many.call_args_list[0].kwargs.get("where", {}).get("team_id") == "team-abc" + assert result.summary.total == 2 # summary counts all team guardrails @pytest.mark.asyncio @@ -1360,4 +1373,187 @@ async def test_reject_guardrail_submission_success(mocker): assert result["status"] == "rejected" mock_prisma.db.litellm_guardrailstable.update.assert_called_once() call_data = mock_prisma.db.litellm_guardrailstable.update.call_args[1]["data"] - assert call_data["status"] == "rejected" \ No newline at end of file + assert call_data["status"] == "rejected" + + +# --- Tests for review fixes --- + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "api_base,expected_detail", + [ + ("file:///etc/passwd", "http or https scheme"), + ("ftp://internal.host/data", "http or https scheme"), + ("javascript:alert(1)", "http or https scheme"), + ("://missing-scheme", "http or https scheme"), + ("https://", "valid hostname"), + ], + ids=[ + "file_scheme", + "ftp_scheme", + "javascript_scheme", + "no_scheme", + "no_hostname", + ], +) +async def test_register_guardrail_rejects_bad_api_base(mocker, api_base, expected_detail): + """Register returns 400 when api_base has invalid scheme or missing hostname.""" + mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock()) + req = RegisterGuardrailRequest( + guardrail_name="bad-url-guard", + litellm_params={ + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "api_base": api_base, + }, + ) + user = UserAPIKeyAuth(user_id="u1", user_email="a@b.com", team_id="team-1") + + with pytest.raises(HTTPException) as exc_info: + await register_guardrail(req, user) + assert exc_info.value.status_code == 400 + assert expected_detail in exc_info.value.detail + + +@pytest.mark.asyncio +async def test_register_guardrail_accepts_valid_https_url(mocker): + """Register accepts valid https api_base URLs.""" + mock_prisma = mocker.Mock() + mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=None) + created_row = mocker.Mock( + guardrail_id="valid-url-123", + guardrail_name="valid-guard", + status="pending_review", + submitted_at=datetime.now(), + ) + mock_prisma.db.litellm_guardrailstable.create = AsyncMock(return_value=created_row) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) + + req = RegisterGuardrailRequest( + guardrail_name="valid-guard", + litellm_params={ + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "api_base": "https://guardrails.example.com/v1/check", + }, + ) + user = UserAPIKeyAuth(user_id="u1", user_email="a@b.com", team_id="team-1") + + result = await register_guardrail(req, user) + assert result.guardrail_id == "valid-url-123" + assert result.status == "pending_review" + + +@pytest.mark.asyncio +async def test_approve_guardrail_init_failure_returns_warning(mocker): + """Approve returns a warning field when in-memory initialization fails.""" + mock_prisma = mocker.Mock() + row = mocker.Mock( + guardrail_id="warn-me", + guardrail_name="fragile-guard", + status="pending_review", + litellm_params={ + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "api_base": "https://g.com", + }, + guardrail_info={}, + ) + mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.db.litellm_guardrailstable.update = AsyncMock() + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) + + mock_handler = mocker.Mock() + mock_handler.initialize_guardrail = mocker.Mock( + side_effect=Exception("missing dependency") + ) + mocker.patch( + "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", + mock_handler, + ) + user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + result = await approve_guardrail_submission("warn-me", user) + + assert result["status"] == "active" + assert "warning" in result + assert "failed to initialize" in result["warning"].lower() + assert "missing dependency" in result["warning"] + + +@pytest.mark.asyncio +async def test_approve_guardrail_no_warning_on_success(mocker): + """Approve does NOT include a warning field when init succeeds.""" + mock_prisma = mocker.Mock() + row = mocker.Mock( + guardrail_id="ok-guard", + guardrail_name="good-guard", + status="pending_review", + litellm_params={ + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "api_base": "https://g.com", + }, + guardrail_info={}, + ) + mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.db.litellm_guardrailstable.update = AsyncMock() + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) + + mock_handler = mocker.Mock() + mock_handler.initialize_guardrail = mocker.Mock() # no exception + mocker.patch( + "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", + mock_handler, + ) + user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + result = await approve_guardrail_submission("ok-guard", user) + + assert result["status"] == "active" + assert "warning" not in result + + +@pytest.mark.asyncio +async def test_list_submissions_single_db_query(mocker): + """List submissions makes exactly one find_many call (no redundant query).""" + mock_prisma = mocker.Mock() + find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_guardrailstable.find_many = find_many + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) + user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + await list_guardrail_submissions(user_api_key_dict=user) + + assert find_many.call_count == 1 + + +@pytest.mark.asyncio +async def test_list_submissions_summary_counts_unaffected_by_filters(mocker): + """Summary counts reflect all team guardrails regardless of status filter.""" + mock_prisma = mocker.Mock() + pending_row = mocker.Mock( + guardrail_id="p1", guardrail_name="p", status="pending_review", + team_id="t1", litellm_params={}, guardrail_info={}, + submitted_at=None, reviewed_at=None, + created_at=datetime.now(), updated_at=datetime.now(), + ) + active_row = mocker.Mock( + guardrail_id="a1", guardrail_name="a", status="active", + team_id="t1", litellm_params={}, guardrail_info={}, + submitted_at=None, reviewed_at=None, + created_at=datetime.now(), updated_at=datetime.now(), + ) + all_rows = [pending_row, active_row] + mock_prisma.db.litellm_guardrailstable.find_many = AsyncMock(return_value=all_rows) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) + user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + # Filter to only pending, but summary should still show both + result = await list_guardrail_submissions(status="pending_review", user_api_key_dict=user) + + assert len(result.submissions) == 1 # filtered + assert result.summary.total == 2 # unfiltered + assert result.summary.pending_review == 1 + assert result.summary.active == 1 \ No newline at end of file