mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: resolve team id's on projects
This commit is contained in:
parent
8a3aa4d31c
commit
51ef8a96f0
3 changed files with 202 additions and 31 deletions
|
|
@ -7,19 +7,19 @@ repos:
|
|||
language: system
|
||||
types: [python]
|
||||
files: ^(litellm/|litellm_proxy_extras/|enterprise/)
|
||||
- id: isort
|
||||
name: isort
|
||||
entry: isort
|
||||
language: system
|
||||
types: [python]
|
||||
files: (litellm/|litellm_proxy_extras/|enterprise/).*\.py
|
||||
exclude: ^litellm/__init__.py$
|
||||
- id: black
|
||||
name: black
|
||||
entry: poetry run black
|
||||
language: system
|
||||
types: [python]
|
||||
files: (litellm/|litellm_proxy_extras/).*\.py
|
||||
# - id: isort
|
||||
# name: isort
|
||||
# entry: isort
|
||||
# language: system
|
||||
# types: [python]
|
||||
# files: (litellm/|litellm_proxy_extras/|enterprise/).*\.py
|
||||
# exclude: ^litellm/__init__.py$
|
||||
# - id: black
|
||||
# name: black
|
||||
# entry: poetry run black
|
||||
# language: system
|
||||
# types: [python]
|
||||
# files: (litellm/|litellm_proxy_extras/).*\.py
|
||||
- repo: https://github.com/pycqa/flake8
|
||||
rev: 7.0.0 # The version of flake8 to use
|
||||
hooks:
|
||||
|
|
|
|||
|
|
@ -188,6 +188,7 @@ async def _run_project_checks(
|
|||
skip_budget_checks: bool,
|
||||
valid_token: Optional[UserAPIKeyAuth],
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
team_object: Optional[LiteLLM_TeamTable] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Run all project-level checks: blocked, model access, budget, soft budget.
|
||||
|
|
@ -204,11 +205,28 @@ async def _run_project_checks(
|
|||
|
||||
# 2.2 If project can call model
|
||||
if _model and len(project_object.models) > 0:
|
||||
can_project_access_model(
|
||||
model=_model,
|
||||
project_object=project_object,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
resolved_models = list(project_object.models)
|
||||
|
||||
if SpecialModelNames.all_proxy_models.value in resolved_models:
|
||||
pass # all-proxy-models grants access to everything, skip the check
|
||||
elif SpecialModelNames.all_team_models.value in resolved_models:
|
||||
if team_object is not None and team_object.models:
|
||||
resolved_models = list(team_object.models)
|
||||
else:
|
||||
resolved_models = []
|
||||
if len(resolved_models) > 0:
|
||||
_can_object_call_model(
|
||||
model=_model,
|
||||
llm_router=llm_router,
|
||||
models=resolved_models,
|
||||
object_type="project",
|
||||
)
|
||||
else:
|
||||
can_project_access_model(
|
||||
model=_model,
|
||||
project_object=project_object,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
if not skip_budget_checks:
|
||||
# 3.0.2. If project is in budget
|
||||
|
|
@ -464,6 +482,7 @@ async def common_checks( # noqa: PLR0915
|
|||
skip_budget_checks=skip_budget_checks,
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_object=team_object,
|
||||
)
|
||||
|
||||
# If this is a free model, skip all budget checks
|
||||
|
|
|
|||
|
|
@ -59,9 +59,9 @@ def reset_constants_module():
|
|||
# Reload modules before test
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
|
||||
|
||||
yield
|
||||
|
||||
|
||||
# Reload modules after test to clean up
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
|
|
@ -154,9 +154,9 @@ def test_experimental_ui_token_ignores_litellm_ui_session_duration(
|
|||
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
|
||||
now = get_utc_datetime()
|
||||
# Must be ~10 min, NOT 24h. If LITELLM_UI_SESSION_DURATION were incorrectly used, this would fail.
|
||||
assert expires <= now + timedelta(minutes=11), (
|
||||
"Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION"
|
||||
)
|
||||
assert expires <= now + timedelta(
|
||||
minutes=11
|
||||
), "Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION"
|
||||
|
||||
|
||||
def test_get_experimental_ui_login_jwt_auth_token_invalid(
|
||||
|
|
@ -290,13 +290,15 @@ def test_get_cli_jwt_auth_token_custom_expiration(
|
|||
|
||||
# Set custom expiration to 48 hours
|
||||
monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "48")
|
||||
|
||||
|
||||
# Reload the constants module to pick up the new env var
|
||||
importlib.reload(constants)
|
||||
# Also reload auth_checks to pick up the new constant value
|
||||
importlib.reload(auth_checks)
|
||||
|
||||
token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
|
||||
|
||||
token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
||||
valid_sso_user_defined_values
|
||||
)
|
||||
|
||||
# Decrypt and verify token contents
|
||||
decrypted_token = decrypt_value_helper(
|
||||
|
|
@ -312,7 +314,6 @@ def test_get_cli_jwt_auth_token_custom_expiration(
|
|||
assert expires <= get_utc_datetime() + timedelta(hours=48, minutes=1)
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_internal_user_params_with_get_user_object(monkeypatch):
|
||||
"""Test that default_internal_user_params is used when creating a new user via get_user_object"""
|
||||
|
|
@ -433,7 +434,9 @@ async def test_get_user_object_upsert_includes_user_email():
|
|||
mock_prisma_client.db.litellm_usertable.create.assert_called_once()
|
||||
creation_args = mock_prisma_client.db.litellm_usertable.create.call_args[1]["data"]
|
||||
|
||||
assert "user_email" in creation_args, "user_email should be included when upserting a new user"
|
||||
assert (
|
||||
"user_email" in creation_args
|
||||
), "user_email should be included when upserting a new user"
|
||||
assert creation_args["user_email"] == "test@example.com"
|
||||
assert creation_args["user_id"] == "new_test_user"
|
||||
|
||||
|
|
@ -460,7 +463,9 @@ def test_log_budget_lookup_failure_skips_user_not_found():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock)
|
||||
@patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock
|
||||
)
|
||||
async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeypatch):
|
||||
"""
|
||||
Test that _get_team_db_check correctly calls the `new_team` function
|
||||
|
|
@ -494,8 +499,12 @@ async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeyp
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock)
|
||||
async def test_get_team_db_check_does_not_call_new_team_if_exists(mock_new_team, monkeypatch):
|
||||
@patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock
|
||||
)
|
||||
async def test_get_team_db_check_does_not_call_new_team_if_exists(
|
||||
mock_new_team, monkeypatch
|
||||
):
|
||||
"""
|
||||
Test that _get_team_db_check does NOT call the `new_team` function
|
||||
if the team already exists in the database.
|
||||
|
|
@ -1629,3 +1638,146 @@ async def test_custom_auth_common_checks_opt_in():
|
|||
parent_otel_span=None,
|
||||
)
|
||||
mock_common.assert_called_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests for _run_project_checks: special model name resolution
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_project(models: list) -> "LiteLLM_ProjectTableCachedObj":
|
||||
from litellm.proxy._types import LiteLLM_ProjectTableCachedObj
|
||||
|
||||
return LiteLLM_ProjectTableCachedObj(
|
||||
project_id="proj-1",
|
||||
created_by="test",
|
||||
updated_by="test",
|
||||
models=models,
|
||||
)
|
||||
|
||||
|
||||
def _make_team(models: list) -> LiteLLM_TeamTable:
|
||||
return LiteLLM_TeamTable(
|
||||
team_id="team-1",
|
||||
models=models,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_project_checks_all_team_models_allowed():
|
||||
"""Project with all-team-models should allow a model that the team allows."""
|
||||
from litellm.proxy.auth.auth_checks import _run_project_checks
|
||||
|
||||
project = _make_project(["all-team-models"])
|
||||
team = _make_team(["openai/gpt-5", "openai/gpt-4o"])
|
||||
proxy_logging = MagicMock()
|
||||
|
||||
await _run_project_checks(
|
||||
project_object=project,
|
||||
_model="openai/gpt-5",
|
||||
llm_router=None,
|
||||
skip_budget_checks=True,
|
||||
valid_token=None,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
team_object=team,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_project_checks_all_team_models_denied():
|
||||
"""Project with all-team-models should deny a model the team does not allow."""
|
||||
from litellm.proxy.auth.auth_checks import _run_project_checks
|
||||
|
||||
project = _make_project(["all-team-models"])
|
||||
team = _make_team(["openai/gpt-4o"])
|
||||
proxy_logging = MagicMock()
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _run_project_checks(
|
||||
project_object=project,
|
||||
_model="openai/gpt-5",
|
||||
llm_router=None,
|
||||
skip_budget_checks=True,
|
||||
valid_token=None,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
team_object=team,
|
||||
)
|
||||
assert "project" in exc_info.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_project_checks_all_team_models_no_team():
|
||||
"""Project with all-team-models but no team object should skip the model check."""
|
||||
from litellm.proxy.auth.auth_checks import _run_project_checks
|
||||
|
||||
project = _make_project(["all-team-models"])
|
||||
proxy_logging = MagicMock()
|
||||
|
||||
await _run_project_checks(
|
||||
project_object=project,
|
||||
_model="openai/gpt-5",
|
||||
llm_router=None,
|
||||
skip_budget_checks=True,
|
||||
valid_token=None,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
team_object=None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_project_checks_all_proxy_models():
|
||||
"""Project with all-proxy-models should allow any model."""
|
||||
from litellm.proxy.auth.auth_checks import _run_project_checks
|
||||
|
||||
project = _make_project(["all-proxy-models"])
|
||||
proxy_logging = MagicMock()
|
||||
|
||||
await _run_project_checks(
|
||||
project_object=project,
|
||||
_model="anything/any-model",
|
||||
llm_router=None,
|
||||
skip_budget_checks=True,
|
||||
valid_token=None,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
team_object=None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_project_checks_explicit_models_allowed():
|
||||
"""Project with explicit model list should allow a listed model."""
|
||||
from litellm.proxy.auth.auth_checks import _run_project_checks
|
||||
|
||||
project = _make_project(["openai/gpt-5"])
|
||||
proxy_logging = MagicMock()
|
||||
|
||||
await _run_project_checks(
|
||||
project_object=project,
|
||||
_model="openai/gpt-5",
|
||||
llm_router=None,
|
||||
skip_budget_checks=True,
|
||||
valid_token=None,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
team_object=None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_project_checks_explicit_models_denied():
|
||||
"""Project with explicit model list should deny an unlisted model."""
|
||||
from litellm.proxy.auth.auth_checks import _run_project_checks
|
||||
|
||||
project = _make_project(["openai/gpt-4o"])
|
||||
proxy_logging = MagicMock()
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _run_project_checks(
|
||||
project_object=project,
|
||||
_model="openai/gpt-5",
|
||||
llm_router=None,
|
||||
skip_budget_checks=True,
|
||||
valid_token=None,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
team_object=None,
|
||||
)
|
||||
assert "project" in exc_info.value.message
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue