mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
Merge pull request #5099 from BerriAI/litellm_personal_user_budgets
fix(user_api_key_auth.py): respect team budgets over user budget, if key belongs to team
This commit is contained in:
commit
e1610d37b9
14 changed files with 256 additions and 219 deletions
|
|
@ -173,3 +173,23 @@ export PROXY_LOGOUT_URL="https://www.google.com"
|
|||
<Image img={require('../../img/ui_logout.png')} style={{ width: '400px', height: 'auto' }} />
|
||||
|
||||
|
||||
### Set max budget for internal users
|
||||
|
||||
Automatically apply budget per internal user when they sign up
|
||||
|
||||
```yaml
|
||||
litellm_settings:
|
||||
max_internal_user_budget: 10
|
||||
```
|
||||
|
||||
This sets a max budget of $10 USD for internal users when they sign up.
|
||||
|
||||
This budget only applies to personal keys created by that user - seen under `Default Team` on the UI.
|
||||
|
||||
<Image img={require('../../img/max_budget_for_internal_users.png')} style={{ width: '500px', height: 'auto' }} />
|
||||
|
||||
This budget does not apply to keys created under non-default teams.
|
||||
|
||||
### Set max budget for teams
|
||||
|
||||
[**Go Here**](./team_budgets.md)
|
||||
|
|
@ -53,6 +53,12 @@ UI_PASSWORD=langchain # password to sign in on UI
|
|||
|
||||
On accessing the LiteLLM UI, you will be prompted to enter your username, password
|
||||
|
||||
## Invite-other users
|
||||
|
||||
Allow others to create/delete their own keys.
|
||||
|
||||
[**Go Here**](./self_serve.md)
|
||||
|
||||
## ✨ Enterprise Features
|
||||
|
||||
Features here are behind a commercial license in our `/enterprise` folder. [**See Code**](https://github.com/BerriAI/litellm/tree/main/enterprise)
|
||||
|
|
|
|||
BIN
docs/my-website/img/max_budget_for_internal_users.png
Normal file
BIN
docs/my-website/img/max_budget_for_internal_users.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 130 KiB |
|
|
@ -260,6 +260,7 @@ upperbound_key_generate_params: Optional[LiteLLM_UpperboundKeyGenerateParams] =
|
|||
default_user_params: Optional[Dict] = None
|
||||
default_team_settings: Optional[List] = None
|
||||
max_user_budget: Optional[float] = None
|
||||
max_internal_user_budget: Optional[float] = None
|
||||
max_end_user_budget: Optional[float] = None
|
||||
#### REQUEST PRIORITIZATION ####
|
||||
priority_reservation: Optional[Dict[str, float]] = None
|
||||
|
|
|
|||
|
|
@ -730,10 +730,15 @@ LITELLM_EXCEPTION_TYPES = [
|
|||
|
||||
|
||||
class BudgetExceededError(Exception):
|
||||
def __init__(self, current_cost, max_budget):
|
||||
def __init__(
|
||||
self, current_cost: float, max_budget: float, message: Optional[str] = None
|
||||
):
|
||||
self.current_cost = current_cost
|
||||
self.max_budget = max_budget
|
||||
message = f"Budget has been exceeded! Current cost: {current_cost}, Max budget: {max_budget}"
|
||||
message = (
|
||||
message
|
||||
or f"Budget has been exceeded! Current cost: {current_cost}, Max budget: {max_budget}"
|
||||
)
|
||||
self.message = message
|
||||
super().__init__(message)
|
||||
|
||||
|
|
|
|||
|
|
@ -1338,6 +1338,7 @@ class LiteLLM_UserTable(LiteLLMBase):
|
|||
models: list = []
|
||||
tpm_limit: Optional[int] = None
|
||||
rpm_limit: Optional[int] = None
|
||||
user_role: Optional[str] = None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
|
|
|
|||
|
|
@ -55,11 +55,11 @@ def common_checks(
|
|||
1. If team is blocked
|
||||
2. If team can call model
|
||||
3. If team is in budget
|
||||
5. If user passed in (JWT or key.user_id) - is in budget
|
||||
4. If end_user (either via JWT or 'user' passed to /chat/completions, /embeddings endpoint) is in budget
|
||||
5. [OPTIONAL] If 'enforce_end_user' enabled - did developer pass in 'user' param for openai endpoints
|
||||
6. [OPTIONAL] If 'litellm.max_budget' is set (>0), is proxy under budget
|
||||
7. [OPTIONAL] If guardrails modified - is request allowed to change this
|
||||
4. If user passed in (JWT or key.user_id) - is in budget
|
||||
5. If end_user (either via JWT or 'user' passed to /chat/completions, /embeddings endpoint) is in budget
|
||||
6. [OPTIONAL] If 'enforce_end_user' enabled - did developer pass in 'user' param for openai endpoints
|
||||
7. [OPTIONAL] If 'litellm.max_budget' is set (>0), is proxy under budget
|
||||
8. [OPTIONAL] If guardrails modified - is request allowed to change this
|
||||
"""
|
||||
_model = request_body.get("model", None)
|
||||
if team_object is not None and team_object.blocked is True:
|
||||
|
|
@ -88,21 +88,34 @@ def common_checks(
|
|||
and team_object.spend is not None
|
||||
and team_object.spend > team_object.max_budget
|
||||
):
|
||||
raise Exception(
|
||||
f"Team={team_object.team_id} over budget. Spend={team_object.spend}, Budget={team_object.max_budget}"
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=team_object.spend,
|
||||
max_budget=team_object.max_budget,
|
||||
message=f"Team={team_object.team_id} over budget. Spend={team_object.spend}, Budget={team_object.max_budget}",
|
||||
)
|
||||
if user_object is not None and user_object.max_budget is not None:
|
||||
# 4. If user is in budget
|
||||
## 4.1 check personal budget, if personal key
|
||||
if (
|
||||
(team_object is None or team_object.team_id is None)
|
||||
and user_object is not None
|
||||
and user_object.max_budget is not None
|
||||
):
|
||||
user_budget = user_object.max_budget
|
||||
if user_budget > user_object.spend:
|
||||
raise Exception(
|
||||
f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_object.spend}, Budget={user_budget}"
|
||||
if user_budget < user_object.spend:
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=user_object.spend,
|
||||
max_budget=user_budget,
|
||||
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_object.spend}, Budget={user_budget}",
|
||||
)
|
||||
## 4.2 check team member budget, if team key
|
||||
# 5. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
|
||||
if end_user_object is not None and end_user_object.litellm_budget_table is not None:
|
||||
end_user_budget = end_user_object.litellm_budget_table.max_budget
|
||||
if end_user_budget is not None and end_user_object.spend > end_user_budget:
|
||||
raise Exception(
|
||||
f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}"
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=end_user_object.spend,
|
||||
max_budget=end_user_budget,
|
||||
message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}",
|
||||
)
|
||||
# 6. [OPTIONAL] If 'enforce_user_param' enabled - did developer pass in 'user' param for openai endpoints
|
||||
if (
|
||||
|
|
|
|||
|
|
@ -581,8 +581,8 @@ async def user_api_key_auth(
|
|||
"allowed_model_region"
|
||||
)
|
||||
|
||||
user_id_information: Optional[List] = None
|
||||
if valid_token is not None:
|
||||
user_obj: Optional[LiteLLM_UserTable] = None
|
||||
# Got Valid Token from Cache, DB
|
||||
# Run checks for
|
||||
# 1. If token can call model
|
||||
|
|
@ -650,114 +650,17 @@ async def user_api_key_auth(
|
|||
valid_token=valid_token,
|
||||
)
|
||||
|
||||
# Check 2. If user_id for this token is in budget
|
||||
# Check 2. If user_id for this token is in budget - done in common_checks()
|
||||
if valid_token.user_id is not None:
|
||||
user_id_list = [valid_token.user_id]
|
||||
for id in user_id_list:
|
||||
value = user_api_key_cache.get_cache(key=id)
|
||||
if value is not None:
|
||||
if user_id_information is None:
|
||||
user_id_information = []
|
||||
user_id_information.append(value)
|
||||
if user_id_information is None or (
|
||||
isinstance(user_id_information, list)
|
||||
and len(user_id_information) < 1
|
||||
):
|
||||
if prisma_client is not None:
|
||||
user_id_information = await prisma_client.get_data(
|
||||
user_id_list=[
|
||||
valid_token.user_id,
|
||||
],
|
||||
table_name="user",
|
||||
query_type="find_all",
|
||||
)
|
||||
if user_id_information is not None:
|
||||
for _id in user_id_information:
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=_id["user_id"],
|
||||
value=_id,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"user_id_information: {user_id_information}"
|
||||
user_obj = await get_user_object(
|
||||
user_id=valid_token.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
if user_id_information is not None:
|
||||
if isinstance(user_id_information, list):
|
||||
## Check if user in budget
|
||||
for _user in user_id_information:
|
||||
if _user is None:
|
||||
continue
|
||||
assert isinstance(_user, dict)
|
||||
# check if user is admin #
|
||||
|
||||
# Token exists, not expired now check if its in budget for the user
|
||||
user_max_budget = _user.get("max_budget", None)
|
||||
user_current_spend = _user.get("spend", None)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"user_id: {_user.get('user_id', None)}; user_max_budget: {user_max_budget}; user_current_spend: {user_current_spend}"
|
||||
)
|
||||
|
||||
if (
|
||||
user_max_budget is not None
|
||||
and user_current_spend is not None
|
||||
):
|
||||
call_info = CallInfo(
|
||||
token=valid_token.token,
|
||||
spend=user_current_spend,
|
||||
max_budget=user_max_budget,
|
||||
user_id=_user.get("user_id", None),
|
||||
user_email=_user.get("user_email", None),
|
||||
key_alias=valid_token.key_alias,
|
||||
)
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.budget_alerts(
|
||||
type="user_budget",
|
||||
user_info=call_info,
|
||||
)
|
||||
)
|
||||
|
||||
_user_id = _user.get("user_id", None)
|
||||
if user_current_spend > user_max_budget:
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=user_current_spend,
|
||||
max_budget=user_max_budget,
|
||||
)
|
||||
else:
|
||||
# Token exists, not expired now check if its in budget for the user
|
||||
user_max_budget = getattr(
|
||||
user_id_information, "max_budget", None
|
||||
)
|
||||
user_current_spend = getattr(user_id_information, "spend", None)
|
||||
|
||||
if (
|
||||
user_max_budget is not None
|
||||
and user_current_spend is not None
|
||||
):
|
||||
call_info = CallInfo(
|
||||
token=valid_token.token,
|
||||
spend=user_current_spend,
|
||||
max_budget=user_max_budget,
|
||||
user_id=getattr(user_id_information, "user_id", None),
|
||||
user_email=getattr(
|
||||
user_id_information, "user_email", None
|
||||
),
|
||||
key_alias=valid_token.key_alias,
|
||||
)
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.budget_alerts(
|
||||
type="user_budget",
|
||||
user_info=call_info,
|
||||
)
|
||||
)
|
||||
|
||||
if user_current_spend > user_max_budget:
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=user_current_spend,
|
||||
max_budget=user_max_budget,
|
||||
)
|
||||
|
||||
# Check 3. Check if user is in their team budget
|
||||
if valid_token.team_member_spend is not None:
|
||||
if prisma_client is not None:
|
||||
|
|
@ -829,11 +732,8 @@ async def user_api_key_auth(
|
|||
|
||||
user_email: Optional[str] = None
|
||||
# Check if the token has any user id information
|
||||
if user_id_information is not None and len(user_id_information) > 0:
|
||||
specific_user_id_information = user_id_information[0]
|
||||
_user_email = specific_user_id_information.get("user_email", None)
|
||||
if _user_email is not None:
|
||||
user_email = str(_user_email)
|
||||
if user_obj is not None:
|
||||
user_email = user_obj.user_email
|
||||
|
||||
call_info = CallInfo(
|
||||
token=valid_token.token,
|
||||
|
|
@ -983,7 +883,7 @@ async def user_api_key_auth(
|
|||
_ = common_checks(
|
||||
request_body=request_data,
|
||||
team_object=_team_obj,
|
||||
user_object=None,
|
||||
user_object=user_obj,
|
||||
end_user_object=_end_user_object,
|
||||
general_settings=general_settings,
|
||||
global_proxy_spend=global_proxy_spend,
|
||||
|
|
@ -1007,9 +907,9 @@ async def user_api_key_auth(
|
|||
if _end_user_object is not None:
|
||||
valid_token_dict.update(end_user_params)
|
||||
|
||||
_user_role = _get_user_role(user_id_information=user_id_information)
|
||||
_user_role = _get_user_role(user_obj=user_obj)
|
||||
|
||||
if not _is_user_proxy_admin(user_id_information): # if non-admin
|
||||
if not _is_user_proxy_admin(user_obj=user_obj): # if non-admin
|
||||
if is_llm_api_route(route=route):
|
||||
pass
|
||||
elif is_llm_api_route(route=request["route"].name):
|
||||
|
|
@ -1091,14 +991,9 @@ async def user_api_key_auth(
|
|||
else:
|
||||
user_role = "unknown"
|
||||
user_id = "unknown"
|
||||
if (
|
||||
user_id_information is not None
|
||||
and isinstance(user_id_information, list)
|
||||
and len(user_id_information) > 0
|
||||
):
|
||||
_user = user_id_information[0]
|
||||
user_role = _user.get("user_role", "unknown")
|
||||
user_id = _user.get("user_id", "unknown")
|
||||
if user_obj is not None:
|
||||
user_role = user_obj.user_role or "unknown"
|
||||
user_id = user_obj.user_id or "unknown"
|
||||
raise Exception(
|
||||
f"Only proxy admin can be used to generate, delete, update info for new keys/users/teams. Route={route}. Your role={user_role}. Your user_id={user_id}"
|
||||
)
|
||||
|
|
@ -1144,9 +1039,7 @@ async def user_api_key_auth(
|
|||
# Do something if the current route starts with any of the allowed routes
|
||||
pass
|
||||
else:
|
||||
if user_id_information is not None and _is_user_proxy_admin(
|
||||
user_id_information
|
||||
):
|
||||
if user_obj is not None and _is_user_proxy_admin(user_obj=user_obj):
|
||||
return UserAPIKeyAuth(
|
||||
api_key=api_key,
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
|
|
@ -1172,7 +1065,7 @@ async def user_api_key_auth(
|
|||
raise Exception("Invalid proxy server token passed")
|
||||
if valid_token_dict is not None:
|
||||
return _return_user_api_key_auth_obj(
|
||||
user_id_information=user_id_information,
|
||||
user_obj=user_obj,
|
||||
api_key=api_key,
|
||||
parent_otel_span=parent_otel_span,
|
||||
valid_token_dict=valid_token_dict,
|
||||
|
|
@ -1219,17 +1112,16 @@ async def user_api_key_auth(
|
|||
|
||||
|
||||
def _return_user_api_key_auth_obj(
|
||||
user_id_information: Optional[list],
|
||||
user_obj: Optional[LiteLLM_UserTable],
|
||||
api_key: str,
|
||||
parent_otel_span: Optional[Span],
|
||||
valid_token_dict: dict,
|
||||
route: str,
|
||||
) -> UserAPIKeyAuth:
|
||||
retrieved_user_role = (
|
||||
_get_user_role(user_id_information=user_id_information)
|
||||
or LitellmUserRoles.INTERNAL_USER
|
||||
_get_user_role(user_obj=user_obj) or LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
if user_id_information is not None and _is_user_proxy_admin(user_id_information):
|
||||
if user_obj is not None and _is_user_proxy_admin(user_obj=user_obj):
|
||||
return UserAPIKeyAuth(
|
||||
api_key=api_key,
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
|
|
@ -1270,30 +1162,19 @@ def _has_user_setup_sso():
|
|||
return sso_setup
|
||||
|
||||
|
||||
def _is_user_proxy_admin(user_id_information: Optional[list]):
|
||||
if user_id_information is None:
|
||||
def _is_user_proxy_admin(user_obj: Optional[LiteLLM_UserTable]):
|
||||
if user_obj is None:
|
||||
return False
|
||||
|
||||
if len(user_id_information) == 0 or user_id_information[0] is None:
|
||||
return False
|
||||
|
||||
_user = user_id_information[0]
|
||||
if (
|
||||
_user.get("user_role", None) is not None
|
||||
and _user.get("user_role") == LitellmUserRoles.PROXY_ADMIN.value
|
||||
user_obj.user_role is not None
|
||||
and user_obj.user_role == LitellmUserRoles.PROXY_ADMIN.value
|
||||
):
|
||||
return True
|
||||
|
||||
# if user_id_information contains litellm-proxy-budget
|
||||
# get first user_id that is not litellm-proxy-budget
|
||||
for user in user_id_information:
|
||||
if user.get("user_id") != "litellm-proxy-budget":
|
||||
_user = user
|
||||
break
|
||||
|
||||
if (
|
||||
_user.get("user_role", None) is not None
|
||||
and _user.get("user_role") == LitellmUserRoles.PROXY_ADMIN.value
|
||||
user_obj.user_role is not None
|
||||
and user_obj.user_role == LitellmUserRoles.PROXY_ADMIN.value
|
||||
):
|
||||
return True
|
||||
|
||||
|
|
@ -1301,29 +1182,20 @@ def _is_user_proxy_admin(user_id_information: Optional[list]):
|
|||
|
||||
|
||||
def _get_user_role(
|
||||
user_id_information: Optional[list],
|
||||
) -> Optional[
|
||||
Literal[
|
||||
LitellmUserRoles.PROXY_ADMIN,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||||
LitellmUserRoles.INTERNAL_USER,
|
||||
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
|
||||
LitellmUserRoles.TEAM,
|
||||
LitellmUserRoles.CUSTOMER,
|
||||
]
|
||||
]:
|
||||
if user_id_information is None:
|
||||
user_obj: Optional[LiteLLM_UserTable],
|
||||
) -> Optional[LitellmUserRoles]:
|
||||
if user_obj is None:
|
||||
return None
|
||||
|
||||
if len(user_id_information) == 0 or user_id_information[0] is None:
|
||||
return None
|
||||
_user = user_obj
|
||||
|
||||
_user = user_id_information[0]
|
||||
_user_role = _user.user_role
|
||||
try:
|
||||
role = LitellmUserRoles(_user_role)
|
||||
except ValueError:
|
||||
return LitellmUserRoles.INTERNAL_USER
|
||||
|
||||
_user_role = _user.get("user_role")
|
||||
if _user_role in list(LitellmUserRoles.__annotations__.keys()):
|
||||
return _user_role
|
||||
return LitellmUserRoles.INTERNAL_USER
|
||||
return role
|
||||
|
||||
|
||||
def _check_valid_ip(allowed_ips: Optional[List[str]], request: Request) -> bool:
|
||||
|
|
|
|||
|
|
@ -87,6 +87,10 @@ async def new_user(
|
|||
"user" # only create a user, don't create key if 'auto_create_key' set to False
|
||||
)
|
||||
|
||||
if "max_budget" in data_json and data_json["max_budget"] is None:
|
||||
if litellm.max_internal_user_budget is not None:
|
||||
data_json["max_budget"] = litellm.max_internal_user_budget
|
||||
|
||||
response = await generate_key_helper_fn(request_type="user", **data_json)
|
||||
|
||||
# Admin UI Logic
|
||||
|
|
|
|||
|
|
@ -938,6 +938,7 @@ def test_completion_function_plus_image(model):
|
|||
}
|
||||
]
|
||||
|
||||
try:
|
||||
response = completion(
|
||||
model=model,
|
||||
messages=[image_message],
|
||||
|
|
@ -949,8 +950,6 @@ def test_completion_function_plus_image(model):
|
|||
print(response)
|
||||
except litellm.InternalServerError:
|
||||
pass
|
||||
except Exception as e:
|
||||
pytest.fail(f"error occurred: {str(e)}")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
|
|||
|
|
@ -504,7 +504,7 @@ def test_call_with_user_over_budget(prisma_client):
|
|||
asyncio.run(test())
|
||||
except Exception as e:
|
||||
error_detail = e.message
|
||||
assert "Budget has been exceeded" in error_detail
|
||||
assert "ExceededBudget:" in error_detail
|
||||
assert isinstance(e, ProxyException)
|
||||
assert e.type == ProxyErrorTypes.budget_exceeded
|
||||
print(vars(e))
|
||||
|
|
@ -607,7 +607,7 @@ def test_call_with_end_user_over_budget(prisma_client):
|
|||
# use generated key to auth in
|
||||
result = await user_api_key_auth(request=request, api_key=bearer_token)
|
||||
print("result from user auth with new key", result)
|
||||
pytest.fail(f"This should have failed!. They key crossed it's budget")
|
||||
pytest.fail("This should have failed!. They key crossed it's budget")
|
||||
|
||||
asyncio.run(test())
|
||||
except Exception as e:
|
||||
|
|
@ -779,12 +779,12 @@ def test_call_with_user_over_budget_stream(prisma_client):
|
|||
# use generated key to auth in
|
||||
result = await user_api_key_auth(request=request, api_key=bearer_token)
|
||||
print("result from user auth with new key", result)
|
||||
pytest.fail(f"This should have failed!. They key crossed it's budget")
|
||||
pytest.fail("This should have failed!. They key crossed it's budget")
|
||||
|
||||
asyncio.run(test())
|
||||
except Exception as e:
|
||||
error_detail = e.message
|
||||
assert "Budget has been exceeded" in error_detail
|
||||
assert "ExceededBudget:" in error_detail
|
||||
assert isinstance(e, ProxyException)
|
||||
assert e.type == ProxyErrorTypes.budget_exceeded
|
||||
print(vars(e))
|
||||
|
|
@ -2511,7 +2511,6 @@ async def test_update_user_role(prisma_client):
|
|||
Tests if we update user role, incorrect values are not stored in cache
|
||||
-> create a user with role == INTERNAL_USER
|
||||
-> access an Admin only route -> expect to fail
|
||||
|
||||
-> update user role to == PROXY_ADMIN
|
||||
-> access an Admin only route -> expect to succeed
|
||||
"""
|
||||
|
|
@ -2556,6 +2555,7 @@ async def test_update_user_role(prisma_client):
|
|||
await asyncio.sleep(2)
|
||||
|
||||
# use generated key to auth in
|
||||
print("\n\nMAKING NEW REQUEST WITH UPDATED USER ROLE\n\n")
|
||||
result = await user_api_key_auth(request=request, api_key=api_key)
|
||||
print("result from user auth with new key", result)
|
||||
|
||||
|
|
|
|||
|
|
@ -800,3 +800,39 @@ async def test_get_team_redis(client_no_auth):
|
|||
pass
|
||||
|
||||
mock_client.assert_called_once()
|
||||
|
||||
|
||||
import random
|
||||
import uuid
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, NewUserRequest, UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
|
||||
from litellm.tests.test_key_generate_prisma import prisma_client
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_user_default_budget(prisma_client):
|
||||
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
setattr(litellm, "max_internal_user_budget", 10)
|
||||
await litellm.proxy.proxy_server.prisma_client.connect()
|
||||
user = f"ishaan {uuid.uuid4().hex}"
|
||||
request = NewUserRequest(user_id=user) # create a key with no budget
|
||||
with patch.object(
|
||||
litellm.proxy.proxy_server.prisma_client, "insert_data", new=AsyncMock()
|
||||
) as mock_client:
|
||||
await new_user(
|
||||
request,
|
||||
)
|
||||
|
||||
mock_client.assert_called()
|
||||
|
||||
print(f"mock_client.call_args: {mock_client.call_args}")
|
||||
print("mock_client.call_args.kwargs: {}".format(mock_client.call_args.kwargs))
|
||||
|
||||
assert (
|
||||
mock_client.call_args.kwargs["data"]["max_budget"]
|
||||
== litellm.max_internal_user_budget
|
||||
)
|
||||
|
|
|
|||
|
|
@ -96,23 +96,87 @@ async def test_check_blocked_team():
|
|||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"user_role", ["app_user", "internal_user", "proxy_admin_viewer"]
|
||||
"user_role, expected_role",
|
||||
[
|
||||
("app_user", "internal_user"),
|
||||
("internal_user", "internal_user"),
|
||||
("proxy_admin_viewer", "proxy_admin_viewer"),
|
||||
],
|
||||
)
|
||||
def test_returned_user_api_key_auth(user_role):
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
def test_returned_user_api_key_auth(user_role, expected_role):
|
||||
from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles
|
||||
from litellm.proxy.auth.user_api_key_auth import _return_user_api_key_auth_obj
|
||||
|
||||
user_id_information = [{"user_role": user_role}]
|
||||
|
||||
new_obj = _return_user_api_key_auth_obj(
|
||||
user_id_information,
|
||||
user_obj=LiteLLM_UserTable(
|
||||
user_role=user_role, user_id="", max_budget=None, user_email=""
|
||||
),
|
||||
api_key="hello-world",
|
||||
parent_otel_span=None,
|
||||
valid_token_dict={},
|
||||
route="/chat/completion",
|
||||
)
|
||||
|
||||
if user_role in list(LitellmUserRoles.__annotations__.keys()):
|
||||
assert new_obj.user_role == user_role
|
||||
else:
|
||||
assert new_obj.user_role == "internal_user"
|
||||
assert new_obj.user_role == expected_role
|
||||
|
||||
|
||||
@pytest.mark.parametrize("key_ownership", ["user_key", "team_key"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_personal_budgets(key_ownership):
|
||||
"""
|
||||
Set a personal budget on a user
|
||||
|
||||
- have it only apply when key belongs to user -> raises BudgetExceededError
|
||||
- if key belongs to team, have key respect team budget -> allows call to go through
|
||||
"""
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
from litellm.proxy._types import LiteLLM_UserTable, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.proxy_server import hash_token, user_api_key_cache
|
||||
|
||||
_user_id = "1234"
|
||||
user_key = "sk-12345678"
|
||||
|
||||
if key_ownership == "user_key":
|
||||
valid_token = UserAPIKeyAuth(
|
||||
token=hash_token(user_key),
|
||||
last_refreshed_at=time.time(),
|
||||
user_id=_user_id,
|
||||
spend=20,
|
||||
)
|
||||
elif key_ownership == "team_key":
|
||||
valid_token = UserAPIKeyAuth(
|
||||
token=hash_token(user_key),
|
||||
last_refreshed_at=time.time(),
|
||||
user_id=_user_id,
|
||||
team_id="my-special-team",
|
||||
team_max_budget=100,
|
||||
spend=20,
|
||||
)
|
||||
await asyncio.sleep(1)
|
||||
user_obj = LiteLLM_UserTable(
|
||||
user_id=_user_id, spend=11, max_budget=10, user_email=""
|
||||
)
|
||||
user_api_key_cache.set_cache(key=hash_token(user_key), value=valid_token)
|
||||
user_api_key_cache.set_cache(key="{}".format(_user_id), value=user_obj)
|
||||
|
||||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", "hello-world")
|
||||
|
||||
request = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
|
||||
try:
|
||||
await user_api_key_auth(request=request, api_key="Bearer " + user_key)
|
||||
|
||||
if key_ownership == "user_key":
|
||||
pytest.fail("Expected this call to fail. User is over limit.")
|
||||
except Exception:
|
||||
if key_ownership == "team_key":
|
||||
pytest.fail("Expected this call to work. Key is below team budget.")
|
||||
|
|
|
|||
|
|
@ -45,32 +45,33 @@ interface ViewUserSpendProps {
|
|||
const ViewUserSpend: React.FC<ViewUserSpendProps> = ({ userID, userRole, accessToken, userSpend, selectedTeam }) => {
|
||||
console.log(`userSpend: ${userSpend}`)
|
||||
let [spend, setSpend] = useState(userSpend !== null ? userSpend : 0.0);
|
||||
const [maxBudget, setMaxBudget] = useState(0.0);
|
||||
const [maxBudget, setMaxBudget] = useState(selectedTeam ? selectedTeam.max_budget : null);
|
||||
console.log(`maxBudget: ${maxBudget}, selectedTeam.max_budget: ${selectedTeam.max_budget}, selectedTeam: ${JSON.stringify(selectedTeam)}`)
|
||||
const [userModels, setUserModels] = useState([]);
|
||||
useEffect(() => {
|
||||
const fetchData = async () => {
|
||||
if (!accessToken || !userID || !userRole) {
|
||||
return;
|
||||
}
|
||||
if (userRole === "Admin" && userSpend == null) {
|
||||
try {
|
||||
const globalSpend = await getTotalSpendCall(accessToken);
|
||||
if (globalSpend) {
|
||||
if (globalSpend.spend) {
|
||||
setSpend(globalSpend.spend);
|
||||
} else {
|
||||
setSpend(0.0);
|
||||
}
|
||||
if (globalSpend.max_budget) {
|
||||
setMaxBudget(globalSpend.max_budget);
|
||||
} else {
|
||||
setMaxBudget(0.0);
|
||||
}
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("Error fetching global spend data:", error);
|
||||
}
|
||||
}
|
||||
// if (userRole === "Admin" && userSpend == null) {
|
||||
// try {
|
||||
// const globalSpend = await getTotalSpendCall(accessToken);
|
||||
// if (globalSpend) {
|
||||
// if (globalSpend.spend) {
|
||||
// setSpend(globalSpend.spend);
|
||||
// } else {
|
||||
// setSpend(0.0);
|
||||
// }
|
||||
// if (globalSpend.max_budget) {
|
||||
// setMaxBudget(globalSpend.max_budget);
|
||||
// } else {
|
||||
// setMaxBudget(null);
|
||||
// }
|
||||
// }
|
||||
// } catch (error) {
|
||||
// console.error("Error fetching global spend data:", error);
|
||||
// }
|
||||
// }
|
||||
};
|
||||
const fetchUserModels = async () => {
|
||||
try {
|
||||
|
|
@ -102,6 +103,11 @@ const ViewUserSpend: React.FC<ViewUserSpendProps> = ({ userID, userRole, accessT
|
|||
setSpend(userSpend)
|
||||
}
|
||||
}, [userSpend])
|
||||
useEffect(() => {
|
||||
if (selectedTeam && selectedTeam.max_budget !== maxBudget) {
|
||||
setMaxBudget(selectedTeam.max_budget);
|
||||
}
|
||||
}, [selectedTeam, maxBudget]);
|
||||
|
||||
// logic to decide what models to display
|
||||
let modelsToDisplay = [];
|
||||
|
|
@ -127,14 +133,24 @@ const ViewUserSpend: React.FC<ViewUserSpendProps> = ({ userID, userRole, accessT
|
|||
console.log(`spend in view user spend: ${spend}`)
|
||||
return (
|
||||
<div className="flex items-center">
|
||||
<div className="flex justify-between gap-x-6">
|
||||
<div>
|
||||
<p className="text-tremor-default text-tremor-content dark:text-dark-tremor-content">
|
||||
Total Spend{" "}
|
||||
Total Spend
|
||||
</p>
|
||||
<p className="text-2xl text-tremor-content-strong dark:text-dark-tremor-content-strong font-semibold">
|
||||
${roundedSpend}
|
||||
</p>
|
||||
</div>
|
||||
<div>
|
||||
<p className="text-tremor-default text-tremor-content dark:text-dark-tremor-content">
|
||||
Max Budget
|
||||
</p>
|
||||
<p className="text-2xl text-tremor-content-strong dark:text-dark-tremor-content-strong font-semibold">
|
||||
{displayMaxBudget}
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
{/* <div className="ml-auto">
|
||||
<Accordion>
|
||||
<AccordionHeader><Text>Team Models</Text></AccordionHeader>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue