fix(model_management_endpoints.py): allow team admin to update model … (#10539)

* fix(model_management_endpoints.py): allow team admin to update model via `/model/{model_id}/update` route

Fixes ui regression where team admin could not modify their own models

* fix(provider_specific_fields.tsx): style fix

* fix(table.tsx): allow expanding multiple rows

* fix(organization_endpoints.py): more robust check if user can give org model access

handle when user has models=["all-proxy-models"]

* fix(organization_endpoints.py): enable proxy admin with 'all-proxy-model' access to create new org with specific models

Fixes LIT-135

* fix: fix linting error

* fix: fix ui linting error

* fix(index.tsx): fix linting errors
This commit is contained in:
Krish Dholakia 2025-05-03 19:34:35 -07:00 • committed by GitHub
parent 42a91bae6b
commit 880c2a736b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 52 additions and 71 deletions

File diff suppressed because one or more lines are too long

View file

@ -479,6 +479,7 @@ class LiteLLMRoutes(enum.Enum):
"/model/update",
"/model/delete",
"/user/daily/activity",
"/model/{model_id}/update",
] # routes that manage their own allowed/disallowed logic
## Org Admin Routes ##

View file

@ -154,6 +154,7 @@ async def patch_model(
from litellm.proxy.proxy_server import (
litellm_proxy_admin_name,
llm_router,
premium_user,
prisma_client,
store_model_in_db,
)
@ -193,6 +194,12 @@ async def patch_model(
param=None,
)
await ModelManagementAuthChecks.can_user_make_model_call(
model_params=db_model,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
premium_user=premium_user,
)
# Create update dictionary only for provided fields
update_data = update_db_model(db_model=db_model, updated_patch=patch_data)

View file

@ -18,6 +18,7 @@ from fastapi import APIRouter, Depends, HTTPException, Request, status
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import *
from litellm.proxy.auth.auth_checks import can_user_call_model
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.budget_management_endpoints import (
new_budget,
@ -103,10 +104,17 @@ async def new_organization(
```
"""
from litellm.proxy.proxy_server import litellm_proxy_admin_name, prisma_client
from litellm.proxy.proxy_server import (
litellm_proxy_admin_name,
llm_router,
prisma_client,
)
if prisma_client is None:
raise HTTPException(status_code=500, detail={"error": "No db connected"})
raise HTTPException(
status_code=500,
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
if (
user_api_key_dict.user_role is None
@ -119,6 +127,22 @@ async def new_organization(
},
)
if llm_router is None:
raise HTTPException(
status_code=500, detail={"error": CommonProxyErrors.no_llm_router.value}
)
user_object_correct_type: Optional[LiteLLM_UserTable] = None
if user_api_key_dict.user_id is not None:
try:
user_object = await prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_api_key_dict.user_id}
)
user_object_correct_type = LiteLLM_UserTable(**user_object.model_dump())
except Exception:
pass
if data.budget_id is None:
"""
Every organization needs a budget attached.
@ -157,14 +181,11 @@ async def new_organization(
"error": "User not allowed to give access to all models. Select models you want org to have access to."
},
)
for m in data.models:
if m not in user_api_key_dict.models:
raise HTTPException(
status_code=400,
detail={
"error": f"User not allowed to give access to model={m}. Models you have access to = {user_api_key_dict.models}"
},
)
await can_user_call_model(
m, llm_router=llm_router, user_object=user_object_correct_type
)
organization_row = LiteLLM_OrganizationTable(
**data.json(exclude_none=True),

View file

@ -24,7 +24,7 @@ sys.path.insert(
) # Adds the parent directory to the system path
import asyncio
import logging
from unittest.mock import MagicMock
import pytest
import litellm
@ -141,6 +141,7 @@ async def test_create_new_user_in_organization(prisma_client, user_role):
master_key = "sk-1234"
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
setattr(litellm.proxy.proxy_server, "master_key", master_key)
setattr(litellm.proxy.proxy_server, "llm_router", MagicMock())
await litellm.proxy.proxy_server.prisma_client.connect()
@ -151,6 +152,7 @@ async def test_create_new_user_in_organization(prisma_client, user_role):
organization_alias=f"new-org-{uuid.uuid4()}",
),
user_api_key_dict=UserAPIKeyAuth(
user_id=created_user_id,
user_role=LitellmUserRoles.PROXY_ADMIN,
),
)
@ -203,6 +205,7 @@ async def test_org_admin_create_team_permissions(prisma_client):
master_key = "sk-1234"
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
setattr(litellm.proxy.proxy_server, "master_key", master_key)
setattr(litellm.proxy.proxy_server, "llm_router", MagicMock())
await litellm.proxy.proxy_server.prisma_client.connect()
@ -274,6 +277,7 @@ async def test_org_admin_create_user_permissions(prisma_client):
master_key = "sk-1234"
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
setattr(litellm.proxy.proxy_server, "master_key", master_key)
setattr(litellm.proxy.proxy_server, "llm_router", MagicMock())
await litellm.proxy.proxy_server.prisma_client.connect()
@ -345,6 +349,7 @@ async def test_org_admin_create_user_team_wrong_org_permissions(prisma_client):
master_key = "sk-1234"
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
setattr(litellm.proxy.proxy_server, "master_key", master_key)
setattr(litellm.proxy.proxy_server, "llm_router", MagicMock())
await litellm.proxy.proxy_server.prisma_client.connect()
created_user_id = f"new-user-{uuid.uuid4()}"

View file

@ -435,11 +435,9 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
{/* Special case for Vertex Credentials help text */}
{field.key === "vertex_credentials" && (
<Row>
<Col span={10}></Col>
<Col span={10}>
<Col>
<Text className="mb-3 mt-1">
Give litellm a gcp service account(.json file), so it
can make the relevant calls
Give a gcp service account(.json file)
</Text>
</Col>
</Row>

View file

@ -86,8 +86,8 @@ export const SessionView: React.FC<SessionViewProps> = ({ sessionId, logs, onBac
data={logs}
renderSubComponent={RequestViewer}
getRowCanExpand={() => true}
expandedRequestId={expandedRequestId}
onRowExpand={setExpandedRequestId}
loadingMessage="Loading logs..."
noDataMessage="No logs found"
/>
</div>
</div>

View file

@ -691,8 +691,6 @@ export default function SpendLogsTable({
data={filteredData}
renderSubComponent={RequestViewer}
getRowCanExpand={() => true}
onRowExpand={handleRowExpand}
expandedRequestId={expandedRequestId}
/>
</div>
</>

View file

@ -1,4 +1,4 @@
import { Fragment, useEffect } from "react";
import { Fragment } from "react";
import {
ColumnDef,
flexRender,
@ -23,21 +23,16 @@ interface DataTableProps<TData, TValue> {
renderSubComponent: (props: { row: Row<TData> }) => React.ReactElement;
getRowCanExpand: (row: Row<TData>) => boolean;
isLoading?: boolean;
expandedRequestId?: string | null;
onRowExpand?: (requestId: string | null) => void;
setSelectedKeyIdInfoView?: (keyId: string | null) => void;
loadingMessage?: string;
noDataMessage?: string;
}
export function DataTable<TData extends { request_id: string }, TValue>({
export function DataTable<TData, TValue>({
data = [],
columns,
getRowCanExpand,
renderSubComponent,
isLoading = false,
expandedRequestId,
onRowExpand,
loadingMessage = "🚅 Loading logs...",
noDataMessage = "No logs found",
}: DataTableProps<TData, TValue>) {
@ -47,53 +42,10 @@ export function DataTable<TData extends { request_id: string }, TValue>({
getRowCanExpand,
getCoreRowModel: getCoreRowModel(),
getExpandedRowModel: getExpandedRowModel(),
state: {
expanded: expandedRequestId
? data.reduce((acc, row, index) => {
if (row.request_id === expandedRequestId) {
acc[index] = true;
}
return acc;
}, {} as Record<string, boolean>)
: {},
},
onExpandedChange: (updater) => {
if (!onRowExpand) return;
// Get current expanded state
const currentExpanded = expandedRequestId
? data.reduce((acc, row, index) => {
if (row.request_id === expandedRequestId) {
acc[index] = true;
}
return acc;
}, {} as Record<string, boolean>)
: {};
// Calculate new expanded state
const newExpanded = typeof updater === 'function'
? updater(currentExpanded)
: updater;
// If empty, it means we're closing the expanded row
if (Object.keys(newExpanded).length === 0) {
onRowExpand(null);
return;
}
// Find the request_id of the expanded row
const expandedIndex = Object.keys(newExpanded)[0];
const expandedRow = expandedIndex !== undefined ? data[parseInt(expandedIndex)] : null;
// Call the onRowExpand callback with the request_id
onRowExpand(expandedRow ? expandedRow.request_id : null);
},
});
// No need for the useEffect here as we're handling everything in onExpandedChange
return (
<div className="rounded-lg custom-border">
<div className="rounded-lg custom-border table-wrapper">
<Table className="[&_td]:py-0.5 [&_th]:py-1">
<TableHead>
{table.getHeaderGroups().map((headerGroup) => (
@ -160,4 +112,4 @@ export function DataTable<TData extends { request_id: string }, TValue>({
</Table>
</div>
);
}
}