mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(route_checks.py): support team metadata
This commit is contained in:
parent
82d7a7248e
commit
5aa5a3b425
2 changed files with 315 additions and 25 deletions
|
|
@ -472,14 +472,25 @@ class RouteChecks:
|
|||
Supports both exact match and prefix match.
|
||||
"""
|
||||
metadata = user_api_key_dict.metadata
|
||||
if metadata is None:
|
||||
team_metadata = user_api_key_dict.team_metadata or {}
|
||||
if metadata is None and team_metadata is None:
|
||||
return False
|
||||
if "allowed_passthrough_routes" not in metadata:
|
||||
if (
|
||||
"allowed_passthrough_routes" not in metadata
|
||||
and "allowed_passthrough_routes" not in team_metadata
|
||||
):
|
||||
return False
|
||||
if metadata["allowed_passthrough_routes"] is None:
|
||||
if (
|
||||
metadata.get("allowed_passthrough_routes") is None
|
||||
and team_metadata.get("allowed_passthrough_routes") is None
|
||||
):
|
||||
return False
|
||||
|
||||
allowed_passthrough_routes = metadata["allowed_passthrough_routes"]
|
||||
allowed_passthrough_routes = (
|
||||
metadata.get("allowed_passthrough_routes")
|
||||
or team_metadata.get("allowed_passthrough_routes")
|
||||
or []
|
||||
)
|
||||
|
||||
# Check if route matches any allowed passthrough route (exact or prefix match)
|
||||
for allowed_route in allowed_passthrough_routes:
|
||||
|
|
|
|||
|
|
@ -130,22 +130,24 @@ def test_virtual_key_allowed_routes_with_litellm_routes_member_name_denied():
|
|||
assert "Only allowed to call routes: ['info_routes']" in str(exc_info.value)
|
||||
assert "Tried to call route: /chat/completions" in str(exc_info.value)
|
||||
|
||||
@pytest.mark.parametrize("route", [
|
||||
"/anthropic/v1/messages",
|
||||
"/anthropic/v1/count_tokens",
|
||||
"/gemini/v1/models",
|
||||
"/gemini/countTokens",
|
||||
])
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route",
|
||||
[
|
||||
"/anthropic/v1/messages",
|
||||
"/anthropic/v1/count_tokens",
|
||||
"/gemini/v1/models",
|
||||
"/gemini/countTokens",
|
||||
],
|
||||
)
|
||||
def test_virtual_key_llm_api_route_includes_passthrough_prefix(route):
|
||||
"""
|
||||
Virtual key with llm_api_routes should allow passthrough routes like /anthropic/v1/messages
|
||||
|
||||
|
||||
Relevant issue: https://github.com/BerriAI/litellm/issues/14017
|
||||
"""
|
||||
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user", allowed_routes=["llm_api_routes"]
|
||||
)
|
||||
valid_token = UserAPIKeyAuth(user_id="test_user", allowed_routes=["llm_api_routes"])
|
||||
|
||||
result = RouteChecks.is_virtual_key_allowed_to_call_route(
|
||||
route=route, valid_token=valid_token
|
||||
|
|
@ -232,7 +234,7 @@ def test_virtual_key_allowed_routes_with_no_member_names_only_explicit():
|
|||
|
||||
def test_anthropic_count_tokens_route_is_llm_api_route():
|
||||
"""Test that /v1/messages/count_tokens is recognized as an LLM API route for Anthropic"""
|
||||
|
||||
|
||||
# Test the core anthropic routes
|
||||
assert RouteChecks.is_llm_api_route("/v1/messages") is True
|
||||
assert RouteChecks.is_llm_api_route("/v1/messages/count_tokens") is True
|
||||
|
|
@ -240,11 +242,11 @@ def test_anthropic_count_tokens_route_is_llm_api_route():
|
|||
|
||||
def test_anthropic_count_tokens_route_accessible_to_internal_users():
|
||||
"""Test that internal users can access the Anthropic count_tokens route"""
|
||||
|
||||
|
||||
# Test that the route is recognized as an LLM API route (which means it's accessible to internal users)
|
||||
# This is the core check that was failing in the original issue
|
||||
assert RouteChecks.is_llm_api_route("/v1/messages/count_tokens") is True
|
||||
|
||||
|
||||
# Also test that the regular messages route still works
|
||||
assert RouteChecks.is_llm_api_route("/v1/messages") is True
|
||||
|
||||
|
|
@ -252,7 +254,7 @@ def test_anthropic_count_tokens_route_accessible_to_internal_users():
|
|||
def test_virtual_key_llm_api_routes_allows_registered_pass_through_endpoints():
|
||||
"""
|
||||
Test that virtual keys with llm_api_routes permission can access registered pass-through endpoints.
|
||||
|
||||
|
||||
This tests the scenario where a pass-through endpoint is registered from the DB
|
||||
(e.g., /azure-assistant) and a virtual key with llm_api_routes permission should be able to access
|
||||
both the exact path and subpaths (e.g., /azure-assistant/openai/assistants).
|
||||
|
|
@ -272,7 +274,7 @@ def test_virtual_key_llm_api_routes_allows_registered_pass_through_endpoints():
|
|||
"type": "subpath",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._registered_pass_through_routes",
|
||||
mock_registered_routes,
|
||||
|
|
@ -282,21 +284,21 @@ def test_virtual_key_llm_api_routes_allows_registered_pass_through_endpoints():
|
|||
user_id="test_user",
|
||||
allowed_routes=["llm_api_routes"],
|
||||
)
|
||||
|
||||
|
||||
# Test exact match for registered pass-through endpoint
|
||||
result1 = RouteChecks.is_virtual_key_allowed_to_call_route(
|
||||
route="/azure-assistant",
|
||||
valid_token=valid_token,
|
||||
)
|
||||
assert result1 is True
|
||||
|
||||
|
||||
# Test subpath for registered pass-through endpoint with subpath type
|
||||
result2 = RouteChecks.is_virtual_key_allowed_to_call_route(
|
||||
route="/custom-endpoint/openai/assistants",
|
||||
valid_token=valid_token,
|
||||
)
|
||||
assert result2 is True
|
||||
|
||||
|
||||
# Test exact match for subpath type
|
||||
result3 = RouteChecks.is_virtual_key_allowed_to_call_route(
|
||||
route="/custom-endpoint",
|
||||
|
|
@ -319,7 +321,7 @@ def test_virtual_key_without_llm_api_routes_cannot_access_pass_through():
|
|||
"type": "exact",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._registered_pass_through_routes",
|
||||
mock_registered_routes,
|
||||
|
|
@ -329,12 +331,289 @@ def test_virtual_key_without_llm_api_routes_cannot_access_pass_through():
|
|||
user_id="test_user",
|
||||
allowed_routes=["info_routes"],
|
||||
)
|
||||
|
||||
|
||||
# Test that access is denied
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
RouteChecks.is_virtual_key_allowed_to_call_route(
|
||||
route="/azure-assistant",
|
||||
valid_token=valid_token,
|
||||
)
|
||||
|
||||
|
||||
assert "Virtual key is not allowed to call this route" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_key_metadata_exact_match():
|
||||
"""Test that key metadata allowed_passthrough_routes allows exact match"""
|
||||
|
||||
# Create a UserAPIKeyAuth with allowed_passthrough_routes in metadata
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={"allowed_passthrough_routes": ["/custom-endpoint"]},
|
||||
)
|
||||
|
||||
# Test exact match
|
||||
result = RouteChecks.check_passthrough_route_access(
|
||||
route="/custom-endpoint",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_key_metadata_prefix_match():
|
||||
"""Test that key metadata allowed_passthrough_routes allows prefix match"""
|
||||
|
||||
# Create a UserAPIKeyAuth with allowed_passthrough_routes in metadata
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={"allowed_passthrough_routes": ["/custom-endpoint"]},
|
||||
)
|
||||
|
||||
# Test prefix match
|
||||
result = RouteChecks.check_passthrough_route_access(
|
||||
route="/custom-endpoint/v1/chat/completions",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_key_metadata_no_match():
|
||||
"""Test that key metadata allowed_passthrough_routes denies non-matching routes"""
|
||||
|
||||
# Create a UserAPIKeyAuth with allowed_passthrough_routes in metadata
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={"allowed_passthrough_routes": ["/custom-endpoint"]},
|
||||
)
|
||||
|
||||
# Test non-matching route
|
||||
result = RouteChecks.check_passthrough_route_access(
|
||||
route="/other-endpoint",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result is False
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_team_metadata_exact_match():
|
||||
"""Test that team metadata allowed_passthrough_routes allows exact match"""
|
||||
|
||||
# Create a UserAPIKeyAuth with allowed_passthrough_routes in team_metadata
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={},
|
||||
team_metadata={"allowed_passthrough_routes": ["/team-endpoint"]},
|
||||
)
|
||||
|
||||
# Test exact match
|
||||
result = RouteChecks.check_passthrough_route_access(
|
||||
route="/team-endpoint",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_team_metadata_prefix_match():
|
||||
"""Test that team metadata allowed_passthrough_routes allows prefix match"""
|
||||
|
||||
# Create a UserAPIKeyAuth with allowed_passthrough_routes in team_metadata
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={},
|
||||
team_metadata={"allowed_passthrough_routes": ["/team-endpoint"]},
|
||||
)
|
||||
|
||||
# Test prefix match
|
||||
result = RouteChecks.check_passthrough_route_access(
|
||||
route="/team-endpoint/v1/messages",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_team_metadata_no_match():
|
||||
"""Test that team metadata allowed_passthrough_routes denies non-matching routes"""
|
||||
|
||||
# Create a UserAPIKeyAuth with allowed_passthrough_routes in team_metadata
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={},
|
||||
team_metadata={"allowed_passthrough_routes": ["/team-endpoint"]},
|
||||
)
|
||||
|
||||
# Test non-matching route
|
||||
result = RouteChecks.check_passthrough_route_access(
|
||||
route="/other-endpoint",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result is False
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_key_metadata_takes_precedence():
|
||||
"""Test that key metadata takes precedence over team metadata"""
|
||||
|
||||
# Create a UserAPIKeyAuth with different allowed_passthrough_routes in both metadata
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={"allowed_passthrough_routes": ["/key-endpoint"]},
|
||||
team_metadata={"allowed_passthrough_routes": ["/team-endpoint"]},
|
||||
)
|
||||
|
||||
# Test that key endpoint is allowed
|
||||
result1 = RouteChecks.check_passthrough_route_access(
|
||||
route="/key-endpoint",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
# Test that team endpoint is NOT allowed (key metadata takes precedence)
|
||||
result2 = RouteChecks.check_passthrough_route_access(
|
||||
route="/team-endpoint",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result1 is True
|
||||
assert result2 is False
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_no_metadata():
|
||||
"""Test that route is denied when metadata and team_metadata don't have allowed_passthrough_routes"""
|
||||
|
||||
# Create a UserAPIKeyAuth without allowed_passthrough_routes
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
)
|
||||
|
||||
# Test that route is denied
|
||||
result = RouteChecks.check_passthrough_route_access(
|
||||
route="/any-endpoint",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result is False
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_no_allowed_passthrough_routes_key():
|
||||
"""Test that route is denied when allowed_passthrough_routes is not in metadata"""
|
||||
|
||||
# Create a UserAPIKeyAuth with metadata but no allowed_passthrough_routes
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={"other_field": "value"},
|
||||
team_metadata={},
|
||||
)
|
||||
|
||||
# Test that route is denied
|
||||
result = RouteChecks.check_passthrough_route_access(
|
||||
route="/any-endpoint",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result is False
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_allowed_passthrough_routes_is_none():
|
||||
"""Test that route is denied when allowed_passthrough_routes is None"""
|
||||
|
||||
# Create a UserAPIKeyAuth with allowed_passthrough_routes set to None
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={"allowed_passthrough_routes": None},
|
||||
team_metadata={"allowed_passthrough_routes": None},
|
||||
)
|
||||
|
||||
# Test that route is denied
|
||||
result = RouteChecks.check_passthrough_route_access(
|
||||
route="/any-endpoint",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result is False
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_multiple_routes():
|
||||
"""Test that multiple allowed_passthrough_routes work correctly"""
|
||||
|
||||
# Create a UserAPIKeyAuth with multiple allowed_passthrough_routes
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={
|
||||
"allowed_passthrough_routes": [
|
||||
"/endpoint-1",
|
||||
"/endpoint-2",
|
||||
"/endpoint-3",
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
# Test that all allowed routes work
|
||||
result1 = RouteChecks.check_passthrough_route_access(
|
||||
route="/endpoint-1/v1/chat",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
result2 = RouteChecks.check_passthrough_route_access(
|
||||
route="/endpoint-2",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
result3 = RouteChecks.check_passthrough_route_access(
|
||||
route="/endpoint-3/completions",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
# Test that non-allowed route fails
|
||||
result4 = RouteChecks.check_passthrough_route_access(
|
||||
route="/endpoint-4",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result1 is True
|
||||
assert result2 is True
|
||||
assert result3 is True
|
||||
assert result4 is False
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_prevents_false_prefix_match():
|
||||
"""Test that prefix matching doesn't allow false matches like /endpoint vs /endpoint-2"""
|
||||
|
||||
# Create a UserAPIKeyAuth with allowed_passthrough_routes
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={"allowed_passthrough_routes": ["/endpoint"]},
|
||||
)
|
||||
|
||||
# Test that /endpoint-2 is NOT allowed (not a valid prefix match)
|
||||
result = RouteChecks.check_passthrough_route_access(
|
||||
route="/endpoint-2",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result is False
|
||||
|
||||
# Test that /endpoint/something IS allowed (valid prefix match)
|
||||
result2 = RouteChecks.check_passthrough_route_access(
|
||||
route="/endpoint/something",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result2 is True
|
||||
|
||||
|
||||
def test_check_passthrough_route_access_empty_list():
|
||||
"""Test that empty allowed_passthrough_routes list denies all routes"""
|
||||
|
||||
# Create a UserAPIKeyAuth with empty allowed_passthrough_routes
|
||||
valid_token = UserAPIKeyAuth(
|
||||
user_id="test_user",
|
||||
metadata={"allowed_passthrough_routes": []},
|
||||
)
|
||||
|
||||
# Test that route is denied
|
||||
result = RouteChecks.check_passthrough_route_access(
|
||||
route="/any-endpoint",
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
assert result is False
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue