mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
42a91bae6b
commit
880c2a736b
9 changed files with 52 additions and 71 deletions
File diff suppressed because one or more lines are too long
|
|
@ -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 ##
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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()}"
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -691,8 +691,6 @@ export default function SpendLogsTable({
|
|||
data={filteredData}
|
||||
renderSubComponent={RequestViewer}
|
||||
getRowCanExpand={() => true}
|
||||
onRowExpand={handleRowExpand}
|
||||
expandedRequestId={expandedRequestId}
|
||||
/>
|
||||
</div>
|
||||
</>
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue