mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
Merge pull request #5120 from BerriAI/litellm_sso_max_internal_budget_fixes
feat: set max_internal_budget for user w/ sso
This commit is contained in:
commit
529178c82a
6 changed files with 101 additions and 38 deletions
|
|
@ -1744,3 +1744,12 @@ class ProxyErrorTypes(str, enum.Enum):
|
|||
auth_error = "auth_error"
|
||||
internal_server_error = "internal_server_error"
|
||||
bad_request_error = "bad_request_error"
|
||||
|
||||
|
||||
class SSOUserDefinedValues(TypedDict):
|
||||
models: List[str]
|
||||
user_id: str
|
||||
user_email: Optional[str]
|
||||
user_role: Optional[str]
|
||||
max_budget: Optional[float]
|
||||
budget_duration: Optional[str]
|
||||
|
|
|
|||
|
|
@ -87,12 +87,16 @@ async def new_user(
|
|||
"user" # only create a user, don't create key if 'auto_create_key' set to False
|
||||
)
|
||||
|
||||
is_internal_user = False
|
||||
if data.user_role == LitellmUserRoles.INTERNAL_USER:
|
||||
is_internal_user = True
|
||||
|
||||
if "max_budget" in data_json and data_json["max_budget"] is None:
|
||||
if litellm.max_internal_user_budget is not None:
|
||||
if is_internal_user and litellm.max_internal_user_budget is not None:
|
||||
data_json["max_budget"] = litellm.max_internal_user_budget
|
||||
|
||||
if "budget_duration" in data_json and data_json["budget_duration"] is None:
|
||||
if litellm.internal_user_budget_duration is not None:
|
||||
if is_internal_user and litellm.internal_user_budget_duration is not None:
|
||||
data_json["budget_duration"] = litellm.internal_user_budget_duration
|
||||
|
||||
response = await generate_key_helper_fn(request_type="user", **data_json)
|
||||
|
|
|
|||
|
|
@ -283,7 +283,6 @@ except Exception as e:
|
|||
pass
|
||||
|
||||
server_root_path = os.getenv("SERVER_ROOT_PATH", "")
|
||||
print("server root path: ", server_root_path) # noqa
|
||||
_license_check = LicenseCheck()
|
||||
premium_user: bool = _license_check.is_premium()
|
||||
ui_link = f"{server_root_path}/ui/"
|
||||
|
|
@ -8629,8 +8628,13 @@ async def auth_callback(request: Request):
|
|||
_last_name = getattr(result, "last_name", "") or ""
|
||||
user_id = _first_name + _last_name
|
||||
|
||||
if user_email is not None and (user_id is None or len(user_id) == 0):
|
||||
user_id = user_email
|
||||
|
||||
user_info = None
|
||||
user_id_models: List = []
|
||||
max_internal_user_budget = litellm.max_internal_user_budget
|
||||
internal_user_budget_duration = litellm.internal_user_budget_duration
|
||||
|
||||
# User might not be already created on first generation of key
|
||||
# But if it is, we want their models preferences
|
||||
|
|
@ -8642,10 +8646,13 @@ async def auth_callback(request: Request):
|
|||
"spend": 0,
|
||||
"team_id": "litellm-dashboard",
|
||||
}
|
||||
user_defined_values = {
|
||||
user_defined_values: SSOUserDefinedValues = {
|
||||
"models": user_id_models,
|
||||
"user_id": user_id,
|
||||
"user_email": user_email,
|
||||
"max_budget": max_internal_user_budget,
|
||||
"user_role": None,
|
||||
"budget_duration": internal_user_budget_duration,
|
||||
}
|
||||
_user_id_from_sso = user_id
|
||||
try:
|
||||
|
|
@ -8657,10 +8664,16 @@ async def auth_callback(request: Request):
|
|||
)
|
||||
if user_info is not None:
|
||||
user_defined_values = {
|
||||
"models": getattr(user_info, "models", []),
|
||||
"models": getattr(user_info, "models", user_id_models),
|
||||
"user_id": getattr(user_info, "user_id", user_id),
|
||||
"user_email": getattr(user_info, "user_id", user_email),
|
||||
"user_role": getattr(user_info, "user_role", None),
|
||||
"max_budget": getattr(
|
||||
user_info, "max_budget", max_internal_user_budget
|
||||
),
|
||||
"budget_duration": getattr(
|
||||
user_info, "budget_duration", internal_user_budget_duration
|
||||
),
|
||||
}
|
||||
user_role = getattr(user_info, "user_role", None)
|
||||
|
||||
|
|
@ -8674,6 +8687,12 @@ async def auth_callback(request: Request):
|
|||
"user_id": user_id,
|
||||
"user_email": getattr(user_info, "user_id", user_email),
|
||||
"user_role": getattr(user_info, "user_role", None),
|
||||
"max_budget": getattr(
|
||||
user_info, "max_budget", max_internal_user_budget
|
||||
),
|
||||
"budget_duration": getattr(
|
||||
user_info, "budget_duration", internal_user_budget_duration
|
||||
),
|
||||
}
|
||||
user_role = getattr(user_info, "user_role", None)
|
||||
|
||||
|
|
@ -8690,10 +8709,38 @@ async def auth_callback(request: Request):
|
|||
"user_email": litellm.default_user_params.get(
|
||||
"user_email", user_email
|
||||
),
|
||||
"user_role": litellm.default_user_params.get("user_role", None),
|
||||
"max_budget": litellm.default_user_params.get(
|
||||
"max_budget", max_internal_user_budget
|
||||
),
|
||||
"budget_duration": litellm.default_user_params.get(
|
||||
"budget_duration", internal_user_budget_duration
|
||||
),
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
pass
|
||||
|
||||
is_internal_user = False
|
||||
if (
|
||||
user_defined_values["user_role"] is not None
|
||||
and user_defined_values["user_role"] == LitellmUserRoles.INTERNAL_USER.value
|
||||
):
|
||||
is_internal_user = True
|
||||
if (
|
||||
is_internal_user is True
|
||||
and user_defined_values["max_budget"] is None
|
||||
and litellm.max_internal_user_budget is not None
|
||||
):
|
||||
user_defined_values["max_budget"] = litellm.max_internal_user_budget
|
||||
|
||||
if (
|
||||
is_internal_user is True
|
||||
and user_defined_values["budget_duration"] is None
|
||||
and litellm.internal_user_budget_duration is not None
|
||||
):
|
||||
user_defined_values["budget_duration"] = litellm.internal_user_budget_duration
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"user_defined_values for creating ui key: {user_defined_values}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -811,15 +811,22 @@ from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
|
|||
from litellm.tests.test_key_generate_prisma import prisma_client
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"user_role",
|
||||
[LitellmUserRoles.INTERNAL_USER.value, LitellmUserRoles.PROXY_ADMIN.value],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_user_default_budget(prisma_client):
|
||||
async def test_create_user_default_budget(prisma_client, user_role):
|
||||
|
||||
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)
|
||||
setattr(litellm, "internal_user_budget_duration", "5m")
|
||||
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
|
||||
request = NewUserRequest(
|
||||
user_id=user, user_role=user_role
|
||||
) # create a key with no budget
|
||||
with patch.object(
|
||||
litellm.proxy.proxy_server.prisma_client, "insert_data", new=AsyncMock()
|
||||
) as mock_client:
|
||||
|
|
@ -832,7 +839,16 @@ async def test_create_user_default_budget(prisma_client):
|
|||
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
|
||||
)
|
||||
if user_role == LitellmUserRoles.INTERNAL_USER.value:
|
||||
assert (
|
||||
mock_client.call_args.kwargs["data"]["max_budget"]
|
||||
== litellm.max_internal_user_budget
|
||||
)
|
||||
assert (
|
||||
mock_client.call_args.kwargs["data"]["budget_duration"]
|
||||
== litellm.internal_user_budget_duration
|
||||
)
|
||||
|
||||
else:
|
||||
assert mock_client.call_args.kwargs["data"]["max_budget"] is None
|
||||
assert mock_client.call_args.kwargs["data"]["budget_duration"] is None
|
||||
|
|
|
|||
|
|
@ -315,6 +315,7 @@ const UserDashboard: React.FC<UserDashboardProps> = ({
|
|||
<ViewUserSpend
|
||||
userID={userID}
|
||||
userRole={userRole}
|
||||
userMaxBudget={userSpendData?.max_budget || null}
|
||||
accessToken={accessToken}
|
||||
userSpend={teamSpend}
|
||||
selectedTeam={selectedTeam ? selectedTeam : null}
|
||||
|
|
|
|||
|
|
@ -40,39 +40,30 @@ interface ViewUserSpendProps {
|
|||
userRole: string | null;
|
||||
accessToken: string | null;
|
||||
userSpend: number | null;
|
||||
userMaxBudget: number | null;
|
||||
selectedTeam: any | null;
|
||||
}
|
||||
const ViewUserSpend: React.FC<ViewUserSpendProps> = ({ userID, userRole, accessToken, userSpend, selectedTeam }) => {
|
||||
const ViewUserSpend: React.FC<ViewUserSpendProps> = ({ userID, userRole, accessToken, userSpend, userMaxBudget, selectedTeam }) => {
|
||||
console.log(`userSpend: ${userSpend}`)
|
||||
let [spend, setSpend] = useState(userSpend !== null ? userSpend : 0.0);
|
||||
const [maxBudget, setMaxBudget] = useState(selectedTeam ? selectedTeam.max_budget : null);
|
||||
const team_budget = selectedTeam ? selectedTeam.max_budget : "unlimited";
|
||||
console.log(`maxBudget: ${maxBudget}, selectedTeam.max_budget: ${team_budget}, selectedTeam: ${JSON.stringify(selectedTeam)}`)
|
||||
useEffect(() => {
|
||||
console.log(`Updating team max budget for default team - userMaxBudget - ${userMaxBudget}, selectedTeam.team_alias - ${selectedTeam.team_alias}`)
|
||||
if (selectedTeam) {
|
||||
if (selectedTeam.team_alias === "Default Team") {
|
||||
setMaxBudget(userMaxBudget);
|
||||
} else {
|
||||
setMaxBudget(selectedTeam.max_budget);
|
||||
}
|
||||
}
|
||||
}, [selectedTeam, userMaxBudget]);
|
||||
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(null);
|
||||
// }
|
||||
// }
|
||||
// } catch (error) {
|
||||
// console.error("Error fetching global spend data:", error);
|
||||
// }
|
||||
// }
|
||||
};
|
||||
const fetchUserModels = async () => {
|
||||
try {
|
||||
|
|
@ -104,11 +95,6 @@ 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 = [];
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue