diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html deleted file mode 100644 index 8272d90eace..00000000000 --- a/litellm/proxy/_experimental/out/onboarding.html +++ /dev/null @@ -1 +0,0 @@ -LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 3698ca7053b..ce2639b3b35 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 ## diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 42dd903e796..06abbf9a229 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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) diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index 799b5dbec77..50a75b18474 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -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), diff --git a/tests/proxy_admin_ui_tests/test_role_based_access.py b/tests/proxy_admin_ui_tests/test_role_based_access.py index ff73143bf44..a1a3a406da2 100644 --- a/tests/proxy_admin_ui_tests/test_role_based_access.py +++ b/tests/proxy_admin_ui_tests/test_role_based_access.py @@ -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()}" diff --git a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx index 8b00795ac6a..7146ae2b995 100644 --- a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx +++ b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx @@ -435,11 +435,9 @@ const ProviderSpecificFields: React.FC = ({ {/* Special case for Vertex Credentials help text */} {field.key === "vertex_credentials" && ( - - + - Give litellm a gcp service account(.json file), so it - can make the relevant calls + Give a gcp service account(.json file) diff --git a/ui/litellm-dashboard/src/components/view_logs/SessionView.tsx b/ui/litellm-dashboard/src/components/view_logs/SessionView.tsx index 873d2f5791a..062c9f643a9 100644 --- a/ui/litellm-dashboard/src/components/view_logs/SessionView.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/SessionView.tsx @@ -86,8 +86,8 @@ export const SessionView: React.FC = ({ sessionId, logs, onBac data={logs} renderSubComponent={RequestViewer} getRowCanExpand={() => true} - expandedRequestId={expandedRequestId} - onRowExpand={setExpandedRequestId} + loadingMessage="Loading logs..." + noDataMessage="No logs found" /> diff --git a/ui/litellm-dashboard/src/components/view_logs/index.tsx b/ui/litellm-dashboard/src/components/view_logs/index.tsx index 111d373137a..8a4e280cbfe 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.tsx @@ -691,8 +691,6 @@ export default function SpendLogsTable({ data={filteredData} renderSubComponent={RequestViewer} getRowCanExpand={() => true} - onRowExpand={handleRowExpand} - expandedRequestId={expandedRequestId} /> diff --git a/ui/litellm-dashboard/src/components/view_logs/table.tsx b/ui/litellm-dashboard/src/components/view_logs/table.tsx index 51131f0dd7d..43e35486707 100644 --- a/ui/litellm-dashboard/src/components/view_logs/table.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/table.tsx @@ -1,4 +1,4 @@ -import { Fragment, useEffect } from "react"; +import { Fragment } from "react"; import { ColumnDef, flexRender, @@ -23,21 +23,16 @@ interface DataTableProps { renderSubComponent: (props: { row: Row }) => React.ReactElement; getRowCanExpand: (row: Row) => boolean; isLoading?: boolean; - expandedRequestId?: string | null; - onRowExpand?: (requestId: string | null) => void; - setSelectedKeyIdInfoView?: (keyId: string | null) => void; loadingMessage?: string; noDataMessage?: string; } -export function DataTable({ +export function DataTable({ data = [], columns, getRowCanExpand, renderSubComponent, isLoading = false, - expandedRequestId, - onRowExpand, loadingMessage = "🚅 Loading logs...", noDataMessage = "No logs found", }: DataTableProps) { @@ -47,53 +42,10 @@ export function DataTable({ 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) - : {}, - }, - 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) - : {}; - - // 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 ( -
+
{table.getHeaderGroups().map((headerGroup) => ( @@ -160,4 +112,4 @@ export function DataTable({
); -} +} \ No newline at end of file