chore(proxy): make endpoint-override credential checks provider-specific and scope alias deletion to the owning team

Requiring 'any non-empty credential' on a destination change was insufficient for AWS: Bedrock/SageMaker sign with ambient credentials (instance role / IRSA web identity) regardless of api_key, so a team admin could repoint aws_sts_endpoint/aws_bedrock_runtime_endpoint and include an unrelated api_key to have ambient AWS credentials sent to their host. The check now requires a credential the overridden destination's provider actually consumes (AWS credential for AWS endpoints, api_key-class for generic), shared between the model-management write paths and /health/test_connection via a single destination_credential_policy module so the two cannot drift.

Separately, delete_team_model_alias removed every alias row matching the public model name; since a public name is not unique across teams, one team admin could delete another team's alias. It now only updates the owning team's row. Adds regression tests for the AWS-endpoint provider-specific check (model management and test_connection) and for cross-team alias isolation.
This commit is contained in:
user 2026-05-31 10:24:00 +00:00
parent 5e57f966dc
commit a66bb372bd
No known key found for this signature in database
5 changed files with 282 additions and 164 deletions

View file

@ -0,0 +1,66 @@
"""Single source of truth for the destination-override credential policy.
A non-admin must not redirect a model's endpoint and have the proxy's stored or
ambient credentials sent there. The credential that actually authenticates a
request depends on the destination's provider, so requiring "any credential" is
insufficient: an api_key does not stop Bedrock/SageMaker from signing with
ambient AWS credentials. Shared by the model-management write paths and
/health/test_connection so the two policies cannot drift.
"""
from typing import Tuple
# Caller-controllable endpoint URLs that retarget the outbound request. Mirrors
# the endpoint-redirect subset of auth_utils._BANNED_REQUEST_BODY_PARAMS.
URL_DESTINATION_FIELDS: Tuple[str, ...] = (
"api_base",
"base_url",
"aws_bedrock_runtime_endpoint",
"aws_sts_endpoint",
"sagemaker_base_url",
"s3_endpoint_url",
"deployment_url",
)
# Credential fields that must never be silently re-pointed at a new endpoint.
CREDENTIAL_FIELDS: Tuple[str, ...] = (
"api_key",
"litellm_credential_name",
"aws_secret_access_key",
"aws_session_token",
"aws_web_identity_token",
"azure_ad_token",
"vertex_credentials",
)
_AWS_CREDENTIALS: Tuple[str, ...] = (
"aws_secret_access_key",
"aws_session_token",
"aws_web_identity_token",
)
# api_key-class credentials consumed by OpenAI-compatible / generic providers.
_INLINE_CREDENTIALS: Tuple[str, ...] = (
"api_key",
"litellm_credential_name",
"azure_ad_token",
"vertex_credentials",
)
# AWS providers sign with ambient credentials (instance role / IRSA web identity)
# when no explicit AWS credential is given, so only an AWS credential makes a
# redirect of these fields self-contained.
_CONSUMED_CREDENTIALS = {
"aws_bedrock_runtime_endpoint": _AWS_CREDENTIALS,
"aws_sts_endpoint": _AWS_CREDENTIALS,
"sagemaker_base_url": _AWS_CREDENTIALS,
"s3_endpoint_url": _AWS_CREDENTIALS,
}
def consumed_credentials_for(destination_field: str) -> Tuple[str, ...]:
"""Credentials the destination's provider actually authenticates with.
Supplying a credential outside this set does not make redirecting the field
safe: the provider would fall back to the proxy's ambient/stored credentials
and send them to the new endpoint.
"""
return _CONSUMED_CREDENTIALS.get(destination_field, _INLINE_CREDENTIALS)

View file

@ -27,6 +27,15 @@ from litellm.proxy._types import (
WebhookEvent,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.destination_credential_policy import (
CREDENTIAL_FIELDS as _HEALTH_CREDENTIAL_FIELDS,
)
from litellm.proxy.common_utils.destination_credential_policy import (
URL_DESTINATION_FIELDS as _HEALTH_URL_FIELDS,
)
from litellm.proxy.common_utils.destination_credential_policy import (
consumed_credentials_for,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.health_check import (
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS,
@ -78,29 +87,6 @@ def _reject_os_environ_references(params: dict) -> None:
stack.append(value)
_HEALTH_CREDENTIAL_FIELDS = (
"api_key",
"litellm_credential_name",
"aws_secret_access_key",
"aws_session_token",
"aws_web_identity_token",
"azure_ad_token",
"vertex_credentials",
)
# Caller-controlled endpoint URLs: overriding one re-points the request (and any
# credential) at a new host, so it is SSRF-validated. Mirrors the endpoint-redirect
# subset of auth_utils._BANNED_REQUEST_BODY_PARAMS.
_HEALTH_URL_FIELDS = (
"api_base",
"base_url",
"aws_bedrock_runtime_endpoint",
"aws_sts_endpoint",
"sagemaker_base_url",
"s3_endpoint_url",
"deployment_url",
)
def _assert_non_admin_destination_override_is_safe(
config_litellm_params: dict,
request_litellm_params: dict,
@ -110,9 +96,10 @@ def _assert_non_admin_destination_override_is_safe(
When a non-admin overrides the connection target (api_base/base_url/provider
endpoint), every credential the resolved deployment stores must be re-supplied
by the request (non-empty), and the request must carry a non-empty credential
at all. Otherwise the downstream provider sends a stored credential, or falls
back to an environment secret (e.g. OPENAI_API_KEY), to the caller-chosen
by the request (non-empty), and each overridden URL must be accompanied by a
non-empty credential its provider actually consumes. Otherwise the downstream
provider sends a stored credential, or falls back to an ambient/environment
secret (e.g. OPENAI_API_KEY or ambient AWS credentials), to the caller-chosen
address — including the case where the caller supplies an unrelated credential
field to leave the stored one riding along. The overridden URL is also
SSRF-validated so it cannot point at an internal or cloud-metadata host.
@ -123,7 +110,10 @@ def _assert_non_admin_destination_override_is_safe(
if is_proxy_admin(user_api_key_dict):
return
if not any(request_litellm_params.get(field) for field in _HEALTH_URL_FIELDS):
overridden_urls = [
field for field in _HEALTH_URL_FIELDS if request_litellm_params.get(field)
]
if not overridden_urls:
return
for field in _HEALTH_CREDENTIAL_FIELDS:
if config_litellm_params.get(field) and not request_litellm_params.get(field):
@ -133,15 +123,15 @@ def _assert_non_admin_destination_override_is_safe(
"error": f"Re-provide {field} when overriding the connection target; a stored credential cannot be sent to a caller-chosen endpoint."
},
)
if not any(
request_litellm_params.get(field) for field in _HEALTH_CREDENTIAL_FIELDS
):
raise HTTPException(
status_code=400,
detail={
"error": "Supply your own api_key when overriding the connection target (e.g. api_base); otherwise a stored or environment credential could be sent to it."
},
)
for field in overridden_urls:
consumed = consumed_credentials_for(field)
if not any(request_litellm_params.get(cred) for cred in consumed):
raise HTTPException(
status_code=400,
detail={
"error": f"Overriding {field} requires supplying a credential it uses ({', '.join(consumed)}); otherwise a stored or ambient credential could be sent to it."
},
)
if not getattr(litellm, "user_url_validation", False):
return
for field in _HEALTH_URL_FIELDS:

View file

@ -43,6 +43,11 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy.common_utils.destination_credential_policy import (
CREDENTIAL_FIELDS,
URL_DESTINATION_FIELDS,
consumed_credentials_for,
)
from litellm.proxy.common_utils.resource_ownership import is_proxy_admin
from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin
from litellm.proxy.management_endpoints.team_endpoints import (
@ -69,31 +74,13 @@ from litellm.utils import get_utc_datetime
router = APIRouter()
# Credential fields that must never be silently re-pointed at a new endpoint,
# and the routing fields that change where they are sent.
_CREDENTIAL_LITELLM_PARAMS = (
"api_key",
"litellm_credential_name",
"aws_secret_access_key",
"aws_session_token",
"aws_web_identity_token",
"azure_ad_token",
"vertex_credentials",
)
# Caller-controlled endpoint URLs that must be SSRF-validated on write. Mirrors
# the endpoint-redirect subset of auth_utils._BANNED_REQUEST_BODY_PARAMS so a
# stored config can't reach an internal host that the request-time guard blocks.
_URL_LITELLM_PARAMS = (
"api_base",
"base_url",
"aws_bedrock_runtime_endpoint",
"aws_sts_endpoint",
"sagemaker_base_url",
"s3_endpoint_url",
"deployment_url",
)
# Destination fields whose change must drop an inherited credential (URLs above
# plus the provider selector, which isn't a URL so isn't SSRF-validated).
# Credential and endpoint-URL field sets, plus the provider-specific credential
# policy, live in one shared module so the write paths and /health/test_connection
# cannot drift apart.
_CREDENTIAL_LITELLM_PARAMS = CREDENTIAL_FIELDS
_URL_LITELLM_PARAMS = URL_DESTINATION_FIELDS
# Destination fields whose change requires a fresh credential (URLs above plus the
# provider selector, which isn't a URL so isn't SSRF-validated).
_DESTINATION_LITELLM_PARAMS = _URL_LITELLM_PARAMS + ("custom_llm_provider",)
@ -144,30 +131,35 @@ def _assert_credential_resupplied_on_destination_change(
) -> None:
"""Changing a model's destination (api_base/base_url/provider endpoint) must
not send a credential to the new endpoint that the caller did not themselves
provide. Two things are required. First, the patch must carry a non-empty
credential at all: even a model stored without one falls back to the proxy's
environment secret (e.g. OPENAI_API_KEY) at call time, which would then be
sent to the new endpoint. Second, every credential the model already had must
be re-supplied non-empty, so an inherited secret is never kept alongside an
unrelated new one. A non-empty re-supplied value is the caller's own
credential going to their own endpoint, which is allowed."""
destination_changed = any(
field in patch_plaintext and patch_plaintext[field] != db_plaintext(field)
provide. Two things are required. First, each changed destination must be
accompanied by a non-empty credential its provider actually consumes: an
api_key does not make an AWS endpoint change safe, because Bedrock/SageMaker
sign with ambient AWS credentials regardless, and even a model with no stored
credential falls back to the proxy's environment secret at call time. Second,
every credential the model already had must be re-supplied non-empty, so an
inherited secret is never kept alongside an unrelated new one. A non-empty
re-supplied value is the caller's own credential going to their own endpoint,
which is allowed."""
changed_destinations = [
field
for field in _DESTINATION_LITELLM_PARAMS
)
if not destination_changed:
if field in patch_plaintext and patch_plaintext[field] != db_plaintext(field)
]
if not changed_destinations:
return
if not any(patch_plaintext.get(field) for field in _CREDENTIAL_LITELLM_PARAMS):
raise ProxyException(
message=(
"Provide a non-empty credential (e.g. api_key) when changing the "
"model endpoint; otherwise the proxy's environment credential "
"would be used at the new destination."
),
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_400_BAD_REQUEST,
param="api_key",
)
for field in changed_destinations:
consumed = consumed_credentials_for(field)
if not any(patch_plaintext.get(cred) for cred in consumed):
raise ProxyException(
message=(
f"Changing {field} requires supplying a non-empty credential it "
f"uses ({', '.join(consumed)}); otherwise the proxy's stored or "
"ambient credential would be sent to the new destination."
),
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_400_BAD_REQUEST,
param=field,
)
for field in _CREDENTIAL_LITELLM_PARAMS:
if db_plaintext(field) and not patch_plaintext.get(field):
raise ProxyException(
@ -1105,6 +1097,7 @@ async def delete_model(
model_params.model_info.team_public_model_name
or model_params.model_name
),
team_id=model_params.model_info.team_id,
prisma_client=prisma_client,
)
@ -1204,12 +1197,15 @@ async def delete_model(
async def delete_team_model_alias(
public_model_name: str,
team_id: str,
prisma_client: PrismaClient,
) -> List[Tuple[str, str]]:
"""
Delete a team model alias
Iterate through all team model aliases and delete the one that matches the model_id
Only the owning team's alias is touched: a public model name is not unique
across teams, so without scoping by team_id one team admin could delete
another team's alias that happens to point at the same public name.
Returns:
- List of team id + model alias pairs that were removed
@ -1220,6 +1216,8 @@ async def delete_team_model_alias(
tasks = []
removed_model_aliases = []
for team_model_alias in team_model_aliases:
if team_model_alias.team is None or team_model_alias.team.team_id != team_id:
continue
model_aliases = team_model_alias.model_aliases # {"alias": "public model name"}
id = team_model_alias.id

View file

@ -1991,6 +1991,24 @@ def test_non_admin_destination_override_guard():
guard({"api_key": "sk-victim"}, {"model": "gpt-4o"}, non_admin)
# Proxy admin is exempt even when overriding without a credential.
guard({"api_key": "sk-victim"}, {"api_base": "https://attacker.example"}, admin)
# An AWS endpoint override is NOT satisfied by an api_key (Bedrock/SageMaker
# sign with ambient AWS credentials regardless), so it must be rejected.
with pytest.raises(HTTPException):
guard(
{},
{"aws_bedrock_runtime_endpoint": "https://attacker", "api_key": "x"},
non_admin,
)
with patch.object(litellm, "user_url_validation", False):
# Supplying an actual AWS credential is accepted.
guard(
{},
{
"aws_bedrock_runtime_endpoint": "https://attacker",
"aws_secret_access_key": "ak",
},
non_admin,
)
# Non-admin override to an internal/metadata IP -> SSRF-rejected.
with patch.object(litellm, "user_url_validation", True):
with pytest.raises(HTTPException):

View file

@ -293,105 +293,76 @@ class MockPrismaWrapper:
class TestDeleteTeamModelAlias:
@staticmethod
def _alias_row(row_id, model_aliases, team_id):
from types import SimpleNamespace
return SimpleNamespace(
id=row_id,
model_aliases=model_aliases,
team=SimpleNamespace(team_id=team_id),
)
@pytest.mark.asyncio
async def test_delete_team_model_alias_success(self):
"""Test successful deletion of a team model alias"""
"""Only the owning team's alias is removed, even when another team's row
points at the same public model name."""
from litellm.proxy.management_endpoints.model_management_endpoints import (
delete_team_model_alias,
)
# Setup test data
model_aliases_list = [
{
"id": 1,
"model_aliases": {
"alias1": "public_model_1",
"alias2": "public_model_2",
},
"updated_by": "test_user",
"created_by": "test_user",
},
{
"id": 2,
"model_aliases": {
"alias3": "public_model_3",
"alias4": "public_model_1",
},
"updated_by": "test_user",
"created_by": "test_user",
}, # public_model_1 appears twice
rows = [
self._alias_row(
1, {"alias1": "public_model_1", "alias2": "public_model_2"}, "team1"
),
# team2 also aliases public_model_1; it must NOT be touched.
self._alias_row(2, {"alias4": "public_model_1"}, "team2"),
]
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_modeltable = AsyncMock()
mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=rows)
mock_prisma.db.litellm_modeltable.update = AsyncMock()
# Create mock prisma client
mock_prisma = MockPrismaClient(team_exists=True)
mock_prisma.db = MockPrismaWrapper(model_aliases_list)
# Call the function
await delete_team_model_alias(
public_model_name="public_model_1", prisma_client=mock_prisma
removed = await delete_team_model_alias(
public_model_name="public_model_1",
team_id="team1",
prisma_client=mock_prisma,
)
# Verify results
mock_db = mock_prisma.db.litellm_modeltable
assert (
len(mock_db.update_calls) == 2
) # Should have 2 update calls since public_model_1 appears twice
# Verify first update
first_update = mock_db.update_calls[0]
assert first_update["where"] == {"id": 1}
assert json.loads(first_update["data"]["model_aliases"]) == {
update = mock_prisma.db.litellm_modeltable.update
assert update.call_count == 1 # only team1's row
assert update.call_args.kwargs["where"] == {"id": 1}
assert json.loads(update.call_args.kwargs["data"]["model_aliases"]) == {
"alias2": "public_model_2"
}
# Verify second update
second_update = mock_db.update_calls[1]
assert second_update["where"] == {"id": 2}
assert json.loads(second_update["data"]["model_aliases"]) == {
"alias3": "public_model_3"
}
assert removed == [("team1", "alias1")]
@pytest.mark.asyncio
async def test_delete_team_model_alias_no_matches(self):
"""Test deletion when no matching model alias exists"""
"""No updates when the owning team has no alias for the public name."""
from litellm.proxy.management_endpoints.model_management_endpoints import (
delete_team_model_alias,
)
# Setup test data with no matching model
model_aliases_list = [
{
"id": 1,
"model_aliases": {
"alias1": "public_model_1",
"alias2": "public_model_2",
},
"updated_by": "test_user",
"created_by": "test_user",
},
{
"id": 2,
"model_aliases": {
"alias3": "public_model_3",
"alias4": "public_model_4",
},
"updated_by": "test_user",
"created_by": "test_user",
},
rows = [
self._alias_row(
1, {"alias1": "public_model_1", "alias2": "public_model_2"}, "team1"
),
]
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_modeltable = AsyncMock()
mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=rows)
mock_prisma.db.litellm_modeltable.update = AsyncMock()
# Create mock prisma client
mock_prisma = MockPrismaClient(team_exists=True)
mock_prisma.db = MockPrismaWrapper(model_aliases_list)
# Call the function with non-existent model
await delete_team_model_alias(
public_model_name="non_existent_model", prisma_client=mock_prisma
public_model_name="non_existent_model",
team_id="team1",
prisma_client=mock_prisma,
)
# Verify no updates were made
mock_db = mock_prisma.db.litellm_modeltable
assert len(mock_db.update_calls) == 0
assert mock_prisma.db.litellm_modeltable.update.call_count == 0
class TestClearCache:
@ -2345,8 +2316,9 @@ class TestModelMgmtAuthzHardening:
captured = {}
async def fake_delete_alias(public_model_name, prisma_client):
async def fake_delete_alias(public_model_name, team_id, prisma_client):
captured["public_model_name"] = public_model_name
captured["team_id"] = team_id
return []
_PS = "litellm.proxy.proxy_server"
@ -2469,9 +2441,10 @@ class TestModelMgmtAuthzHardening:
def test_destination_change_partial_resupply_is_rejected(self):
from litellm.proxy._types import ProxyException
# Patch redirects api_base AND supplies an UNRELATED credential
# (aws_secret_access_key); the inherited api_key is still not re-supplied,
# so the change must be rejected rather than letting api_key ride along.
# Patch redirects api_base but supplies only an UNRELATED credential
# (aws_secret_access_key, which an OpenAI-compatible api_base does not
# consume), so the change must be rejected rather than letting the stored
# api_key (or the env fallback) ride along to the new endpoint.
db = {"api_base": "https://real.example", "api_key": "sk-inherited"}
with pytest.raises(ProxyException) as e:
self._assert_resupply(
@ -2481,7 +2454,7 @@ class TestModelMgmtAuthzHardening:
},
db,
)
assert e.value.param == "api_key"
assert e.value.param == "api_base"
# --- Veria sweep follow-ups: authoritative pricing set, base_model, extra URLs ---
@ -2558,3 +2531,76 @@ class TestModelMgmtAuthzHardening:
self._assert_resupply(
{"aws_sts_endpoint": "https://sts.attacker", "api_key": "x"}, db
)
def test_unrelated_credential_does_not_satisfy_aws_endpoint_override(self):
from litellm.proxy._types import ProxyException
# AWS providers sign with ambient credentials regardless of api_key, so an
# api_key must NOT satisfy an AWS endpoint override (it would leave ambient
# AWS creds to be SigV4-signed to the attacker host). Model has no stored
# AWS credential (ambient).
db = {"aws_bedrock_runtime_endpoint": "https://real.aws"}
with pytest.raises(ProxyException) as e:
self._assert_resupply(
{
"aws_bedrock_runtime_endpoint": "https://attacker.example",
"api_key": "sk-unrelated",
},
db,
)
assert e.value.param == "aws_bedrock_runtime_endpoint"
# Supplying an actual AWS credential is accepted.
self._assert_resupply(
{
"aws_bedrock_runtime_endpoint": "https://attacker.example",
"aws_secret_access_key": "ak",
},
db,
)
# Conversely, an unrelated AWS credential does not satisfy a generic
# api_base override (OpenAI ignores it and falls back to OPENAI_API_KEY).
with pytest.raises(ProxyException):
self._assert_resupply(
{"api_base": "https://attacker", "aws_secret_access_key": "ak"},
{"api_base": "https://real"},
)
@pytest.mark.asyncio
async def test_delete_team_model_alias_only_touches_owning_team(self):
from types import SimpleNamespace
from litellm.proxy.management_endpoints.model_management_endpoints import (
delete_team_model_alias,
)
# Two teams have an alias pointing at the same public model name.
row_a = SimpleNamespace(
id="A",
model_aliases={"gpt4": "shared-public"},
team=SimpleNamespace(team_id="teamA"),
)
row_b = SimpleNamespace(
id="B",
model_aliases={"gpt4": "shared-public"},
team=SimpleNamespace(team_id="teamB"),
)
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_modeltable = AsyncMock()
mock_prisma.db.litellm_modeltable.find_many = AsyncMock(
return_value=[row_a, row_b]
)
mock_prisma.db.litellm_modeltable.update = AsyncMock()
removed = await delete_team_model_alias(
public_model_name="shared-public",
team_id="teamA",
prisma_client=mock_prisma,
)
# Only the owning team's alias row is updated; teamB's is untouched.
assert mock_prisma.db.litellm_modeltable.update.call_count == 1
assert mock_prisma.db.litellm_modeltable.update.call_args.kwargs["where"] == {
"id": "A"
}
assert removed == [("teamA", "gpt4")]