mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
5e57f966dc
commit
a66bb372bd
5 changed files with 282 additions and 164 deletions
66
litellm/proxy/common_utils/destination_credential_policy.py
Normal file
66
litellm/proxy/common_utils/destination_credential_policy.py
Normal 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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue