Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 30 additions & 0 deletions enterprise/litellm_enterprise/proxy/auth/user_api_key_auth.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
from typing import Any, Optional

from fastapi import Request

from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth


async def enterprise_custom_auth(
request: Request, api_key: str, user_custom_auth: Any
) -> Optional[UserAPIKeyAuth]:
from litellm_enterprise.proxy.proxy_server import custom_auth_settings

if custom_auth_settings is None:
return None

if custom_auth_settings["mode"] == "on":
return await user_custom_auth(request, api_key)
elif custom_auth_settings["mode"] == "off":
return None
elif custom_auth_settings["mode"] == "auto":
try:
return await user_custom_auth(request, api_key)
except Exception as e:
verbose_proxy_logger.debug(
f"Error in custom auth, checking litellm auth: {e}"
)
return None
else:
raise ValueError(f"Invalid mode: {custom_auth_settings['mode']}")
25 changes: 25 additions & 0 deletions enterprise/litellm_enterprise/proxy/proxy_server.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
from typing import Optional

from litellm_enterprise.types.proxy.proxy_server import CustomAuthSettings

custom_auth_settings: Optional[CustomAuthSettings] = None


class EnterpriseProxyConfig:
async def load_custom_auth_settings(
self, general_settings: dict
) -> CustomAuthSettings:
print(f"General settings: {general_settings}")
custom_auth_settings = general_settings.get("custom_auth_settings", None)
print(f"Custom auth settings: {custom_auth_settings}")
if custom_auth_settings is not None:
custom_auth_settings = CustomAuthSettings(
mode=custom_auth_settings.get("mode"),
)
print(f"Custom auth settings: {custom_auth_settings}")
return custom_auth_settings

async def load_enterprise_config(self, general_settings: dict) -> None:
global custom_auth_settings
custom_auth_settings = await self.load_custom_auth_settings(general_settings)
return None
5 changes: 5 additions & 0 deletions enterprise/litellm_enterprise/types/proxy/proxy_server.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
from typing import Literal, TypedDict


class CustomAuthSettings(TypedDict):
mode: Literal["on", "off", "auto"]
6 changes: 3 additions & 3 deletions litellm/proxy/_new_secret_config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -78,9 +78,9 @@ model_list:
prompt_label: "latest"
api_key: os.environ/OPENAI_API_KEY

litellm_settings:
callbacks: ["langfuse"]

general_settings:
store_model_in_db: true
store_prompts_in_spend_logs: true
custom_auth: custom_auth_auto.user_api_key_auth
custom_auth_settings:
mode: "auto"
2 changes: 2 additions & 0 deletions litellm/proxy/_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -938,6 +938,7 @@ class NewCustomerRequest(BudgetNewRequest):
alias: Optional[str] = None # human-friendly alias
blocked: bool = False # allow/disallow requests for this end-user
budget_id: Optional[str] = None # give either a budget_id or max_budget
spend: Optional[float] = None
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
Expand Down Expand Up @@ -1222,6 +1223,7 @@ class TeamRequest(LiteLLMPydanticObjectBase):
class LiteLLM_BudgetTable(LiteLLMPydanticObjectBase):
"""Represents user-controllable params for a LiteLLM_BudgetTable record"""

budget_id: Optional[str] = None
soft_budget: Optional[float] = None
max_budget: Optional[float] = None
max_parallel_requests: Optional[int] = None
Expand Down
18 changes: 17 additions & 1 deletion litellm/proxy/auth/user_api_key_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,16 @@
from litellm.secret_managers.main import get_secret_bool
from litellm.types.services import ServiceTypes

try:
from litellm_enterprise.proxy.auth.user_api_key_auth import (
enterprise_custom_auth as _enterprise_custom_auth,
)

enterprise_custom_auth: Optional[Callable] = _enterprise_custom_auth
except ImportError as e:
verbose_proxy_logger.debug(f"Error in enterprise custom auth: {e}")
enterprise_custom_auth = None

user_api_key_service_logger_obj = ServiceLogging() # used for tracking latency on OTEL

custom_litellm_key_header = APIKeyHeader(
Expand Down Expand Up @@ -346,7 +356,13 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
)

### USER-DEFINED AUTH FUNCTION ###
if user_custom_auth is not None:
if enterprise_custom_auth is not None:
response = await enterprise_custom_auth(
request=request, api_key=api_key, user_custom_auth=user_custom_auth
)
if response is not None:
return UserAPIKeyAuth.model_validate(response)
elif user_custom_auth is not None:
response = await user_custom_auth(request=request, api_key=api_key) # type: ignore
return UserAPIKeyAuth.model_validate(response)

Expand Down
18 changes: 18 additions & 0 deletions litellm/proxy/custom_auth_auto.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
"""
Example custom auth function.

This will allow all keys starting with "my-custom-key" to pass through.
"""
from fastapi import Request

from litellm.proxy._types import UserAPIKeyAuth


async def user_api_key_auth(request: Request, api_key: str) -> UserAPIKeyAuth:
try:
if api_key.startswith("my-custom-key"):
return UserAPIKeyAuth(api_key=api_key)
else:
raise Exception("Invalid API key")
except Exception:
raise Exception("Invalid API key")
63 changes: 58 additions & 5 deletions litellm/proxy/management_endpoints/customer_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -185,6 +185,7 @@ async def new_end_user(
- model_max_budget: Optional[dict] - [Not Implemented Yet] Specify max budget for a given model. Example: {"openai/gpt-4o-mini": {"max_budget": 100.0, "budget_duration": "1d"}}
- max_parallel_requests: Optional[int] - [Not Implemented Yet] Specify max parallel requests for a given customer.
- soft_budget: Optional[float] - [Not Implemented Yet] Get alerts when customer crosses given budget, doesn't block requests.
- spend: Optional[float] - Specify initial spend for a given customer.


- Allow specifying allowed regions
Expand Down Expand Up @@ -424,13 +425,65 @@ async def update_end_user(
): # models default to [], spend defaults to 0, we should not reset these values
non_default_values[k] = v

## ADD USER, IF NEW ##
## Get end user table data ##
end_user_table_data = await prisma_client.db.litellm_endusertable.find_first(
where={"user_id": data.user_id}, include={"litellm_budget_table": True}
)

if end_user_table_data is None:
raise HTTPException(
status_code=400,
detail={
"error": "End User Id={} does not exist in db".format(data.user_id)
},
)

end_user_table_data_typed = LiteLLM_EndUserTable(
**end_user_table_data.model_dump()
)

## Get budget table data ##
end_user_budget_table = end_user_table_data_typed.litellm_budget_table

## Get all params for budget table ##
budget_table_data = {}
update_end_user_table_data = {}
for k, v in non_default_values.items():
if k in LiteLLM_BudgetTable.model_fields.keys():
budget_table_data[k] = v

if k in LiteLLM_EndUserTable.model_fields.keys():
update_end_user_table_data[k] = v

## Check if budget id is set ##
if budget_table_data:
if end_user_budget_table is None:
## Create new budget ##
budget_table_data_record = (
await prisma_client.db.litellm_budgettable.create(
data=budget_table_data, include={"litellm_endusertable": True}
)
)

update_end_user_table_data[
"budget_id"
] = budget_table_data_record.budget_id
else:
## Update existing budget ##
budget_table_data_record = (
await prisma_client.db.litellm_budgettable.update(
where={"budget_id": end_user_budget_table.budget_id},
data=budget_table_data,
)
)

## Update user table, with update params + new budget id (if set) ##
verbose_proxy_logger.debug("/customer/update: Received data = %s", data)
if data.user_id is not None and len(data.user_id) > 0:
non_default_values["user_id"] = data.user_id # type: ignore
update_end_user_table_data["user_id"] = data.user_id # type: ignore
verbose_proxy_logger.debug("In update customer, user_id condition block.")
response = await prisma_client.db.litellm_endusertable.update(
where={"user_id": data.user_id}, data=non_default_values # type: ignore
where={"user_id": data.user_id}, data=update_end_user_table_data, include={"litellm_budget_table": True} # type: ignore
)
if response is None:
raise ValueError(
Expand All @@ -444,13 +497,13 @@ async def update_end_user(
raise ValueError(f"user_id is required, passed user_id = {data.user_id}")

# update based on remaining passed in values

except Exception as e:
verbose_proxy_logger.error(
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.update_end_user(): Exception occured - {}".format(
str(e)
)
)
verbose_proxy_logger.debug(traceback.format_exc())
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "detail", f"Internal Server Error({str(e)})"),
Expand Down
7 changes: 6 additions & 1 deletion litellm/proxy/proxy_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -395,10 +395,12 @@ def generate_feedback_box():
# Import enterprise routes
try:
from litellm_enterprise.proxy.enterprise_routes import router as _enterprise_router
from litellm_enterprise.proxy.proxy_server import EnterpriseProxyConfig

enterprise_router = _enterprise_router
enterprise_proxy_config: Optional[EnterpriseProxyConfig] = EnterpriseProxyConfig()
except ImportError:
pass
enterprise_proxy_config = None
###################

server_root_path = os.getenv("SERVER_ROOT_PATH", "")
Expand Down Expand Up @@ -1863,6 +1865,9 @@ async def load_config( # noqa: PLR0915
value=custom_sso, config_file_path=config_file_path
)

if enterprise_proxy_config is not None:
await enterprise_proxy_config.load_enterprise_config(general_settings)

## pass through endpoints
if general_settings.get("pass_through_endpoints", None) is not None:
await initialize_pass_through_endpoints(
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
from unittest.mock import AsyncMock, patch

import pytest
from fastapi import FastAPI, HTTPException
from fastapi.testclient import TestClient

from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_EndUserTable,
LitellmUserRoles,
)
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
from litellm.proxy.management_endpoints.customer_endpoints import router
from litellm.proxy.proxy_server import ProxyException

app = FastAPI()
app.include_router(router)
client = TestClient(app)


@pytest.fixture
def mock_prisma_client():
with patch("litellm.proxy.proxy_server.prisma_client") as mock:
yield mock


@pytest.fixture
def mock_user_api_key_auth():
with patch("litellm.proxy.proxy_server.user_api_key_auth") as mock:
mock.return_value = UserAPIKeyAuth(
user_id="test-user", user_role=LitellmUserRoles.PROXY_ADMIN
)
yield mock


def test_update_customer_success(mock_prisma_client, mock_user_api_key_auth):
# Mock the database responses
mock_end_user = LiteLLM_EndUserTable(
user_id="test-user-1", alias="Test User", blocked=False
)
updated_mock_end_user = LiteLLM_EndUserTable(
user_id="test-user-1", alias="Updated Test User", blocked=False
)

# Mock the find_first response
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(
return_value=mock_end_user
)

# Mock the update response
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(
return_value=updated_mock_end_user
)

# Test data
test_data = {"user_id": "test-user-1", "alias": "Updated Test User"}

# Make the request
response = client.post(
"/customer/update", json=test_data, headers={"Authorization": "Bearer test-key"}
)

# Assert response
assert response.status_code == 200
assert response.json()["user_id"] == "test-user-1"
assert response.json()["alias"] == "Updated Test User"


def test_update_customer_not_found(mock_prisma_client, mock_user_api_key_auth):
# Mock the database response to return None (user not found)
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=None)

# Test data
test_data = {"user_id": "non-existent-user", "alias": "Test User"}

# Make the request
try:
response = client.post(
"/customer/update",
json=test_data,
headers={"Authorization": "Bearer test-key"},
)
except Exception as e:
print(e, type(e))
assert isinstance(e, ProxyException)
assert int(e.code) == 400
assert "End User Id=non-existent-user does not exist in db" in e.message