mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
chore(typing): clear 2.7k basedpyright Any errors across 15 hotspot files
Replace Any-typed seams with real types in the files carrying the highest reportAny/reportExplicitAny density: typed Prisma read helpers in the MCP db layer and verification token repository, TypedDicts for OAuth credential payloads and aggregated spend rows, a DailySpendRecord protocol for the daily activity endpoints, and concrete request/response types in the volcengine, openai evals, azure batches, azure_ai count_tokens, and ocr transformation modules. Modernize touched annotations to PEP 604/585 forms. No casts, no type: ignore, no noqa, no new Any annotations, no behavior changes. Whole-tree basedpyright: reportAny 27,005 -> 24,427, reportExplicitAny 7,439 -> 7,280, no rule increased anywhere. Budgets ratcheted: basedpyright -2,869, ruff-strict -1,505, type-discipline -167.
This commit is contained in:
parent
24123269cc
commit
fbfb63c948
24 changed files with 2292 additions and 1576 deletions
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 37484
|
||||
"limit": 34906
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2704
|
||||
"limit": 2701
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 330
|
||||
|
|
@ -12,7 +12,7 @@
|
|||
"limit": 516
|
||||
},
|
||||
"reportCallIssue": {
|
||||
"limit": 124
|
||||
"limit": 123
|
||||
},
|
||||
"reportConstantRedefinition": {
|
||||
"limit": 59
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 42
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 10389
|
||||
"limit": 10230
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 11
|
||||
|
|
@ -54,10 +54,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5900
|
||||
"limit": 5893
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15903
|
||||
"limit": 15886
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 41
|
||||
|
|
@ -99,31 +99,31 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 45894
|
||||
"limit": 45870
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 113
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 40539
|
||||
"limit": 40525
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 20403
|
||||
"limit": 20384
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 32141
|
||||
"limit": 32099
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 177
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 1025
|
||||
"limit": 1023
|
||||
},
|
||||
"reportUnnecessaryContains": {
|
||||
"limit": 7
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 1209
|
||||
"limit": 1206
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 165
|
||||
|
|
|
|||
|
|
@ -11,7 +11,8 @@ Endpoints for /project operations
|
|||
#### PROJECT MANAGEMENT ####
|
||||
|
||||
import json
|
||||
from typing import List, Optional, Union
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
|
||||
|
|
@ -25,15 +26,24 @@ from litellm.proxy.management_helpers.utils import (
|
|||
)
|
||||
from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_models
|
||||
from prisma.actions import LiteLLM_TeamTableActions
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _team_table(prisma_client: PrismaClient) -> "LiteLLM_TeamTableActions[prisma_models.LiteLLM_TeamTable]":
|
||||
team_table: LiteLLM_TeamTableActions[prisma_models.LiteLLM_TeamTable] = prisma_client.db.litellm_teamtable
|
||||
return team_table
|
||||
|
||||
|
||||
async def _check_user_permission_for_project(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
team_id: Optional[str],
|
||||
team_id: str | None,
|
||||
prisma_client: PrismaClient,
|
||||
require_admin: bool = False,
|
||||
team_object: Optional[LiteLLM_TeamTable] = None,
|
||||
team_object: LiteLLM_TeamTable | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if user has permission to manage a project.
|
||||
|
|
@ -57,9 +67,7 @@ async def _check_user_permission_for_project(
|
|||
|
||||
team = team_object
|
||||
if team is None:
|
||||
team = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": team_id}
|
||||
)
|
||||
team = await _team_table(prisma_client).find_unique(where={"team_id": team_id})
|
||||
|
||||
if team and team.admins:
|
||||
return user_api_key_dict.user_id in team.admins
|
||||
|
|
@ -70,9 +78,9 @@ async def _check_user_permission_for_project(
|
|||
async def _validate_team_exists(
|
||||
team_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
):
|
||||
) -> "prisma_models.LiteLLM_TeamTable":
|
||||
"""Validate that a team exists. Returns the team row."""
|
||||
team = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
team = await _team_table(prisma_client).find_unique(
|
||||
where={"team_id": team_id},
|
||||
)
|
||||
|
||||
|
|
@ -89,7 +97,7 @@ async def _validate_team_exists(
|
|||
|
||||
def _check_team_project_limits(
|
||||
team_object: LiteLLM_TeamTable,
|
||||
data: Union[NewProjectRequest, UpdateProjectRequest],
|
||||
data: NewProjectRequest | UpdateProjectRequest,
|
||||
) -> None:
|
||||
"""
|
||||
Check that project limits respect its parent Team's limits.
|
||||
|
|
@ -108,16 +116,12 @@ def _check_team_project_limits(
|
|||
if data.max_budget is not None and data.max_budget < 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"max_budget cannot be negative. Received: {data.max_budget}"
|
||||
},
|
||||
detail={"error": f"max_budget cannot be negative. Received: {data.max_budget}"},
|
||||
)
|
||||
if data.soft_budget is not None and data.soft_budget < 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"
|
||||
},
|
||||
detail={"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"},
|
||||
)
|
||||
|
||||
# --- soft_budget < max_budget ---
|
||||
|
|
@ -131,7 +135,7 @@ def _check_team_project_limits(
|
|||
)
|
||||
|
||||
# --- Validate project models are a subset of team models ---
|
||||
project_models = getattr(data, "models", None)
|
||||
project_models = data.models
|
||||
team_models = team_object.models or []
|
||||
if project_models and len(team_models) > 0:
|
||||
# If team has 'all-proxy-models', skip validation as it allows all models
|
||||
|
|
@ -148,11 +152,7 @@ def _check_team_project_limits(
|
|||
# --- Validate project max_budget <= team max_budget ---
|
||||
# Team stores budget fields directly (max_budget, tpm_limit, rpm_limit)
|
||||
# unlike Project which uses a separate LiteLLM_BudgetTable relation
|
||||
if (
|
||||
data.max_budget is not None
|
||||
and team_object.max_budget is not None
|
||||
and data.max_budget > team_object.max_budget
|
||||
):
|
||||
if data.max_budget is not None and team_object.max_budget is not None and data.max_budget > team_object.max_budget:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
|
|
@ -161,11 +161,7 @@ def _check_team_project_limits(
|
|||
)
|
||||
|
||||
# --- Validate project tpm_limit <= team tpm_limit ---
|
||||
if (
|
||||
data.tpm_limit is not None
|
||||
and team_object.tpm_limit is not None
|
||||
and data.tpm_limit > team_object.tpm_limit
|
||||
):
|
||||
if data.tpm_limit is not None and team_object.tpm_limit is not None and data.tpm_limit > team_object.tpm_limit:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
|
|
@ -174,11 +170,7 @@ def _check_team_project_limits(
|
|||
)
|
||||
|
||||
# --- Validate project rpm_limit <= team rpm_limit ---
|
||||
if (
|
||||
data.rpm_limit is not None
|
||||
and team_object.rpm_limit is not None
|
||||
and data.rpm_limit > team_object.rpm_limit
|
||||
):
|
||||
if data.rpm_limit is not None and team_object.rpm_limit is not None and data.rpm_limit > team_object.rpm_limit:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
|
|
@ -189,19 +181,19 @@ def _check_team_project_limits(
|
|||
|
||||
async def _create_budget_for_project(
|
||||
data: NewProjectRequest,
|
||||
user_id: Optional[str],
|
||||
user_id: str | None,
|
||||
litellm_proxy_admin_name: str,
|
||||
prisma_client: PrismaClient,
|
||||
) -> str:
|
||||
"""Create a budget for the project and return budget_id."""
|
||||
budget_params = LiteLLM_BudgetTable.model_fields.keys()
|
||||
_json_data = data.json(exclude_none=True)
|
||||
_json_data: Mapping[str, object] = data.json(exclude_none=True)
|
||||
_budget_data = {k: v for k, v in _json_data.items() if k in budget_params}
|
||||
budget_row = LiteLLM_BudgetTable(**_budget_data)
|
||||
budget_row = LiteLLM_BudgetTable.model_validate(_budget_data)
|
||||
|
||||
new_budget = prisma_client.jsonify_object(budget_row.json(exclude_none=True))
|
||||
|
||||
_budget = await prisma_client.db.litellm_budgettable.create(
|
||||
_budget: prisma_models.LiteLLM_BudgetTable = await prisma_client.db.litellm_budgettable.create(
|
||||
data={
|
||||
**new_budget,
|
||||
"created_by": user_id or litellm_proxy_admin_name,
|
||||
|
|
@ -214,8 +206,8 @@ async def _create_budget_for_project(
|
|||
|
||||
async def _set_project_object_permission(
|
||||
data: NewProjectRequest,
|
||||
prisma_client: Optional[PrismaClient],
|
||||
) -> Optional[str]:
|
||||
prisma_client: PrismaClient | None,
|
||||
) -> str | None:
|
||||
"""
|
||||
Creates the LiteLLM_ObjectPermissionTable record for the project.
|
||||
Returns the object_permission_id if created, otherwise None.
|
||||
|
|
@ -224,7 +216,7 @@ async def _set_project_object_permission(
|
|||
return None
|
||||
|
||||
if data.object_permission is not None:
|
||||
created_object_permission = (
|
||||
created_object_permission: prisma_models.LiteLLM_ObjectPermissionTable = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=data.object_permission.model_dump(exclude_none=True),
|
||||
)
|
||||
|
|
@ -344,8 +336,7 @@ async def new_project(
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Only premium users can add tags to projects. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
"error": "Only premium users can add tags to projects. " + CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -353,8 +344,7 @@ async def new_project(
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Project management is an enterprise feature. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
"error": "Project management is an enterprise feature. " + CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -375,13 +365,11 @@ async def new_project(
|
|||
)
|
||||
|
||||
# Validate team exists and get team object with budget
|
||||
team_object = await _validate_team_exists(
|
||||
team_id=data.team_id, prisma_client=prisma_client
|
||||
)
|
||||
team_object = await _validate_team_exists(team_id=data.team_id, prisma_client=prisma_client)
|
||||
|
||||
# Validate project limits against team limits
|
||||
_check_team_project_limits(
|
||||
team_object=LiteLLM_TeamTable(**team_object.model_dump()),
|
||||
team_object=LiteLLM_TeamTable.model_validate(team_object.model_dump()),
|
||||
data=data,
|
||||
)
|
||||
|
||||
|
|
@ -391,7 +379,7 @@ async def new_project(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
team_id=data.team_id,
|
||||
prisma_client=prisma_client,
|
||||
team_object=LiteLLM_TeamTable(**team_object.model_dump()),
|
||||
team_object=LiteLLM_TeamTable.model_validate(team_object.model_dump()),
|
||||
)
|
||||
|
||||
if not has_permission:
|
||||
|
|
@ -449,17 +437,13 @@ async def new_project(
|
|||
value=getattr(data, field),
|
||||
)
|
||||
|
||||
new_project_row = prisma_client.jsonify_object(
|
||||
project_row.json(exclude_none=True)
|
||||
)
|
||||
new_project_row = prisma_client.jsonify_object(project_row.json(exclude_none=True))
|
||||
|
||||
# Remove budget fields (following organization_endpoints.py pattern)
|
||||
new_project_row = _remove_budget_fields_from_project_data(new_project_row)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"new_project_row: {json.dumps(new_project_row, indent=2)}"
|
||||
)
|
||||
response = await prisma_client.db.litellm_projecttable.create(
|
||||
verbose_proxy_logger.info(f"new_project_row: {json.dumps(new_project_row, indent=2)}")
|
||||
response: prisma_models.LiteLLM_ProjectTable = await prisma_client.db.litellm_projecttable.create(
|
||||
data={
|
||||
**new_project_row, # type: ignore
|
||||
},
|
||||
|
|
@ -469,9 +453,7 @@ async def new_project(
|
|||
return response
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.project_endpoints.new_project(): Exception occured - {}".format(
|
||||
str(e)
|
||||
)
|
||||
"litellm.proxy.management_endpoints.project_endpoints.new_project(): Exception occured - {}".format(str(e))
|
||||
)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
|
@ -539,8 +521,7 @@ async def update_project(
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Only premium users can add tags to projects. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
"error": "Only premium users can add tags to projects. " + CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -548,8 +529,7 @@ async def update_project(
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Project management is an enterprise feature. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
"error": "Project management is an enterprise feature. " + CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -576,9 +556,9 @@ async def update_project(
|
|||
)
|
||||
|
||||
# Fetch existing project
|
||||
existing_project = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
where={"project_id": data.project_id}
|
||||
)
|
||||
existing_project: (
|
||||
prisma_models.LiteLLM_ProjectTable | None
|
||||
) = await prisma_client.db.litellm_projecttable.find_unique(where={"project_id": data.project_id})
|
||||
|
||||
if existing_project is None:
|
||||
raise ProxyException(
|
||||
|
|
@ -595,9 +575,7 @@ async def update_project(
|
|||
target_team_id = data.team_id or existing_project.team_id
|
||||
target_team_obj = None
|
||||
if target_team_id is not None:
|
||||
target_team_obj = await _validate_team_exists(
|
||||
team_id=target_team_id, prisma_client=prisma_client
|
||||
)
|
||||
target_team_obj = await _validate_team_exists(team_id=target_team_id, prisma_client=prisma_client)
|
||||
|
||||
has_permission = await _check_user_permission_for_project(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -620,32 +598,26 @@ async def update_project(
|
|||
team_id=data.team_id,
|
||||
prisma_client=prisma_client,
|
||||
team_object=(
|
||||
LiteLLM_TeamTable(**target_team_obj.model_dump())
|
||||
if target_team_obj
|
||||
else None
|
||||
LiteLLM_TeamTable.model_validate(target_team_obj.model_dump()) if target_team_obj else None
|
||||
),
|
||||
)
|
||||
if not can_assign_to_target:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Cannot reassign project to a team you are not an admin of"
|
||||
},
|
||||
detail={"error": "Cannot reassign project to a team you are not an admin of"},
|
||||
)
|
||||
|
||||
# Validate project limits against team limits
|
||||
if target_team_obj is not None:
|
||||
_check_team_project_limits(
|
||||
team_object=LiteLLM_TeamTable(**target_team_obj.model_dump()),
|
||||
team_object=LiteLLM_TeamTable.model_validate(target_team_obj.model_dump()),
|
||||
data=data,
|
||||
)
|
||||
|
||||
# Prepare update data
|
||||
update_data = data.json(exclude_none=True, exclude={"project_id"})
|
||||
update_data = prisma_client.jsonify_object(update_data)
|
||||
update_data["updated_by"] = (
|
||||
user_api_key_dict.user_id or litellm_proxy_admin_name
|
||||
)
|
||||
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
|
||||
|
||||
# Handle budget updates
|
||||
budget_fields = LiteLLM_BudgetTable.model_fields.keys()
|
||||
|
|
@ -671,21 +643,17 @@ async def update_project(
|
|||
if existing_project.object_permission_id:
|
||||
# Update existing permission
|
||||
await prisma_client.db.litellm_objectpermissiontable.update(
|
||||
where={
|
||||
"object_permission_id": existing_project.object_permission_id
|
||||
},
|
||||
where={"object_permission_id": existing_project.object_permission_id},
|
||||
data=object_permission_data,
|
||||
)
|
||||
else:
|
||||
# Create new permission
|
||||
created_permission = (
|
||||
created_permission: prisma_models.LiteLLM_ObjectPermissionTable = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=object_permission_data,
|
||||
)
|
||||
)
|
||||
update_data["object_permission_id"] = (
|
||||
created_permission.object_permission_id
|
||||
)
|
||||
update_data["object_permission_id"] = created_permission.object_permission_id
|
||||
|
||||
# Handle metadata fields
|
||||
for field in LiteLLM_ManagementEndpoint_MetadataFields:
|
||||
|
|
@ -698,7 +666,7 @@ async def update_project(
|
|||
update_data = _remove_budget_fields_from_project_data(update_data)
|
||||
|
||||
# Update project
|
||||
updated_project = await prisma_client.db.litellm_projecttable.update(
|
||||
updated_project: prisma_models.LiteLLM_ProjectTable | None = await prisma_client.db.litellm_projecttable.update(
|
||||
where={"project_id": data.project_id},
|
||||
data=update_data,
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
|
|
@ -718,7 +686,7 @@ async def update_project(
|
|||
"/project/delete",
|
||||
tags=["project management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[LiteLLM_ProjectTable],
|
||||
response_model=list[LiteLLM_ProjectTable],
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def delete_project(
|
||||
|
|
@ -749,8 +717,7 @@ async def delete_project(
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Project management is an enterprise feature. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
"error": "Project management is an enterprise feature. " + CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -778,9 +745,7 @@ async def delete_project(
|
|||
|
||||
for project_id in data.project_ids:
|
||||
# Check if project exists
|
||||
existing_project = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
where={"project_id": project_id}
|
||||
)
|
||||
existing_project = await prisma_client.db.litellm_projecttable.find_unique(where={"project_id": project_id})
|
||||
|
||||
if existing_project is None:
|
||||
raise ProxyException(
|
||||
|
|
@ -791,11 +756,9 @@ async def delete_project(
|
|||
)
|
||||
|
||||
# Check if there are any keys associated with this project
|
||||
associated_keys = (
|
||||
await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={"project_id": project_id}
|
||||
)
|
||||
)
|
||||
associated_keys: Sequence[
|
||||
prisma_models.LiteLLM_VerificationToken
|
||||
] = await prisma_client.db.litellm_verificationtoken.find_many(where={"project_id": project_id})
|
||||
|
||||
if len(associated_keys) > 0:
|
||||
raise ProxyException(
|
||||
|
|
@ -806,9 +769,9 @@ async def delete_project(
|
|||
)
|
||||
|
||||
# Delete the project
|
||||
deleted_project = await prisma_client.db.litellm_projecttable.delete(
|
||||
where={"project_id": project_id}
|
||||
)
|
||||
deleted_project: (
|
||||
prisma_models.LiteLLM_ProjectTable | None
|
||||
) = await prisma_client.db.litellm_projecttable.delete(where={"project_id": project_id})
|
||||
|
||||
deleted_projects.append(deleted_project)
|
||||
|
||||
|
|
@ -854,7 +817,7 @@ async def project_info(
|
|||
)
|
||||
|
||||
# Fetch project
|
||||
project = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
project: prisma_models.LiteLLM_ProjectTable | None = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
where={"project_id": project_id},
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
)
|
||||
|
|
@ -872,17 +835,11 @@ async def project_info(
|
|||
is_team_member = False
|
||||
|
||||
if project.team_id and user_api_key_dict.user_id:
|
||||
team = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": project.team_id}
|
||||
)
|
||||
team = await _team_table(prisma_client).find_unique(where={"team_id": project.team_id})
|
||||
if team:
|
||||
caller_user_id = user_api_key_dict.user_id
|
||||
for m in team.members_with_roles or []:
|
||||
m_user_id = (
|
||||
m.get("user_id")
|
||||
if isinstance(m, dict)
|
||||
else getattr(m, "user_id", None)
|
||||
)
|
||||
m_user_id = m.get("user_id") if isinstance(m, dict) else getattr(m, "user_id", None)
|
||||
if m_user_id == caller_user_id:
|
||||
is_team_member = True
|
||||
break
|
||||
|
|
@ -896,9 +853,7 @@ async def project_info(
|
|||
return project
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.project_endpoints.project_info(): Exception occured - {}".format(
|
||||
str(e)
|
||||
)
|
||||
"litellm.proxy.management_endpoints.project_endpoints.project_info(): Exception occured - {}".format(str(e))
|
||||
)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
|
@ -907,7 +862,7 @@ async def project_info(
|
|||
"/project/list",
|
||||
tags=["project management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[LiteLLM_ProjectTable],
|
||||
response_model=list[LiteLLM_ProjectTable],
|
||||
)
|
||||
async def list_projects(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
|
|
@ -932,21 +887,19 @@ async def list_projects(
|
|||
|
||||
# If proxy admin, get all projects
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
projects = await prisma_client.db.litellm_projecttable.find_many(
|
||||
projects: Sequence[
|
||||
prisma_models.LiteLLM_ProjectTable
|
||||
] = await prisma_client.db.litellm_projecttable.find_many(
|
||||
include={"litellm_budget_table": True, "object_permission": True}
|
||||
)
|
||||
else:
|
||||
# Look up the user's team memberships via the reverse-index on
|
||||
# LiteLLM_UserTable.teams (maintained by team_member_add alongside
|
||||
# members_with_roles). This avoids a full scan of all team rows.
|
||||
user_record = await prisma_client.db.litellm_usertable.find_unique(
|
||||
user_record: prisma_models.LiteLLM_UserTable | None = await prisma_client.db.litellm_usertable.find_unique(
|
||||
where={"user_id": user_api_key_dict.user_id},
|
||||
)
|
||||
user_team_ids = (
|
||||
user_record.teams
|
||||
if user_record is not None and user_record.teams
|
||||
else []
|
||||
)
|
||||
user_team_ids: Sequence[str] = user_record.teams if user_record is not None and user_record.teams else []
|
||||
|
||||
projects = await prisma_client.db.litellm_projecttable.find_many(
|
||||
where={"team_id": {"in": user_team_ids}},
|
||||
|
|
|
|||
|
|
@ -2,7 +2,8 @@
|
|||
Azure Batches API Handler
|
||||
"""
|
||||
|
||||
from typing import Any, Coroutine, Optional, Union, cast
|
||||
from collections.abc import Coroutine
|
||||
from typing import cast
|
||||
|
||||
import httpx
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
|
|
@ -33,32 +34,30 @@ class AzureBatchesAPI(BaseAzureLLM):
|
|||
async def acreate_batch(
|
||||
self,
|
||||
create_batch_data: CreateBatchRequest,
|
||||
azure_client: Union[AsyncAzureOpenAI, AsyncOpenAI],
|
||||
azure_client: AsyncAzureOpenAI | AsyncOpenAI,
|
||||
) -> LiteLLMBatch:
|
||||
response = await azure_client.batches.create(**create_batch_data) # type: ignore[arg-type]
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
def create_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
create_batch_data: CreateBatchRequest,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
max_retries: Optional[int],
|
||||
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]:
|
||||
azure_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = (
|
||||
self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
api_version: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
) -> LiteLLMBatch | Coroutine[object, object, LiteLLMBatch]:
|
||||
azure_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
if azure_client is None:
|
||||
raise ValueError(
|
||||
|
|
@ -73,38 +72,36 @@ class AzureBatchesAPI(BaseAzureLLM):
|
|||
return self.acreate_batch( # type: ignore
|
||||
create_batch_data=create_batch_data, azure_client=azure_client
|
||||
)
|
||||
response = cast(Union[AzureOpenAI, OpenAI], azure_client).batches.create(**create_batch_data) # type: ignore[arg-type]
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
response = cast(AzureOpenAI | OpenAI, azure_client).batches.create(**create_batch_data) # type: ignore[arg-type]
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
async def aretrieve_batch(
|
||||
self,
|
||||
retrieve_batch_data: RetrieveBatchRequest,
|
||||
client: Union[AsyncAzureOpenAI, AsyncOpenAI],
|
||||
client: AsyncAzureOpenAI | AsyncOpenAI,
|
||||
) -> LiteLLMBatch:
|
||||
response = await client.batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type]
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
def retrieve_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
retrieve_batch_data: RetrieveBatchRequest,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
max_retries: Optional[int],
|
||||
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
api_version: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
):
|
||||
azure_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = (
|
||||
self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
azure_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
if azure_client is None:
|
||||
raise ValueError(
|
||||
|
|
@ -119,38 +116,36 @@ class AzureBatchesAPI(BaseAzureLLM):
|
|||
return self.aretrieve_batch( # type: ignore
|
||||
retrieve_batch_data=retrieve_batch_data, client=azure_client
|
||||
)
|
||||
response = cast(Union[AzureOpenAI, OpenAI], azure_client).batches.retrieve(**retrieve_batch_data)
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
response = cast(AzureOpenAI | OpenAI, azure_client).batches.retrieve(**retrieve_batch_data)
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
async def acancel_batch(
|
||||
self,
|
||||
cancel_batch_data: CancelBatchRequest,
|
||||
client: Union[AsyncAzureOpenAI, AsyncOpenAI],
|
||||
client: AsyncAzureOpenAI | AsyncOpenAI,
|
||||
) -> LiteLLMBatch:
|
||||
response = await client.batches.cancel(**cancel_batch_data)
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
def cancel_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
cancel_batch_data: CancelBatchRequest,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
max_retries: Optional[int],
|
||||
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
api_version: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
):
|
||||
azure_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = (
|
||||
self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
azure_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
if azure_client is None:
|
||||
raise ValueError(
|
||||
|
|
@ -172,13 +167,13 @@ class AzureBatchesAPI(BaseAzureLLM):
|
|||
"Azure client is not an instance of AzureOpenAI or OpenAI. Make sure you passed a sync client."
|
||||
)
|
||||
response = azure_client.batches.cancel(**cancel_batch_data)
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
async def alist_batches(
|
||||
self,
|
||||
client: Union[AsyncAzureOpenAI, AsyncOpenAI],
|
||||
after: Optional[str] = None,
|
||||
limit: Optional[int] = None,
|
||||
client: AsyncAzureOpenAI | AsyncOpenAI,
|
||||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
):
|
||||
response = await client.batches.list(after=after, limit=limit) # type: ignore
|
||||
return response
|
||||
|
|
@ -186,25 +181,23 @@ class AzureBatchesAPI(BaseAzureLLM):
|
|||
def list_batches(
|
||||
self,
|
||||
_is_async: bool,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
max_retries: Optional[int],
|
||||
after: Optional[str] = None,
|
||||
limit: Optional[int] = None,
|
||||
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
api_version: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
):
|
||||
azure_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = (
|
||||
self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
azure_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
if azure_client is None:
|
||||
raise ValueError(
|
||||
|
|
|
|||
|
|
@ -4,8 +4,6 @@ Azure AI Anthropic CountTokens API transformation logic.
|
|||
Extends the base Anthropic CountTokens transformation with Azure authentication.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION
|
||||
from litellm.llms.anthropic.count_tokens.transformation import (
|
||||
AnthropicCountTokensConfig,
|
||||
|
|
@ -25,8 +23,8 @@ class AzureAIAnthropicCountTokensConfig(AnthropicCountTokensConfig):
|
|||
def get_required_headers(
|
||||
self,
|
||||
api_key: str,
|
||||
litellm_params: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, str]:
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
Get the required headers for the Azure AI Anthropic CountTokens API.
|
||||
|
||||
|
|
@ -53,7 +51,7 @@ class AzureAIAnthropicCountTokensConfig(AnthropicCountTokensConfig):
|
|||
if "api_key" not in litellm_params:
|
||||
litellm_params["api_key"] = api_key
|
||||
|
||||
litellm_params_obj = GenericLiteLLMParams(**litellm_params)
|
||||
litellm_params_obj = GenericLiteLLMParams.model_validate(litellm_params)
|
||||
|
||||
# Get Azure auth headers (api-key or Authorization)
|
||||
azure_headers = BaseAzureLLM._base_validate_azure_environment(headers={}, litellm_params=litellm_params_obj)
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
OpenAI Evals API configuration and transformations
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
from collections.abc import Mapping
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -31,6 +31,10 @@ from litellm.types.router import GenericLiteLLMParams
|
|||
from litellm.types.utils import LlmProviders
|
||||
|
||||
|
||||
def _parsed_response_json(raw_response: httpx.Response) -> Mapping[str, object]:
|
||||
return raw_response.json()
|
||||
|
||||
|
||||
class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
||||
"""OpenAI-specific Evals API configuration"""
|
||||
|
||||
|
|
@ -38,7 +42,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
def custom_llm_provider(self) -> LlmProviders:
|
||||
return LlmProviders.OPENAI
|
||||
|
||||
def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict:
|
||||
def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict:
|
||||
"""Add OpenAI-specific headers"""
|
||||
import litellm
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
|
@ -61,9 +65,9 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_base: str | None,
|
||||
endpoint: str,
|
||||
eval_id: Optional[str] = None,
|
||||
eval_id: str | None = None,
|
||||
) -> str:
|
||||
"""Get complete URL for OpenAI Evals API"""
|
||||
if api_base is None:
|
||||
|
|
@ -79,7 +83,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
create_request: CreateEvalRequest,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Dict:
|
||||
) -> dict:
|
||||
"""Transform create eval request for OpenAI"""
|
||||
verbose_logger.debug("Transforming create eval request: %s", create_request)
|
||||
|
||||
|
|
@ -94,17 +98,17 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Eval:
|
||||
"""Transform OpenAI response to Eval object"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming create eval response: %s", response_json)
|
||||
|
||||
return Eval(**response_json)
|
||||
return Eval.model_validate(response_json)
|
||||
|
||||
def transform_list_evals_request(
|
||||
self,
|
||||
list_params: ListEvalsParams,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform list evals request for OpenAI"""
|
||||
api_base = "https://api.openai.com"
|
||||
if litellm_params and litellm_params.api_base:
|
||||
|
|
@ -113,7 +117,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
url = self.get_complete_url(api_base=api_base, endpoint="evals")
|
||||
|
||||
# Build query parameters
|
||||
query_params: Dict[str, Any] = {}
|
||||
query_params: dict[str, object] = {}
|
||||
if "limit" in list_params and list_params["limit"]:
|
||||
query_params["limit"] = list_params["limit"]
|
||||
if "after" in list_params and list_params["after"]:
|
||||
|
|
@ -138,10 +142,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ListEvalsResponse:
|
||||
"""Transform OpenAI response to ListEvalsResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming list evals response: %s", response_json)
|
||||
|
||||
return ListEvalsResponse(**response_json)
|
||||
return ListEvalsResponse.model_validate(response_json)
|
||||
|
||||
def transform_get_eval_request(
|
||||
self,
|
||||
|
|
@ -149,7 +153,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform get eval request for OpenAI"""
|
||||
url = self.get_complete_url(api_base=api_base, endpoint="evals", eval_id=eval_id)
|
||||
|
||||
|
|
@ -163,10 +167,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Eval:
|
||||
"""Transform OpenAI response to Eval object"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming get eval response: %s", response_json)
|
||||
|
||||
return Eval(**response_json)
|
||||
return Eval.model_validate(response_json)
|
||||
|
||||
def transform_update_eval_request(
|
||||
self,
|
||||
|
|
@ -175,7 +179,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict, Dict]:
|
||||
) -> tuple[str, dict, dict]:
|
||||
"""Transform update eval request for OpenAI"""
|
||||
url = self.get_complete_url(api_base=api_base, endpoint="evals", eval_id=eval_id)
|
||||
|
||||
|
|
@ -192,10 +196,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Eval:
|
||||
"""Transform OpenAI response to Eval object"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming update eval response: %s", response_json)
|
||||
|
||||
return Eval(**response_json)
|
||||
return Eval.model_validate(response_json)
|
||||
|
||||
def transform_delete_eval_request(
|
||||
self,
|
||||
|
|
@ -203,7 +207,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform delete eval request for OpenAI"""
|
||||
url = self.get_complete_url(api_base=api_base, endpoint="evals", eval_id=eval_id)
|
||||
|
||||
|
|
@ -217,10 +221,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> DeleteEvalResponse:
|
||||
"""Transform OpenAI response to DeleteEvalResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming delete eval response: %s", response_json)
|
||||
|
||||
return DeleteEvalResponse(**response_json)
|
||||
return DeleteEvalResponse.model_validate(response_json)
|
||||
|
||||
def transform_cancel_eval_request(
|
||||
self,
|
||||
|
|
@ -228,12 +232,12 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict, Dict]:
|
||||
) -> tuple[str, dict, dict]:
|
||||
"""Transform cancel eval request for OpenAI"""
|
||||
url = f"{self.get_complete_url(api_base=api_base, endpoint='evals', eval_id=eval_id)}/cancel"
|
||||
|
||||
# Empty body for cancel request
|
||||
request_body: Dict[str, Any] = {}
|
||||
request_body: dict[str, object] = {}
|
||||
|
||||
verbose_logger.debug("Cancel eval request - URL: %s", url)
|
||||
|
||||
|
|
@ -245,10 +249,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> CancelEvalResponse:
|
||||
"""Transform OpenAI response to CancelEvalResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming cancel eval response: %s", response_json)
|
||||
|
||||
return CancelEvalResponse(**response_json)
|
||||
return CancelEvalResponse.model_validate(response_json)
|
||||
|
||||
# Run API Transformations
|
||||
def transform_create_run_request(
|
||||
|
|
@ -257,7 +261,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
create_request: CreateRunRequest,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform create run request for OpenAI"""
|
||||
api_base = "https://api.openai.com"
|
||||
if litellm_params and litellm_params.api_base:
|
||||
|
|
@ -279,10 +283,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Run:
|
||||
"""Transform OpenAI response to Run object"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming create run response: %s", response_json)
|
||||
|
||||
return Run(**response_json)
|
||||
return Run.model_validate(response_json)
|
||||
|
||||
def transform_list_runs_request(
|
||||
self,
|
||||
|
|
@ -290,7 +294,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
list_params: ListRunsParams,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform list runs request for OpenAI"""
|
||||
api_base = "https://api.openai.com"
|
||||
if litellm_params and litellm_params.api_base:
|
||||
|
|
@ -300,7 +304,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
url = f"{api_base}/v1/evals/{encoded_eval_id}/runs"
|
||||
|
||||
# Build query parameters
|
||||
query_params: Dict[str, Any] = {}
|
||||
query_params: dict[str, object] = {}
|
||||
if "limit" in list_params and list_params["limit"]:
|
||||
query_params["limit"] = list_params["limit"]
|
||||
if "after" in list_params and list_params["after"]:
|
||||
|
|
@ -323,10 +327,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ListRunsResponse:
|
||||
"""Transform OpenAI response to ListRunsResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming list runs response: %s", response_json)
|
||||
|
||||
return ListRunsResponse(**response_json)
|
||||
return ListRunsResponse.model_validate(response_json)
|
||||
|
||||
def transform_get_run_request(
|
||||
self,
|
||||
|
|
@ -335,7 +339,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform get run request for OpenAI"""
|
||||
encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id")
|
||||
encoded_run_id = encode_url_path_segment(run_id, field_name="run_id")
|
||||
|
|
@ -351,10 +355,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Run:
|
||||
"""Transform OpenAI response to Run object"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming get run response: %s", response_json)
|
||||
|
||||
return Run(**response_json)
|
||||
return Run.model_validate(response_json)
|
||||
|
||||
def transform_cancel_run_request(
|
||||
self,
|
||||
|
|
@ -363,14 +367,14 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict, Dict]:
|
||||
) -> tuple[str, dict, dict]:
|
||||
"""Transform cancel run request for OpenAI"""
|
||||
encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id")
|
||||
encoded_run_id = encode_url_path_segment(run_id, field_name="run_id")
|
||||
url = f"{api_base}/v1/evals/{encoded_eval_id}/runs/{encoded_run_id}/cancel"
|
||||
|
||||
# Empty body for cancel request
|
||||
request_body: Dict[str, Any] = {}
|
||||
request_body: dict[str, object] = {}
|
||||
|
||||
verbose_logger.debug("Cancel run request - URL: %s", url)
|
||||
|
||||
|
|
@ -382,10 +386,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> CancelRunResponse:
|
||||
"""Transform OpenAI response to CancelRunResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming cancel run response: %s", response_json)
|
||||
|
||||
return CancelRunResponse(**response_json)
|
||||
return CancelRunResponse.model_validate(response_json)
|
||||
|
||||
def transform_delete_run_request(
|
||||
self,
|
||||
|
|
@ -394,14 +398,14 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict, Dict]:
|
||||
) -> tuple[str, dict, dict]:
|
||||
"""Transform delete run request for OpenAI"""
|
||||
encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id")
|
||||
encoded_run_id = encode_url_path_segment(run_id, field_name="run_id")
|
||||
url = f"{api_base}/v1/evals/{encoded_eval_id}/runs/{encoded_run_id}"
|
||||
|
||||
# Empty body for delete request
|
||||
request_body: Dict[str, Any] = {}
|
||||
request_body: dict[str, object] = {}
|
||||
|
||||
verbose_logger.debug("Delete run request - URL: %s", url)
|
||||
|
||||
|
|
@ -413,7 +417,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> RunDeleteResponse:
|
||||
"""Transform OpenAI response to RunDeleteResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming delete run response: %s", response_json)
|
||||
|
||||
return RunDeleteResponse(**response_json)
|
||||
return RunDeleteResponse.model_validate(response_json)
|
||||
|
|
|
|||
|
|
@ -1,11 +1,9 @@
|
|||
from collections.abc import Callable, Mapping, Sequence
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Tuple,
|
||||
Protocol,
|
||||
Union,
|
||||
get_args,
|
||||
get_origin,
|
||||
|
|
@ -17,10 +15,10 @@ from pydantic import fields as pyd_fields
|
|||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
_safe_convert_created_field,
|
||||
)
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import (
|
||||
|
|
@ -47,8 +45,15 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class _EventModelClass(Protocol):
|
||||
@property
|
||||
def model_fields(self) -> Mapping[str, pyd_fields.FieldInfo]: ...
|
||||
|
||||
def model_validate(self, obj: Mapping[str, object]) -> ResponsesAPIStreamingResponse: ...
|
||||
|
||||
|
||||
class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
_SUPPORTED_OPTIONAL_PARAMS: List[str] = [
|
||||
_SUPPORTED_OPTIONAL_PARAMS: list[str] = [
|
||||
# Doc-listed knobs
|
||||
"instructions",
|
||||
"max_output_tokens",
|
||||
|
|
@ -89,9 +94,7 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
supported.remove("metadata")
|
||||
return supported
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> VolcEngineError:
|
||||
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> VolcEngineError:
|
||||
typed_headers: httpx.Headers = headers if isinstance(headers, httpx.Headers) else httpx.Headers(headers or {})
|
||||
return VolcEngineError(
|
||||
status_code=status_code,
|
||||
|
|
@ -99,14 +102,14 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
headers=typed_headers,
|
||||
)
|
||||
|
||||
def validate_environment(self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]) -> dict:
|
||||
def validate_environment(self, headers: dict, model: str, litellm_params: GenericLiteLLMParams | None) -> dict:
|
||||
"""
|
||||
Build auth headers for Volcengine Responses API.
|
||||
"""
|
||||
if litellm_params is None:
|
||||
litellm_params = GenericLiteLLMParams()
|
||||
elif isinstance(litellm_params, dict):
|
||||
litellm_params = GenericLiteLLMParams(**litellm_params)
|
||||
litellm_params = GenericLiteLLMParams.model_validate(litellm_params)
|
||||
|
||||
api_key = (
|
||||
litellm_params.api_key
|
||||
|
|
@ -122,7 +125,7 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_base: str | None,
|
||||
litellm_params: dict,
|
||||
) -> str:
|
||||
"""
|
||||
|
|
@ -149,7 +152,7 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
response_api_optional_params: ResponsesAPIOptionalRequestParams,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> Dict:
|
||||
) -> dict:
|
||||
"""
|
||||
Volcengine Responses API aligns with OpenAI parameters.
|
||||
Remove parameters not supported by the public docs.
|
||||
|
|
@ -173,11 +176,11 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
def transform_responses_api_request(
|
||||
self,
|
||||
model: str,
|
||||
input: Union[str, ResponseInputParam],
|
||||
response_api_optional_request_params: Dict,
|
||||
input: str | ResponseInputParam,
|
||||
response_api_optional_request_params: dict,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Dict:
|
||||
) -> dict:
|
||||
"""
|
||||
Volcengine rejects any undocumented fields (including extra_body). Fail fast
|
||||
with clear errors and re-filter with the documented whitelist before delegating
|
||||
|
|
@ -210,7 +213,7 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
def transform_streaming_response(
|
||||
self,
|
||||
model: str,
|
||||
parsed_chunk: dict,
|
||||
parsed_chunk: Mapping[str, object],
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ResponsesAPIStreamingResponse:
|
||||
"""
|
||||
|
|
@ -222,18 +225,19 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
if isinstance(chunk, dict):
|
||||
resp = chunk.get("response")
|
||||
if isinstance(resp, dict) and "output" not in resp:
|
||||
resp_items: Mapping[str, object] = resp
|
||||
patched_chunk = dict(chunk)
|
||||
patched_resp = dict(resp)
|
||||
patched_resp = dict(resp_items)
|
||||
patched_resp["output"] = []
|
||||
patched_chunk["response"] = patched_resp
|
||||
chunk = patched_chunk
|
||||
|
||||
event_type = str(chunk.get("type")) if isinstance(chunk, dict) else None
|
||||
event_pydantic_model = OpenAIResponsesAPIConfig.get_event_model_class(event_type=event_type)
|
||||
event_pydantic_model: _EventModelClass = OpenAIResponsesAPIConfig.get_event_model_class(event_type=event_type)
|
||||
|
||||
patched_chunk = self._fill_missing_fields(chunk, event_pydantic_model)
|
||||
|
||||
return event_pydantic_model(**patched_chunk)
|
||||
return event_pydantic_model.model_validate(patched_chunk)
|
||||
|
||||
def transform_response_api_response(
|
||||
self,
|
||||
|
|
@ -246,7 +250,7 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
original_response=raw_response.text,
|
||||
additional_args={"complete_input_dict": {}},
|
||||
)
|
||||
raw_response_json = raw_response.json()
|
||||
raw_response_json = self._parsed_response_body(raw_response)
|
||||
if "created_at" in raw_response_json:
|
||||
raw_response_json["created_at"] = _safe_convert_created_field(raw_response_json["created_at"])
|
||||
except Exception:
|
||||
|
|
@ -256,10 +260,11 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
processed_headers = process_response_headers(raw_response_headers)
|
||||
|
||||
try:
|
||||
response = ResponsesAPIResponse(**raw_response_json)
|
||||
response = ResponsesAPIResponse.model_validate(raw_response_json)
|
||||
except Exception:
|
||||
verbose_logger.debug("Volcengine Responses API: falling back to model_construct for response parsing.")
|
||||
response = ResponsesAPIResponse.model_construct(**raw_response_json)
|
||||
construct_response: Callable[..., ResponsesAPIResponse] = ResponsesAPIResponse.model_construct
|
||||
response = construct_response(**raw_response_json)
|
||||
|
||||
response._hidden_params["additional_headers"] = processed_headers
|
||||
response._hidden_params["headers"] = raw_response_headers
|
||||
|
|
@ -274,10 +279,10 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
encoded_response_id = encode_url_path_segment(response_id, field_name="response_id")
|
||||
url = f"{api_base}/{encoded_response_id}"
|
||||
data: Dict = {}
|
||||
data: dict = {}
|
||||
return url, data
|
||||
|
||||
def transform_delete_response_api_response(
|
||||
|
|
@ -286,16 +291,17 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> DeleteResponseResult:
|
||||
try:
|
||||
raw_response_json = raw_response.json()
|
||||
raw_response_json = self._parsed_response_body(raw_response)
|
||||
except Exception:
|
||||
raise VolcEngineError(message=raw_response.text, status_code=raw_response.status_code)
|
||||
try:
|
||||
return DeleteResponseResult(**raw_response_json)
|
||||
return DeleteResponseResult.model_validate(raw_response_json)
|
||||
except Exception:
|
||||
verbose_logger.debug(
|
||||
"Volcengine Responses API: falling back to model_construct for delete response parsing."
|
||||
)
|
||||
return DeleteResponseResult.model_construct(**raw_response_json)
|
||||
construct_delete_result: Callable[..., DeleteResponseResult] = DeleteResponseResult.model_construct
|
||||
return construct_delete_result(**raw_response_json)
|
||||
|
||||
#########################################################
|
||||
########## GET RESPONSE API TRANSFORMATION ###############
|
||||
|
|
@ -306,10 +312,10 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
encoded_response_id = encode_url_path_segment(response_id, field_name="response_id")
|
||||
url = f"{api_base}/{encoded_response_id}"
|
||||
data: Dict = {}
|
||||
data: dict = {}
|
||||
return url, data
|
||||
|
||||
def transform_get_response_api_response(
|
||||
|
|
@ -318,14 +324,14 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ResponsesAPIResponse:
|
||||
try:
|
||||
raw_response_json = raw_response.json()
|
||||
raw_response_json = self._parsed_response_body(raw_response)
|
||||
except Exception:
|
||||
raise VolcEngineError(message=raw_response.text, status_code=raw_response.status_code)
|
||||
|
||||
raw_response_headers = dict(raw_response.headers)
|
||||
processed_headers = process_response_headers(raw_response_headers)
|
||||
|
||||
response = ResponsesAPIResponse(**raw_response_json)
|
||||
response = ResponsesAPIResponse.model_validate(raw_response_json)
|
||||
response._hidden_params["additional_headers"] = processed_headers
|
||||
response._hidden_params["headers"] = raw_response_headers
|
||||
return response
|
||||
|
|
@ -339,15 +345,15 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
after: Optional[str] = None,
|
||||
before: Optional[str] = None,
|
||||
include: Optional[List[str]] = None,
|
||||
after: str | None = None,
|
||||
before: str | None = None,
|
||||
include: list[str] | None = None,
|
||||
limit: int = 20,
|
||||
order: Literal["asc", "desc"] = "desc",
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
encoded_response_id = encode_url_path_segment(response_id, field_name="response_id")
|
||||
url = f"{api_base}/{encoded_response_id}/input_items"
|
||||
params: Dict[str, Any] = {}
|
||||
params: dict[str, str | int] = {}
|
||||
if after is not None:
|
||||
params["after"] = after
|
||||
if before is not None:
|
||||
|
|
@ -364,9 +370,9 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Dict:
|
||||
) -> dict:
|
||||
try:
|
||||
return raw_response.json()
|
||||
return self._parsed_response_body(raw_response)
|
||||
except Exception:
|
||||
raise VolcEngineError(message=raw_response.text, status_code=raw_response.status_code)
|
||||
|
||||
|
|
@ -379,10 +385,10 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
encoded_response_id = encode_url_path_segment(response_id, field_name="response_id")
|
||||
url = f"{api_base}/{encoded_response_id}/cancel"
|
||||
data: Dict = {}
|
||||
data: dict = {}
|
||||
return url, data
|
||||
|
||||
def transform_cancel_response_api_response(
|
||||
|
|
@ -391,23 +397,23 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ResponsesAPIResponse:
|
||||
try:
|
||||
raw_response_json = raw_response.json()
|
||||
raw_response_json = self._parsed_response_body(raw_response)
|
||||
except Exception:
|
||||
raise VolcEngineError(message=raw_response.text, status_code=raw_response.status_code)
|
||||
|
||||
raw_response_headers = dict(raw_response.headers)
|
||||
processed_headers = process_response_headers(raw_response_headers)
|
||||
|
||||
response = ResponsesAPIResponse(**raw_response_json)
|
||||
response = ResponsesAPIResponse.model_validate(raw_response_json)
|
||||
response._hidden_params["additional_headers"] = processed_headers
|
||||
response._hidden_params["headers"] = raw_response_headers
|
||||
return response
|
||||
|
||||
def should_fake_stream(
|
||||
self,
|
||||
model: Optional[str],
|
||||
stream: Optional[bool],
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
model: str | None,
|
||||
stream: bool | None,
|
||||
custom_llm_provider: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Volcengine Responses API supports native streaming; never fall back to fake stream.
|
||||
|
|
@ -415,7 +421,24 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
return False
|
||||
|
||||
@staticmethod
|
||||
def _fill_missing_fields(chunk: Any, event_model: Any) -> Dict[str, Any]:
|
||||
def _parsed_response_body(raw_response: httpx.Response) -> dict[str, object]:
|
||||
return raw_response.json()
|
||||
|
||||
@staticmethod
|
||||
def _annotation_origin(annotation: object) -> object:
|
||||
return get_origin(annotation)
|
||||
|
||||
@staticmethod
|
||||
def _annotation_args(annotation: object) -> tuple[object, ...]:
|
||||
return get_args(annotation)
|
||||
|
||||
@staticmethod
|
||||
def _field_annotation(field: pyd_fields.FieldInfo) -> object:
|
||||
annotation: object = field.annotation
|
||||
return annotation
|
||||
|
||||
@staticmethod
|
||||
def _fill_missing_fields(chunk: Mapping[str, object], event_model: object | None) -> Mapping[str, object]:
|
||||
"""
|
||||
Heuristically fill missing required fields with safe defaults based on the
|
||||
event model's field annotations. This keeps parsing tolerant of providers that
|
||||
|
|
@ -424,31 +447,37 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
if not isinstance(chunk, dict) or event_model is None:
|
||||
return chunk
|
||||
|
||||
patched: Dict[str, Any] = dict(chunk)
|
||||
fields_map = getattr(event_model, "model_fields", {}) or {}
|
||||
patched = dict(chunk)
|
||||
fields_map: Mapping[str, pyd_fields.FieldInfo] = getattr(event_model, "model_fields", {}) or {}
|
||||
|
||||
for name, field in fields_map.items():
|
||||
if name in patched:
|
||||
patched[name] = VolcEngineResponsesAPIConfig._maybe_fill_nested(patched[name], field.annotation)
|
||||
patched[name] = VolcEngineResponsesAPIConfig._maybe_fill_nested(
|
||||
patched[name], VolcEngineResponsesAPIConfig._field_annotation(field)
|
||||
)
|
||||
continue
|
||||
|
||||
# Explicit default or factory
|
||||
if field.default is not pyd_fields.PydanticUndefined and field.default is not None:
|
||||
patched[name] = field.default
|
||||
field_default: object = field.default
|
||||
if field_default is not pyd_fields.PydanticUndefined and field_default is not None:
|
||||
patched[name] = field_default
|
||||
continue
|
||||
if field.default_factory is not None and field.default_factory is not pyd_fields.PydanticUndefined:
|
||||
patched[name] = field.default_factory()
|
||||
default_factory: Callable[..., object] | None = field.default_factory
|
||||
if default_factory is not None and default_factory is not pyd_fields.PydanticUndefined:
|
||||
patched[name] = default_factory()
|
||||
continue
|
||||
|
||||
# Heuristic defaults for missing required fields
|
||||
patched[name] = VolcEngineResponsesAPIConfig._default_for_annotation(field.annotation)
|
||||
patched[name] = VolcEngineResponsesAPIConfig._default_for_annotation(
|
||||
VolcEngineResponsesAPIConfig._field_annotation(field)
|
||||
)
|
||||
|
||||
return patched
|
||||
|
||||
@staticmethod
|
||||
def _default_for_annotation(annotation: Any) -> Any:
|
||||
origin = get_origin(annotation)
|
||||
args = get_args(annotation)
|
||||
def _default_for_annotation(annotation: object) -> object:
|
||||
origin = VolcEngineResponsesAPIConfig._annotation_origin(annotation)
|
||||
args = VolcEngineResponsesAPIConfig._annotation_args(annotation)
|
||||
|
||||
if annotation is int:
|
||||
return 0
|
||||
|
|
@ -456,7 +485,7 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
return []
|
||||
if origin is Union:
|
||||
# Prefer empty list when any option is a list
|
||||
if any((arg is list or get_origin(arg) is list) for arg in args):
|
||||
if any((arg is list or VolcEngineResponsesAPIConfig._annotation_origin(arg) is list) for arg in args):
|
||||
return []
|
||||
if type(None) in args:
|
||||
return None
|
||||
|
|
@ -467,53 +496,51 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
def _maybe_fill_nested(value: Any, annotation: Any) -> Any:
|
||||
def _maybe_fill_nested(value: object, annotation: object) -> object:
|
||||
"""
|
||||
Recursively fill nested dict/list structures based on the annotated model.
|
||||
"""
|
||||
model_cls = VolcEngineResponsesAPIConfig._pick_model_class(annotation, value)
|
||||
args = get_args(annotation)
|
||||
args = VolcEngineResponsesAPIConfig._annotation_args(annotation)
|
||||
|
||||
if isinstance(value, dict) and model_cls is not None:
|
||||
return VolcEngineResponsesAPIConfig._fill_missing_fields(value, model_cls)
|
||||
nested_items: Mapping[str, object] = value
|
||||
return VolcEngineResponsesAPIConfig._fill_missing_fields(nested_items, model_cls)
|
||||
|
||||
if isinstance(value, list):
|
||||
# Attempt to fill list elements if we know the element annotation
|
||||
elem_ann: Any = args[0] if args else None
|
||||
elem_ann: object = args[0] if args else None
|
||||
if elem_ann is not None:
|
||||
return [VolcEngineResponsesAPIConfig._maybe_fill_nested(v, elem_ann) for v in value]
|
||||
nested_elements: Sequence[object] = value
|
||||
return [VolcEngineResponsesAPIConfig._maybe_fill_nested(v, elem_ann) for v in nested_elements]
|
||||
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _pick_model_class(annotation: Any, value: Any) -> Optional[Any]:
|
||||
def _pick_model_class(annotation: object, value: object) -> object | None:
|
||||
"""
|
||||
Choose the best-matching Pydantic model class for a nested dict.
|
||||
"""
|
||||
candidates: List[Any] = []
|
||||
origin = get_origin(annotation)
|
||||
|
||||
if hasattr(annotation, "model_fields"):
|
||||
candidates.append(annotation)
|
||||
if origin is Union:
|
||||
for arg in get_args(annotation):
|
||||
if hasattr(arg, "model_fields"):
|
||||
candidates.append(arg)
|
||||
origin = VolcEngineResponsesAPIConfig._annotation_origin(annotation)
|
||||
union_args = VolcEngineResponsesAPIConfig._annotation_args(annotation) if origin is Union else ()
|
||||
candidates = tuple(candidate for candidate in (annotation, *union_args) if hasattr(candidate, "model_fields"))
|
||||
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
# Try to match by literal "type" field when available
|
||||
if isinstance(value, dict):
|
||||
v_type = value.get("type")
|
||||
value_items: Mapping[str, object] = value
|
||||
v_type = value_items.get("type")
|
||||
for candidate in candidates:
|
||||
try:
|
||||
type_field = candidate.model_fields.get("type")
|
||||
candidate_fields: Mapping[str, pyd_fields.FieldInfo] = getattr(candidate, "model_fields")
|
||||
type_field = candidate_fields.get("type")
|
||||
if type_field is None:
|
||||
continue
|
||||
literal_ann = type_field.annotation
|
||||
if get_origin(literal_ann) is Literal:
|
||||
literal_values = get_args(literal_ann)
|
||||
literal_ann = VolcEngineResponsesAPIConfig._field_annotation(type_field)
|
||||
if VolcEngineResponsesAPIConfig._annotation_origin(literal_ann) is Literal:
|
||||
literal_values = VolcEngineResponsesAPIConfig._annotation_args(literal_ann)
|
||||
if v_type in literal_values:
|
||||
return candidate
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -7,9 +7,10 @@ import base64
|
|||
import mimetypes
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Callable, Coroutine, Mapping
|
||||
from dataclasses import dataclass
|
||||
from io import IOBase
|
||||
from typing import Any, Callable, Coroutine, Union, cast
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -42,7 +43,7 @@ class _PreparedOCRRequest:
|
|||
provider_config: BaseOCRConfig
|
||||
optional_params: dict[str, object]
|
||||
litellm_params: dict[str, object]
|
||||
effective_timeout: Union[float, httpx.Timeout]
|
||||
effective_timeout: float | httpx.Timeout
|
||||
litellm_logging_obj: LiteLLMLoggingObj
|
||||
|
||||
|
||||
|
|
@ -63,13 +64,13 @@ _RUST_OCR_PROVIDERS = {
|
|||
|
||||
def _prepare_ocr_request(
|
||||
model: str,
|
||||
document: dict[str, Any],
|
||||
document: Mapping[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
timeout: Union[float, httpx.Timeout] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, Any] | None,
|
||||
kwargs: dict[str, Any],
|
||||
extra_headers: dict[str, object] | None,
|
||||
kwargs: dict[str, object],
|
||||
) -> _PreparedOCRRequest:
|
||||
litellm_logging_obj = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj"))
|
||||
litellm_call_id = cast(str | None, kwargs.get("litellm_call_id", None))
|
||||
|
|
@ -120,7 +121,7 @@ def _prepare_ocr_request(
|
|||
|
||||
verbose_logger.debug(f"OCR call - model: {model}, provider: {custom_llm_provider}")
|
||||
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params = GenericLiteLLMParams.model_validate(kwargs)
|
||||
|
||||
supported_params = ocr_provider_config.get_supported_ocr_params(model=model)
|
||||
non_default_params = {}
|
||||
|
|
@ -155,7 +156,7 @@ def _prepare_ocr_request(
|
|||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=cast(dict[str, object] | None, extra_headers),
|
||||
extra_headers=extra_headers,
|
||||
provider_config=ocr_provider_config,
|
||||
optional_params=cast(dict[str, object], optional_params),
|
||||
litellm_params=dict(litellm_params),
|
||||
|
|
@ -305,13 +306,13 @@ async def _run_rust_aocr(
|
|||
@client
|
||||
async def aocr(
|
||||
model: str,
|
||||
document: dict[str, Any],
|
||||
document: Mapping[str, object],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
timeout: Union[float, httpx.Timeout] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
**kwargs,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
**kwargs: object,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Async OCR function.
|
||||
|
|
@ -567,14 +568,14 @@ def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str,
|
|||
@client
|
||||
def ocr(
|
||||
model: str,
|
||||
document: dict[str, Any],
|
||||
document: Mapping[str, object],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
timeout: Union[float, httpx.Timeout] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
**kwargs,
|
||||
) -> Union[OCRResponse, Coroutine[Any, Any, OCRResponse]]:
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
**kwargs: object,
|
||||
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
"""
|
||||
Synchronous OCR function.
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -11,7 +11,7 @@ collaborators acquire their globals per call, mirroring v1's lazy-import pattern
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Mapping
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -54,7 +54,7 @@ ServerLookup = Callable[[str], "MCPServer | None"]
|
|||
StoreBuilder = Callable[[ServerLookup], tuple[InvalidatableOAuthTokenStore, bool]]
|
||||
|
||||
|
||||
async def _read_credential(user_id: str, server_id: str) -> dict[str, object] | None:
|
||||
async def _read_credential(user_id: str, server_id: str) -> Mapping[str, object] | None:
|
||||
from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415
|
||||
get_user_oauth_credential,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -10,14 +10,14 @@ injected, so the DB/decoding plumbing stays testable and out of this seam.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
OAuthToken,
|
||||
)
|
||||
|
||||
CredentialReader = Callable[[str, str], Awaitable["dict[str, object] | None"]]
|
||||
CredentialReader = Callable[[str, str], Awaitable["Mapping[str, object] | None"]]
|
||||
|
||||
|
||||
def _iso_to_epoch(expires_at: str) -> float | None:
|
||||
|
|
@ -39,7 +39,7 @@ def _to_scopes(raw: object) -> tuple[str, ...]:
|
|||
return ()
|
||||
|
||||
|
||||
def _to_oauth_token(payload: dict[str, object]) -> OAuthToken | None:
|
||||
def _to_oauth_token(payload: Mapping[str, object]) -> OAuthToken | None:
|
||||
access_token = payload.get("access_token")
|
||||
if not isinstance(access_token, str):
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -1,18 +1,11 @@
|
|||
import asyncio
|
||||
import importlib
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from datetime import datetime
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Awaitable,
|
||||
Callable,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Set,
|
||||
Tuple,
|
||||
Union,
|
||||
)
|
||||
|
||||
import httpx
|
||||
|
|
@ -38,6 +31,9 @@ from litellm.proxy._experimental.mcp_server.utils import (
|
|||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.utils import CallTypes
|
||||
|
|
@ -97,12 +93,12 @@ if MCP_AVAILABLE:
|
|||
########################################################
|
||||
############ MCP Server REST API Routes #################
|
||||
async def _safe_fire_mcp_tool_call_logging(
|
||||
logging_obj: Optional[Any],
|
||||
logging_obj: Any | None,
|
||||
result: Any,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
request_data: Optional[Mapping[str, object]] = None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
request_data: Mapping[str, object] | None = None,
|
||||
) -> None:
|
||||
if logging_obj is None:
|
||||
return
|
||||
|
|
@ -134,7 +130,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
async def _handle_virtual_mcp_tool(
|
||||
request: Request,
|
||||
data: Dict[str, Any],
|
||||
data: dict[str, Any],
|
||||
tool_name: str,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> Any:
|
||||
|
|
@ -212,9 +208,9 @@ if MCP_AVAILABLE:
|
|||
|
||||
def _get_server_auth_header(
|
||||
server,
|
||||
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
|
||||
mcp_auth_header: Optional[str],
|
||||
) -> Optional[Union[Dict[str, str], str]]:
|
||||
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
|
||||
mcp_auth_header: str | None,
|
||||
) -> dict[str, str] | str | None:
|
||||
"""Helper function to get server-specific auth header with case-insensitive matching."""
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
lookup_mcp_server_auth_in_headers,
|
||||
|
|
@ -230,7 +226,7 @@ if MCP_AVAILABLE:
|
|||
return server_auth
|
||||
return mcp_auth_header
|
||||
|
||||
def _is_v1_resolved_oauth2_server(server: Optional[MCPServer]) -> bool:
|
||||
def _is_v1_resolved_oauth2_server(server: MCPServer | None) -> bool:
|
||||
"""Whether this server's per-user OAuth2 token is still resolved by v1.
|
||||
|
||||
A server the v2 resolver owns reads its stored token from the resolver at connect
|
||||
|
|
@ -246,7 +242,7 @@ if MCP_AVAILABLE:
|
|||
return False
|
||||
return to_server_spec(server) is None
|
||||
|
||||
def _v1_resolved_oauth2_server_ids(allowed_server_ids: List[str]) -> Set[str]:
|
||||
def _v1_resolved_oauth2_server_ids(allowed_server_ids: list[str]) -> set[str]:
|
||||
"""Return the subset of *allowed_server_ids* whose per-user OAuth2 token is still
|
||||
resolved by v1.
|
||||
|
||||
|
|
@ -260,10 +256,10 @@ if MCP_AVAILABLE:
|
|||
}
|
||||
|
||||
async def _get_user_oauth_extra_headers(
|
||||
server,
|
||||
server: MCPServer,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prefetched_creds: Optional[Dict[str, Dict[str, Any]]] = None,
|
||||
) -> Optional[Dict[str, str]]:
|
||||
prefetched_creds: dict[str, "OAuthCredentialPayload"] | None = None,
|
||||
) -> dict[str, str] | None:
|
||||
"""
|
||||
For OAuth2 servers, look up the user's stored access token and return it
|
||||
as extra_headers {"Authorization": "Bearer <token>"} so that it reaches
|
||||
|
|
@ -315,7 +311,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
async def _prefetch_user_oauth_creds(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> Dict[str, Dict[str, Any]]:
|
||||
) -> dict[str, "OAuthCredentialPayload"]:
|
||||
"""Fetch all OAuth2 credentials for the user in a single DB query.
|
||||
|
||||
Returns a dict keyed by server_id. Used to avoid N+1 DB queries when
|
||||
|
|
@ -379,8 +375,8 @@ if MCP_AVAILABLE:
|
|||
|
||||
def _resolve_mcp_server_id_for_rest(
|
||||
server_id: str,
|
||||
allowed_server_ids: Union[Set[str], List[str]],
|
||||
client_ip: Optional[str] = None,
|
||||
allowed_server_ids: set[str] | list[str],
|
||||
client_ip: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Map REST ``server_id`` (UUID, server_name, or alias) to canonical server_id.
|
||||
|
|
@ -400,7 +396,7 @@ if MCP_AVAILABLE:
|
|||
request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
server_id: str,
|
||||
) -> Tuple[List[MCPServer], str]:
|
||||
) -> tuple[list[MCPServer], str]:
|
||||
"""
|
||||
Resolve allowed MCP servers for a tool call with IP filtering.
|
||||
|
||||
|
|
@ -471,7 +467,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
|
||||
# Build allowed_mcp_servers list (only include allowed servers)
|
||||
allowed_mcp_servers: List[MCPServer] = []
|
||||
allowed_mcp_servers: list[MCPServer] = []
|
||||
for allowed_server_id in allowed_server_ids_set:
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(allowed_server_id)
|
||||
if server is not None:
|
||||
|
|
@ -482,9 +478,9 @@ if MCP_AVAILABLE:
|
|||
async def _get_tools_for_single_server(
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
apply_tool_filters: bool = True,
|
||||
):
|
||||
"""Helper function to get tools for a single server.
|
||||
|
|
@ -530,7 +526,7 @@ if MCP_AVAILABLE:
|
|||
async def _resolve_allowed_mcp_servers_for_tool_call(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
server_id: str,
|
||||
) -> List[MCPServer]:
|
||||
) -> list[MCPServer]:
|
||||
"""Resolve allowed MCP servers for the given user and validate server_id access."""
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
allowed_server_ids_set = set()
|
||||
|
|
@ -545,7 +541,7 @@ if MCP_AVAILABLE:
|
|||
"message": f"The key is not allowed to access server {server_id}",
|
||||
},
|
||||
)
|
||||
allowed_mcp_servers: List[MCPServer] = []
|
||||
allowed_mcp_servers: list[MCPServer] = []
|
||||
for allowed_server_id in allowed_server_ids_set:
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(allowed_server_id)
|
||||
if server is not None:
|
||||
|
|
@ -554,10 +550,10 @@ if MCP_AVAILABLE:
|
|||
|
||||
async def _list_tools_for_single_server(
|
||||
server_id: str,
|
||||
allowed_server_ids: List[str],
|
||||
rest_client_ip: Optional[str],
|
||||
allowed_server_ids: list[str],
|
||||
rest_client_ip: str | None,
|
||||
mcp_server_auth_headers: dict,
|
||||
mcp_auth_header: Optional[str],
|
||||
mcp_auth_header: str | None,
|
||||
raw_headers_from_request: dict,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
apply_tool_filters: bool = True,
|
||||
|
|
@ -644,12 +640,12 @@ if MCP_AVAILABLE:
|
|||
"message": "Successfully retrieved tools",
|
||||
}
|
||||
|
||||
def _as_query_str(value: Any) -> Optional[str]:
|
||||
def _as_query_str(value: Any) -> str | None:
|
||||
"""Coerce an Optional[str] Query param to str|None, dropping unresolved FastAPI defaults."""
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
async def _resolve_toolset_scope(
|
||||
toolset_name: Optional[str],
|
||||
toolset_name: str | None,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> UserAPIKeyAuth:
|
||||
"""Resolve ``toolset_name`` to its scoped ``UserAPIKeyAuth``, or return unchanged."""
|
||||
|
|
@ -670,11 +666,9 @@ if MCP_AVAILABLE:
|
|||
@router.get("/tools/list", dependencies=[Depends(user_api_key_auth)])
|
||||
async def list_tool_rest_api(
|
||||
request: Request,
|
||||
server_id: Optional[str] = Query(None, description="The server id to list tools for"),
|
||||
mcp_server_name: Optional[str] = Query(
|
||||
None, description="Filter tools to a single MCP server by name or alias"
|
||||
),
|
||||
toolset_name: Optional[str] = Query(None, description="Filter tools to a single toolset by name"),
|
||||
server_id: str | None = Query(None, description="The server id to list tools for"),
|
||||
mcp_server_name: str | None = Query(None, description="Filter tools to a single MCP server by name or alias"),
|
||||
toolset_name: str | None = Query(None, description="Filter tools to a single toolset by name"),
|
||||
include_disabled_tools: bool = Query(
|
||||
False,
|
||||
description=(
|
||||
|
|
@ -981,7 +975,7 @@ if MCP_AVAILABLE:
|
|||
) = await _resolve_allowed_mcp_servers_with_ip_filter(request, user_api_key_dict, server_id)
|
||||
|
||||
# Look up per-user OAuth headers for this server (mirrors list_tool_rest_api).
|
||||
user_oauth_extra_headers: Optional[Dict[str, str]] = None
|
||||
user_oauth_extra_headers: dict[str, str] | None = None
|
||||
target_server = next(
|
||||
(s for s in allowed_mcp_servers if s.server_id == canonical_server_id),
|
||||
None,
|
||||
|
|
@ -1094,18 +1088,18 @@ if MCP_AVAILABLE:
|
|||
(client_id, client_secret, scopes) — any value may be ``None``.
|
||||
"""
|
||||
creds = request.credentials if isinstance(request.credentials, dict) else {}
|
||||
client_id: Optional[str] = creds.get("client_id")
|
||||
client_secret: Optional[str] = creds.get("client_secret")
|
||||
client_id: str | None = creds.get("client_id")
|
||||
client_secret: str | None = creds.get("client_secret")
|
||||
scopes_raw = creds.get("scopes")
|
||||
scopes: Optional[List[str]] = scopes_raw if isinstance(scopes_raw, list) else None
|
||||
scopes: list[str] | None = scopes_raw if isinstance(scopes_raw, list) else None
|
||||
return client_id, client_secret, scopes
|
||||
|
||||
async def _execute_with_mcp_client(
|
||||
request: NewMCPServerRequest,
|
||||
operation: Callable[..., Awaitable[Any]],
|
||||
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
|
||||
oauth2_headers: Optional[Dict[str, str]] = None,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
mcp_auth_header: str | dict[str, str] | None = None,
|
||||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Create a temporary MCP client from *request*, run *operation*, and return the result.
|
||||
|
|
@ -1128,7 +1122,7 @@ if MCP_AVAILABLE:
|
|||
try:
|
||||
client_id, client_secret, scopes = _extract_credentials(request)
|
||||
|
||||
_oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = request.oauth2_flow or (
|
||||
_oauth2_flow: Literal["client_credentials", "authorization_code"] | None = request.oauth2_flow or (
|
||||
"client_credentials" if client_id and client_secret and request.token_url else None
|
||||
)
|
||||
# client_credentials requires token_url to fetch a token; without it the
|
||||
|
|
@ -1244,7 +1238,7 @@ if MCP_AVAILABLE:
|
|||
spec = await load_openapi_spec_async(spec_path)
|
||||
paths = spec.get("paths", {})
|
||||
components = spec.get("components", {})
|
||||
tools: List[dict] = []
|
||||
tools: list[dict] = []
|
||||
used_names: set = set()
|
||||
for path, path_item in paths.items():
|
||||
for method in ("get", "post", "put", "delete", "patch"):
|
||||
|
|
@ -1351,7 +1345,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
headers = request.headers
|
||||
|
||||
mcp_auth_header: Optional[str] = None
|
||||
mcp_auth_header: str | None = None
|
||||
if new_mcp_server_request.auth_type in {
|
||||
MCPAuth.api_key,
|
||||
MCPAuth.bearer_token,
|
||||
|
|
@ -1365,7 +1359,7 @@ if MCP_AVAILABLE:
|
|||
# Authorization doubles as the admission fallback (LITELLM_API_KEY_HEADER_NAME_SECONDARY):
|
||||
# when the primary x-litellm-api-key header is absent, the Authorization value is the
|
||||
# caller's LiteLLM key, not an upstream token, and must never be forwarded upstream.
|
||||
oauth2_headers: Optional[Dict[str, str]] = None
|
||||
oauth2_headers: dict[str, str] | None = None
|
||||
if new_mcp_server_request.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES and headers.get(
|
||||
MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY
|
||||
):
|
||||
|
|
@ -1376,8 +1370,8 @@ if MCP_AVAILABLE:
|
|||
return await session.list_tools()
|
||||
|
||||
list_tools_response = await client.run_with_session(_list_tools_session_operation)
|
||||
list_tools_result: List[MCPTool] = list_tools_response.tools
|
||||
model_dumped_tools: List[dict] = [tool.model_dump() for tool in list_tools_result]
|
||||
list_tools_result: list[MCPTool] = list_tools_response.tools
|
||||
model_dumped_tools: list[dict] = [tool.model_dump() for tool in list_tools_result]
|
||||
return {
|
||||
"tools": model_dumped_tools,
|
||||
"error": None,
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,7 +1,8 @@
|
|||
import hashlib
|
||||
import json
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Protocol, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
|
@ -13,9 +14,81 @@ from litellm.repositories.table_repositories import AgentsRepository
|
|||
from litellm.types.agents import AgentConfig, AgentResponse, PatchAgentRequest
|
||||
|
||||
|
||||
class AgentObjectPermissionRecord(Protocol):
|
||||
def model_dump(self) -> dict[str, object]: ...
|
||||
|
||||
def dict(self) -> dict[str, object]: ...
|
||||
|
||||
|
||||
class AgentRecordDump(TypedDict):
|
||||
agent_id: str
|
||||
agent_name: str
|
||||
litellm_params: dict[str, object] | None
|
||||
agent_card_params: dict[str, object]
|
||||
static_headers: dict[str, str] | None
|
||||
extra_headers: list[str] | None
|
||||
object_permission: dict[str, object] | None
|
||||
spend: float
|
||||
tpm_limit: int | None
|
||||
rpm_limit: int | None
|
||||
session_tpm_limit: int | None
|
||||
session_rpm_limit: int | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
created_by: str | None
|
||||
updated_by: str | None
|
||||
|
||||
|
||||
class AgentRecord(Protocol):
|
||||
agent_id: str
|
||||
agent_name: str
|
||||
object_permission_id: str | None
|
||||
object_permission: AgentObjectPermissionRecord | None
|
||||
spend: float
|
||||
|
||||
def model_dump(self) -> AgentRecordDump: ...
|
||||
|
||||
def __iter__(self) -> Iterator[tuple[str, object]]: ...
|
||||
|
||||
|
||||
class AgentTableClient(Protocol):
|
||||
async def create(
|
||||
self,
|
||||
data: Mapping[str, object],
|
||||
include: Mapping[str, bool] | None = None,
|
||||
) -> AgentRecord: ...
|
||||
|
||||
async def find_unique(
|
||||
self,
|
||||
where: Mapping[str, object],
|
||||
include: Mapping[str, bool] | None = None,
|
||||
) -> AgentRecord | None: ...
|
||||
|
||||
async def find_many(
|
||||
self,
|
||||
where: Mapping[str, object] | None = None,
|
||||
order: Mapping[str, str] | None = None,
|
||||
include: Mapping[str, bool] | None = None,
|
||||
) -> Sequence[AgentRecord]: ...
|
||||
|
||||
async def update(
|
||||
self,
|
||||
where: Mapping[str, object],
|
||||
data: Mapping[str, object],
|
||||
include: Mapping[str, bool] | None = None,
|
||||
) -> AgentRecord: ...
|
||||
|
||||
async def delete(self, where: Mapping[str, object]) -> AgentRecord: ...
|
||||
|
||||
|
||||
def agents_table(prisma_client: PrismaClient) -> AgentTableClient:
|
||||
table: AgentTableClient = AgentsRepository(prisma_client).table
|
||||
return table
|
||||
|
||||
|
||||
class AgentRegistry:
|
||||
def __init__(self):
|
||||
self.agent_list: List[AgentResponse] = []
|
||||
self.agent_list: list[AgentResponse] = []
|
||||
|
||||
def reset_agent_list(self):
|
||||
self.agent_list = []
|
||||
|
|
@ -26,13 +99,13 @@ class AgentRegistry:
|
|||
def deregister_agent(self, agent_name: str):
|
||||
self.agent_list = [agent for agent in self.agent_list if agent.agent_name != agent_name]
|
||||
|
||||
def get_agent_list(self, agent_names: Optional[List[str]] = None):
|
||||
def get_agent_list(self, agent_names: Sequence[str] | None = None):
|
||||
if agent_names is not None:
|
||||
return [agent for agent in self.agent_list if agent.agent_name in agent_names]
|
||||
return self.agent_list
|
||||
|
||||
def get_public_agent_list(self) -> List[AgentResponse]:
|
||||
public_agent_list: List[AgentResponse] = []
|
||||
def get_public_agent_list(self) -> list[AgentResponse]:
|
||||
public_agent_list: list[AgentResponse] = []
|
||||
if litellm.public_agent_groups is None:
|
||||
return public_agent_list
|
||||
for agent in self.agent_list:
|
||||
|
|
@ -43,7 +116,7 @@ class AgentRegistry:
|
|||
def _create_agent_id(self, agent_config: AgentConfig) -> str:
|
||||
return hashlib.sha256(json.dumps(agent_config, sort_keys=True).encode()).hexdigest()
|
||||
|
||||
def load_agents_from_config(self, agent_config: Optional[List[AgentConfig]] = None):
|
||||
def load_agents_from_config(self, agent_config: Sequence[AgentConfig] | None = None):
|
||||
if agent_config is None:
|
||||
return None
|
||||
|
||||
|
|
@ -63,8 +136,8 @@ class AgentRegistry:
|
|||
|
||||
def load_agents_from_db_and_config(
|
||||
self,
|
||||
agent_config: Optional[List[AgentConfig]] = None,
|
||||
db_agents: Optional[List[Dict[str, Any]]] = None,
|
||||
agent_config: Sequence[AgentConfig] | None = None,
|
||||
db_agents: list[dict[str, Any]] | None = None,
|
||||
):
|
||||
self.reset_agent_list()
|
||||
|
||||
|
|
@ -96,7 +169,7 @@ class AgentRegistry:
|
|||
agent: AgentConfig,
|
||||
prisma_client: PrismaClient,
|
||||
created_by: str,
|
||||
agent_id: Optional[str] = None,
|
||||
agent_id: str | None = None,
|
||||
) -> AgentResponse:
|
||||
"""
|
||||
Add an agent to the database.
|
||||
|
|
@ -126,18 +199,18 @@ class AgentRegistry:
|
|||
agent_card_params: str = safe_dumps(agent_card_params_dict)
|
||||
|
||||
# Handle object_permission (MCP tool access for agent)
|
||||
object_permission_id: Optional[str] = None
|
||||
object_permission_id: str | None = None
|
||||
if agent.get("object_permission") is not None:
|
||||
agent_copy = dict(agent)
|
||||
object_permission_id = await handle_update_object_permission_common(agent_copy, None, prisma_client)
|
||||
|
||||
# Serialize static_headers
|
||||
static_headers_obj = agent.get("static_headers")
|
||||
static_headers_val: Optional[str] = safe_dumps(dict(static_headers_obj)) if static_headers_obj else None
|
||||
static_headers_val: str | None = safe_dumps(dict(static_headers_obj)) if static_headers_obj else None
|
||||
|
||||
extra_headers_val: Optional[List[str]] = agent.get("extra_headers")
|
||||
extra_headers_val = agent.get("extra_headers")
|
||||
|
||||
create_data: Dict[str, Any] = {
|
||||
create_data: dict[str, object] = {
|
||||
"agent_name": agent_name,
|
||||
"litellm_params": litellm_params,
|
||||
"agent_card_params": agent_card_params,
|
||||
|
|
@ -166,7 +239,7 @@ class AgentRegistry:
|
|||
create_data[rate_field] = _val
|
||||
|
||||
# Create agent in DB
|
||||
created_agent = await AgentsRepository(prisma_client).table.create(
|
||||
created_agent = await agents_table(prisma_client).create(
|
||||
data=create_data,
|
||||
include={"object_permission": True},
|
||||
)
|
||||
|
|
@ -181,12 +254,12 @@ class AgentRegistry:
|
|||
except Exception as e:
|
||||
raise Exception(f"Error adding agent to DB: {str(e)}")
|
||||
|
||||
async def delete_agent_from_db(self, agent_id: str, prisma_client: PrismaClient) -> Dict[str, Any]:
|
||||
async def delete_agent_from_db(self, agent_id: str, prisma_client: PrismaClient) -> Mapping[str, object]:
|
||||
"""
|
||||
Delete an agent from the database
|
||||
"""
|
||||
try:
|
||||
deleted_agent = await AgentsRepository(prisma_client).table.delete(where={"agent_id": agent_id})
|
||||
deleted_agent = await agents_table(prisma_client).delete(where={"agent_id": agent_id})
|
||||
return dict(deleted_agent)
|
||||
except Exception as e:
|
||||
raise Exception(f"Error deleting agent from DB: {str(e)}")
|
||||
|
|
@ -221,7 +294,7 @@ class AgentRegistry:
|
|||
raise Exception(f"Agent with ID {agent_id} not found")
|
||||
|
||||
augment_agent = {**existing_agent, **agent}
|
||||
update_data: Dict[str, Any] = {}
|
||||
update_data: dict[str, Any] = {}
|
||||
if augment_agent.get("agent_name"):
|
||||
update_data["agent_name"] = augment_agent.get("agent_name")
|
||||
if augment_agent.get("litellm_params"):
|
||||
|
|
@ -254,7 +327,7 @@ class AgentRegistry:
|
|||
if object_permission_id is not None:
|
||||
update_data["object_permission_id"] = object_permission_id
|
||||
# Patch agent in DB
|
||||
patched_agent = await AgentsRepository(prisma_client).table.update(
|
||||
patched_agent = await agents_table(prisma_client).update(
|
||||
where={"agent_id": agent_id},
|
||||
data={
|
||||
**update_data,
|
||||
|
|
@ -307,9 +380,9 @@ class AgentRegistry:
|
|||
static_headers_val_u: str = (
|
||||
safe_dumps(dict(static_headers_obj_u)) if static_headers_obj_u is not None else safe_dumps({})
|
||||
)
|
||||
extra_headers_val_u: List[str] = agent.get("extra_headers") or []
|
||||
extra_headers_val_u = agent.get("extra_headers") or []
|
||||
|
||||
update_data: Dict[str, Any] = {
|
||||
update_data: dict[str, object] = {
|
||||
"agent_name": agent_name,
|
||||
"litellm_params": litellm_params,
|
||||
"agent_card_params": agent_card_params,
|
||||
|
|
@ -330,7 +403,7 @@ class AgentRegistry:
|
|||
update_data[rate_field] = _val
|
||||
|
||||
if agent.get("object_permission") is not None:
|
||||
existing_agent = await AgentsRepository(prisma_client).table.find_unique(where={"agent_id": agent_id})
|
||||
existing_agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
|
||||
existing_object_permission_id = (
|
||||
existing_agent.object_permission_id if existing_agent is not None else None
|
||||
)
|
||||
|
|
@ -344,7 +417,7 @@ class AgentRegistry:
|
|||
update_data["object_permission_id"] = object_permission_id
|
||||
|
||||
# Update agent in DB
|
||||
updated_agent = await AgentsRepository(prisma_client).table.update(
|
||||
updated_agent = await agents_table(prisma_client).update(
|
||||
where={"agent_id": agent_id},
|
||||
data=update_data,
|
||||
include={"object_permission": True},
|
||||
|
|
@ -363,17 +436,17 @@ class AgentRegistry:
|
|||
@staticmethod
|
||||
async def get_all_agents_from_db(
|
||||
prisma_client: PrismaClient,
|
||||
) -> List[Dict[str, Any]]:
|
||||
) -> list[dict[str, object]]:
|
||||
"""
|
||||
Get all agents from the database
|
||||
"""
|
||||
try:
|
||||
agents_from_db = await AgentsRepository(prisma_client).table.find_many(
|
||||
agents_from_db = await agents_table(prisma_client).find_many(
|
||||
order={"created_at": "desc"},
|
||||
include={"object_permission": True},
|
||||
)
|
||||
|
||||
agents: List[Dict[str, Any]] = []
|
||||
agents: list[dict[str, object]] = []
|
||||
for agent in agents_from_db:
|
||||
agent_dict = dict(agent)
|
||||
# object_permission is eagerly loaded via include above
|
||||
|
|
@ -391,7 +464,7 @@ class AgentRegistry:
|
|||
def get_agent_by_id(
|
||||
self,
|
||||
agent_id: str,
|
||||
) -> Optional[AgentResponse]:
|
||||
) -> AgentResponse | None:
|
||||
"""
|
||||
Get an agent by its ID from the database
|
||||
"""
|
||||
|
|
@ -404,7 +477,7 @@ class AgentRegistry:
|
|||
except Exception as e:
|
||||
raise Exception(f"Error getting agent from DB: {str(e)}")
|
||||
|
||||
def get_agent_by_name(self, agent_name: str) -> Optional[AgentResponse]:
|
||||
def get_agent_by_name(self, agent_name: str) -> AgentResponse | None:
|
||||
"""
|
||||
Get an agent by its name from the database
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -11,9 +11,11 @@ Follows the A2A Spec.
|
|||
import asyncio
|
||||
import os
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TypedDict
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from typing_extensions import Required
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -30,6 +32,7 @@ from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
|
|||
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
|
||||
from litellm.proxy.utils import get_custom_url
|
||||
from litellm.types.agents import (
|
||||
AgentCard,
|
||||
AgentConfig,
|
||||
AgentKeySummary,
|
||||
AgentMakePublicResponse,
|
||||
|
|
@ -49,7 +52,7 @@ def _proxy_base_url(http_request: Request) -> str:
|
|||
return get_custom_url(str(http_request.base_url), route=None)
|
||||
|
||||
|
||||
def _validate_protocol_version(upstream_card: Mapping[str, Any] | None) -> None:
|
||||
def _validate_protocol_version(upstream_card: AgentCard | None) -> None:
|
||||
"""Reject an agent card pinning an unsupported A2A protocol version."""
|
||||
version = upstream_card.get("protocolVersion") if upstream_card else None
|
||||
if version is not None and normalize_protocol_version(version) is None:
|
||||
|
|
@ -63,12 +66,12 @@ def _validate_protocol_version(upstream_card: Mapping[str, Any] | None) -> None:
|
|||
|
||||
|
||||
def _build_merged_agent_card(
|
||||
upstream_card: Mapping[str, Any] | None,
|
||||
upstream_card: AgentCard | None,
|
||||
*,
|
||||
agent_id: str,
|
||||
http_request: Request,
|
||||
agent_name: str | None = None,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""Apply the LiteLLM-fronting merge to ``upstream_card`` for ``agent_id``."""
|
||||
proxy_base = _proxy_base_url(http_request)
|
||||
_validate_protocol_version(upstream_card)
|
||||
|
|
@ -88,7 +91,7 @@ def _build_merged_agent_card(
|
|||
router = APIRouter()
|
||||
|
||||
|
||||
async def _attach_keys_to_agents(agents: list[AgentResponse], prisma_client) -> None:
|
||||
async def _attach_keys_to_agents(agents: Sequence[AgentResponse], prisma_client) -> None:
|
||||
"""Attach each agent's virtual keys, derived from the key table's agent_id
|
||||
foreign key. Mirrors how spend is joined into the agent response so the UI
|
||||
never has to cross-reference a full key dump client-side. Only non-secret
|
||||
|
|
@ -113,7 +116,7 @@ async def _attach_keys_to_agents(agents: list[AgentResponse], prisma_client) ->
|
|||
|
||||
|
||||
def _redact_sensitive_agent_fields(
|
||||
agents: list[AgentResponse],
|
||||
agents: Sequence[AgentResponse],
|
||||
) -> list[AgentResponse]:
|
||||
"""
|
||||
Return copies of the given agents with sensitive configuration fields
|
||||
|
|
@ -156,9 +159,15 @@ AGENT_HEALTH_CHECK_TIMEOUT_SECONDS = float(os.environ.get("LITELLM_AGENT_HEALTH_
|
|||
AGENT_HEALTH_CHECK_GATHER_TIMEOUT_SECONDS = float(os.environ.get("LITELLM_AGENT_HEALTH_CHECK_GATHER_TIMEOUT", "30.0"))
|
||||
|
||||
|
||||
class _AgentHealthResult(TypedDict, total=False):
|
||||
agent_id: Required[str]
|
||||
healthy: Required[bool]
|
||||
error: str
|
||||
|
||||
|
||||
async def _check_agent_url_health(
|
||||
agent: AgentResponse,
|
||||
) -> Dict[str, Any]:
|
||||
) -> _AgentHealthResult:
|
||||
"""
|
||||
Perform a GET request against the agent's URL and return the health result.
|
||||
|
||||
|
|
@ -194,7 +203,7 @@ async def _check_agent_url_health(
|
|||
"/v1/agents",
|
||||
tags=["[beta] A2A Agents"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[AgentResponse],
|
||||
response_model=list[AgentResponse],
|
||||
)
|
||||
async def get_agents(
|
||||
request: Request,
|
||||
|
|
@ -230,7 +239,7 @@ async def get_agents(
|
|||
)
|
||||
|
||||
try:
|
||||
returned_agents: List[AgentResponse] = []
|
||||
returned_agents: list[AgentResponse] = []
|
||||
|
||||
# Admin users get all agents
|
||||
if (
|
||||
|
|
@ -256,7 +265,7 @@ async def get_agents(
|
|||
if prisma_client is not None:
|
||||
agent_ids = [agent.agent_id for agent in returned_agents]
|
||||
if agent_ids:
|
||||
db_agents = await AgentsRepository(prisma_client).table.find_many(
|
||||
db_agents = await agents_table(prisma_client).find_many(
|
||||
where={"agent_id": {"in": agent_ids}},
|
||||
)
|
||||
spend_map = {a.agent_id: a.spend for a in db_agents}
|
||||
|
|
@ -285,7 +294,7 @@ async def get_agents(
|
|||
agents_with_url = [agent for agent in returned_agents if (agent.agent_card_params or {}).get("url")]
|
||||
agents_without_url = [agent for agent in returned_agents if not (agent.agent_card_params or {}).get("url")]
|
||||
try:
|
||||
health_results = await asyncio.wait_for(
|
||||
health_results: Sequence[_AgentHealthResult] = await asyncio.wait_for(
|
||||
asyncio.gather(*[_check_agent_url_health(agent) for agent in agents_with_url]),
|
||||
timeout=AGENT_HEALTH_CHECK_GATHER_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
|
@ -317,10 +326,12 @@ async def get_agents(
|
|||
|
||||
#### CRUD ENDPOINTS FOR AGENTS ####
|
||||
|
||||
from litellm.proxy.agent_endpoints.agent_registry import (
|
||||
agents_table,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.agent_registry import (
|
||||
global_agent_registry as AGENT_REGISTRY,
|
||||
)
|
||||
from litellm.repositories.table_repositories import AgentsRepository
|
||||
|
||||
|
||||
@router.post(
|
||||
|
|
@ -487,7 +498,7 @@ async def get_agent_by_id(
|
|||
try:
|
||||
agent = AGENT_REGISTRY.get_agent_by_id(agent_id=agent_id)
|
||||
if agent is None:
|
||||
agent_row = await AgentsRepository(prisma_client).table.find_unique(
|
||||
agent_row = await agents_table(prisma_client).find_unique(
|
||||
where={"agent_id": agent_id},
|
||||
include={"object_permission": True},
|
||||
)
|
||||
|
|
@ -501,7 +512,7 @@ async def get_agent_by_id(
|
|||
agent = AgentResponse(**agent_dict) # type: ignore
|
||||
else:
|
||||
# Agent found in memory — refresh spend from DB
|
||||
db_row = await AgentsRepository(prisma_client).table.find_unique(where={"agent_id": agent_id})
|
||||
db_row = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
|
||||
if db_row is not None:
|
||||
agent.spend = db_row.spend
|
||||
|
||||
|
|
@ -578,7 +589,7 @@ async def update_agent(
|
|||
|
||||
try:
|
||||
# Check if agent exists
|
||||
existing_agent = await AgentsRepository(prisma_client).table.find_unique(where={"agent_id": agent_id})
|
||||
existing_agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
|
||||
if existing_agent is not None:
|
||||
existing_agent = dict(existing_agent)
|
||||
|
||||
|
|
@ -680,7 +691,7 @@ async def patch_agent(
|
|||
|
||||
try:
|
||||
# Check if agent exists
|
||||
existing_agent = await AgentsRepository(prisma_client).table.find_unique(where={"agent_id": agent_id})
|
||||
existing_agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
|
||||
if existing_agent is not None:
|
||||
existing_agent = dict(existing_agent)
|
||||
|
||||
|
|
@ -767,9 +778,9 @@ async def delete_agent(
|
|||
|
||||
try:
|
||||
# Check if agent exists
|
||||
existing_agent = await AgentsRepository(prisma_client).table.find_unique(where={"agent_id": agent_id})
|
||||
existing_agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
|
||||
if existing_agent is not None:
|
||||
existing_agent = dict[Any, Any](existing_agent)
|
||||
existing_agent = dict[str, object](existing_agent)
|
||||
|
||||
if existing_agent is None:
|
||||
raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found in DB.")
|
||||
|
|
@ -849,7 +860,7 @@ async def make_agent_public(
|
|||
agent = AGENT_REGISTRY.get_agent_by_id(agent_id=agent_id)
|
||||
if agent is None:
|
||||
# check if agent exists in DB
|
||||
agent = await AgentsRepository(prisma_client).table.find_unique(where={"agent_id": agent_id})
|
||||
agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
|
||||
if agent is not None:
|
||||
agent = AgentResponse(**agent.model_dump()) # type: ignore
|
||||
|
||||
|
|
@ -966,7 +977,7 @@ async def make_agents_public(
|
|||
agent = AGENT_REGISTRY.get_agent_by_id(agent_id=agent_id)
|
||||
if agent is None:
|
||||
# check if agent exists in DB
|
||||
agent = await AgentsRepository(prisma_client).table.find_unique(where={"agent_id": agent_id})
|
||||
agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
|
||||
if agent is not None:
|
||||
agent = AgentResponse(**agent.model_dump()) # type: ignore
|
||||
|
||||
|
|
@ -1031,7 +1042,7 @@ async def get_agent_daily_activity(
|
|||
)
|
||||
|
||||
agent_ids_list = agent_ids.split(",") if agent_ids else None
|
||||
exclude_agent_ids_list: List[str] | None = None
|
||||
exclude_agent_ids_list: list[str] | None = None
|
||||
if exclude_agent_ids:
|
||||
exclude_agent_ids_list = exclude_agent_ids.split(",") if exclude_agent_ids else None
|
||||
|
||||
|
|
@ -1044,7 +1055,7 @@ async def get_agent_daily_activity(
|
|||
)
|
||||
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
|
||||
|
||||
where_condition: Dict[str, Any] = {}
|
||||
where_condition: dict[str, object] = {}
|
||||
if not _user_has_admin_view(user_api_key_dict):
|
||||
permitted_agent_ids = await AgentRequestHandler.get_allowed_agents(user_api_key_auth=user_api_key_dict)
|
||||
# `get_allowed_agents` returns an empty list when the caller's key
|
||||
|
|
@ -1058,7 +1069,7 @@ async def get_agent_daily_activity(
|
|||
if user_api_key_dict.user_id is None:
|
||||
permitted_agent_ids = []
|
||||
else:
|
||||
owned_records = await AgentsRepository(prisma_client).table.find_many(
|
||||
owned_records = await agents_table(prisma_client).find_many(
|
||||
where={"created_by": user_api_key_dict.user_id}
|
||||
)
|
||||
permitted_agent_ids = [a.agent_id for a in owned_records]
|
||||
|
|
@ -1093,8 +1104,10 @@ async def get_agent_daily_activity(
|
|||
if agent_ids_list:
|
||||
where_condition["agent_id"] = {"in": list(agent_ids_list)}
|
||||
|
||||
agent_records = await AgentsRepository(prisma_client).table.find_many(where=where_condition)
|
||||
agent_metadata = {agent.agent_id: {"agent_name": agent.agent_name} for agent in agent_records}
|
||||
agent_records = await agents_table(prisma_client).find_many(where=where_condition)
|
||||
agent_metadata: Mapping[str, dict[str, object]] = {
|
||||
agent.agent_id: {"agent_name": agent.agent_name} for agent in agent_records
|
||||
}
|
||||
|
||||
return await get_daily_activity(
|
||||
prisma_client=prisma_client,
|
||||
|
|
|
|||
|
|
@ -4,11 +4,13 @@ GET /guardrails/usage/overview, /guardrails/usage/detail/:id, /guardrails/usage/
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, Literal, Union, overload
|
||||
|
||||
from fastapi import APIRouter, Depends, Query
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
|
@ -21,12 +23,65 @@ from litellm.repositories.table_repositories import (
|
|||
SpendLogsRepository,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_models
|
||||
from prisma import types as prisma_types
|
||||
from prisma.actions import LiteLLM_GuardrailsTableActions, LiteLLM_PolicyTableActions
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.types.guardrails import Guardrail
|
||||
|
||||
_DbOrConfigGuardrail = Union[prisma_models.LiteLLM_GuardrailsTable, Guardrail]
|
||||
_DailyMetricsRow = Union[prisma_models.LiteLLM_DailyGuardrailMetrics, prisma_models.LiteLLM_DailyPolicyMetrics]
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _guardrails_table(
|
||||
prisma_client: "PrismaClient",
|
||||
) -> "LiteLLM_GuardrailsTableActions[prisma_models.LiteLLM_GuardrailsTable]":
|
||||
guardrails_table: LiteLLM_GuardrailsTableActions[prisma_models.LiteLLM_GuardrailsTable] = GuardrailsRepository(
|
||||
prisma_client
|
||||
).table
|
||||
return guardrails_table
|
||||
|
||||
|
||||
def _policies_table(
|
||||
prisma_client: "PrismaClient",
|
||||
) -> "LiteLLM_PolicyTableActions[prisma_models.LiteLLM_PolicyTable]":
|
||||
policies_table: LiteLLM_PolicyTableActions[prisma_models.LiteLLM_PolicyTable] = PolicyRepository(
|
||||
prisma_client
|
||||
).table
|
||||
return policies_table
|
||||
|
||||
|
||||
# --- Response models ---
|
||||
|
||||
|
||||
class UsageChartPoint(TypedDict):
|
||||
date: str
|
||||
passed: int
|
||||
blocked: int
|
||||
score: NotRequired[float | None]
|
||||
|
||||
|
||||
class _MetricTotals(TypedDict):
|
||||
requests: int
|
||||
passed: int
|
||||
blocked: int
|
||||
flagged: int
|
||||
|
||||
|
||||
class _PrevPeriodCounts(TypedDict):
|
||||
req: int
|
||||
blocked: int
|
||||
|
||||
|
||||
class _DailyPassBlocked(TypedDict):
|
||||
passed: int
|
||||
blocked: int
|
||||
|
||||
|
||||
class UsageOverviewRow(BaseModel):
|
||||
id: str
|
||||
name: str
|
||||
|
|
@ -34,15 +89,15 @@ class UsageOverviewRow(BaseModel):
|
|||
provider: str
|
||||
requestsEvaluated: int
|
||||
failRate: float
|
||||
avgScore: Optional[float]
|
||||
avgLatency: Optional[float]
|
||||
avgScore: float | None
|
||||
avgLatency: float | None
|
||||
status: str # healthy | warning | critical
|
||||
trend: str # up | down | stable
|
||||
|
||||
|
||||
class UsageOverviewResponse(BaseModel):
|
||||
rows: List[UsageOverviewRow]
|
||||
chart: List[Dict[str, Any]] # [{ date, passed, blocked }]
|
||||
rows: list[UsageOverviewRow]
|
||||
chart: list[UsageChartPoint] # [{ date, passed, blocked }]
|
||||
totalRequests: int
|
||||
totalBlocked: int
|
||||
passRate: float
|
||||
|
|
@ -55,28 +110,28 @@ class UsageDetailResponse(BaseModel):
|
|||
provider: str
|
||||
requestsEvaluated: int
|
||||
failRate: float
|
||||
avgScore: Optional[float]
|
||||
avgLatency: Optional[float]
|
||||
avgScore: float | None
|
||||
avgLatency: float | None
|
||||
status: str
|
||||
trend: str
|
||||
description: Optional[str]
|
||||
time_series: List[Dict[str, Any]]
|
||||
description: str | None
|
||||
time_series: list[UsageChartPoint]
|
||||
|
||||
|
||||
class UsageLogEntry(BaseModel):
|
||||
id: str
|
||||
timestamp: str
|
||||
action: str # blocked | passed | flagged
|
||||
score: Optional[float]
|
||||
latency_ms: Optional[float]
|
||||
model: Optional[str]
|
||||
input_snippet: Optional[str]
|
||||
output_snippet: Optional[str]
|
||||
reason: Optional[str]
|
||||
score: float | None
|
||||
latency_ms: float | None
|
||||
model: str | None
|
||||
input_snippet: str | None
|
||||
output_snippet: str | None
|
||||
reason: str | None
|
||||
|
||||
|
||||
class UsageLogsResponse(BaseModel):
|
||||
logs: List[UsageLogEntry]
|
||||
logs: list[UsageLogEntry]
|
||||
total: int
|
||||
page: int
|
||||
page_size: int
|
||||
|
|
@ -101,10 +156,10 @@ def _trend_from_comparison(current_fail: float, previous_fail: float) -> str:
|
|||
return "stable"
|
||||
|
||||
|
||||
def _aggregate_daily_metrics(metrics: Any, id_attr: str) -> Dict[str, Dict[str, Any]]:
|
||||
agg: Dict[str, Dict[str, Any]] = {}
|
||||
def _aggregate_daily_metrics(metrics: "Sequence[_DailyMetricsRow]", id_attr: str) -> Mapping[str, _MetricTotals]:
|
||||
agg: dict[str, _MetricTotals] = {}
|
||||
for m in metrics:
|
||||
gid = getattr(m, id_attr)
|
||||
gid: str = getattr(m, id_attr)
|
||||
if gid not in agg:
|
||||
agg[gid] = {"requests": 0, "passed": 0, "blocked": 0, "flagged": 0}
|
||||
agg[gid]["requests"] += int(m.requests_evaluated or 0)
|
||||
|
|
@ -114,10 +169,10 @@ def _aggregate_daily_metrics(metrics: Any, id_attr: str) -> Dict[str, Dict[str,
|
|||
return agg
|
||||
|
||||
|
||||
def _prev_fail_rates(metrics_prev: Any, id_attr: str) -> Dict[str, float]:
|
||||
prev_agg_raw: Dict[str, Dict[str, int]] = {}
|
||||
def _prev_fail_rates(metrics_prev: "Sequence[_DailyMetricsRow]", id_attr: str) -> Mapping[str, float]:
|
||||
prev_agg_raw: dict[str, _PrevPeriodCounts] = {}
|
||||
for m in metrics_prev:
|
||||
gid = getattr(m, id_attr)
|
||||
gid: str = getattr(m, id_attr)
|
||||
r, b = int(m.requests_evaluated or 0), int(m.blocked_count or 0)
|
||||
if gid not in prev_agg_raw:
|
||||
prev_agg_raw[gid] = {"req": 0, "blocked": 0}
|
||||
|
|
@ -126,8 +181,8 @@ def _prev_fail_rates(metrics_prev: Any, id_attr: str) -> Dict[str, float]:
|
|||
return {gid: (100.0 * v["blocked"] / v["req"]) if v["req"] else 0.0 for gid, v in prev_agg_raw.items()}
|
||||
|
||||
|
||||
def _chart_from_metrics(metrics: Any) -> List[Dict[str, Any]]:
|
||||
chart_by_date: Dict[str, Dict[str, int]] = {}
|
||||
def _chart_from_metrics(metrics: "Sequence[_DailyMetricsRow]") -> list[UsageChartPoint]:
|
||||
chart_by_date: dict[str, _DailyPassBlocked] = {}
|
||||
for m in metrics:
|
||||
d = m.date
|
||||
if d not in chart_by_date:
|
||||
|
|
@ -137,14 +192,26 @@ def _chart_from_metrics(metrics: Any) -> List[Dict[str, Any]]:
|
|||
return [{"date": d, "passed": v["passed"], "blocked": v["blocked"]} for d, v in sorted(chart_by_date.items())]
|
||||
|
||||
|
||||
def _get_guardrail_field(g: Any, field: str) -> Any:
|
||||
_GuardrailStrField = Literal["guardrail_id", "guardrail_name"]
|
||||
_GuardrailObjectField = Literal["litellm_params", "guardrail_info"]
|
||||
|
||||
|
||||
@overload
|
||||
def _get_guardrail_field(g: "_DbOrConfigGuardrail", field: _GuardrailStrField) -> str | None: ...
|
||||
|
||||
|
||||
@overload
|
||||
def _get_guardrail_field(g: "_DbOrConfigGuardrail", field: _GuardrailObjectField) -> object: ...
|
||||
|
||||
|
||||
def _get_guardrail_field(g: "_DbOrConfigGuardrail", field: _GuardrailStrField | _GuardrailObjectField) -> object:
|
||||
"""Read `field` off a guardrail whether it's a Prisma row (attr) or a dict/TypedDict (key)."""
|
||||
if isinstance(g, dict):
|
||||
return g.get(field)
|
||||
return getattr(g, field, None)
|
||||
|
||||
|
||||
def _to_dict(value: Any) -> Dict[str, Any]:
|
||||
def _to_dict(value: object) -> dict[str, Any]:
|
||||
"""Coerce a pydantic model (e.g. LitellmParams) / dict value into a plain dict."""
|
||||
if isinstance(value, BaseModel):
|
||||
return value.model_dump(exclude_none=True)
|
||||
|
|
@ -153,7 +220,7 @@ def _to_dict(value: Any) -> Dict[str, Any]:
|
|||
return {}
|
||||
|
||||
|
||||
def _get_guardrail_attrs(g: Any) -> tuple[Any, str]:
|
||||
def _get_guardrail_attrs(g: "_DbOrConfigGuardrail") -> tuple[Any, str]:
|
||||
"""Get (guardrail_id, display_name) from guardrail - handles Prisma model or dict."""
|
||||
gid = _get_guardrail_field(g, "guardrail_id")
|
||||
name = _get_guardrail_field(g, "guardrail_name")
|
||||
|
|
@ -161,18 +228,18 @@ def _get_guardrail_attrs(g: Any) -> tuple[Any, str]:
|
|||
|
||||
|
||||
def _guardrail_overview_rows(
|
||||
guardrails: Any,
|
||||
agg: Dict[str, Dict[str, Any]],
|
||||
prev_agg: Dict[str, float],
|
||||
) -> List[UsageOverviewRow]:
|
||||
rows: List[UsageOverviewRow] = []
|
||||
covered_keys: set = set()
|
||||
guardrails: "Sequence[_DbOrConfigGuardrail]",
|
||||
agg: Mapping[str, _MetricTotals],
|
||||
prev_agg: Mapping[str, float],
|
||||
) -> list[UsageOverviewRow]:
|
||||
rows: list[UsageOverviewRow] = []
|
||||
covered_keys: set[str] = set()
|
||||
for g in guardrails:
|
||||
gid, display_name = _get_guardrail_attrs(g)
|
||||
# Metrics are keyed by logical name from spend log metadata; guardrails table uses UUID
|
||||
lookup_keys = [k for k in (display_name, gid) if k]
|
||||
lookup_keys: Sequence[str] = [k for k in (display_name, gid) if k]
|
||||
covered_keys.update(lookup_keys)
|
||||
a = {"requests": 0, "passed": 0, "blocked": 0, "flagged": 0}
|
||||
a: _MetricTotals = {"requests": 0, "passed": 0, "blocked": 0, "flagged": 0}
|
||||
for k in lookup_keys:
|
||||
if k in agg:
|
||||
a = agg[k]
|
||||
|
|
@ -229,11 +296,11 @@ def _guardrail_overview_rows(
|
|||
|
||||
|
||||
def _policy_overview_rows(
|
||||
policies: Any,
|
||||
agg: Dict[str, Dict[str, Any]],
|
||||
prev_agg: Dict[str, float],
|
||||
) -> List[UsageOverviewRow]:
|
||||
rows: List[UsageOverviewRow] = []
|
||||
policies: "Sequence[prisma_models.LiteLLM_PolicyTable]",
|
||||
agg: Mapping[str, _MetricTotals],
|
||||
prev_agg: Mapping[str, float],
|
||||
) -> list[UsageOverviewRow]:
|
||||
rows: list[UsageOverviewRow] = []
|
||||
for p in policies:
|
||||
pid = p.policy_id
|
||||
a = agg.get(pid, {"requests": 0, "passed": 0, "blocked": 0, "flagged": 0})
|
||||
|
|
@ -264,8 +331,8 @@ def _policy_overview_rows(
|
|||
response_model=UsageOverviewResponse,
|
||||
)
|
||||
async def guardrails_usage_overview(
|
||||
start_date: Optional[str] = Query(None, description="YYYY-MM-DD"),
|
||||
end_date: Optional[str] = Query(None, description="YYYY-MM-DD"),
|
||||
start_date: str | None = Query(None, description="YYYY-MM-DD"),
|
||||
end_date: str | None = Query(None, description="YYYY-MM-DD"),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Return guardrail performance overview for the dashboard."""
|
||||
|
|
@ -281,23 +348,23 @@ async def guardrails_usage_overview(
|
|||
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
|
||||
|
||||
try:
|
||||
db_guardrails = await GuardrailsRepository(prisma_client).table.find_many()
|
||||
db_guardrails = await _guardrails_table(prisma_client).find_many()
|
||||
seen_ids = {gid for g in db_guardrails if (gid := _get_guardrail_field(g, "guardrail_id")) is not None}
|
||||
config_guardrails = [
|
||||
g for g in IN_MEMORY_GUARDRAIL_HANDLER.list_config_guardrails() if g.get("guardrail_id") not in seen_ids
|
||||
]
|
||||
guardrails: List[Any] = [*db_guardrails, *config_guardrails]
|
||||
guardrails: Sequence[_DbOrConfigGuardrail] = [*db_guardrails, *config_guardrails]
|
||||
|
||||
# Daily metrics in range
|
||||
metrics = await DailyGuardrailMetricsRepository(prisma_client).table.find_many(
|
||||
where={"date": {"gte": start, "lte": end}}
|
||||
)
|
||||
metrics: Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics] = await DailyGuardrailMetricsRepository(
|
||||
prisma_client
|
||||
).table.find_many(where={"date": {"gte": start, "lte": end}})
|
||||
|
||||
# Previous period for trend
|
||||
start_prev = (datetime.strptime(start, "%Y-%m-%d") - timedelta(days=7)).strftime("%Y-%m-%d")
|
||||
metrics_prev = await DailyGuardrailMetricsRepository(prisma_client).table.find_many(
|
||||
where={"date": {"gte": start_prev, "lt": start}}
|
||||
)
|
||||
metrics_prev: Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics] = await DailyGuardrailMetricsRepository(
|
||||
prisma_client
|
||||
).table.find_many(where={"date": {"gte": start_prev, "lt": start}})
|
||||
|
||||
agg = _aggregate_daily_metrics(metrics, "guardrail_id")
|
||||
prev_agg = _prev_fail_rates(metrics_prev, "guardrail_id")
|
||||
|
|
@ -327,8 +394,8 @@ async def guardrails_usage_overview(
|
|||
)
|
||||
async def guardrails_usage_detail(
|
||||
guardrail_id: str,
|
||||
start_date: Optional[str] = Query(None),
|
||||
end_date: Optional[str] = Query(None),
|
||||
start_date: str | None = Query(None),
|
||||
end_date: str | None = Query(None),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Return single guardrail usage metrics and time series."""
|
||||
|
|
@ -345,7 +412,7 @@ async def guardrails_usage_detail(
|
|||
|
||||
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
|
||||
|
||||
guardrail: Any = await GuardrailsRepository(prisma_client).table.find_unique(where={"guardrail_id": guardrail_id})
|
||||
guardrail = await _guardrails_table(prisma_client).find_unique(where={"guardrail_id": guardrail_id})
|
||||
if guardrail is None:
|
||||
guardrail = IN_MEMORY_GUARDRAIL_HANDLER.get_config_guardrail_by_id(guardrail_id=guardrail_id)
|
||||
if guardrail is None:
|
||||
|
|
@ -357,13 +424,17 @@ async def guardrails_usage_detail(
|
|||
logical_id = _get_guardrail_field(guardrail, "guardrail_name")
|
||||
metric_ids = [i for i in (logical_id, guardrail_id) if i]
|
||||
|
||||
metrics = await DailyGuardrailMetricsRepository(prisma_client).table.find_many(
|
||||
metrics: Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics] = await DailyGuardrailMetricsRepository(
|
||||
prisma_client
|
||||
).table.find_many(
|
||||
where={
|
||||
"guardrail_id": {"in": metric_ids},
|
||||
"date": {"gte": start, "lte": end},
|
||||
}
|
||||
)
|
||||
metrics_prev = await DailyGuardrailMetricsRepository(prisma_client).table.find_many(
|
||||
metrics_prev: Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics] = await DailyGuardrailMetricsRepository(
|
||||
prisma_client
|
||||
).table.find_many(
|
||||
where={
|
||||
"guardrail_id": {"in": metric_ids},
|
||||
"date": {"lt": start},
|
||||
|
|
@ -380,14 +451,14 @@ async def guardrails_usage_detail(
|
|||
trend = _trend_from_comparison(fail_rate, prev_fail)
|
||||
|
||||
# Aggregate by date in case metrics exist under both UUID and logical name
|
||||
ts_by_date: Dict[str, Dict[str, Any]] = {}
|
||||
ts_by_date: dict[str, _DailyPassBlocked] = {}
|
||||
for m in metrics:
|
||||
d = m.date
|
||||
if d not in ts_by_date:
|
||||
ts_by_date[d] = {"passed": 0, "blocked": 0}
|
||||
ts_by_date[d]["passed"] += int(m.passed_count or 0)
|
||||
ts_by_date[d]["blocked"] += int(m.blocked_count or 0)
|
||||
time_series = [
|
||||
time_series: list[UsageChartPoint] = [
|
||||
{"date": d, "passed": v["passed"], "blocked": v["blocked"], "score": None}
|
||||
for d, v in sorted(ts_by_date.items())
|
||||
]
|
||||
|
|
@ -412,18 +483,18 @@ async def guardrails_usage_detail(
|
|||
|
||||
|
||||
def _build_usage_logs_where(
|
||||
guardrail_ids: Optional[List[str]],
|
||||
policy_id: Optional[str],
|
||||
start_date: Optional[str],
|
||||
end_date: Optional[str],
|
||||
) -> Dict[str, Any]:
|
||||
where: Dict[str, Any] = {}
|
||||
guardrail_ids: list[str] | None,
|
||||
policy_id: str | None,
|
||||
start_date: str | None,
|
||||
end_date: str | None,
|
||||
) -> "prisma_types.LiteLLM_SpendLogGuardrailIndexWhereInput":
|
||||
where: prisma_types.LiteLLM_SpendLogGuardrailIndexWhereInput = {}
|
||||
if guardrail_ids:
|
||||
where["guardrail_id"] = {"in": guardrail_ids} if len(guardrail_ids) > 1 else guardrail_ids[0]
|
||||
if policy_id:
|
||||
where["policy_id"] = policy_id
|
||||
if start_date or end_date:
|
||||
st_filter: Dict[str, Any] = {}
|
||||
st_filter: prisma_types.DateTimeFilter = {}
|
||||
if start_date:
|
||||
sd = start_date.replace("Z", "+00:00").strip()
|
||||
if "T" not in sd:
|
||||
|
|
@ -438,7 +509,9 @@ def _build_usage_logs_where(
|
|||
return where
|
||||
|
||||
|
||||
def _usage_log_entry_from_row(r: Any, sl: Any, action_filter: Optional[str]) -> Optional[UsageLogEntry]:
|
||||
def _usage_log_entry_from_row(
|
||||
r: "prisma_models.LiteLLM_SpendLogGuardrailIndex", sl: Any, action_filter: str | None
|
||||
) -> UsageLogEntry | None:
|
||||
meta = sl.metadata
|
||||
if isinstance(meta, str):
|
||||
try:
|
||||
|
|
@ -488,7 +561,7 @@ def _usage_log_entry_from_row(r: Any, sl: Any, action_filter: Optional[str]) ->
|
|||
)
|
||||
|
||||
|
||||
def _snippet(text: Any, max_len: int = 200) -> Optional[str]:
|
||||
def _snippet(text: Any, max_len: int = 200) -> str | None:
|
||||
if text is None:
|
||||
return None
|
||||
if isinstance(text, str):
|
||||
|
|
@ -510,7 +583,7 @@ def _snippet(text: Any, max_len: int = 200) -> Optional[str]:
|
|||
return result
|
||||
|
||||
|
||||
def _input_snippet_for_log(sl: Any) -> Optional[str]:
|
||||
def _input_snippet_for_log(sl: "prisma_models.LiteLLM_SpendLogs") -> str | None:
|
||||
"""Snippet for request input: prefer messages, fall back to proxy_server_request (same as drawer)."""
|
||||
out = _snippet(sl.messages)
|
||||
if out:
|
||||
|
|
@ -541,13 +614,13 @@ def _input_snippet_for_log(sl: Any) -> Optional[str]:
|
|||
response_model=UsageLogsResponse,
|
||||
)
|
||||
async def guardrails_usage_logs(
|
||||
guardrail_id: Optional[str] = Query(None),
|
||||
policy_id: Optional[str] = Query(None),
|
||||
guardrail_id: str | None = Query(None),
|
||||
policy_id: str | None = Query(None),
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(50, ge=1, le=100),
|
||||
action: Optional[str] = Query(None),
|
||||
start_date: Optional[str] = Query(None),
|
||||
end_date: Optional[str] = Query(None),
|
||||
action: str | None = Query(None),
|
||||
start_date: str | None = Query(None),
|
||||
end_date: str | None = Query(None),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Return paginated run logs for a guardrail (or policy) from SpendLogs via index."""
|
||||
|
|
@ -562,13 +635,11 @@ async def guardrails_usage_logs(
|
|||
try:
|
||||
# Index rows may store either guardrail_id (UUID) or guardrail_name from metadata.
|
||||
# Query by both so we match regardless of which was written.
|
||||
effective_guardrail_ids: List[str] = [guardrail_id] if guardrail_id else []
|
||||
effective_guardrail_ids: list[str] = [guardrail_id] if guardrail_id else []
|
||||
if guardrail_id:
|
||||
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
|
||||
|
||||
guardrail: Any = await GuardrailsRepository(prisma_client).table.find_unique(
|
||||
where={"guardrail_id": guardrail_id}
|
||||
)
|
||||
guardrail = await _guardrails_table(prisma_client).find_unique(where={"guardrail_id": guardrail_id})
|
||||
if guardrail is None:
|
||||
guardrail = IN_MEMORY_GUARDRAIL_HANDLER.get_config_guardrail_by_id(guardrail_id=guardrail_id)
|
||||
if guardrail:
|
||||
|
|
@ -577,19 +648,23 @@ async def guardrails_usage_logs(
|
|||
effective_guardrail_ids.append(logical_name)
|
||||
|
||||
where = _build_usage_logs_where(effective_guardrail_ids or None, policy_id, start_date, end_date)
|
||||
index_rows = await SpendLogGuardrailIndexRepository(prisma_client).table.find_many(
|
||||
index_rows: Sequence[prisma_models.LiteLLM_SpendLogGuardrailIndex] = await SpendLogGuardrailIndexRepository(
|
||||
prisma_client
|
||||
).table.find_many(
|
||||
where=where,
|
||||
order={"start_time": "desc"},
|
||||
skip=(page - 1) * page_size,
|
||||
take=page_size + 1,
|
||||
)
|
||||
total = await SpendLogGuardrailIndexRepository(prisma_client).table.count(where=where)
|
||||
total: int = await SpendLogGuardrailIndexRepository(prisma_client).table.count(where=where)
|
||||
request_ids = [r.request_id for r in index_rows[:page_size]]
|
||||
if not request_ids:
|
||||
return UsageLogsResponse(logs=[], total=total, page=page, page_size=page_size)
|
||||
spend_logs = await SpendLogsRepository(prisma_client).table.find_many(where={"request_id": {"in": request_ids}})
|
||||
spend_logs: Sequence[prisma_models.LiteLLM_SpendLogs] = await SpendLogsRepository(
|
||||
prisma_client
|
||||
).table.find_many(where={"request_id": {"in": request_ids}})
|
||||
log_by_id = {s.request_id: s for s in spend_logs}
|
||||
logs_out: List[UsageLogEntry] = []
|
||||
logs_out: list[UsageLogEntry] = []
|
||||
for r in index_rows[:page_size]:
|
||||
sl = log_by_id.get(r.request_id)
|
||||
if not sl:
|
||||
|
|
@ -614,8 +689,8 @@ async def guardrails_usage_logs(
|
|||
response_model=UsageOverviewResponse,
|
||||
)
|
||||
async def policies_usage_overview(
|
||||
start_date: Optional[str] = Query(None, description="YYYY-MM-DD"),
|
||||
end_date: Optional[str] = Query(None, description="YYYY-MM-DD"),
|
||||
start_date: str | None = Query(None, description="YYYY-MM-DD"),
|
||||
end_date: str | None = Query(None, description="YYYY-MM-DD"),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Return policy performance overview for the dashboard."""
|
||||
|
|
@ -629,11 +704,13 @@ async def policies_usage_overview(
|
|||
start = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d")
|
||||
|
||||
try:
|
||||
policies = await PolicyRepository(prisma_client).table.find_many()
|
||||
metrics = await DailyPolicyMetricsRepository(prisma_client).table.find_many(
|
||||
where={"date": {"gte": start, "lte": end}}
|
||||
)
|
||||
metrics_prev = await DailyPolicyMetricsRepository(prisma_client).table.find_many(
|
||||
policies = await _policies_table(prisma_client).find_many()
|
||||
metrics: Sequence[prisma_models.LiteLLM_DailyPolicyMetrics] = await DailyPolicyMetricsRepository(
|
||||
prisma_client
|
||||
).table.find_many(where={"date": {"gte": start, "lte": end}})
|
||||
metrics_prev: Sequence[prisma_models.LiteLLM_DailyPolicyMetrics] = await DailyPolicyMetricsRepository(
|
||||
prisma_client
|
||||
).table.find_many(
|
||||
where={
|
||||
"date": {
|
||||
"gte": (datetime.strptime(start, "%Y-%m-%d") - timedelta(days=7)).strftime("%Y-%m-%d"),
|
||||
|
|
|
|||
|
|
@ -1,9 +1,15 @@
|
|||
import asyncio
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Awaitable, Callable, Dict, List, Optional, Set, Tuple, Union
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Protocol,
|
||||
Union,
|
||||
)
|
||||
|
||||
from fastapi import HTTPException, status
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
|
|
@ -16,6 +22,7 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
|||
BreakdownMetrics,
|
||||
DailySpendData,
|
||||
DailySpendMetadata,
|
||||
GroupedData,
|
||||
KeyMetadata,
|
||||
KeyMetricWithMetadata,
|
||||
MetricWithMetadata,
|
||||
|
|
@ -23,8 +30,16 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
|||
SpendMetrics,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma.models import (
|
||||
LiteLLM_DeletedVerificationToken as PrismaDeletedVerificationToken,
|
||||
)
|
||||
from prisma.models import (
|
||||
LiteLLM_VerificationToken as PrismaVerificationToken,
|
||||
)
|
||||
|
||||
# Mapping from Prisma accessor names to actual PostgreSQL table names.
|
||||
_PRISMA_TO_PG_TABLE: Dict[str, str] = {
|
||||
_PRISMA_TO_PG_TABLE: Mapping[str, str] = {
|
||||
"litellm_dailyuserspend": "LiteLLM_DailyUserSpend",
|
||||
"litellm_dailyteamspend": "LiteLLM_DailyTeamSpend",
|
||||
"litellm_dailyorganizationspend": "LiteLLM_DailyOrganizationSpend",
|
||||
|
|
@ -34,7 +49,98 @@ _PRISMA_TO_PG_TABLE: Dict[str, str] = {
|
|||
}
|
||||
|
||||
|
||||
def update_metrics(existing_metrics: SpendMetrics, record: Any) -> SpendMetrics:
|
||||
class DailySpendRecord(Protocol):
|
||||
@property
|
||||
def date(self) -> str: ...
|
||||
|
||||
@property
|
||||
def api_key(self) -> str: ...
|
||||
|
||||
@property
|
||||
def model(self) -> str | None: ...
|
||||
|
||||
@property
|
||||
def model_group(self) -> str | None: ...
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> str | None: ...
|
||||
|
||||
@property
|
||||
def mcp_namespaced_tool_name(self) -> str | None: ...
|
||||
|
||||
@property
|
||||
def endpoint(self) -> str | None: ...
|
||||
|
||||
@property
|
||||
def prompt_tokens(self) -> int: ...
|
||||
|
||||
@property
|
||||
def completion_tokens(self) -> int: ...
|
||||
|
||||
@property
|
||||
def spend(self) -> float: ...
|
||||
|
||||
@property
|
||||
def cache_read_input_tokens(self) -> int: ...
|
||||
|
||||
@property
|
||||
def cache_creation_input_tokens(self) -> int: ...
|
||||
|
||||
@property
|
||||
def compression_saved_tokens(self) -> int: ...
|
||||
|
||||
@property
|
||||
def compression_savings_spend(self) -> float: ...
|
||||
|
||||
@property
|
||||
def prompt_caching_savings_spend(self) -> float: ...
|
||||
|
||||
@property
|
||||
def api_requests(self) -> int: ...
|
||||
|
||||
@property
|
||||
def successful_requests(self) -> int: ...
|
||||
|
||||
@property
|
||||
def failed_requests(self) -> int: ...
|
||||
|
||||
|
||||
class _KeyMetadataDict(TypedDict, total=False):
|
||||
key_alias: str | None
|
||||
team_id: str | None
|
||||
|
||||
|
||||
_WhereValue = Union[str, dict[str, object]]
|
||||
|
||||
|
||||
class _AggregatedSpendData(TypedDict):
|
||||
results: list[DailySpendData]
|
||||
totals: SpendMetrics
|
||||
|
||||
|
||||
class _GroupingSetsRow(SimpleNamespace):
|
||||
date: str
|
||||
api_key: str | None
|
||||
model: str | None
|
||||
model_group: str | None
|
||||
custom_llm_provider: str | None
|
||||
mcp_namespaced_tool_name: str | None
|
||||
endpoint: str | None
|
||||
group_level: int
|
||||
spend: float | None
|
||||
prompt_tokens: int | None
|
||||
completion_tokens: int | None
|
||||
cache_read_input_tokens: int | None
|
||||
cache_creation_input_tokens: int | None
|
||||
compression_saved_tokens: int | None
|
||||
compression_savings_spend: float | None
|
||||
prompt_caching_savings_spend: float | None
|
||||
api_requests: int | None
|
||||
successful_requests: int | None
|
||||
failed_requests: int | None
|
||||
|
||||
|
||||
def update_metrics(existing_metrics: SpendMetrics, record: DailySpendRecord) -> SpendMetrics:
|
||||
"""Update metrics with new record data.
|
||||
|
||||
Rollup rows can carry None for numeric fields when SUM() spans zero rows
|
||||
|
|
@ -58,7 +164,7 @@ def update_metrics(existing_metrics: SpendMetrics, record: Any) -> SpendMetrics:
|
|||
return existing_metrics
|
||||
|
||||
|
||||
def _is_user_agent_tag(tag: Optional[str]) -> bool:
|
||||
def _is_user_agent_tag(tag: str | None) -> bool:
|
||||
"""Determine whether a tag should be treated as a User-Agent tag."""
|
||||
if not tag:
|
||||
return False
|
||||
|
|
@ -66,15 +172,15 @@ def _is_user_agent_tag(tag: Optional[str]) -> bool:
|
|||
return normalized_tag.startswith("user-agent:") or normalized_tag.startswith("user agent:")
|
||||
|
||||
|
||||
def compute_tag_metadata_totals(records: List[Any]) -> SpendMetrics:
|
||||
def compute_tag_metadata_totals(records: Sequence[DailySpendRecord]) -> SpendMetrics:
|
||||
"""
|
||||
Deduplicate spend metrics for tags using request_id, ignoring User-Agent prefixed tags.
|
||||
|
||||
Each unique request_id contributes at most one record (the tag with max spend) to metadata.
|
||||
"""
|
||||
deduped_records: Dict[str, Any] = {}
|
||||
deduped_records: dict[str, DailySpendRecord] = {}
|
||||
for record in records:
|
||||
request_id = getattr(record, "request_id", None)
|
||||
request_id: str | None = getattr(record, "request_id", None)
|
||||
if not request_id:
|
||||
continue
|
||||
|
||||
|
|
@ -94,12 +200,12 @@ def compute_tag_metadata_totals(records: List[Any]) -> SpendMetrics:
|
|||
|
||||
def update_breakdown_metrics(
|
||||
breakdown: BreakdownMetrics,
|
||||
record: Any,
|
||||
model_metadata: Dict[str, Dict[str, Any]],
|
||||
provider_metadata: Dict[str, Dict[str, Any]],
|
||||
api_key_metadata: Dict[str, Dict[str, Any]],
|
||||
entity_id_field: Optional[str] = None,
|
||||
entity_metadata_field: Optional[Dict[str, dict]] = None,
|
||||
record: DailySpendRecord,
|
||||
model_metadata: Mapping[str, dict[str, object]],
|
||||
provider_metadata: Mapping[str, dict[str, object]],
|
||||
api_key_metadata: Mapping[str, _KeyMetadataDict],
|
||||
entity_id_field: str | None = None,
|
||||
entity_metadata_field: Mapping[str, dict[str, object]] | None = None,
|
||||
) -> BreakdownMetrics:
|
||||
"""Updates breakdown metrics for a single record using the existing update_metrics function"""
|
||||
|
||||
|
|
@ -269,23 +375,27 @@ def update_breakdown_metrics(
|
|||
|
||||
async def get_api_key_metadata(
|
||||
prisma_client: PrismaClient,
|
||||
api_keys: Set[str],
|
||||
) -> Dict[str, Dict[str, Any]]:
|
||||
api_keys: set[str],
|
||||
) -> dict[str, _KeyMetadataDict]:
|
||||
"""Get api key metadata, falling back to deleted keys table for keys not found in active table.
|
||||
|
||||
This ensures that key_alias and team_id are preserved in historical activity logs
|
||||
even after a key is deleted or regenerated.
|
||||
"""
|
||||
key_records = await VerificationTokenRepository(prisma_client).table.find_many(
|
||||
key_records: list[PrismaVerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many(
|
||||
where={"token": {"in": list(api_keys)}}
|
||||
)
|
||||
result = {k.token: {"key_alias": k.key_alias, "team_id": k.team_id} for k in key_records}
|
||||
result: dict[str, _KeyMetadataDict] = {
|
||||
k.token: {"key_alias": k.key_alias, "team_id": k.team_id} for k in key_records
|
||||
}
|
||||
|
||||
# For any keys not found in the active table, check the deleted keys table
|
||||
missing_keys = api_keys - set(result.keys())
|
||||
if missing_keys:
|
||||
try:
|
||||
deleted_key_records = await DeletedVerificationTokenRepository(prisma_client).table.find_many(
|
||||
deleted_key_records: list[PrismaDeletedVerificationToken] = await DeletedVerificationTokenRepository(
|
||||
prisma_client
|
||||
).table.find_many(
|
||||
where={"token": {"in": list(missing_keys)}},
|
||||
order={"deleted_at": "desc"},
|
||||
)
|
||||
|
|
@ -309,8 +419,8 @@ async def get_api_key_metadata(
|
|||
def _adjust_dates_for_timezone(
|
||||
start_date: str,
|
||||
end_date: str,
|
||||
timezone_offset_minutes: Optional[int],
|
||||
) -> Tuple[str, str]:
|
||||
timezone_offset_minutes: int | None,
|
||||
) -> tuple[str, str]:
|
||||
"""
|
||||
Pass-through for the local date range; the timezone offset is intentionally ignored here.
|
||||
|
||||
|
|
@ -335,19 +445,19 @@ def _adjust_dates_for_timezone(
|
|||
def _build_where_conditions(
|
||||
*,
|
||||
entity_id_field: str,
|
||||
entity_id: Optional[Union[str, List[str]]],
|
||||
entity_id: str | list[str] | None,
|
||||
start_date: str,
|
||||
end_date: str,
|
||||
model: Optional[str],
|
||||
api_key: Optional[Union[str, List[str]]],
|
||||
exclude_entity_ids: Optional[List[str]] = None,
|
||||
timezone_offset_minutes: Optional[int] = None,
|
||||
) -> Dict[str, Any]:
|
||||
model: str | None,
|
||||
api_key: str | list[str] | None,
|
||||
exclude_entity_ids: list[str] | None = None,
|
||||
timezone_offset_minutes: int | None = None,
|
||||
) -> dict[str, "_WhereValue"]:
|
||||
"""Build prisma where clause for daily activity queries."""
|
||||
# Adjust dates for timezone if provided
|
||||
adjusted_start, adjusted_end = _adjust_dates_for_timezone(start_date, end_date, timezone_offset_minutes)
|
||||
|
||||
where_conditions: Dict[str, Any] = {
|
||||
where_conditions: dict[str, _WhereValue] = {
|
||||
"date": {
|
||||
"gte": adjusted_start,
|
||||
"lte": adjusted_end,
|
||||
|
|
@ -369,7 +479,7 @@ def _build_where_conditions(
|
|||
where_conditions[entity_id_field] = {"equals": entity_id}
|
||||
|
||||
if exclude_entity_ids:
|
||||
current = where_conditions.get(entity_id_field, {})
|
||||
current: _WhereValue = where_conditions.get(entity_id_field, {})
|
||||
if isinstance(current, str):
|
||||
current = {"equals": current}
|
||||
current["not"] = {"in": exclude_entity_ids}
|
||||
|
|
@ -382,14 +492,14 @@ def _build_aggregated_sql_query(
|
|||
*,
|
||||
table_name: str,
|
||||
entity_id_field: str,
|
||||
entity_id: Optional[Union[str, List[str]]],
|
||||
entity_id: str | list[str] | None,
|
||||
start_date: str,
|
||||
end_date: str,
|
||||
model: Optional[str],
|
||||
api_key: Optional[str],
|
||||
exclude_entity_ids: Optional[List[str]] = None,
|
||||
timezone_offset_minutes: Optional[int] = None,
|
||||
) -> Tuple[str, List[Any]]:
|
||||
model: str | None,
|
||||
api_key: str | None,
|
||||
exclude_entity_ids: list[str] | None = None,
|
||||
timezone_offset_minutes: int | None = None,
|
||||
) -> tuple[str, list[str]]:
|
||||
"""Build a parameterized SQL GROUP BY query for aggregated daily activity.
|
||||
|
||||
Groups by (date, api_key, model, model_group, custom_llm_provider,
|
||||
|
|
@ -406,8 +516,8 @@ def _build_aggregated_sql_query(
|
|||
|
||||
adjusted_start, adjusted_end = _adjust_dates_for_timezone(start_date, end_date, timezone_offset_minutes)
|
||||
|
||||
sql_conditions: List[str] = []
|
||||
sql_params: List[Any] = []
|
||||
sql_conditions: list[str] = []
|
||||
sql_params: list[str] = []
|
||||
p = 1 # parameter index (1-based for PostgreSQL $N placeholders)
|
||||
|
||||
# Date range (always present)
|
||||
|
|
@ -506,17 +616,17 @@ def _build_aggregated_sql_query(
|
|||
|
||||
def _aggregate_spend_records_sync(
|
||||
*,
|
||||
records: List[Any],
|
||||
api_key_metadata: Dict[str, Dict[str, Any]],
|
||||
entity_id_field: Optional[str],
|
||||
entity_metadata_field: Optional[Dict[str, dict]],
|
||||
) -> Dict[str, Any]:
|
||||
model_metadata: Dict[str, Dict[str, Any]] = {}
|
||||
provider_metadata: Dict[str, Dict[str, Any]] = {}
|
||||
records: Sequence[DailySpendRecord],
|
||||
api_key_metadata: Mapping[str, _KeyMetadataDict],
|
||||
entity_id_field: str | None,
|
||||
entity_metadata_field: Mapping[str, dict[str, object]] | None,
|
||||
) -> _AggregatedSpendData:
|
||||
model_metadata: dict[str, dict[str, object]] = {}
|
||||
provider_metadata: dict[str, dict[str, object]] = {}
|
||||
|
||||
results: List[DailySpendData] = []
|
||||
results: list[DailySpendData] = []
|
||||
total_metrics = SpendMetrics()
|
||||
grouped_data: Dict[str, Dict[str, Any]] = {}
|
||||
grouped_data: dict[str, GroupedData] = {}
|
||||
|
||||
for record in records:
|
||||
date_str = record.date
|
||||
|
|
@ -557,18 +667,18 @@ def _aggregate_spend_records_sync(
|
|||
async def _aggregate_spend_records(
|
||||
*,
|
||||
prisma_client: PrismaClient,
|
||||
records: List[Any],
|
||||
entity_id_field: Optional[str],
|
||||
entity_metadata_field: Optional[Dict[str, dict]],
|
||||
) -> Dict[str, Any]:
|
||||
records: Sequence[DailySpendRecord],
|
||||
entity_id_field: str | None,
|
||||
entity_metadata_field: Mapping[str, dict[str, object]] | None,
|
||||
) -> _AggregatedSpendData:
|
||||
"""Aggregate rows into DailySpendData list and total metrics.
|
||||
|
||||
The per-row loop is offloaded to a worker thread via asyncio.to_thread so
|
||||
a large result set doesn't peg the event loop.
|
||||
"""
|
||||
api_keys: Set[str] = {record.api_key for record in records if record.api_key}
|
||||
api_keys: set[str] = {record.api_key for record in records if record.api_key}
|
||||
|
||||
api_key_metadata: Dict[str, Dict[str, Any]] = {}
|
||||
api_key_metadata: dict[str, _KeyMetadataDict] = {}
|
||||
if api_keys:
|
||||
api_key_metadata = await get_api_key_metadata(prisma_client, api_keys)
|
||||
|
||||
|
|
@ -603,7 +713,7 @@ _GROUP_DATE_ENDPOINT = 62 # 0b0111110
|
|||
_GROUP_DATE_ENDPOINT_API_KEY = 30 # 0b0011110
|
||||
|
||||
|
||||
def _record_to_spend_metrics(record: Any) -> SpendMetrics:
|
||||
def _record_to_spend_metrics(record: _GroupingSetsRow) -> SpendMetrics:
|
||||
"""Build a SpendMetrics directly from one already-aggregated rollup row.
|
||||
|
||||
SUM() over zero rows is SQL NULL, so rollup rows (notably the grand-total
|
||||
|
|
@ -627,16 +737,16 @@ def _record_to_spend_metrics(record: Any) -> SpendMetrics:
|
|||
)
|
||||
|
||||
|
||||
def _key_metadata(api_key_metadata: Dict[str, Dict[str, Any]], api_key: str) -> KeyMetadata:
|
||||
def _key_metadata(api_key_metadata: Mapping[str, _KeyMetadataDict], api_key: str) -> KeyMetadata:
|
||||
meta = api_key_metadata.get(api_key, {})
|
||||
return KeyMetadata(key_alias=meta.get("key_alias"), team_id=meta.get("team_id"))
|
||||
|
||||
|
||||
def _aggregate_grouping_sets_records_sync(
|
||||
*,
|
||||
records: List[Any],
|
||||
api_key_metadata: Dict[str, Dict[str, Any]],
|
||||
) -> Dict[str, Any]:
|
||||
records: Sequence[_GroupingSetsRow],
|
||||
api_key_metadata: Mapping[str, _KeyMetadataDict],
|
||||
) -> _AggregatedSpendData:
|
||||
"""Build the response from rollup rows produced by the GROUPING SETS query.
|
||||
|
||||
Each row carries a `group_level` bitmask (from Postgres GROUPING()) that
|
||||
|
|
@ -645,16 +755,16 @@ def _aggregate_grouping_sets_records_sync(
|
|||
summing in Python and no nested update_metrics calls.
|
||||
"""
|
||||
total_metrics = SpendMetrics()
|
||||
grouped_data: Dict[str, Dict[str, Any]] = {}
|
||||
grouped_data: dict[str, GroupedData] = {}
|
||||
|
||||
def ensure_date(date_str: str) -> Dict[str, Any]:
|
||||
bucket = grouped_data.get(date_str)
|
||||
def ensure_date(date_str: str) -> GroupedData:
|
||||
bucket: GroupedData | None = grouped_data.get(date_str)
|
||||
if bucket is None:
|
||||
bucket = {"metrics": SpendMetrics(), "breakdown": BreakdownMetrics()}
|
||||
grouped_data[date_str] = bucket
|
||||
return bucket
|
||||
|
||||
def assign_metric_with_metadata(target: Dict[str, MetricWithMetadata], key: str, metrics: SpendMetrics) -> None:
|
||||
def assign_metric_with_metadata(target: dict[str, MetricWithMetadata], key: str, metrics: SpendMetrics) -> None:
|
||||
existing = target.get(key)
|
||||
if existing is None:
|
||||
target[key] = MetricWithMetadata(metrics=metrics, metadata={})
|
||||
|
|
@ -662,7 +772,7 @@ def _aggregate_grouping_sets_records_sync(
|
|||
existing.metrics = metrics
|
||||
|
||||
def assign_api_key_breakdown(
|
||||
target: Dict[str, MetricWithMetadata],
|
||||
target: dict[str, MetricWithMetadata],
|
||||
parent_key: str,
|
||||
api_key: str,
|
||||
metrics: SpendMetrics,
|
||||
|
|
@ -753,12 +863,12 @@ def _aggregate_grouping_sets_records_sync(
|
|||
async def _aggregate_grouping_sets_records(
|
||||
*,
|
||||
prisma_client: PrismaClient,
|
||||
records: List[Any],
|
||||
) -> Dict[str, Any]:
|
||||
records: Sequence[_GroupingSetsRow],
|
||||
) -> _AggregatedSpendData:
|
||||
"""Async wrapper: fetch api_key_metadata, then dispatch on a worker thread."""
|
||||
api_keys: Set[str] = {r.api_key for r in records if r.api_key}
|
||||
api_keys: set[str] = {r.api_key for r in records if r.api_key}
|
||||
|
||||
api_key_metadata: Dict[str, Dict[str, Any]] = {}
|
||||
api_key_metadata: dict[str, _KeyMetadataDict] = {}
|
||||
if api_keys:
|
||||
api_key_metadata = await get_api_key_metadata(prisma_client, api_keys)
|
||||
|
||||
|
|
@ -770,21 +880,22 @@ async def _aggregate_grouping_sets_records(
|
|||
|
||||
|
||||
async def get_daily_activity(
|
||||
prisma_client: Optional[PrismaClient],
|
||||
prisma_client: PrismaClient | None,
|
||||
table_name: str,
|
||||
entity_id_field: str,
|
||||
entity_id: Optional[Union[str, List[str]]],
|
||||
entity_metadata_field: Optional[Dict[str, dict]],
|
||||
start_date: Optional[str],
|
||||
end_date: Optional[str],
|
||||
model: Optional[str],
|
||||
api_key: Optional[Union[str, List[str]]],
|
||||
entity_id: str | list[str] | None,
|
||||
entity_metadata_field: Mapping[str, dict[str, object]] | None,
|
||||
start_date: str | None,
|
||||
end_date: str | None,
|
||||
model: str | None,
|
||||
api_key: str | list[str] | None,
|
||||
page: int,
|
||||
page_size: int,
|
||||
exclude_entity_ids: Optional[List[str]] = None,
|
||||
metadata_metrics_func: Optional[Callable[[List[Any]], SpendMetrics]] = None,
|
||||
timezone_offset_minutes: Optional[int] = None,
|
||||
resolve_entity_metadata: Optional[Callable[[list[Any]], Awaitable[dict[str, dict]]]] = None,
|
||||
exclude_entity_ids: list[str] | None = None,
|
||||
metadata_metrics_func: Callable[[Sequence[DailySpendRecord]], SpendMetrics] | None = None,
|
||||
timezone_offset_minutes: int | None = None,
|
||||
resolve_entity_metadata: Callable[[Sequence[DailySpendRecord]], Awaitable[dict[str, dict[str, object]]]]
|
||||
| None = None,
|
||||
) -> SpendAnalyticsPaginatedResponse:
|
||||
"""Common function to get daily activity for any entity type.
|
||||
|
||||
|
|
@ -819,7 +930,7 @@ async def get_daily_activity(
|
|||
)
|
||||
|
||||
# Get total count for pagination
|
||||
total_count = await getattr(prisma_client.db, table_name).count(where=where_conditions)
|
||||
total_count: int = await getattr(prisma_client.db, table_name).count(where=where_conditions)
|
||||
|
||||
# Fetch paginated results.
|
||||
# ``date`` alone is not a unique sort key -- a busy tenant has many
|
||||
|
|
@ -831,7 +942,7 @@ async def get_daily_activity(
|
|||
# total. Adding ``id`` (the row's UUID primary key, present on both
|
||||
# LiteLLM_DailyUserSpend and LiteLLM_DailyTeamSpend) as a tiebreaker
|
||||
# gives every page a stable cursor (#30164).
|
||||
daily_spend_data = await getattr(prisma_client.db, table_name).find_many(
|
||||
daily_spend_data: Sequence[DailySpendRecord] = await getattr(prisma_client.db, table_name).find_many(
|
||||
where=where_conditions,
|
||||
order=[
|
||||
{"date": "desc"},
|
||||
|
|
@ -889,17 +1000,17 @@ async def get_daily_activity(
|
|||
|
||||
|
||||
async def get_daily_activity_aggregated(
|
||||
prisma_client: Optional[PrismaClient],
|
||||
prisma_client: PrismaClient | None,
|
||||
table_name: str,
|
||||
entity_id_field: str,
|
||||
entity_id: Optional[Union[str, List[str]]],
|
||||
entity_metadata_field: Optional[Dict[str, dict]],
|
||||
start_date: Optional[str],
|
||||
end_date: Optional[str],
|
||||
model: Optional[str],
|
||||
api_key: Optional[str],
|
||||
exclude_entity_ids: Optional[List[str]] = None,
|
||||
timezone_offset_minutes: Optional[int] = None,
|
||||
entity_id: str | list[str] | None,
|
||||
entity_metadata_field: Mapping[str, dict[str, object]] | None,
|
||||
start_date: str | None,
|
||||
end_date: str | None,
|
||||
model: str | None,
|
||||
api_key: str | None,
|
||||
exclude_entity_ids: list[str] | None = None,
|
||||
timezone_offset_minutes: int | None = None,
|
||||
) -> SpendAnalyticsPaginatedResponse:
|
||||
"""Aggregated variant that returns the full result set (no pagination).
|
||||
|
||||
|
|
@ -939,7 +1050,7 @@ async def get_daily_activity_aggregated(
|
|||
if rows is None:
|
||||
rows = []
|
||||
|
||||
records = [SimpleNamespace(**row) for row in rows]
|
||||
records = [_GroupingSetsRow(**row) for row in rows]
|
||||
|
||||
# The grouping-sets dispatcher places each row directly in its bucket
|
||||
# using the row's GROUPING() bitmask. No Python-side summing needed.
|
||||
|
|
|
|||
|
|
@ -15,8 +15,9 @@ These are members of a Team on LiteLLM
|
|||
import asyncio
|
||||
import json
|
||||
import traceback
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Optional, Union, cast
|
||||
from typing import Any, Optional, cast
|
||||
|
||||
import fastapi
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
|
||||
|
|
@ -29,6 +30,7 @@ from litellm.proxy.auth.auth_checks import get_team_object, get_user_object
|
|||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventHooks
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import (
|
||||
DailySpendRecord,
|
||||
get_daily_activity,
|
||||
get_daily_activity_aggregated,
|
||||
)
|
||||
|
|
@ -59,17 +61,17 @@ from litellm.repositories.verification_token_repository import (
|
|||
from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
||||
SpendAnalyticsPaginatedResponse,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.scim_v2 import (
|
||||
SCIM_ENTERPRISE_METADATA_KEY,
|
||||
SCIM_ENTITLEMENTS_METADATA_KEY,
|
||||
SCIM_ROLES_METADATA_KEY,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
|
||||
BulkUpdateUserRequest,
|
||||
BulkUpdateUserResponse,
|
||||
UserListResponse,
|
||||
UserUpdateResult,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.scim_v2 import (
|
||||
SCIM_ENTERPRISE_METADATA_KEY,
|
||||
SCIM_ENTITLEMENTS_METADATA_KEY,
|
||||
SCIM_ROLES_METADATA_KEY,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.proxy_server import PrismaClient
|
||||
|
|
@ -127,11 +129,11 @@ def _update_internal_new_user_params(data_json: dict, data: NewUserRequest) -> d
|
|||
|
||||
async def _check_duplicate_user_field(
|
||||
field_name: str,
|
||||
field_value: Optional[str],
|
||||
field_value: str | None,
|
||||
prisma_client: Any,
|
||||
*,
|
||||
case_insensitive: bool = False,
|
||||
label: Optional[str] = None,
|
||||
label: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Helper function to check if a field already exists in the user table.
|
||||
|
|
@ -167,7 +169,7 @@ async def _check_duplicate_user_field(
|
|||
)
|
||||
|
||||
|
||||
async def _check_duplicate_user_email(user_email: Optional[str], prisma_client: Any) -> None:
|
||||
async def _check_duplicate_user_email(user_email: str | None, prisma_client: Any) -> None:
|
||||
"""
|
||||
Helper function to check if a user email already exists in the database.
|
||||
"""
|
||||
|
|
@ -180,7 +182,7 @@ async def _check_duplicate_user_email(user_email: Optional[str], prisma_client:
|
|||
)
|
||||
|
||||
|
||||
async def _check_duplicate_user_id(user_id: Optional[str], prisma_client: Any) -> None:
|
||||
async def _check_duplicate_user_id(user_id: str | None, prisma_client: Any) -> None:
|
||||
"""
|
||||
Helper function to check if a user id already exists in the database.
|
||||
"""
|
||||
|
|
@ -194,7 +196,7 @@ async def _check_duplicate_user_id(user_id: Optional[str], prisma_client: Any) -
|
|||
|
||||
async def _add_user_to_organizations(
|
||||
user_id: str,
|
||||
organizations: List[str],
|
||||
organizations: list[str],
|
||||
prisma_client: "PrismaClient",
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
):
|
||||
|
|
@ -231,8 +233,8 @@ async def _add_user_to_team(
|
|||
user_id: str,
|
||||
team_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
user_email: Optional[str] = None,
|
||||
max_budget_in_team: Optional[float] = None,
|
||||
user_email: str | None = None,
|
||||
max_budget_in_team: float | None = None,
|
||||
user_role: Literal["user", "admin"] = "user",
|
||||
):
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_member_add
|
||||
|
|
@ -280,7 +282,7 @@ async def _add_user_to_team(
|
|||
raise e
|
||||
|
||||
|
||||
def check_if_default_team_set() -> Optional[Union[List[str], List[NewUserRequestTeam]]]:
|
||||
def check_if_default_team_set() -> list[str] | list[NewUserRequestTeam] | None:
|
||||
if litellm.default_internal_user_params is None:
|
||||
return None
|
||||
teams = litellm.default_internal_user_params.get("teams")
|
||||
|
|
@ -306,9 +308,9 @@ def check_if_default_team_set() -> Optional[Union[List[str], List[NewUserRequest
|
|||
|
||||
async def add_new_user_to_default_team(
|
||||
user_id: str,
|
||||
user_email: Optional[str],
|
||||
user_email: str | None,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
teams: Union[List[str], List[NewUserRequestTeam]],
|
||||
teams: list[str] | list[NewUserRequestTeam],
|
||||
prisma_client: "PrismaClient",
|
||||
):
|
||||
tasks = []
|
||||
|
|
@ -459,7 +461,7 @@ async def new_user(
|
|||
teams = data.teams
|
||||
if teams is None:
|
||||
teams = check_if_default_team_set()
|
||||
organization_ids = cast(Optional[List[str]], data_json.pop("organizations", None))
|
||||
organization_ids = cast(list[str] | None, data_json.pop("organizations", None))
|
||||
|
||||
response = await generate_key_helper_fn(request_type="user", **data_json)
|
||||
# Admin UI Logic
|
||||
|
|
@ -484,7 +486,7 @@ async def new_user(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
user_id = cast(Optional[str], response.get("user_id", None))
|
||||
user_id = cast(str | None, response.get("user_id", None))
|
||||
|
||||
if organization_ids is not None and user_id is not None:
|
||||
await _add_user_to_organizations(
|
||||
|
|
@ -560,9 +562,9 @@ async def ui_get_available_role(
|
|||
|
||||
|
||||
def get_team_from_list(
|
||||
team_list: Optional[Union[List[LiteLLM_TeamTable], List[TeamListResponseObject]]],
|
||||
team_list: list[LiteLLM_TeamTable] | list[TeamListResponseObject] | None,
|
||||
team_id: str,
|
||||
) -> Optional[Union[LiteLLM_TeamTable, LiteLLM_TeamMembership]]:
|
||||
) -> LiteLLM_TeamTable | LiteLLM_TeamMembership | None:
|
||||
if team_list is None:
|
||||
return None
|
||||
|
||||
|
|
@ -584,12 +586,12 @@ def _is_valid_user_id(user_id: str) -> bool:
|
|||
return True
|
||||
|
||||
|
||||
def get_user_id_from_request(request: Request) -> Optional[str]:
|
||||
def get_user_id_from_request(request: Request) -> str | None:
|
||||
"""
|
||||
Get the user id from the request
|
||||
"""
|
||||
# Get the raw query string and parse it properly to handle + characters
|
||||
user_id: Optional[str] = None
|
||||
user_id: str | None = None
|
||||
query_string = str(request.url.query)
|
||||
if "user_id=" in query_string:
|
||||
# Extract the user_id value from the raw query string
|
||||
|
|
@ -605,14 +607,14 @@ def get_user_id_from_request(request: Request) -> Optional[str]:
|
|||
return user_id
|
||||
|
||||
|
||||
def _normalize_user_info_user_id(request: Request, user_id: Optional[str]) -> Optional[str]:
|
||||
def _normalize_user_info_user_id(request: Request, user_id: str | None) -> str | None:
|
||||
"""Normalize URL-decoded user_id while preserving '+' characters."""
|
||||
if user_id is not None and " " in user_id:
|
||||
return get_user_id_from_request(request=request)
|
||||
return user_id
|
||||
|
||||
|
||||
def _enforce_user_info_access(user_id: Optional[str], user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
def _enforce_user_info_access(user_id: str | None, user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
"""Re-validate that the caller may read the resolved ``user_id`` after
|
||||
URL-decoding has been finalized.
|
||||
|
||||
|
|
@ -645,10 +647,10 @@ def _enforce_user_info_access(user_id: Optional[str], user_api_key_dict: UserAPI
|
|||
|
||||
async def _get_user_info_teams(
|
||||
prisma_client: Any,
|
||||
user_id: Optional[str],
|
||||
user_info: Optional[Any],
|
||||
user_id: str | None,
|
||||
user_info: Any | None,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> tuple[list[Any], Optional[list[Any]]]:
|
||||
) -> tuple[list[Any], list[Any] | None]:
|
||||
"""Fetch and merge teams from membership + user.teams field."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import list_team
|
||||
|
||||
|
|
@ -667,7 +669,7 @@ async def _get_user_info_teams(
|
|||
team_list = teams_1
|
||||
team_id_list = [team.team_id for team in teams_1]
|
||||
|
||||
teams_2: Optional[list[Any]] = None
|
||||
teams_2: list[Any] | None = None
|
||||
target_team_ids = getattr(user_info, "teams", None)
|
||||
|
||||
if target_team_ids and isinstance(target_team_ids, list):
|
||||
|
|
@ -701,8 +703,8 @@ _SCIM_DIRECTORY_METADATA_KEYS = frozenset(
|
|||
|
||||
|
||||
def _redact_scim_enterprise_metadata(
|
||||
metadata: Optional[Dict[str, Any]],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
metadata: dict[str, Any] | None,
|
||||
) -> dict[str, Any] | None:
|
||||
"""SCIM enterprise attributes, entitlements, and roles are persisted in user
|
||||
metadata so reporting can group on them, but they are directory-only fields
|
||||
that generic user-info endpoints must not surface; SCIM clients read them
|
||||
|
|
@ -713,11 +715,11 @@ def _redact_scim_enterprise_metadata(
|
|||
|
||||
|
||||
def _build_user_info_response(
|
||||
user_id: Optional[str],
|
||||
user_info: Optional[Any],
|
||||
keys: Optional[List[LiteLLM_VerificationToken]],
|
||||
user_id: str | None,
|
||||
user_info: Any | None,
|
||||
keys: list[LiteLLM_VerificationToken] | None,
|
||||
team_list: list[Any],
|
||||
teams_1: Optional[list[Any]],
|
||||
teams_1: list[Any] | None,
|
||||
) -> UserInfoResponse:
|
||||
"""Create UserInfoResponse while filtering sensitive fields."""
|
||||
if user_info is None and keys is not None:
|
||||
|
|
@ -749,7 +751,7 @@ def _build_user_info_response(
|
|||
@management_endpoint_wrapper
|
||||
async def user_info(
|
||||
request: Request,
|
||||
user_id: Optional[str] = fastapi.Query(default=None, description="User ID in the request parameters"),
|
||||
user_id: str | None = fastapi.Query(default=None, description="User ID in the request parameters"),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
|
|
@ -886,7 +888,7 @@ async def _check_user_info_v2_access(
|
|||
@management_endpoint_wrapper
|
||||
async def user_info_v2(
|
||||
request: Request,
|
||||
user_id: Optional[str] = fastapi.Query(default=None, description="User ID in the request parameters"),
|
||||
user_id: str | None = fastapi.Query(default=None, description="User ID in the request parameters"),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
|
|
@ -996,7 +998,7 @@ async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth):
|
|||
|
||||
verbose_proxy_logger.debug("results_keys: %s", results)
|
||||
|
||||
_keys_in_db: List = results[0]["keys"] or []
|
||||
_keys_in_db: list = results[0]["keys"] or []
|
||||
# cast all keys to LiteLLM_VerificationToken
|
||||
keys_in_db = []
|
||||
for key in _keys_in_db:
|
||||
|
|
@ -1005,7 +1007,7 @@ async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth):
|
|||
keys_in_db.append(LiteLLM_VerificationToken(**key))
|
||||
|
||||
# cast all teams to LiteLLM_TeamTable
|
||||
_teams_in_db: List = results[0]["teams"] or []
|
||||
_teams_in_db: list = results[0]["teams"] or []
|
||||
_teams_in_db = [LiteLLM_TeamTable(**team) for team in _teams_in_db]
|
||||
_teams_in_db.sort(key=lambda x: getattr(x, "team_alias", "") or "")
|
||||
returned_keys = _process_keys_for_user_info(keys=keys_in_db, all_teams=_teams_in_db)
|
||||
|
|
@ -1032,8 +1034,8 @@ async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth):
|
|||
|
||||
|
||||
def _process_keys_for_user_info(
|
||||
keys: Optional[List[LiteLLM_VerificationToken]],
|
||||
all_teams: Optional[Union[List[LiteLLM_TeamTable], List[TeamListResponseObject]]],
|
||||
keys: list[LiteLLM_VerificationToken] | None,
|
||||
all_teams: list[LiteLLM_TeamTable] | list[TeamListResponseObject] | None,
|
||||
):
|
||||
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
|
||||
from litellm.proxy.proxy_server import general_settings, litellm_master_key_hash
|
||||
|
|
@ -1073,9 +1075,7 @@ def _process_keys_for_user_info(
|
|||
return returned_keys
|
||||
|
||||
|
||||
def _update_internal_user_params(
|
||||
data_json: dict, data: Union[UpdateUserRequest, UpdateUserRequestNoUserIDorEmail]
|
||||
) -> dict:
|
||||
def _update_internal_user_params(data_json: dict, data: UpdateUserRequest | UpdateUserRequestNoUserIDorEmail) -> dict:
|
||||
non_default_values = {}
|
||||
fields_set = data.fields_set() if hasattr(data, "fields_set") else set()
|
||||
|
||||
|
|
@ -1124,11 +1124,11 @@ def _update_internal_user_params(
|
|||
|
||||
|
||||
async def _schedule_user_update_audit_log(
|
||||
response: Dict[str, Any],
|
||||
existing_user_row: Optional[BaseModel],
|
||||
litellm_changed_by: Optional[str],
|
||||
response: dict[str, Any],
|
||||
existing_user_row: BaseModel | None,
|
||||
litellm_changed_by: str | None,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
litellm_proxy_admin_name: Optional[str],
|
||||
litellm_proxy_admin_name: str | None,
|
||||
) -> None:
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
|
|
@ -1156,7 +1156,7 @@ async def _schedule_user_update_audit_log(
|
|||
def _check_user_update_authz(
|
||||
user_request: UpdateUserRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
existing_user_row: Optional[BaseModel],
|
||||
existing_user_row: BaseModel | None,
|
||||
) -> None:
|
||||
"""Authorization checks for /user/update — raises HTTPException on failure."""
|
||||
if user_request.user_role is not None and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value:
|
||||
|
|
@ -1201,8 +1201,8 @@ async def _invalidate_user_spend_counter_if_changed(
|
|||
async def _update_single_user_helper(
|
||||
user_request: UpdateUserRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
litellm_changed_by: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
litellm_changed_by: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Helper function to update a single user.
|
||||
Used by both user_update and bulk_user_update endpoints.
|
||||
|
|
@ -1226,7 +1226,7 @@ async def _update_single_user_helper(
|
|||
non_default_values = _update_internal_user_params(data_json=data_json, data=user_request)
|
||||
_hash_password_in_dict(non_default_values)
|
||||
|
||||
existing_user_row: Optional[BaseModel] = None
|
||||
existing_user_row: BaseModel | None = None
|
||||
if user_request.user_id:
|
||||
existing_user_row = await UserRepository(prisma_client).table.find_first(
|
||||
where={"user_id": user_request.user_id}
|
||||
|
|
@ -1261,7 +1261,7 @@ async def _update_single_user_helper(
|
|||
)
|
||||
|
||||
existing_metadata = (
|
||||
cast(Dict, getattr(existing_user_row, "metadata", {}) or {}) if existing_user_row is not None else {}
|
||||
cast(dict, getattr(existing_user_row, "metadata", {}) or {}) if existing_user_row is not None else {}
|
||||
)
|
||||
|
||||
non_default_values = prepare_metadata_fields(
|
||||
|
|
@ -1274,7 +1274,7 @@ async def _update_single_user_helper(
|
|||
validate_finite_spend(non_default_values.get("spend"))
|
||||
|
||||
# Perform the update
|
||||
response: Optional[Dict[str, Any]] = None
|
||||
response: dict[str, Any] | None = None
|
||||
|
||||
if user_request.user_id and len(user_request.user_id) > 0:
|
||||
non_default_values["user_id"] = user_request.user_id
|
||||
|
|
@ -1434,11 +1434,11 @@ async def user_update(
|
|||
|
||||
|
||||
async def bulk_update_processed_users(
|
||||
users_to_update: List[UpdateUserRequest],
|
||||
users_to_update: list[UpdateUserRequest],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
litellm_changed_by: Optional[str] = None,
|
||||
litellm_changed_by: str | None = None,
|
||||
) -> BulkUpdateUserResponse:
|
||||
results: List[UserUpdateResult] = []
|
||||
results: list[UserUpdateResult] = []
|
||||
successful_updates = 0
|
||||
failed_updates = 0
|
||||
|
||||
|
|
@ -1502,7 +1502,7 @@ async def bulk_update_processed_users(
|
|||
async def bulk_user_update(
|
||||
data: BulkUpdateUserRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
litellm_changed_by: Optional[str] = Header(
|
||||
litellm_changed_by: str | None = Header(
|
||||
None,
|
||||
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
|
||||
),
|
||||
|
|
@ -1578,7 +1578,7 @@ async def bulk_user_update(
|
|||
)
|
||||
|
||||
# Determine the list of users to update
|
||||
users_to_update: Union[List[UpdateUserRequest], List[UpdateUserRequestNoUserIDorEmail]] = []
|
||||
users_to_update: list[UpdateUserRequest] | list[UpdateUserRequestNoUserIDorEmail] = []
|
||||
|
||||
if data.all_users and data.user_updates:
|
||||
# Only proxy admins can update all users at once
|
||||
|
|
@ -1616,7 +1616,7 @@ async def bulk_user_update(
|
|||
|
||||
successful_updates = 0
|
||||
failed_updates = 0
|
||||
results: List[UserUpdateResult] = []
|
||||
results: list[UserUpdateResult] = []
|
||||
|
||||
try:
|
||||
# Perform bulk database update
|
||||
|
|
@ -1696,7 +1696,7 @@ async def bulk_user_update(
|
|||
)
|
||||
|
||||
return await bulk_update_processed_users(
|
||||
users_to_update=cast(List[UpdateUserRequest], users_to_update),
|
||||
users_to_update=cast(list[UpdateUserRequest], users_to_update),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
)
|
||||
|
|
@ -1704,7 +1704,7 @@ async def bulk_user_update(
|
|||
|
||||
async def get_user_key_counts(
|
||||
prisma_client,
|
||||
user_ids: Optional[List[str]] = None,
|
||||
user_ids: list[str] | None = None,
|
||||
):
|
||||
"""
|
||||
Helper function to get the count of keys for each user using Prisma's count method.
|
||||
|
|
@ -1739,8 +1739,8 @@ async def get_user_key_counts(
|
|||
return result
|
||||
|
||||
|
||||
def _validate_sort_params(sort_by: Optional[str], sort_order: str) -> Optional[Dict[str, str]]:
|
||||
order_by: Dict[str, str] = {}
|
||||
def _validate_sort_params(sort_by: str | None, sort_order: str) -> dict[str, str] | None:
|
||||
order_by: dict[str, str] = {}
|
||||
|
||||
if sort_by is None:
|
||||
return None
|
||||
|
|
@ -1773,11 +1773,11 @@ def _validate_sort_params(sort_by: Optional[str], sort_order: str) -> Optional[D
|
|||
|
||||
async def _authorize_user_list_request(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
organization_ids: Optional[str],
|
||||
organization_ids: str | None,
|
||||
prisma_client: Any,
|
||||
user_api_key_cache: Any,
|
||||
proxy_logging_obj: Any,
|
||||
) -> Optional[str]:
|
||||
) -> str | None:
|
||||
"""
|
||||
Authorize the /user/list request and return the (possibly scoped) organization_ids string.
|
||||
|
||||
|
|
@ -1844,19 +1844,19 @@ async def _authorize_user_list_request(
|
|||
response_model=UserListResponse,
|
||||
)
|
||||
async def get_users(
|
||||
role: Optional[str] = fastapi.Query(default=None, description="Filter users by role"),
|
||||
user_ids: Optional[str] = fastapi.Query(default=None, description="Get list of users by user_ids"),
|
||||
sso_user_ids: Optional[str] = fastapi.Query(default=None, description="Get list of users by sso_user_id"),
|
||||
user_email: Optional[str] = fastapi.Query(default=None, description="Filter users by partial email match"),
|
||||
team: Optional[str] = fastapi.Query(default=None, description="Filter users by team id"),
|
||||
role: str | None = fastapi.Query(default=None, description="Filter users by role"),
|
||||
user_ids: str | None = fastapi.Query(default=None, description="Get list of users by user_ids"),
|
||||
sso_user_ids: str | None = fastapi.Query(default=None, description="Get list of users by sso_user_id"),
|
||||
user_email: str | None = fastapi.Query(default=None, description="Filter users by partial email match"),
|
||||
team: str | None = fastapi.Query(default=None, description="Filter users by team id"),
|
||||
page: int = fastapi.Query(default=1, ge=1, description="Page number"),
|
||||
page_size: int = fastapi.Query(default=25, ge=1, le=100, description="Number of items per page"),
|
||||
sort_by: Optional[str] = fastapi.Query(
|
||||
sort_by: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Column to sort by (e.g. 'user_id', 'user_email', 'created_at', 'spend')",
|
||||
),
|
||||
sort_order: str = fastapi.Query(default="asc", description="Sort order ('asc' or 'desc')"),
|
||||
organization_ids: Optional[str] = fastapi.Query(
|
||||
organization_ids: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Filter users by organization membership. Comma-separated list of org IDs.",
|
||||
),
|
||||
|
|
@ -1914,7 +1914,7 @@ async def get_users(
|
|||
skip = (page - 1) * page_size
|
||||
|
||||
# Build where conditions based on provided parameters
|
||||
where_conditions: Dict[str, Any] = {}
|
||||
where_conditions: dict[str, Any] = {}
|
||||
|
||||
if role:
|
||||
where_conditions["user_role"] = role
|
||||
|
|
@ -1958,7 +1958,7 @@ async def get_users(
|
|||
|
||||
# Build order_by conditions
|
||||
|
||||
order_by: Optional[Dict[str, str]] = (
|
||||
order_by: dict[str, str] | None = (
|
||||
_validate_sort_params(sort_by, sort_order) if sort_by is not None and isinstance(sort_by, str) else None
|
||||
)
|
||||
|
||||
|
|
@ -1984,7 +1984,7 @@ async def get_users(
|
|||
total_pages = -(-total_count // page_size) # Ceiling division
|
||||
|
||||
# Prepare response
|
||||
user_list: List[LiteLLM_UserTableWithKeyCount] = []
|
||||
user_list: list[LiteLLM_UserTableWithKeyCount] = []
|
||||
if users is not None:
|
||||
for user in users:
|
||||
user_dump = user.model_dump()
|
||||
|
|
@ -2011,7 +2011,7 @@ async def get_users(
|
|||
async def delete_user(
|
||||
data: DeleteUserRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
litellm_changed_by: Optional[str] = Header(
|
||||
litellm_changed_by: str | None = Header(
|
||||
None,
|
||||
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
|
||||
),
|
||||
|
|
@ -2080,7 +2080,7 @@ async def delete_user(
|
|||
|
||||
# Batch-fetch target memberships once before the per-user loop. Avoids
|
||||
# an N+1 DB call when delete_user is called with a large user_ids list.
|
||||
target_org_ids_by_user: Dict[str, set] = {}
|
||||
target_org_ids_by_user: dict[str, set] = {}
|
||||
if not caller_is_proxy_admin:
|
||||
all_target_memberships = await OrganizationMembershipRepository(prisma_client).table.find_many(
|
||||
where={"user_id": {"in": data.user_ids}}
|
||||
|
|
@ -2156,7 +2156,7 @@ async def delete_user(
|
|||
),
|
||||
)
|
||||
if is_member_in_team:
|
||||
_db_new_team_members: List[dict] = [m.model_dump() for m in new_team_members]
|
||||
_db_new_team_members: list[dict] = [m.model_dump() for m in new_team_members]
|
||||
team.members_with_roles = json.dumps(_db_new_team_members)
|
||||
teams_to_update.append(team)
|
||||
|
||||
|
|
@ -2241,11 +2241,11 @@ async def add_internal_user_to_organization(
|
|||
|
||||
async def _resolve_org_filter_for_user_search(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
team_id: Optional[str],
|
||||
team_id: str | None,
|
||||
prisma_client: Any,
|
||||
user_api_key_cache: Any,
|
||||
proxy_logging_obj: Any,
|
||||
) -> Optional[List[str]]:
|
||||
) -> list[str] | None:
|
||||
"""
|
||||
Return a list of org IDs to filter by, or ``None`` for no filter.
|
||||
|
||||
|
|
@ -2279,7 +2279,7 @@ async def _resolve_org_filter_for_user_search(
|
|||
|
||||
# Collect org IDs from ALL org memberships (any role, not just ORG_ADMIN).
|
||||
# This allows team admins who are org members to search users in their org.
|
||||
member_org_ids: List[str] = []
|
||||
member_org_ids: list[str] = []
|
||||
if caller_user is not None:
|
||||
member_org_ids = [m.organization_id for m in (caller_user.organization_memberships or [])]
|
||||
|
||||
|
|
@ -2311,7 +2311,7 @@ async def _resolve_team_org_filter(
|
|||
prisma_client: Any,
|
||||
user_api_key_cache: Any,
|
||||
proxy_logging_obj: Any,
|
||||
) -> List[str]:
|
||||
) -> list[str]:
|
||||
"""Look up the team and return its org as a filter list, or raise 403."""
|
||||
from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin
|
||||
|
||||
|
|
@ -2351,13 +2351,13 @@ async def _resolve_team_org_filter(
|
|||
dependencies=[Depends(user_api_key_auth)],
|
||||
include_in_schema=False,
|
||||
responses={
|
||||
200: {"model": List[LiteLLM_UserTableFiltered]},
|
||||
200: {"model": list[LiteLLM_UserTableFiltered]},
|
||||
},
|
||||
)
|
||||
async def ui_view_users(
|
||||
user_id: Optional[str] = fastapi.Query(default=None, description="User ID in the request parameters"),
|
||||
user_email: Optional[str] = fastapi.Query(default=None, description="User email in the request parameters"),
|
||||
team_id: Optional[str] = fastapi.Query(
|
||||
user_id: str | None = fastapi.Query(default=None, description="User ID in the request parameters"),
|
||||
user_email: str | None = fastapi.Query(default=None, description="User email in the request parameters"),
|
||||
team_id: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Team ID — used when a team admin searches for users to add to their team",
|
||||
),
|
||||
|
|
@ -2400,7 +2400,7 @@ async def ui_view_users(
|
|||
skip = (page - 1) * page_size
|
||||
|
||||
# Build where conditions based on provided parameters
|
||||
where_conditions: Dict[str, Any] = {}
|
||||
where_conditions: dict[str, Any] = {}
|
||||
|
||||
if user_id:
|
||||
where_conditions["user_id"] = {
|
||||
|
|
@ -2419,7 +2419,7 @@ async def ui_view_users(
|
|||
where_conditions["organization_memberships"] = {"some": {"organization_id": {"in": org_filter_ids}}}
|
||||
|
||||
# Query users with pagination and filters
|
||||
users: Optional[List[BaseModel]] = await UserRepository(prisma_client).table.find_many(
|
||||
users: list[BaseModel] | None = await UserRepository(prisma_client).table.find_many(
|
||||
where=where_conditions,
|
||||
skip=skip,
|
||||
take=page_size,
|
||||
|
|
@ -2441,10 +2441,14 @@ async def ui_view_users(
|
|||
# Using shared metric helper implementations from common_daily_activity
|
||||
|
||||
|
||||
async def _resolve_user_email_metadata(prisma_client: "PrismaClient", records: list[Any]) -> dict[str, dict]:
|
||||
async def _resolve_user_email_metadata(
|
||||
prisma_client: "PrismaClient", records: Sequence[DailySpendRecord]
|
||||
) -> dict[str, dict]:
|
||||
"""Map each user_id on the page to its email/alias so the Usage dashboard can
|
||||
label the 'Spend Per User' chart with the email instead of the raw UUID."""
|
||||
user_ids = {record.user_id for record in records if getattr(record, "user_id", None)}
|
||||
user_ids = {
|
||||
user_id for record in records if isinstance(user_id := getattr(record, "user_id", None), str) and user_id
|
||||
}
|
||||
if not user_ids:
|
||||
return {}
|
||||
users = await UserRepository(prisma_client).table.find_many(where={"user_id": {"in": list(user_ids)}})
|
||||
|
|
@ -2459,29 +2463,29 @@ async def _resolve_user_email_metadata(prisma_client: "PrismaClient", records: l
|
|||
)
|
||||
@management_endpoint_wrapper
|
||||
async def get_user_daily_activity(
|
||||
start_date: Optional[str] = fastapi.Query(
|
||||
start_date: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Start date in YYYY-MM-DD format",
|
||||
),
|
||||
end_date: Optional[str] = fastapi.Query(
|
||||
end_date: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="End date in YYYY-MM-DD format",
|
||||
),
|
||||
model: Optional[str] = fastapi.Query(
|
||||
model: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Filter by specific model",
|
||||
),
|
||||
api_key: Optional[str] = fastapi.Query(
|
||||
api_key: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Filter by specific API key",
|
||||
),
|
||||
user_id: Optional[str] = fastapi.Query(
|
||||
user_id: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Filter by specific user ID. Admins can filter by any user or omit for global view. Non-admins must provide their own user_id.",
|
||||
),
|
||||
page: int = fastapi.Query(default=1, description="Page number for pagination", ge=1),
|
||||
page_size: int = fastapi.Query(default=50, description="Items per page", ge=1, le=1000),
|
||||
timezone: Optional[int] = fastapi.Query(
|
||||
timezone: int | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Timezone offset in minutes from UTC (e.g., 480 for PST). "
|
||||
"Matches JavaScript's Date.getTimezoneOffset() convention.",
|
||||
|
|
@ -2568,27 +2572,27 @@ async def get_user_daily_activity(
|
|||
)
|
||||
@management_endpoint_wrapper
|
||||
async def get_user_daily_activity_aggregated(
|
||||
start_date: Optional[str] = fastapi.Query(
|
||||
start_date: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Start date in YYYY-MM-DD format",
|
||||
),
|
||||
end_date: Optional[str] = fastapi.Query(
|
||||
end_date: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="End date in YYYY-MM-DD format",
|
||||
),
|
||||
model: Optional[str] = fastapi.Query(
|
||||
model: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Filter by specific model",
|
||||
),
|
||||
api_key: Optional[str] = fastapi.Query(
|
||||
api_key: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Filter by specific API key",
|
||||
),
|
||||
user_id: Optional[str] = fastapi.Query(
|
||||
user_id: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Filter by specific user ID. Admins can filter by any user or omit for global view. Non-admins must provide their own user_id.",
|
||||
),
|
||||
timezone: Optional[int] = fastapi.Query(
|
||||
timezone: int | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Timezone offset in minutes from UTC (e.g., 480 for PST). "
|
||||
"Matches JavaScript's Date.getTimezoneOffset() convention.",
|
||||
|
|
|
|||
|
|
@ -19,9 +19,10 @@ import functools
|
|||
import importlib
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Dict, Iterable, List, Literal, Optional, Set
|
||||
from typing import Any, Literal
|
||||
|
||||
from fastapi import (
|
||||
APIRouter,
|
||||
|
|
@ -47,10 +48,10 @@ from litellm._logging import verbose_logger, verbose_proxy_logger
|
|||
from litellm._uuid import uuid
|
||||
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
build_env_var_setup_url,
|
||||
collect_env_var_references,
|
||||
LITELLM_MCP_SERVER_DESCRIPTION,
|
||||
LITELLM_MCP_SERVER_NAME,
|
||||
build_env_var_setup_url,
|
||||
collect_env_var_references,
|
||||
get_server_prefix,
|
||||
parse_admin_env_vars,
|
||||
)
|
||||
|
|
@ -194,7 +195,7 @@ if MCP_AVAILABLE:
|
|||
expires_at: datetime
|
||||
|
||||
def _validate_mcp_server_name_fields(payload: Any) -> None:
|
||||
candidates: List[tuple[str, Optional[str]]] = []
|
||||
candidates: list[tuple[str, str | None]] = []
|
||||
|
||||
server_name = getattr(payload, "server_name", None)
|
||||
alias = getattr(payload, "alias", None)
|
||||
|
|
@ -260,7 +261,7 @@ if MCP_AVAILABLE:
|
|||
general_settings as proxy_general_settings,
|
||||
)
|
||||
|
||||
required_fields: Optional[List[str]] = proxy_general_settings.get("mcp_required_fields")
|
||||
required_fields: list[str] | None = proxy_general_settings.get("mcp_required_fields")
|
||||
if not required_fields:
|
||||
return
|
||||
|
||||
|
|
@ -320,7 +321,7 @@ if MCP_AVAILABLE:
|
|||
return server.server_name
|
||||
return server.server_id
|
||||
|
||||
def _build_mcp_registry_entry_for_server(server: MCPServer, base_url: str) -> Dict[str, Any]:
|
||||
def _build_mcp_registry_entry_for_server(server: MCPServer, base_url: str) -> dict[str, Any]:
|
||||
server_name = _build_mcp_registry_server_name(server)
|
||||
title = server_name
|
||||
description = server_name
|
||||
|
|
@ -344,7 +345,7 @@ if MCP_AVAILABLE:
|
|||
],
|
||||
}
|
||||
|
||||
def _build_builtin_registry_entry(base_url: str) -> Dict[str, Any]:
|
||||
def _build_builtin_registry_entry(base_url: str) -> dict[str, Any]:
|
||||
remote_url = _build_registry_remote_url(base_url, "/mcp")
|
||||
return {
|
||||
"name": LITELLM_MCP_SERVER_NAME,
|
||||
|
|
@ -359,7 +360,7 @@ if MCP_AVAILABLE:
|
|||
],
|
||||
}
|
||||
|
||||
_temporary_mcp_servers: Dict[str, _TemporaryMCPServerEntry] = {}
|
||||
_temporary_mcp_servers: dict[str, _TemporaryMCPServerEntry] = {}
|
||||
|
||||
def _prune_expired_temporary_mcp_servers() -> None:
|
||||
if not _temporary_mcp_servers:
|
||||
|
|
@ -391,7 +392,7 @@ if MCP_AVAILABLE:
|
|||
if cache_backend is None or not hasattr(cache_backend, "async_set_cache"):
|
||||
return
|
||||
|
||||
payload: Dict[str, Any] = server.model_dump(mode="json")
|
||||
payload: dict[str, Any] = server.model_dump(mode="json")
|
||||
payload_json = json.dumps(payload)
|
||||
try:
|
||||
encrypted_payload = encrypt_value_helper(payload_json)
|
||||
|
|
@ -414,7 +415,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
async def _get_temporary_mcp_server_from_redis(
|
||||
server_id: str,
|
||||
) -> Optional[MCPServer]:
|
||||
) -> MCPServer | None:
|
||||
"""
|
||||
Best-effort read from Redis shared cache. Returns None on miss/errors.
|
||||
|
||||
|
|
@ -455,7 +456,7 @@ if MCP_AVAILABLE:
|
|||
return None
|
||||
if not isinstance(loaded, dict):
|
||||
return None
|
||||
payload_dict: Dict[str, Any] = loaded
|
||||
payload_dict: dict[str, Any] = loaded
|
||||
|
||||
try:
|
||||
return MCPServer(**payload_dict)
|
||||
|
|
@ -465,7 +466,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
async def get_cached_temporary_mcp_server(
|
||||
server_id: str,
|
||||
) -> Optional[MCPServer]:
|
||||
) -> MCPServer | None:
|
||||
_prune_expired_temporary_mcp_servers()
|
||||
entry = _temporary_mcp_servers.get(server_id)
|
||||
if entry is None:
|
||||
|
|
@ -520,7 +521,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
def _redact_mcp_credentials_list(
|
||||
mcp_servers: Iterable[LiteLLM_MCPServerTable],
|
||||
) -> List[LiteLLM_MCPServerTable]:
|
||||
) -> list[LiteLLM_MCPServerTable]:
|
||||
return [_redact_mcp_credentials(server) for server in mcp_servers]
|
||||
|
||||
def _user_is_full_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
|
||||
|
|
@ -587,7 +588,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
def _sanitize_mcp_server_list_for_non_admin(
|
||||
mcp_servers: Iterable[LiteLLM_MCPServerTable],
|
||||
) -> List[LiteLLM_MCPServerTable]:
|
||||
) -> list[LiteLLM_MCPServerTable]:
|
||||
return [_sanitize_mcp_server_for_non_admin(s) for s in mcp_servers]
|
||||
|
||||
def _sanitize_mcp_server_for_virtual_key(
|
||||
|
|
@ -644,7 +645,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
def _sanitize_mcp_server_list_for_virtual_key(
|
||||
mcp_servers: Iterable[LiteLLM_MCPServerTable],
|
||||
) -> List[LiteLLM_MCPServerTable]:
|
||||
) -> list[LiteLLM_MCPServerTable]:
|
||||
return [_sanitize_mcp_server_for_virtual_key(server) for server in mcp_servers]
|
||||
|
||||
# (server attribute, credentials key) a session server inherits from the server it derives from.
|
||||
|
|
@ -697,7 +698,7 @@ if MCP_AVAILABLE:
|
|||
except AttributeError:
|
||||
pass
|
||||
|
||||
payload_dict: Dict[str, Any]
|
||||
payload_dict: dict[str, Any]
|
||||
try:
|
||||
payload_dict = payload.model_dump() # type: ignore[attr-defined]
|
||||
except AttributeError:
|
||||
|
|
@ -707,7 +708,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
def _build_temporary_mcp_server_record(
|
||||
payload: NewMCPServerRequest,
|
||||
created_by: Optional[str],
|
||||
created_by: str | None,
|
||||
) -> LiteLLM_MCPServerTable:
|
||||
now = datetime.utcnow()
|
||||
server_id = payload.server_id or str(uuid.uuid4())
|
||||
|
|
@ -848,7 +849,7 @@ if MCP_AVAILABLE:
|
|||
verbose_proxy_logger.debug("MCP registry request from IP=%s", client_ip)
|
||||
|
||||
base_url = get_request_base_url(request)
|
||||
registry_servers: List[Dict[str, Any]] = []
|
||||
registry_servers: list[dict[str, Any]] = []
|
||||
registry_servers.append({"server": _build_builtin_registry_entry(base_url)})
|
||||
|
||||
# Centralized IP-based filtering: external callers only see public servers
|
||||
|
|
@ -881,7 +882,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
async def _get_team_scoped_mcp_server_list(
|
||||
team_id: str,
|
||||
) -> List[LiteLLM_MCPServerTable]:
|
||||
) -> list[LiteLLM_MCPServerTable]:
|
||||
"""
|
||||
Return MCP servers scoped to a team: team's allowed servers + allow_all_keys servers.
|
||||
Used by the Create Key UI to populate the MCP server dropdown.
|
||||
|
|
@ -908,7 +909,7 @@ if MCP_AVAILABLE:
|
|||
return []
|
||||
|
||||
# Collect servers from registry
|
||||
servers: List[LiteLLM_MCPServerTable] = []
|
||||
servers: list[LiteLLM_MCPServerTable] = []
|
||||
for server_id in all_allowed_ids:
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
|
||||
if server is not None:
|
||||
|
|
@ -919,7 +920,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
async def _resolve_accessible_mcp_servers(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> List[LiteLLM_MCPServerTable]:
|
||||
) -> list[LiteLLM_MCPServerTable]:
|
||||
"""The server set the dashboard grid shows (GET /v1/mcp/server, no team
|
||||
filter), returned unredacted. Callers that surface this to a client must
|
||||
apply their own redaction; the per-user env-var status endpoint relies on
|
||||
|
|
@ -932,7 +933,7 @@ if MCP_AVAILABLE:
|
|||
if _get_user_mcp_management_mode() == "view_all" and not _is_restricted_virtual_key_request(user_api_key_dict):
|
||||
return await global_mcp_server_manager.get_all_mcp_servers_unfiltered()
|
||||
|
||||
aggregated: Dict[str, LiteLLM_MCPServerTable] = {}
|
||||
aggregated: dict[str, LiteLLM_MCPServerTable] = {}
|
||||
for auth_context in await build_effective_auth_contexts(user_api_key_dict):
|
||||
for server in await global_mcp_server_manager.get_all_allowed_mcp_servers(user_api_key_auth=auth_context):
|
||||
aggregated.setdefault(server.server_id, server)
|
||||
|
|
@ -942,11 +943,11 @@ if MCP_AVAILABLE:
|
|||
"/server",
|
||||
description="Returns the mcp server list with associated teams",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[LiteLLM_MCPServerTable],
|
||||
response_model=list[LiteLLM_MCPServerTable],
|
||||
)
|
||||
async def fetch_all_mcp_servers(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
team_id: Optional[str] = Query(
|
||||
team_id: str | None = Query(
|
||||
None,
|
||||
description="Filter MCP servers by team scope. When provided, returns only "
|
||||
"servers the team has access to plus globally available (allow_all_keys) servers. "
|
||||
|
|
@ -1048,7 +1049,7 @@ if MCP_AVAILABLE:
|
|||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def health_check_servers(
|
||||
server_ids: Optional[List[str]] = Query(
|
||||
server_ids: list[str] | None = Query(
|
||||
None,
|
||||
description="Server IDs to check. If not provided, checks all accessible servers.",
|
||||
),
|
||||
|
|
@ -1081,7 +1082,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
|
||||
server_status_map: Dict[str, Optional[Literal["healthy", "unhealthy", "unknown"]]] = {}
|
||||
server_status_map: dict[str, Literal["healthy", "unhealthy", "unknown"] | None] = {}
|
||||
for auth_context in auth_contexts:
|
||||
servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_and_teams(
|
||||
user_api_key_auth=auth_context,
|
||||
|
|
@ -1399,7 +1400,7 @@ if MCP_AVAILABLE:
|
|||
async def add_mcp_server(
|
||||
payload: NewMCPServerRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
litellm_changed_by: Optional[str] = Header(
|
||||
litellm_changed_by: str | None = Header(
|
||||
None,
|
||||
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
|
||||
),
|
||||
|
|
@ -1489,7 +1490,7 @@ if MCP_AVAILABLE:
|
|||
async def add_session_mcp_server(
|
||||
payload: NewMCPServerRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
litellm_changed_by: Optional[str] = Header(
|
||||
litellm_changed_by: str | None = Header(
|
||||
None,
|
||||
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
|
||||
),
|
||||
|
|
@ -1647,7 +1648,7 @@ if MCP_AVAILABLE:
|
|||
async def _get_cached_temporary_mcp_server_or_404(
|
||||
server_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request: Optional[Request] = None,
|
||||
request: Request | None = None,
|
||||
) -> MCPServer:
|
||||
server = await get_cached_temporary_mcp_server(server_id)
|
||||
resolved_from_temp_cache = server is not None
|
||||
|
|
@ -1677,7 +1678,7 @@ if MCP_AVAILABLE:
|
|||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={"error": f"Access denied to MCP server {server_id}"},
|
||||
)
|
||||
allowed_server_ids: Set[str] = set()
|
||||
allowed_server_ids: set[str] = set()
|
||||
for auth_context in await build_effective_auth_contexts(user_api_key_dict):
|
||||
allowed_server_ids.update(await global_mcp_server_manager.get_allowed_mcp_servers(auth_context))
|
||||
if server.server_id not in allowed_server_ids:
|
||||
|
|
@ -1696,13 +1697,13 @@ if MCP_AVAILABLE:
|
|||
request: Request,
|
||||
server_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(_mcp_oauth_user_api_key_auth),
|
||||
client_id: Optional[str] = None,
|
||||
client_id: str | None = None,
|
||||
redirect_uri: str = Query(...),
|
||||
state: str = "",
|
||||
code_challenge: Optional[str] = None,
|
||||
code_challenge_method: Optional[str] = None,
|
||||
response_type: Optional[str] = None,
|
||||
scope: Optional[str] = None,
|
||||
code_challenge: str | None = None,
|
||||
code_challenge_method: str | None = None,
|
||||
response_type: str | None = None,
|
||||
scope: str | None = None,
|
||||
):
|
||||
mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
|
|
@ -1756,13 +1757,13 @@ if MCP_AVAILABLE:
|
|||
server_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(_mcp_oauth_user_api_key_auth),
|
||||
grant_type: str = Form(...),
|
||||
code: Optional[str] = Form(None),
|
||||
redirect_uri: Optional[str] = Form(None),
|
||||
client_id: Optional[str] = Form(None),
|
||||
client_secret: Optional[str] = Form(None),
|
||||
code_verifier: Optional[str] = Form(None),
|
||||
refresh_token: Optional[str] = Form(None),
|
||||
scope: Optional[str] = Form(None),
|
||||
code: str | None = Form(None),
|
||||
redirect_uri: str | None = Form(None),
|
||||
client_id: str | None = Form(None),
|
||||
client_secret: str | None = Form(None),
|
||||
code_verifier: str | None = Form(None),
|
||||
refresh_token: str | None = Form(None),
|
||||
scope: str | None = Form(None),
|
||||
):
|
||||
mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
|
|
@ -1844,7 +1845,7 @@ if MCP_AVAILABLE:
|
|||
async def remove_mcp_server(
|
||||
server_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
litellm_changed_by: Optional[str] = Header(
|
||||
litellm_changed_by: str | None = Header(
|
||||
None,
|
||||
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
|
||||
),
|
||||
|
|
@ -2007,7 +2008,7 @@ if MCP_AVAILABLE:
|
|||
# expires_at rather than recomputing it here (which could diverge by
|
||||
# milliseconds or if the storage logic ever adds a grace period).
|
||||
stored = await get_user_oauth_credential(prisma_client, user_id, server_id)
|
||||
expires_at: Optional[str] = stored.get("expires_at") if stored else None
|
||||
expires_at: str | None = stored.get("expires_at") if stored else None
|
||||
return MCPOAuthUserCredentialStatus(
|
||||
server_id=server_id,
|
||||
has_credential=True,
|
||||
|
|
@ -2076,7 +2077,7 @@ if MCP_AVAILABLE:
|
|||
cred = await get_user_oauth_credential(prisma_client, user_id, server_id)
|
||||
if cred is None:
|
||||
return MCPOAuthUserCredentialStatus(server_id=server_id, has_credential=False, is_expired=False)
|
||||
expires_at: Optional[str] = cred.get("expires_at")
|
||||
expires_at: str | None = cred.get("expires_at")
|
||||
is_expired = False
|
||||
if expires_at:
|
||||
try:
|
||||
|
|
@ -2096,7 +2097,7 @@ if MCP_AVAILABLE:
|
|||
"/user-credentials",
|
||||
description="List all OAuth2 MCP credentials stored for the calling user",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[MCPUserCredentialListItem],
|
||||
response_model=list[MCPUserCredentialListItem],
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def list_mcp_user_credentials(
|
||||
|
|
@ -2114,13 +2115,15 @@ if MCP_AVAILABLE:
|
|||
if not oauth_creds:
|
||||
return []
|
||||
# Fetch server metadata for display names — single batch query instead of N+1.
|
||||
server_ids = [c["server_id"] for c in oauth_creds]
|
||||
server_ids = [c["server_id"] for c in oauth_creds if "server_id" in c]
|
||||
servers = {srv.server_id: srv for srv in await get_mcp_servers(prisma_client, server_ids)}
|
||||
items: List[MCPUserCredentialListItem] = []
|
||||
items: list[MCPUserCredentialListItem] = []
|
||||
for cred in oauth_creds:
|
||||
if "server_id" not in cred:
|
||||
continue
|
||||
sid = cred["server_id"]
|
||||
srv = servers.get(sid)
|
||||
expires_at: Optional[str] = cred.get("expires_at")
|
||||
expires_at: str | None = cred.get("expires_at")
|
||||
items.append(
|
||||
MCPUserCredentialListItem(
|
||||
server_id=sid,
|
||||
|
|
@ -2182,7 +2185,7 @@ if MCP_AVAILABLE:
|
|||
def _compute_user_env_var_status(
|
||||
*,
|
||||
server: LiteLLM_MCPServerTable,
|
||||
stored_values: Dict[str, str],
|
||||
stored_values: dict[str, str],
|
||||
) -> MCPUserEnvVarsStatus:
|
||||
"""Build a status object for one server given the user's stored values.
|
||||
|
||||
|
|
@ -2211,7 +2214,7 @@ if MCP_AVAILABLE:
|
|||
user_var_names = {spec["name"] for spec in user_specs}
|
||||
blocking = {name for name in (referenced & user_var_names) if name not in global_values}
|
||||
|
||||
required: List[MCPUserEnvVarSpec] = []
|
||||
required: list[MCPUserEnvVarSpec] = []
|
||||
missing_count = 0
|
||||
for spec in user_specs:
|
||||
name = spec["name"]
|
||||
|
|
@ -2334,12 +2337,12 @@ if MCP_AVAILABLE:
|
|||
description="Per-user MCP env var status across every server the user can access. "
|
||||
"Used by the dashboard to highlight servers with missing per-user vars.",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[MCPUserEnvVarsStatus],
|
||||
response_model=list[MCPUserEnvVarsStatus],
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def list_mcp_user_env_var_status(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
) -> List[MCPUserEnvVarsStatus]:
|
||||
) -> list[MCPUserEnvVarsStatus]:
|
||||
prisma_client = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
|
||||
user_id = user_api_key_dict.user_id or ""
|
||||
if not user_id:
|
||||
|
|
@ -2349,7 +2352,7 @@ if MCP_AVAILABLE:
|
|||
return []
|
||||
server_ids = [s.server_id for s in accessible]
|
||||
stored_bulk = await get_user_env_vars_bulk(prisma_client, user_id, server_ids)
|
||||
statuses: List[MCPUserEnvVarsStatus] = []
|
||||
statuses: list[MCPUserEnvVarsStatus] = []
|
||||
for server in accessible:
|
||||
stored = stored_bulk.get(server.server_id, {})
|
||||
status_obj = _compute_user_env_var_status(server=server, stored_values=stored)
|
||||
|
|
@ -2368,7 +2371,7 @@ if MCP_AVAILABLE:
|
|||
async def edit_mcp_server(
|
||||
payload: UpdateMCPServerRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
litellm_changed_by: Optional[str] = Header(
|
||||
litellm_changed_by: str | None = Header(
|
||||
None,
|
||||
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
|
||||
),
|
||||
|
|
@ -2564,16 +2567,16 @@ if MCP_AVAILABLE:
|
|||
"mcp_registry.json",
|
||||
)
|
||||
|
||||
_mcp_registry_cache: Optional[Dict[str, Any]] = None
|
||||
_mcp_registry_cache: dict[str, Any] | None = None
|
||||
|
||||
def _load_mcp_registry() -> Dict[str, Any]:
|
||||
def _load_mcp_registry() -> dict[str, Any]:
|
||||
"""Load the curated MCP registry from disk. Cached after first read."""
|
||||
global _mcp_registry_cache
|
||||
if _mcp_registry_cache is not None:
|
||||
return _mcp_registry_cache
|
||||
try:
|
||||
with open(_MCP_REGISTRY_PATH, "r") as f:
|
||||
data: Dict[str, Any] = json.load(f)
|
||||
data: dict[str, Any] = json.load(f)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(f"Failed to load MCP registry from {_MCP_REGISTRY_PATH}: {e}")
|
||||
data = {"servers": []}
|
||||
|
|
@ -2586,8 +2589,8 @@ if MCP_AVAILABLE:
|
|||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def discover_mcp_servers(
|
||||
query: Optional[str] = Query(None, description="Search filter for server names and descriptions"),
|
||||
category: Optional[str] = Query(None, description="Filter by category"),
|
||||
query: str | None = Query(None, description="Search filter for server names and descriptions"),
|
||||
category: str | None = Query(None, description="Filter by category"),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
|
|
@ -2641,9 +2644,9 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def _load_openapi_registry() -> Dict[str, Any]:
|
||||
def _load_openapi_registry() -> dict[str, Any]:
|
||||
with open(_OPENAPI_REGISTRY_PATH, "r") as f:
|
||||
data: Dict[str, Any] = json.load(f)
|
||||
data: dict[str, Any] = json.load(f)
|
||||
return data
|
||||
|
||||
@router.get(
|
||||
|
|
@ -2694,7 +2697,7 @@ if MCP_AVAILABLE:
|
|||
async def add_mcp_toolset(
|
||||
payload: NewMCPToolsetRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
litellm_changed_by: Optional[str] = Header(None),
|
||||
litellm_changed_by: str | None = Header(None),
|
||||
):
|
||||
"""Create a named toolset — a curated selection of {server_id, tool_name} pairs."""
|
||||
prisma_client = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
|
||||
|
|
@ -2783,7 +2786,7 @@ if MCP_AVAILABLE:
|
|||
async def edit_mcp_toolset(
|
||||
payload: UpdateMCPToolsetRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
litellm_changed_by: Optional[str] = Header(None),
|
||||
litellm_changed_by: str | None = Header(None),
|
||||
):
|
||||
prisma_client = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
|
||||
if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role:
|
||||
|
|
@ -2833,7 +2836,7 @@ if MCP_AVAILABLE:
|
|||
async def remove_mcp_toolset(
|
||||
toolset_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
litellm_changed_by: Optional[str] = Header(None),
|
||||
litellm_changed_by: str | None = Header(None),
|
||||
):
|
||||
prisma_client = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
|
||||
if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role:
|
||||
|
|
|
|||
|
|
@ -13,7 +13,13 @@ Endpoints for /organization operations
|
|||
|
||||
#### ORGANIZATION MANAGEMENT ####
|
||||
|
||||
from typing import Annotated, Any, Dict, List, Mapping, Optional, Tuple
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Annotated,
|
||||
Protocol,
|
||||
overload,
|
||||
)
|
||||
|
||||
import fastapi
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
|
|
@ -57,9 +63,162 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
|||
)
|
||||
from litellm.utils import _update_dictionary
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from types import TracebackType
|
||||
|
||||
from prisma.models import LiteLLM_BudgetTable as PrismaBudgetTable
|
||||
from prisma.models import (
|
||||
LiteLLM_ObjectPermissionTable as PrismaObjectPermissionTable,
|
||||
)
|
||||
from prisma.models import (
|
||||
LiteLLM_OrganizationMembership as PrismaOrganizationMembership,
|
||||
)
|
||||
from prisma.models import LiteLLM_OrganizationTable as PrismaOrganizationTable
|
||||
from prisma.models import LiteLLM_UserTable as PrismaUserTable
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class _UserTableClient(Protocol):
|
||||
async def find_unique(self, where: Mapping[str, object]) -> "PrismaUserTable | None": ...
|
||||
|
||||
|
||||
class _BudgetTableClient(Protocol):
|
||||
async def create(self, data: Mapping[str, object]) -> "PrismaBudgetTable": ...
|
||||
|
||||
|
||||
class _ObjectPermissionTableClient(Protocol):
|
||||
async def create(self, data: Mapping[str, object]) -> "PrismaObjectPermissionTable": ...
|
||||
|
||||
|
||||
class _OrganizationTableClient(Protocol):
|
||||
async def create(
|
||||
self, data: Mapping[str, object], include: Mapping[str, object] | None = None
|
||||
) -> "PrismaOrganizationTable": ...
|
||||
|
||||
async def find_unique(
|
||||
self, where: Mapping[str, object], include: Mapping[str, object] | None = None
|
||||
) -> "PrismaOrganizationTable | None": ...
|
||||
|
||||
async def find_many(
|
||||
self,
|
||||
where: Mapping[str, object] | None = None,
|
||||
include: Mapping[str, object] | None = None,
|
||||
) -> "Sequence[PrismaOrganizationTable]": ...
|
||||
|
||||
async def update(
|
||||
self,
|
||||
where: Mapping[str, object],
|
||||
data: Mapping[str, object],
|
||||
include: Mapping[str, object] | None = None,
|
||||
) -> "PrismaOrganizationTable": ...
|
||||
|
||||
async def delete(
|
||||
self, where: Mapping[str, object], include: Mapping[str, object] | None = None
|
||||
) -> "PrismaOrganizationTable | None": ...
|
||||
|
||||
|
||||
class _OrganizationMembershipTableClient(Protocol):
|
||||
async def create(self, data: Mapping[str, object]) -> "PrismaOrganizationMembership": ...
|
||||
|
||||
async def find_unique(
|
||||
self, where: Mapping[str, object], include: Mapping[str, object] | None = None
|
||||
) -> "PrismaOrganizationMembership | None": ...
|
||||
|
||||
async def find_many(
|
||||
self, where: Mapping[str, object] | None = None
|
||||
) -> "Sequence[PrismaOrganizationMembership]": ...
|
||||
|
||||
async def update(
|
||||
self, where: Mapping[str, object], data: Mapping[str, object]
|
||||
) -> "PrismaOrganizationMembership": ...
|
||||
|
||||
async def delete(self, where: Mapping[str, object]) -> "PrismaOrganizationMembership | None": ...
|
||||
|
||||
async def delete_many(self, where: Mapping[str, object]) -> int: ...
|
||||
|
||||
|
||||
class _TeamTableClient(Protocol):
|
||||
async def delete_many(self, where: Mapping[str, object]) -> int: ...
|
||||
|
||||
|
||||
class _VerificationTokenTableClient(Protocol):
|
||||
async def delete_many(self, where: Mapping[str, object]) -> int: ...
|
||||
|
||||
|
||||
class _ObjectPermissionTxClient(Protocol):
|
||||
async def upsert(
|
||||
self, where: Mapping[str, object], data: Mapping[str, object]
|
||||
) -> "PrismaObjectPermissionTable": ...
|
||||
|
||||
|
||||
class _BudgetTxClient(Protocol):
|
||||
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> "PrismaBudgetTable | None": ...
|
||||
|
||||
|
||||
class _TransactionTables(Protocol):
|
||||
@property
|
||||
def litellm_objectpermissiontable(self) -> "_ObjectPermissionTxClient": ...
|
||||
|
||||
@property
|
||||
def litellm_budgettable(self) -> "_BudgetTxClient": ...
|
||||
|
||||
@property
|
||||
def litellm_organizationtable(self) -> "_OrganizationTableClient": ...
|
||||
|
||||
|
||||
class _TransactionManager(Protocol):
|
||||
async def __aenter__(self) -> "_TransactionTables": ...
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc_value: BaseException | None,
|
||||
traceback: "TracebackType | None",
|
||||
) -> bool | None: ...
|
||||
|
||||
|
||||
@overload
|
||||
def _table(repository: BudgetRepository) -> "_BudgetTableClient": ...
|
||||
|
||||
|
||||
@overload
|
||||
def _table(repository: ObjectPermissionRepository) -> "_ObjectPermissionTableClient": ...
|
||||
|
||||
|
||||
@overload
|
||||
def _table(repository: OrganizationRepository) -> "_OrganizationTableClient": ...
|
||||
|
||||
|
||||
@overload
|
||||
def _table(repository: OrganizationMembershipRepository) -> "_OrganizationMembershipTableClient": ...
|
||||
|
||||
|
||||
@overload
|
||||
def _table(repository: TeamRepository) -> "_TeamTableClient": ...
|
||||
|
||||
|
||||
@overload
|
||||
def _table(repository: UserRepository) -> "_UserTableClient": ...
|
||||
|
||||
|
||||
@overload
|
||||
def _table(repository: VerificationTokenRepository) -> "_VerificationTokenTableClient": ...
|
||||
|
||||
|
||||
def _table(
|
||||
repository: BudgetRepository
|
||||
| ObjectPermissionRepository
|
||||
| OrganizationRepository
|
||||
| OrganizationMembershipRepository
|
||||
| TeamRepository
|
||||
| UserRepository
|
||||
| VerificationTokenRepository,
|
||||
) -> object:
|
||||
prisma_table: object = repository.table
|
||||
return prisma_table
|
||||
|
||||
|
||||
async def _verify_org_access(
|
||||
organization_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -259,14 +418,15 @@ async def new_organization(
|
|||
detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"},
|
||||
)
|
||||
|
||||
user_object_correct_type: Optional[LiteLLM_UserTable] = None
|
||||
user_object_correct_type: LiteLLM_UserTable | None = None
|
||||
|
||||
if user_api_key_dict.user_id is not None:
|
||||
try:
|
||||
user_object = await UserRepository(prisma_client).table.find_unique(
|
||||
user_object = await _table(UserRepository(prisma_client)).find_unique(
|
||||
where={"user_id": user_api_key_dict.user_id}
|
||||
)
|
||||
user_object_correct_type = LiteLLM_UserTable(**user_object.model_dump())
|
||||
if user_object is not None:
|
||||
user_object_correct_type = LiteLLM_UserTable.model_validate(user_object.model_dump())
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
|
@ -279,19 +439,21 @@ async def new_organization(
|
|||
budget_params = LiteLLM_BudgetTable.model_fields.keys()
|
||||
|
||||
# Only include Budget Params when creating an entry in litellm_budgettable
|
||||
_json_data = data.json(exclude_none=True)
|
||||
_json_data = _STR_OBJECT_DICT_ADAPTER.validate_python(data.json(exclude_none=True))
|
||||
_budget_data = {k: v for k, v in _json_data.items() if k in budget_params}
|
||||
budget_row = LiteLLM_BudgetTable(**_budget_data)
|
||||
budget_row = LiteLLM_BudgetTable.model_validate(_budget_data)
|
||||
|
||||
new_budget = prisma_client.jsonify_object(budget_row.json(exclude_none=True))
|
||||
new_budget = _STR_OBJECT_DICT_ADAPTER.validate_python(
|
||||
prisma_client.jsonify_object(budget_row.json(exclude_none=True))
|
||||
)
|
||||
|
||||
_budget = await BudgetRepository(prisma_client).table.create(
|
||||
_budget = await _table(BudgetRepository(prisma_client)).create(
|
||||
data={
|
||||
**new_budget, # type: ignore
|
||||
**new_budget,
|
||||
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
|
||||
"updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
|
||||
}
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
data.budget_id = _budget.budget_id
|
||||
|
||||
|
|
@ -333,11 +495,13 @@ async def new_organization(
|
|||
value=getattr(data, field),
|
||||
)
|
||||
|
||||
new_organization_row = prisma_client.jsonify_object(organization_row.json(exclude_none=True))
|
||||
new_organization_row = _STR_OBJECT_DICT_ADAPTER.validate_python(
|
||||
prisma_client.jsonify_object(organization_row.json(exclude_none=True))
|
||||
)
|
||||
verbose_proxy_logger.info(f"new_organization_row: {json.dumps(new_organization_row, indent=2)}")
|
||||
response = await OrganizationRepository(prisma_client).table.create(
|
||||
response = await _table(OrganizationRepository(prisma_client)).create(
|
||||
data={
|
||||
**new_organization_row, # type: ignore
|
||||
**new_organization_row,
|
||||
},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
|
|
@ -351,14 +515,14 @@ async def new_organization(
|
|||
tags=["organization management"],
|
||||
)
|
||||
async def get_organization_daily_activity(
|
||||
organization_ids: Optional[str] = None,
|
||||
start_date: Optional[str] = None,
|
||||
end_date: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
organization_ids: str | None = None,
|
||||
start_date: str | None = None,
|
||||
end_date: str | None = None,
|
||||
model: str | None = None,
|
||||
api_key: str | None = None,
|
||||
page: int = 1,
|
||||
page_size: int = 10,
|
||||
exclude_organization_ids: Optional[str] = None,
|
||||
exclude_organization_ids: str | None = None,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
|
|
@ -376,13 +540,13 @@ async def get_organization_daily_activity(
|
|||
|
||||
# Parse comma-separated ids
|
||||
org_ids_list = organization_ids.split(",") if organization_ids else None
|
||||
exclude_org_ids_list: Optional[List[str]] = None
|
||||
exclude_org_ids_list: list[str] | None = None
|
||||
if exclude_organization_ids:
|
||||
exclude_org_ids_list = exclude_organization_ids.split(",") if exclude_organization_ids else None
|
||||
|
||||
# Restrict non-proxy-admins to only organizations where they are org_admin
|
||||
if not _user_has_admin_view(user_api_key_dict):
|
||||
memberships = await OrganizationMembershipRepository(prisma_client).table.find_many(
|
||||
memberships = await _table(OrganizationMembershipRepository(prisma_client)).find_many(
|
||||
where={"user_id": user_api_key_dict.user_id}
|
||||
)
|
||||
admin_org_ids = [m.organization_id for m in memberships if m.user_role == LitellmUserRoles.ORG_ADMIN.value]
|
||||
|
|
@ -399,11 +563,10 @@ async def get_organization_daily_activity(
|
|||
)
|
||||
|
||||
# Fetch organization aliases for metadata
|
||||
where_condition = {}
|
||||
where_condition = _STR_OBJECT_DICT_ADAPTER.validate_python({})
|
||||
if org_ids_list:
|
||||
where_condition["organization_id"] = {"in": list(org_ids_list)}
|
||||
org_aliases = await OrganizationRepository(prisma_client).table.find_many(where=where_condition)
|
||||
org_alias_metadata = {o.organization_id: {"organization_alias": o.organization_alias} for o in org_aliases}
|
||||
org_aliases = await _table(OrganizationRepository(prisma_client)).find_many(where=where_condition)
|
||||
|
||||
# Query daily activity for organizations
|
||||
return await get_daily_activity(
|
||||
|
|
@ -411,7 +574,7 @@ async def get_organization_daily_activity(
|
|||
table_name="litellm_dailyorganizationspend",
|
||||
entity_id_field="organization_id",
|
||||
entity_id=org_ids_list,
|
||||
entity_metadata_field=org_alias_metadata,
|
||||
entity_metadata_field={o.organization_id: {"organization_alias": o.organization_alias} for o in org_aliases},
|
||||
exclude_entity_ids=exclude_org_ids_list,
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
|
|
@ -424,8 +587,8 @@ async def get_organization_daily_activity(
|
|||
|
||||
async def _set_object_permission(
|
||||
data: NewOrganizationRequest,
|
||||
prisma_client: Optional[PrismaClient],
|
||||
) -> Optional[str]:
|
||||
prisma_client: PrismaClient | None,
|
||||
) -> str | None:
|
||||
"""
|
||||
Creates the LiteLLM_ObjectPermissionTable record for the organization.
|
||||
- Handles permissions for vector stores and mcp servers.
|
||||
|
|
@ -436,7 +599,7 @@ async def _set_object_permission(
|
|||
return None
|
||||
|
||||
if data.object_permission is not None:
|
||||
created_object_permission = await ObjectPermissionRepository(prisma_client).table.create(
|
||||
created_object_permission = await _table(ObjectPermissionRepository(prisma_client)).create(
|
||||
data=data.object_permission.model_dump(exclude_none=True),
|
||||
)
|
||||
del data.object_permission
|
||||
|
|
@ -522,10 +685,14 @@ async def update_organization(
|
|||
if updated_organization_row_json.get("metadata") is not None:
|
||||
existing_metadata = existing_organization_row.metadata or {}
|
||||
updated_metadata = updated_organization_row_json.get("metadata", {})
|
||||
merged_metadata = _update_dictionary(existing_dict=existing_metadata.copy(), new_dict=updated_metadata)
|
||||
merged_metadata: Mapping[str, object] = _update_dictionary(
|
||||
existing_dict=existing_metadata.copy(), new_dict=updated_metadata
|
||||
)
|
||||
updated_organization_row_json["metadata"] = merged_metadata
|
||||
|
||||
updated_organization_row = prisma_client.jsonify_object(updated_organization_row_json)
|
||||
updated_organization_row = _STR_OBJECT_DICT_ADAPTER.validate_python(
|
||||
prisma_client.jsonify_object(updated_organization_row_json)
|
||||
)
|
||||
if data.object_permission is not None:
|
||||
updated_organization_row = await handle_update_object_permission(
|
||||
data_json=updated_organization_row,
|
||||
|
|
@ -547,7 +714,7 @@ async def update_organization(
|
|||
for field in LiteLLM_BudgetTable.model_fields.keys():
|
||||
updated_organization_row.pop(field, None)
|
||||
|
||||
response = await OrganizationRepository(prisma_client).table.update(
|
||||
response = await _table(OrganizationRepository(prisma_client)).update(
|
||||
where={"organization_id": data.organization_id},
|
||||
data=updated_organization_row,
|
||||
include={"members": True, "teams": True, "litellm_budget_table": True},
|
||||
|
|
@ -557,9 +724,9 @@ async def update_organization(
|
|||
|
||||
|
||||
async def handle_update_object_permission(
|
||||
data_json: dict,
|
||||
data_json: dict[str, object],
|
||||
existing_organization_row: LiteLLM_OrganizationTable,
|
||||
) -> dict:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Handle the update of object permission for an organization.
|
||||
|
||||
|
|
@ -665,7 +832,7 @@ async def update_organization_v2(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
existing_organization_row = await OrganizationRepository(prisma_client).table.find_unique(
|
||||
existing_organization_row = await _table(OrganizationRepository(prisma_client)).find_unique(
|
||||
where={"organization_id": organization_id},
|
||||
)
|
||||
if existing_organization_row is None:
|
||||
|
|
@ -698,15 +865,18 @@ async def update_organization_v2(
|
|||
else ({"object_permission_id": None} if object_permission_cleared else {})
|
||||
)
|
||||
|
||||
organization_write_data = prisma_client.jsonify_object(
|
||||
{
|
||||
**org_column_updates,
|
||||
**object_permission_write,
|
||||
"updated_by": user_api_key_dict.user_id,
|
||||
}
|
||||
organization_write_data = _STR_OBJECT_DICT_ADAPTER.validate_python(
|
||||
prisma_client.jsonify_object(
|
||||
{
|
||||
**org_column_updates,
|
||||
**object_permission_write,
|
||||
"updated_by": user_api_key_dict.user_id,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
async with prisma_client.db.tx() as tx:
|
||||
tx_manager: _TransactionManager = prisma_client.db.tx()
|
||||
async with tx_manager as tx:
|
||||
if object_permission_upsert is not None:
|
||||
await tx.litellm_objectpermissiontable.upsert(
|
||||
where={"object_permission_id": object_permission_upsert.object_permission_id},
|
||||
|
|
@ -716,11 +886,12 @@ async def update_organization_v2(
|
|||
},
|
||||
)
|
||||
if budget_updates:
|
||||
budget_write_data = _STR_OBJECT_DICT_ADAPTER.validate_python(
|
||||
prisma_client.jsonify_object(dict(build_budget_write_data(budget_updates, user_api_key_dict.user_id)))
|
||||
)
|
||||
await tx.litellm_budgettable.update(
|
||||
where={"budget_id": existing_organization_row.budget_id},
|
||||
data=prisma_client.jsonify_object(
|
||||
dict(build_budget_write_data(budget_updates, user_api_key_dict.user_id))
|
||||
),
|
||||
data=budget_write_data,
|
||||
)
|
||||
response = await tx.litellm_organizationtable.update(
|
||||
where={"organization_id": organization_id},
|
||||
|
|
@ -735,7 +906,7 @@ async def update_organization_v2(
|
|||
"/organization/delete",
|
||||
tags=["organization management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[LiteLLM_OrganizationTableWithMembers],
|
||||
response_model=list[LiteLLM_OrganizationTableWithMembers],
|
||||
)
|
||||
async def delete_organization(
|
||||
data: DeleteOrganizationRequest,
|
||||
|
|
@ -765,15 +936,15 @@ async def delete_organization(
|
|||
deleted_orgs = []
|
||||
for organization_id in data.organization_ids:
|
||||
# delete all teams in the organization
|
||||
await TeamRepository(prisma_client).table.delete_many(where={"organization_id": organization_id})
|
||||
await _table(TeamRepository(prisma_client)).delete_many(where={"organization_id": organization_id})
|
||||
# delete all members in the organization
|
||||
await OrganizationMembershipRepository(prisma_client).table.delete_many(
|
||||
await _table(OrganizationMembershipRepository(prisma_client)).delete_many(
|
||||
where={"organization_id": organization_id}
|
||||
)
|
||||
# delete all keys in the organization
|
||||
await VerificationTokenRepository(prisma_client).table.delete_many(where={"organization_id": organization_id})
|
||||
await _table(VerificationTokenRepository(prisma_client)).delete_many(where={"organization_id": organization_id})
|
||||
# delete the organization
|
||||
deleted_org = await OrganizationRepository(prisma_client).table.delete(
|
||||
deleted_org = await _table(OrganizationRepository(prisma_client)).delete(
|
||||
where={"organization_id": organization_id},
|
||||
include={"members": True, "teams": True, "litellm_budget_table": True},
|
||||
)
|
||||
|
|
@ -791,13 +962,11 @@ async def delete_organization(
|
|||
"/organization/list",
|
||||
tags=["organization management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[LiteLLM_OrganizationTableWithMembers],
|
||||
response_model=list[LiteLLM_OrganizationTableWithMembers],
|
||||
)
|
||||
async def list_organization(
|
||||
org_id: Optional[str] = fastapi.Query(
|
||||
default=None, description="Filter organizations by exact organization_id match"
|
||||
),
|
||||
org_alias: Optional[str] = fastapi.Query(
|
||||
org_id: str | None = fastapi.Query(default=None, description="Filter organizations by exact organization_id match"),
|
||||
org_alias: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Filter organizations by partial organization_alias match. Supports case-insensitive search.",
|
||||
),
|
||||
|
|
@ -836,7 +1005,7 @@ async def list_organization(
|
|||
)
|
||||
|
||||
# Build where conditions based on provided filters
|
||||
where_conditions: Dict[str, Any] = {}
|
||||
where_conditions: dict[str, object] = {}
|
||||
|
||||
if org_id:
|
||||
where_conditions["organization_id"] = org_id
|
||||
|
|
@ -849,13 +1018,13 @@ async def list_organization(
|
|||
|
||||
# if proxy admin or admin viewer - get all orgs (with optional filters)
|
||||
if _user_has_admin_view(user_api_key_dict):
|
||||
response = await OrganizationRepository(prisma_client).table.find_many(
|
||||
response = await _table(OrganizationRepository(prisma_client)).find_many(
|
||||
where=where_conditions if where_conditions else None,
|
||||
include={"litellm_budget_table": True, "members": True, "teams": True},
|
||||
)
|
||||
# if internal user - get orgs they are a member of (with optional filters)
|
||||
else:
|
||||
org_memberships = await OrganizationMembershipRepository(prisma_client).table.find_many(
|
||||
org_memberships = await _table(OrganizationMembershipRepository(prisma_client)).find_many(
|
||||
where={"user_id": user_api_key_dict.user_id}
|
||||
)
|
||||
membership_org_ids = [membership.organization_id for membership in org_memberships]
|
||||
|
|
@ -869,7 +1038,7 @@ async def list_organization(
|
|||
response = []
|
||||
else:
|
||||
where_conditions["organization_id"] = org_id
|
||||
response = await OrganizationRepository(prisma_client).table.find_many(
|
||||
response = await _table(OrganizationRepository(prisma_client)).find_many(
|
||||
where=where_conditions,
|
||||
include={
|
||||
"litellm_budget_table": True,
|
||||
|
|
@ -880,7 +1049,7 @@ async def list_organization(
|
|||
else:
|
||||
# Filter by membership and any additional filters
|
||||
where_conditions["organization_id"] = {"in": membership_org_ids}
|
||||
response = await OrganizationRepository(prisma_client).table.find_many(
|
||||
response = await _table(OrganizationRepository(prisma_client)).find_many(
|
||||
where=where_conditions,
|
||||
include={
|
||||
"litellm_budget_table": True,
|
||||
|
|
@ -920,9 +1089,7 @@ async def info_organization(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
response: Optional[LiteLLM_OrganizationTableWithMembers] = await OrganizationRepository(
|
||||
prisma_client
|
||||
).table.find_unique(
|
||||
response = await _table(OrganizationRepository(prisma_client)).find_unique(
|
||||
where={"organization_id": organization_id},
|
||||
include={
|
||||
"litellm_budget_table": True,
|
||||
|
|
@ -939,7 +1106,7 @@ async def info_organization(
|
|||
if response is None:
|
||||
raise HTTPException(status_code=404, detail={"error": "Organization not found"})
|
||||
|
||||
response_pydantic_obj = LiteLLM_OrganizationTableWithMembers(**response.model_dump())
|
||||
response_pydantic_obj = LiteLLM_OrganizationTableWithMembers.model_validate(response.model_dump())
|
||||
|
||||
return response_pydantic_obj
|
||||
|
||||
|
|
@ -975,7 +1142,7 @@ async def deprecated_info_organization(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
response = await OrganizationRepository(prisma_client).table.find_many(
|
||||
response = await _table(OrganizationRepository(prisma_client)).find_many(
|
||||
where={"organization_id": {"in": data.organizations}},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
|
|
@ -1052,7 +1219,7 @@ async def organization_member_add(
|
|||
)
|
||||
|
||||
# Check if organization exists
|
||||
existing_organization_row = await OrganizationRepository(prisma_client).table.find_unique(
|
||||
existing_organization_row = await _table(OrganizationRepository(prisma_client)).find_unique(
|
||||
where={"organization_id": data.organization_id}
|
||||
)
|
||||
if existing_organization_row is None:
|
||||
|
|
@ -1063,14 +1230,14 @@ async def organization_member_add(
|
|||
},
|
||||
)
|
||||
|
||||
members: List[OrgMember]
|
||||
if isinstance(data.member, List):
|
||||
members: Sequence[OrgMember]
|
||||
if isinstance(data.member, list):
|
||||
members = data.member
|
||||
else:
|
||||
members = [data.member]
|
||||
|
||||
updated_users: List[LiteLLM_UserTable] = []
|
||||
updated_organization_memberships: List[LiteLLM_OrganizationMembershipTable] = []
|
||||
updated_users: list[LiteLLM_UserTable] = []
|
||||
updated_organization_memberships: list[LiteLLM_OrganizationMembershipTable] = []
|
||||
|
||||
for member in members:
|
||||
(
|
||||
|
|
@ -1125,7 +1292,7 @@ async def find_member_if_email(user_email: str, prisma_client: PrismaClient) ->
|
|||
"error": f"Unique user not found for user_email={user_email}. Potential duplicate OR non-existent user_email in LiteLLM_UserTable. Use 'user_id' instead."
|
||||
},
|
||||
)
|
||||
existing_user_email_row_pydantic = LiteLLM_UserTable(**existing_user_email_row.model_dump())
|
||||
existing_user_email_row_pydantic = LiteLLM_UserTable.model_validate(existing_user_email_row.model_dump())
|
||||
return existing_user_email_row_pydantic
|
||||
|
||||
|
||||
|
|
@ -1163,7 +1330,7 @@ async def organization_member_update(
|
|||
)
|
||||
|
||||
# Check if organization exists
|
||||
existing_organization_row = await OrganizationRepository(prisma_client).table.find_unique(
|
||||
existing_organization_row = await _table(OrganizationRepository(prisma_client)).find_unique(
|
||||
where={"organization_id": data.organization_id}
|
||||
)
|
||||
if existing_organization_row is None:
|
||||
|
|
@ -1180,7 +1347,9 @@ async def organization_member_update(
|
|||
data.user_id = existing_user_email_row.user_id
|
||||
|
||||
try:
|
||||
existing_organization_membership = await OrganizationMembershipRepository(prisma_client).table.find_unique(
|
||||
existing_organization_membership = await _table(
|
||||
OrganizationMembershipRepository(prisma_client)
|
||||
).find_unique(
|
||||
where={
|
||||
"user_id_organization_id": {
|
||||
"user_id": data.user_id,
|
||||
|
|
@ -1205,7 +1374,7 @@ async def organization_member_update(
|
|||
# org-scoped operations. An org-admin of any org could otherwise
|
||||
# alter a PROXY_ADMIN user's per-org role, which has downstream
|
||||
# effects on admin UI filtering and scope derivation.
|
||||
target_user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": data.user_id})
|
||||
target_user_row = await _table(UserRepository(prisma_client)).find_unique(where={"user_id": data.user_id})
|
||||
if target_user_row is not None and getattr(target_user_row, "user_role", None) in (
|
||||
LitellmUserRoles.PROXY_ADMIN.value,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value,
|
||||
|
|
@ -1222,7 +1391,7 @@ async def organization_member_update(
|
|||
|
||||
# Update member role
|
||||
if data.role is not None:
|
||||
await OrganizationMembershipRepository(prisma_client).table.update(
|
||||
await _table(OrganizationMembershipRepository(prisma_client)).update(
|
||||
where={
|
||||
"user_id_organization_id": {
|
||||
"user_id": data.user_id,
|
||||
|
|
@ -1245,7 +1414,7 @@ async def organization_member_update(
|
|||
)
|
||||
|
||||
# update organization membership with new budget_id
|
||||
await OrganizationMembershipRepository(prisma_client).table.update(
|
||||
await _table(OrganizationMembershipRepository(prisma_client)).update(
|
||||
where={
|
||||
"user_id_organization_id": {
|
||||
"user_id": data.user_id,
|
||||
|
|
@ -1254,9 +1423,7 @@ async def organization_member_update(
|
|||
},
|
||||
data={"budget_id": budget_id},
|
||||
)
|
||||
final_organization_membership: Optional[BaseModel] = await OrganizationMembershipRepository(
|
||||
prisma_client
|
||||
).table.find_unique(
|
||||
final_organization_membership = await _table(OrganizationMembershipRepository(prisma_client)).find_unique(
|
||||
where={
|
||||
"user_id_organization_id": {
|
||||
"user_id": data.user_id,
|
||||
|
|
@ -1272,8 +1439,8 @@ async def organization_member_update(
|
|||
detail={"error": f"Member not found in organization={data.organization_id} for user_id={data.user_id}"},
|
||||
)
|
||||
|
||||
final_organization_membership_pydantic = LiteLLM_OrganizationMembershipTable(
|
||||
**final_organization_membership.model_dump(exclude_none=True)
|
||||
final_organization_membership_pydantic = LiteLLM_OrganizationMembershipTable.model_validate(
|
||||
final_organization_membership.model_dump(exclude_none=True)
|
||||
)
|
||||
return final_organization_membership_pydantic
|
||||
except Exception as e:
|
||||
|
|
@ -1315,7 +1482,7 @@ async def organization_member_delete(
|
|||
existing_user_email_row = await find_member_if_email(data.user_email, prisma_client)
|
||||
data.user_id = existing_user_email_row.user_id
|
||||
|
||||
member_to_delete = await OrganizationMembershipRepository(prisma_client).table.delete(
|
||||
member_to_delete = await _table(OrganizationMembershipRepository(prisma_client)).delete(
|
||||
where={
|
||||
"user_id_organization_id": {
|
||||
"user_id": data.user_id,
|
||||
|
|
@ -1334,7 +1501,7 @@ async def add_member_to_organization(
|
|||
member: OrgMember,
|
||||
organization_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
) -> Tuple[LiteLLM_UserTable, LiteLLM_OrganizationMembershipTable]:
|
||||
) -> tuple[LiteLLM_UserTable, LiteLLM_OrganizationMembershipTable]:
|
||||
"""
|
||||
Add a member to an organization
|
||||
|
||||
|
|
@ -1344,12 +1511,12 @@ async def add_member_to_organization(
|
|||
"""
|
||||
|
||||
try:
|
||||
user_object: Optional[LiteLLM_UserTable] = None
|
||||
user_object: LiteLLM_UserTable | None = None
|
||||
existing_user_id_row = None
|
||||
existing_user_email_row = None
|
||||
## Check if user exists in LiteLLM_UserTable - user exists - either the user_id or user_email is in LiteLLM_UserTable
|
||||
if member.user_id is not None:
|
||||
existing_user_id_row = await UserRepository(prisma_client).table.find_unique(
|
||||
existing_user_id_row = await _table(UserRepository(prisma_client)).find_unique(
|
||||
where={"user_id": member.user_id}
|
||||
)
|
||||
|
||||
|
|
@ -1374,16 +1541,16 @@ async def add_member_to_organization(
|
|||
|
||||
_returned_user = await prisma_client.insert_data(data=new_user_defaults, table_name="user") # type: ignore
|
||||
if _returned_user is not None:
|
||||
user_object = LiteLLM_UserTable(**_returned_user.model_dump())
|
||||
user_object = LiteLLM_UserTable.model_validate(_returned_user.model_dump())
|
||||
elif existing_user_email_row is not None and len(existing_user_email_row) > 1:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": "Multiple users with this email found in db. Please use 'user_id' instead."},
|
||||
)
|
||||
elif existing_user_email_row is not None:
|
||||
user_object = LiteLLM_UserTable(**existing_user_email_row.model_dump())
|
||||
user_object = LiteLLM_UserTable.model_validate(existing_user_email_row.model_dump())
|
||||
elif existing_user_id_row is not None:
|
||||
user_object = LiteLLM_UserTable(**existing_user_id_row.model_dump())
|
||||
user_object = LiteLLM_UserTable.model_validate(existing_user_id_row.model_dump())
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
|
|
@ -1396,14 +1563,16 @@ async def add_member_to_organization(
|
|||
)
|
||||
|
||||
# Add user to organization
|
||||
_organization_membership = await OrganizationMembershipRepository(prisma_client).table.create(
|
||||
_organization_membership = await _table(OrganizationMembershipRepository(prisma_client)).create(
|
||||
data={
|
||||
"organization_id": organization_id,
|
||||
"user_id": user_object.user_id,
|
||||
"user_role": member.role,
|
||||
}
|
||||
)
|
||||
organization_membership = LiteLLM_OrganizationMembershipTable(**_organization_membership.model_dump())
|
||||
organization_membership = LiteLLM_OrganizationMembershipTable.model_validate(
|
||||
_organization_membership.model_dump()
|
||||
)
|
||||
return user_object, organization_membership
|
||||
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -12,8 +12,14 @@ All /tag management endpoints
|
|||
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Protocol,
|
||||
TypedDict,
|
||||
overload,
|
||||
)
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
|
|
@ -42,16 +48,101 @@ from litellm.types.tag_management import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma.models import LiteLLM_BudgetTable as PrismaBudgetTable
|
||||
from prisma.models import LiteLLM_ProxyModelTable as PrismaProxyModelTable
|
||||
from prisma.models import LiteLLM_TagTable as PrismaTagTable
|
||||
from prisma.models import LiteLLM_VerificationToken as PrismaVerificationToken
|
||||
|
||||
from litellm import Router
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.types.router import Deployment
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class _TagRecord(Protocol):
|
||||
tag_name: str
|
||||
description: str | None
|
||||
models: Sequence[str]
|
||||
model_info: object
|
||||
budget_id: str | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
created_by: str | None
|
||||
litellm_budget_table: "PrismaBudgetTable | None"
|
||||
|
||||
|
||||
class _TagTableClient(Protocol):
|
||||
async def find_unique(self, where: Mapping[str, object]) -> "_TagRecord | None": ...
|
||||
|
||||
async def find_many(
|
||||
self,
|
||||
where: Mapping[str, object] | None = None,
|
||||
include: Mapping[str, object] | None = None,
|
||||
) -> "Sequence[_TagRecord]": ...
|
||||
|
||||
async def create(self, data: Mapping[str, object]) -> "PrismaTagTable": ...
|
||||
|
||||
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> "PrismaTagTable": ...
|
||||
|
||||
async def delete(self, where: Mapping[str, object]) -> "PrismaTagTable | None": ...
|
||||
|
||||
|
||||
class _ModelTableClient(Protocol):
|
||||
async def find_many(self, where: Mapping[str, object] | None = None) -> "Sequence[PrismaProxyModelTable]": ...
|
||||
|
||||
|
||||
class _VerificationTokenTableClient(Protocol):
|
||||
async def find_many(
|
||||
self,
|
||||
where: Mapping[str, object] | None = None,
|
||||
select: Mapping[str, object] | None = None,
|
||||
) -> "Sequence[PrismaVerificationToken]": ...
|
||||
|
||||
|
||||
class _DailyTagSpendGroupByRow(TypedDict):
|
||||
tag: str | None
|
||||
_min: Mapping[str, object]
|
||||
_max: Mapping[str, object]
|
||||
|
||||
|
||||
class _DailyTagSpendTableClient(Protocol):
|
||||
async def group_by(
|
||||
self,
|
||||
by: Sequence[str],
|
||||
where: Mapping[str, object] | None = None,
|
||||
min: Mapping[str, object] | None = None,
|
||||
max: Mapping[str, object] | None = None,
|
||||
) -> "Sequence[_DailyTagSpendGroupByRow]": ...
|
||||
|
||||
|
||||
@overload
|
||||
def _table(repository: DailyTagSpendRepository) -> "_DailyTagSpendTableClient": ...
|
||||
|
||||
|
||||
@overload
|
||||
def _table(repository: ModelRepository) -> "_ModelTableClient": ...
|
||||
|
||||
|
||||
@overload
|
||||
def _table(repository: TagRepository) -> "_TagTableClient": ...
|
||||
|
||||
|
||||
@overload
|
||||
def _table(repository: VerificationTokenRepository) -> "_VerificationTokenTableClient": ...
|
||||
|
||||
|
||||
def _table(
|
||||
repository: DailyTagSpendRepository | ModelRepository | TagRepository | VerificationTokenRepository,
|
||||
) -> object:
|
||||
prisma_table: object = repository.table
|
||||
return prisma_table
|
||||
|
||||
|
||||
async def _get_internal_user_api_keys(
|
||||
prisma_client,
|
||||
prisma_client: "PrismaClient",
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> List[str]:
|
||||
) -> list[str]:
|
||||
user_role = user_api_key_dict.user_role
|
||||
if user_role is None or not user_role.is_internal_user_role:
|
||||
return []
|
||||
|
|
@ -64,7 +155,7 @@ async def _get_internal_user_api_keys(
|
|||
if user_id is None:
|
||||
return sorted(user_api_keys)
|
||||
|
||||
key_records = await VerificationTokenRepository(prisma_client).table.find_many(
|
||||
key_records = await _table(VerificationTokenRepository(prisma_client)).find_many(
|
||||
where={"user_id": user_id},
|
||||
select={"token": True},
|
||||
)
|
||||
|
|
@ -74,9 +165,9 @@ async def _get_internal_user_api_keys(
|
|||
|
||||
|
||||
async def _get_tag_list_scope(
|
||||
prisma_client,
|
||||
prisma_client: "PrismaClient",
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> Optional[Dict[str, dict]]:
|
||||
) -> Mapping[str, Mapping[str, Sequence[str]]] | None:
|
||||
user_role = user_api_key_dict.user_role
|
||||
if user_api_key_has_admin_view(user_api_key_dict) or (user_role is None or not user_role.is_internal_user_role):
|
||||
return None
|
||||
|
|
@ -89,10 +180,10 @@ async def _get_tag_list_scope(
|
|||
|
||||
|
||||
async def _get_tag_daily_activity_api_key_filter(
|
||||
prisma_client,
|
||||
prisma_client: "PrismaClient",
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
requested_api_key: Optional[str],
|
||||
) -> Optional[Union[str, List[str]]]:
|
||||
requested_api_key: str | None,
|
||||
) -> str | list[str] | None:
|
||||
user_role = user_api_key_dict.user_role
|
||||
if user_api_key_has_admin_view(user_api_key_dict) or (user_role is None or not user_role.is_internal_user_role):
|
||||
return requested_api_key
|
||||
|
|
@ -106,17 +197,17 @@ async def _get_tag_daily_activity_api_key_filter(
|
|||
return scoped_api_keys
|
||||
|
||||
|
||||
async def _get_model_names(prisma_client, model_ids: list) -> Dict[str, str]:
|
||||
async def _get_model_names(prisma_client: "PrismaClient", model_ids: Sequence[str]) -> dict[str, str]:
|
||||
"""Helper function to get model names from model IDs"""
|
||||
try:
|
||||
models = await ModelRepository(prisma_client).table.find_many(where={"model_id": {"in": model_ids}})
|
||||
models = await _table(ModelRepository(prisma_client)).find_many(where={"model_id": {"in": model_ids}})
|
||||
return {model.model_id: model.model_name for model in models}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error getting model names: {str(e)}")
|
||||
return {}
|
||||
|
||||
|
||||
async def get_deployments_by_model(model: str, llm_router: "Router") -> List["Deployment"]:
|
||||
async def get_deployments_by_model(model: str, llm_router: "Router") -> list["Deployment"]:
|
||||
"""
|
||||
Get all deployments by model
|
||||
"""
|
||||
|
|
@ -181,7 +272,7 @@ async def new_tag(
|
|||
raise HTTPException(status_code=500, detail=CommonProxyErrors.no_llm_router.value)
|
||||
try:
|
||||
# Check if tag already exists
|
||||
existing_tag = await TagRepository(prisma_client).table.find_unique(where={"tag_name": tag.name})
|
||||
existing_tag = await _table(TagRepository(prisma_client)).find_unique(where={"tag_name": tag.name})
|
||||
if existing_tag is not None:
|
||||
raise HTTPException(status_code=400, detail=f"Tag {tag.name} already exists")
|
||||
|
||||
|
|
@ -198,7 +289,7 @@ async def new_tag(
|
|||
model_info = await _get_model_names(prisma_client, tag.models or [])
|
||||
|
||||
# Create new tag in database
|
||||
new_tag_record = await TagRepository(prisma_client).table.create(
|
||||
new_tag_record = await _table(TagRepository(prisma_client)).create(
|
||||
data={
|
||||
"tag_name": tag.name,
|
||||
"description": tag.description,
|
||||
|
|
@ -321,7 +412,7 @@ async def update_tag(
|
|||
|
||||
try:
|
||||
# Check if tag exists
|
||||
existing_tag = await TagRepository(prisma_client).table.find_unique(where={"tag_name": tag.name})
|
||||
existing_tag = await _table(TagRepository(prisma_client)).find_unique(where={"tag_name": tag.name})
|
||||
if existing_tag is None:
|
||||
raise HTTPException(status_code=404, detail=f"Tag {tag.name} not found")
|
||||
|
||||
|
|
@ -351,7 +442,7 @@ async def update_tag(
|
|||
update_data["budget_id"] = budget_id
|
||||
|
||||
# Update tag in database
|
||||
updated_tag_record = await TagRepository(prisma_client).table.update(
|
||||
updated_tag_record = await _table(TagRepository(prisma_client)).update(
|
||||
where={"tag_name": tag.name},
|
||||
data=update_data,
|
||||
)
|
||||
|
|
@ -398,7 +489,7 @@ async def info_tag(
|
|||
|
||||
try:
|
||||
# Query tags from database with budget info
|
||||
tag_records = await TagRepository(prisma_client).table.find_many(
|
||||
tag_records = await _table(TagRepository(prisma_client)).find_many(
|
||||
where={"tag_name": {"in": data.names}},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
|
|
@ -413,7 +504,7 @@ async def info_tag(
|
|||
requested_tags = {}
|
||||
for tag_record in tag_records:
|
||||
# Parse model_info from JSON
|
||||
model_info = {}
|
||||
model_info: object = {}
|
||||
if tag_record.model_info:
|
||||
if isinstance(tag_record.model_info, str):
|
||||
model_info = json.loads(tag_record.model_info)
|
||||
|
|
@ -441,7 +532,7 @@ async def info_tag(
|
|||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
def _validate_tag_list_date_range(start_date: Optional[str], end_date: Optional[str]) -> None:
|
||||
def _validate_tag_list_date_range(start_date: str | None, end_date: str | None) -> None:
|
||||
"""Require both dates together, and enforce YYYY-MM-DD format with start <= end."""
|
||||
if (start_date is None) != (end_date is None):
|
||||
raise HTTPException(
|
||||
|
|
@ -472,7 +563,7 @@ def _validate_tag_list_date_range(start_date: Optional[str], end_date: Optional[
|
|||
)
|
||||
async def list_tags(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
start_date: Optional[str] = Query(
|
||||
start_date: str | None = Query(
|
||||
None,
|
||||
description=(
|
||||
"Optional start date (YYYY-MM-DD). When provided together with "
|
||||
|
|
@ -480,7 +571,7 @@ async def list_tags(
|
|||
"Stored tags are always returned."
|
||||
),
|
||||
),
|
||||
end_date: Optional[str] = Query(
|
||||
end_date: str | None = Query(
|
||||
None,
|
||||
description="Optional end date (YYYY-MM-DD). Must be given with start_date.",
|
||||
),
|
||||
|
|
@ -506,13 +597,13 @@ async def list_tags(
|
|||
# Prisma's distinct fetches all columns for all rows and deduplicates
|
||||
# in application code, which is extremely slow on large tables.
|
||||
# See: https://www.prisma.io/docs/orm/prisma-client/queries/aggregation-grouping-summarizing#distinct-under-the-hood
|
||||
dynamic_tag_where: Dict[str, Any] = {"tag": {"not": None}}
|
||||
dynamic_tag_where: dict[str, object] = {"tag": {"not": None}}
|
||||
if tag_scope:
|
||||
dynamic_tag_where = {**dynamic_tag_where, **tag_scope}
|
||||
if start_date is not None and end_date is not None:
|
||||
dynamic_tag_where["date"] = {"gte": start_date, "lte": end_date}
|
||||
|
||||
dynamic_tag_rows = await DailyTagSpendRepository(prisma_client).table.group_by(
|
||||
dynamic_tag_rows = await _table(DailyTagSpendRepository(prisma_client)).group_by(
|
||||
by=["tag"],
|
||||
where=dynamic_tag_where,
|
||||
min={"created_at": True},
|
||||
|
|
@ -526,7 +617,7 @@ async def list_tags(
|
|||
stored_tag_where = {"tag_name": {"in": used_tag_names}} if tag_scope is not None else None
|
||||
|
||||
## QUERY STORED TAGS ##
|
||||
tag_records = await TagRepository(prisma_client).table.find_many(
|
||||
tag_records = await _table(TagRepository(prisma_client)).find_many(
|
||||
where=stored_tag_where,
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
|
|
@ -536,7 +627,7 @@ async def list_tags(
|
|||
for tag_record in tag_records:
|
||||
stored_tag_names.add(tag_record.tag_name)
|
||||
# Parse model_info from JSON
|
||||
model_info = {}
|
||||
model_info: object = {}
|
||||
if tag_record.model_info:
|
||||
if isinstance(tag_record.model_info, str):
|
||||
model_info = json.loads(tag_record.model_info)
|
||||
|
|
@ -598,12 +689,12 @@ async def delete_tag(
|
|||
|
||||
try:
|
||||
# Check if tag exists
|
||||
existing_tag = await TagRepository(prisma_client).table.find_unique(where={"tag_name": data.name})
|
||||
existing_tag = await _table(TagRepository(prisma_client)).find_unique(where={"tag_name": data.name})
|
||||
if existing_tag is None:
|
||||
raise HTTPException(status_code=404, detail=f"Tag {data.name} not found")
|
||||
|
||||
# Delete tag from database
|
||||
await TagRepository(prisma_client).table.delete(where={"tag_name": data.name})
|
||||
await _table(TagRepository(prisma_client)).delete(where={"tag_name": data.name})
|
||||
|
||||
return {"message": f"Tag {data.name} deleted successfully"}
|
||||
except Exception as e:
|
||||
|
|
@ -617,11 +708,11 @@ async def delete_tag(
|
|||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def get_tag_daily_activity(
|
||||
tags: Optional[str] = None,
|
||||
start_date: Optional[str] = None,
|
||||
end_date: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
tags: str | None = None,
|
||||
start_date: str | None = None,
|
||||
end_date: str | None = None,
|
||||
model: str | None = None,
|
||||
api_key: str | None = None,
|
||||
page: int = 1,
|
||||
page_size: int = 10,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
|
|
|
|||
|
|
@ -8,8 +8,16 @@ by policy_attachments (see AttachmentRegistry).
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Optional,
|
||||
Protocol,
|
||||
TypedDict,
|
||||
Union,
|
||||
)
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.repositories.table_repositories import PolicyRepository
|
||||
|
|
@ -33,7 +41,89 @@ if TYPE_CHECKING:
|
|||
POLICY_VERSION_ID_PREFIX = "policy_"
|
||||
|
||||
|
||||
def _row_to_policy_db_response(row: Any) -> PolicyDBResponse:
|
||||
class _RawPipelineStep(TypedDict):
|
||||
guardrail: str
|
||||
|
||||
|
||||
class _RawPipelineConfig(TypedDict, total=False):
|
||||
mode: str
|
||||
steps: Sequence[Union[PipelineStep, "_RawPipelineStep"]]
|
||||
|
||||
|
||||
class _PolicyRow(Protocol):
|
||||
policy_id: str
|
||||
policy_name: str
|
||||
version_number: int
|
||||
version_status: str
|
||||
parent_version_id: str | None
|
||||
is_latest: bool
|
||||
published_at: datetime | None
|
||||
production_at: datetime | None
|
||||
inherit: str | None
|
||||
description: str | None
|
||||
guardrails_add: list[str] | None
|
||||
guardrails_remove: list[str] | None
|
||||
condition: dict[str, object] | None
|
||||
pipeline: dict[str, object] | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
created_by: str | None
|
||||
updated_by: str | None
|
||||
|
||||
|
||||
class _PolicyVersionSourceRow(Protocol):
|
||||
policy_id: str
|
||||
policy_name: str
|
||||
version_number: int
|
||||
inherit: str | None
|
||||
description: str | None
|
||||
guardrails_add: Sequence[str] | None
|
||||
guardrails_remove: Sequence[str] | None
|
||||
condition: Mapping[str, object] | str | None
|
||||
pipeline: Mapping[str, object] | str | None
|
||||
|
||||
|
||||
class _PolicyTableClient(Protocol):
|
||||
async def create(self, data: Mapping[str, object]) -> _PolicyRow: ...
|
||||
|
||||
async def find_unique(self, where: Mapping[str, object]) -> _PolicyRow | None: ...
|
||||
|
||||
async def find_many(
|
||||
self,
|
||||
where: Mapping[str, object] | None = None,
|
||||
order: Mapping[str, str] | None = None,
|
||||
) -> Sequence[_PolicyRow]: ...
|
||||
|
||||
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _PolicyRow: ...
|
||||
|
||||
async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
|
||||
|
||||
async def delete(self, where: Mapping[str, object]) -> _PolicyRow | None: ...
|
||||
|
||||
async def delete_many(self, where: Mapping[str, object]) -> int: ...
|
||||
|
||||
|
||||
class _PolicyVersionSourceTableClient(Protocol):
|
||||
async def find_unique(self, where: Mapping[str, object]) -> _PolicyVersionSourceRow | None: ...
|
||||
|
||||
async def find_first(
|
||||
self,
|
||||
where: Mapping[str, object],
|
||||
order: Mapping[str, str] | None = None,
|
||||
) -> _PolicyVersionSourceRow | None: ...
|
||||
|
||||
|
||||
def _policy_table(prisma_client: "PrismaClient") -> _PolicyTableClient:
|
||||
table: _PolicyTableClient = PolicyRepository(prisma_client).table
|
||||
return table
|
||||
|
||||
|
||||
def _policy_version_source_table(prisma_client: "PrismaClient") -> _PolicyVersionSourceTableClient:
|
||||
table: _PolicyVersionSourceTableClient = PolicyRepository(prisma_client).table
|
||||
return table
|
||||
|
||||
|
||||
def _row_to_policy_db_response(row: _PolicyRow) -> PolicyDBResponse:
|
||||
"""Build PolicyDBResponse from a Prisma LiteLLM_PolicyTable row."""
|
||||
return PolicyDBResponse(
|
||||
policy_id=row.policy_id,
|
||||
|
|
@ -71,11 +161,11 @@ class PolicyRegistry:
|
|||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._policies: Dict[str, Policy] = {}
|
||||
self._policies_by_id: Dict[str, Tuple[str, Policy]] = {}
|
||||
self._policies: dict[str, Policy] = {}
|
||||
self._policies_by_id: dict[str, tuple[str, Policy]] = {}
|
||||
self._initialized: bool = False
|
||||
|
||||
def load_policies(self, policies_config: Dict[str, Any]) -> None:
|
||||
def load_policies(self, policies_config: Mapping[str, dict[str, object]]) -> None:
|
||||
"""
|
||||
Load policies from a configuration dictionary.
|
||||
|
||||
|
|
@ -98,7 +188,7 @@ class PolicyRegistry:
|
|||
self._initialized = True
|
||||
verbose_proxy_logger.info(f"Loaded {len(self._policies)} policies")
|
||||
|
||||
def _parse_policy(self, policy_name: str, policy_data: Dict[str, Any]) -> Policy:
|
||||
def _parse_policy(self, policy_name: str, policy_data: dict[str, Any]) -> Policy:
|
||||
"""
|
||||
Parse a policy from raw configuration data.
|
||||
|
||||
|
|
@ -139,13 +229,13 @@ class PolicyRegistry:
|
|||
|
||||
@staticmethod
|
||||
def _parse_pipeline(
|
||||
pipeline_data: Optional[Dict[str, Any]],
|
||||
) -> Optional[GuardrailPipeline]:
|
||||
pipeline_data: Optional["_RawPipelineConfig"],
|
||||
) -> GuardrailPipeline | None:
|
||||
"""Parse a pipeline configuration from raw data."""
|
||||
if pipeline_data is None:
|
||||
return None
|
||||
|
||||
steps_data = pipeline_data.get("steps", [])
|
||||
steps_data: Sequence[PipelineStep | _RawPipelineStep] = pipeline_data.get("steps", [])
|
||||
steps = [PipelineStep(**step_data) if isinstance(step_data, dict) else step_data for step_data in steps_data]
|
||||
|
||||
return GuardrailPipeline(
|
||||
|
|
@ -153,7 +243,7 @@ class PolicyRegistry:
|
|||
steps=steps,
|
||||
)
|
||||
|
||||
def get_policy(self, policy_name: str) -> Optional[Policy]:
|
||||
def get_policy(self, policy_name: str) -> Policy | None:
|
||||
"""
|
||||
Get a policy by name.
|
||||
|
||||
|
|
@ -165,7 +255,7 @@ class PolicyRegistry:
|
|||
"""
|
||||
return self._policies.get(policy_name)
|
||||
|
||||
def get_all_policies(self) -> Dict[str, Policy]:
|
||||
def get_all_policies(self) -> dict[str, Policy]:
|
||||
"""
|
||||
Get all loaded policies.
|
||||
|
||||
|
|
@ -174,7 +264,7 @@ class PolicyRegistry:
|
|||
"""
|
||||
return self._policies.copy()
|
||||
|
||||
def get_policy_names(self) -> List[str]:
|
||||
def get_policy_names(self) -> list[str]:
|
||||
"""
|
||||
Get list of all policy names.
|
||||
|
||||
|
|
@ -247,7 +337,7 @@ class PolicyRegistry:
|
|||
self,
|
||||
policy_request: PolicyCreateRequest,
|
||||
prisma_client: "PrismaClient",
|
||||
created_by: Optional[str] = None,
|
||||
created_by: str | None = None,
|
||||
) -> PolicyDBResponse:
|
||||
"""
|
||||
Add a policy to the database.
|
||||
|
|
@ -263,7 +353,7 @@ class PolicyRegistry:
|
|||
try:
|
||||
now = datetime.now(timezone.utc)
|
||||
# Build data dict; new policy is v1 production
|
||||
data: Dict[str, Any] = {
|
||||
data: dict[str, object] = {
|
||||
"policy_name": policy_request.policy_name,
|
||||
"version_number": 1,
|
||||
"version_status": "production",
|
||||
|
|
@ -289,7 +379,7 @@ class PolicyRegistry:
|
|||
validated_pipeline = GuardrailPipeline(**policy_request.pipeline)
|
||||
data["pipeline"] = json.dumps(validated_pipeline.model_dump())
|
||||
|
||||
created_policy = await PolicyRepository(prisma_client).table.create(data=data)
|
||||
created_policy = await _policy_table(prisma_client).create(data=data)
|
||||
|
||||
# Also add to in-memory registry
|
||||
policy = self._parse_policy(
|
||||
|
|
@ -317,7 +407,7 @@ class PolicyRegistry:
|
|||
policy_id: str,
|
||||
policy_request: PolicyUpdateRequest,
|
||||
prisma_client: "PrismaClient",
|
||||
updated_by: Optional[str] = None,
|
||||
updated_by: str | None = None,
|
||||
) -> PolicyDBResponse:
|
||||
"""
|
||||
Update a policy in the database. Only draft versions can be updated.
|
||||
|
|
@ -335,7 +425,7 @@ class PolicyRegistry:
|
|||
Exception: If policy is not in draft status (only drafts are editable).
|
||||
"""
|
||||
try:
|
||||
existing = await PolicyRepository(prisma_client).table.find_unique(where={"policy_id": policy_id})
|
||||
existing = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id})
|
||||
if existing is None:
|
||||
raise Exception(f"Policy with ID {policy_id} not found")
|
||||
version_status = getattr(existing, "version_status", "production")
|
||||
|
|
@ -343,7 +433,7 @@ class PolicyRegistry:
|
|||
raise Exception(f"Only draft versions can be updated. This policy has status '{version_status}'.")
|
||||
|
||||
# Build update data - only include fields that are set
|
||||
update_data: Dict[str, Any] = {
|
||||
update_data: dict[str, object] = {
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
"updated_by": updated_by,
|
||||
}
|
||||
|
|
@ -364,7 +454,7 @@ class PolicyRegistry:
|
|||
validated_pipeline = GuardrailPipeline(**policy_request.pipeline)
|
||||
update_data["pipeline"] = json.dumps(validated_pipeline.model_dump())
|
||||
|
||||
updated_policy = await PolicyRepository(prisma_client).table.update(
|
||||
updated_policy = await _policy_table(prisma_client).update(
|
||||
where={"policy_id": policy_id},
|
||||
data=update_data,
|
||||
)
|
||||
|
|
@ -380,7 +470,7 @@ class PolicyRegistry:
|
|||
self,
|
||||
policy_id: str,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> Dict[str, Any]:
|
||||
) -> Mapping[str, str]:
|
||||
"""
|
||||
Delete a policy version from the database.
|
||||
|
||||
|
|
@ -395,7 +485,7 @@ class PolicyRegistry:
|
|||
Dict with "message" and optional "warning" if production was deleted.
|
||||
"""
|
||||
try:
|
||||
policy = await PolicyRepository(prisma_client).table.find_unique(where={"policy_id": policy_id})
|
||||
policy = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id})
|
||||
|
||||
if policy is None:
|
||||
raise Exception(f"Policy with ID {policy_id} not found")
|
||||
|
|
@ -404,9 +494,9 @@ class PolicyRegistry:
|
|||
policy_name = policy.policy_name
|
||||
|
||||
# Delete from DB
|
||||
await PolicyRepository(prisma_client).table.delete(where={"policy_id": policy_id})
|
||||
await _policy_table(prisma_client).delete(where={"policy_id": policy_id})
|
||||
|
||||
result: Dict[str, Any] = {"message": f"Policy {policy_id} deleted successfully"}
|
||||
result: dict[str, str] = {"message": f"Policy {policy_id} deleted successfully"}
|
||||
|
||||
# Remove from in-memory registry only if this was the production version
|
||||
if version_status == "production":
|
||||
|
|
@ -425,7 +515,7 @@ class PolicyRegistry:
|
|||
self,
|
||||
policy_id: str,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> Optional[PolicyDBResponse]:
|
||||
) -> PolicyDBResponse | None:
|
||||
"""
|
||||
Get a policy by ID from the database.
|
||||
|
||||
|
|
@ -437,7 +527,7 @@ class PolicyRegistry:
|
|||
PolicyDBResponse if found, None otherwise
|
||||
"""
|
||||
try:
|
||||
policy = await PolicyRepository(prisma_client).table.find_unique(where={"policy_id": policy_id})
|
||||
policy = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id})
|
||||
|
||||
if policy is None:
|
||||
return None
|
||||
|
|
@ -447,7 +537,7 @@ class PolicyRegistry:
|
|||
verbose_proxy_logger.exception(f"Error getting policy from DB: {e}")
|
||||
raise Exception(f"Error getting policy from DB: {str(e)}")
|
||||
|
||||
def get_policy_by_id_for_request(self, policy_id: str) -> Optional[Tuple[str, Policy]]:
|
||||
def get_policy_by_id_for_request(self, policy_id: str) -> tuple[str, Policy] | None:
|
||||
"""
|
||||
Return a policy version by ID from in-memory cache (no DB access).
|
||||
|
||||
|
|
@ -466,8 +556,8 @@ class PolicyRegistry:
|
|||
async def get_all_policies_from_db(
|
||||
self,
|
||||
prisma_client: "PrismaClient",
|
||||
version_status: Optional[str] = None,
|
||||
) -> List[PolicyDBResponse]:
|
||||
version_status: str | None = None,
|
||||
) -> list[PolicyDBResponse]:
|
||||
"""
|
||||
Get all policies from the database, optionally filtered by version_status.
|
||||
|
||||
|
|
@ -480,11 +570,11 @@ class PolicyRegistry:
|
|||
List of PolicyDBResponse objects
|
||||
"""
|
||||
try:
|
||||
where: Dict[str, Any] = {}
|
||||
where: dict[str, str] = {}
|
||||
if version_status is not None:
|
||||
where["version_status"] = version_status
|
||||
|
||||
policies = await PolicyRepository(prisma_client).table.find_many(
|
||||
policies = await _policy_table(prisma_client).find_many(
|
||||
where=where if where else None,
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
|
|
@ -524,7 +614,7 @@ class PolicyRegistry:
|
|||
self.add_policy(policy_response.policy_name, policy)
|
||||
|
||||
self._policies_by_id = {}
|
||||
non_production = await PolicyRepository(prisma_client).table.find_many(
|
||||
non_production = await _policy_table(prisma_client).find_many(
|
||||
where={"version_status": {"in": ["draft", "published"]}},
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
|
|
@ -557,7 +647,7 @@ class PolicyRegistry:
|
|||
self,
|
||||
policy_name: str,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> List[str]:
|
||||
) -> list[str]:
|
||||
"""
|
||||
Resolve all guardrails for a policy from the database.
|
||||
|
||||
|
|
@ -622,7 +712,7 @@ class PolicyRegistry:
|
|||
PolicyVersionListResponse with policy_name and list of versions
|
||||
"""
|
||||
try:
|
||||
rows = await PolicyRepository(prisma_client).table.find_many(
|
||||
rows = await _policy_table(prisma_client).find_many(
|
||||
where={"policy_name": policy_name},
|
||||
order={"version_number": "desc"},
|
||||
)
|
||||
|
|
@ -640,8 +730,8 @@ class PolicyRegistry:
|
|||
self,
|
||||
policy_name: str,
|
||||
prisma_client: "PrismaClient",
|
||||
source_policy_id: Optional[str] = None,
|
||||
created_by: Optional[str] = None,
|
||||
source_policy_id: str | None = None,
|
||||
created_by: str | None = None,
|
||||
) -> PolicyDBResponse:
|
||||
"""
|
||||
Create a new draft version of a policy. Copies all fields from the source.
|
||||
|
|
@ -658,14 +748,16 @@ class PolicyRegistry:
|
|||
"""
|
||||
try:
|
||||
if source_policy_id is not None:
|
||||
source = await PolicyRepository(prisma_client).table.find_unique(where={"policy_id": source_policy_id})
|
||||
source = await _policy_version_source_table(prisma_client).find_unique(
|
||||
where={"policy_id": source_policy_id}
|
||||
)
|
||||
if source is None:
|
||||
raise Exception(f"Source policy {source_policy_id} not found")
|
||||
if source.policy_name != policy_name:
|
||||
raise Exception(f"Source policy name '{source.policy_name}' does not match '{policy_name}'")
|
||||
else:
|
||||
# Find current production version for this policy_name
|
||||
prod = await PolicyRepository(prisma_client).table.find_first(
|
||||
prod = await _policy_version_source_table(prisma_client).find_first(
|
||||
where={
|
||||
"policy_name": policy_name,
|
||||
"version_status": "production",
|
||||
|
|
@ -676,7 +768,7 @@ class PolicyRegistry:
|
|||
source = prod
|
||||
|
||||
# Next version number
|
||||
latest = await PolicyRepository(prisma_client).table.find_first(
|
||||
latest = await _policy_version_source_table(prisma_client).find_first(
|
||||
where={"policy_name": policy_name},
|
||||
order={"version_number": "desc"},
|
||||
)
|
||||
|
|
@ -684,12 +776,12 @@ class PolicyRegistry:
|
|||
|
||||
now = datetime.now(timezone.utc)
|
||||
# Set is_latest=False on all existing versions for this policy_name
|
||||
await PolicyRepository(prisma_client).table.update_many(
|
||||
await _policy_table(prisma_client).update_many(
|
||||
where={"policy_name": policy_name},
|
||||
data={"is_latest": False},
|
||||
)
|
||||
|
||||
data: Dict[str, Any] = {
|
||||
data: dict[str, object] = {
|
||||
"policy_name": policy_name,
|
||||
"version_number": next_num,
|
||||
"version_status": "draft",
|
||||
|
|
@ -714,7 +806,7 @@ class PolicyRegistry:
|
|||
if source.pipeline is not None:
|
||||
data["pipeline"] = json.dumps(source.pipeline) if isinstance(source.pipeline, dict) else source.pipeline
|
||||
|
||||
created = await PolicyRepository(prisma_client).table.create(data=data)
|
||||
created = await _policy_table(prisma_client).create(data=data)
|
||||
return _row_to_policy_db_response(created)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error creating new version: {e}")
|
||||
|
|
@ -725,7 +817,7 @@ class PolicyRegistry:
|
|||
policy_id: str,
|
||||
new_status: str,
|
||||
prisma_client: "PrismaClient",
|
||||
updated_by: Optional[str] = None,
|
||||
updated_by: str | None = None,
|
||||
) -> PolicyDBResponse:
|
||||
"""
|
||||
Update a policy version's status. Valid transitions:
|
||||
|
|
@ -748,7 +840,7 @@ class PolicyRegistry:
|
|||
if new_status not in ("published", "production"):
|
||||
raise Exception(f"Invalid status '{new_status}'. Use 'published' or 'production'.")
|
||||
|
||||
row = await PolicyRepository(prisma_client).table.find_unique(where={"policy_id": policy_id})
|
||||
row = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id})
|
||||
if row is None:
|
||||
raise Exception(f"Policy with ID {policy_id} not found")
|
||||
|
||||
|
|
@ -759,7 +851,7 @@ class PolicyRegistry:
|
|||
if new_status == "published":
|
||||
if current != "draft":
|
||||
raise Exception(f"Only draft versions can be published. Current status: '{current}'.")
|
||||
updated = await PolicyRepository(prisma_client).table.update(
|
||||
updated = await _policy_table(prisma_client).update(
|
||||
where={"policy_id": policy_id},
|
||||
data={
|
||||
"version_status": "published",
|
||||
|
|
@ -780,7 +872,7 @@ class PolicyRegistry:
|
|||
raise Exception("Cannot promote draft directly to production. Publish the version first.")
|
||||
|
||||
# Demote current production to published
|
||||
await PolicyRepository(prisma_client).table.update_many(
|
||||
await _policy_table(prisma_client).update_many(
|
||||
where={
|
||||
"policy_name": policy_name,
|
||||
"version_status": "production",
|
||||
|
|
@ -793,7 +885,7 @@ class PolicyRegistry:
|
|||
)
|
||||
|
||||
# Promote this version to production
|
||||
updated = await PolicyRepository(prisma_client).table.update(
|
||||
updated = await _policy_table(prisma_client).update(
|
||||
where={"policy_id": policy_id},
|
||||
data={
|
||||
"version_status": "production",
|
||||
|
|
@ -843,8 +935,8 @@ class PolicyRegistry:
|
|||
PolicyVersionCompareResponse with both versions and field_diffs
|
||||
"""
|
||||
try:
|
||||
a = await PolicyRepository(prisma_client).table.find_unique(where={"policy_id": policy_id_a})
|
||||
b = await PolicyRepository(prisma_client).table.find_unique(where={"policy_id": policy_id_b})
|
||||
a = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id_a})
|
||||
b = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id_b})
|
||||
if a is None:
|
||||
raise Exception(f"Policy {policy_id_a} not found")
|
||||
if b is None:
|
||||
|
|
@ -854,15 +946,15 @@ class PolicyRegistry:
|
|||
resp_b = _row_to_policy_db_response(b)
|
||||
|
||||
# Compare fields that are part of policy content (not metadata)
|
||||
compare_fields = [
|
||||
compare_fields = (
|
||||
"inherit",
|
||||
"description",
|
||||
"guardrails_add",
|
||||
"guardrails_remove",
|
||||
"condition",
|
||||
"pipeline",
|
||||
]
|
||||
field_diffs: Dict[str, Dict[str, Any]] = {}
|
||||
)
|
||||
field_diffs: dict[str, dict[str, object]] = {}
|
||||
for field in compare_fields:
|
||||
val_a = getattr(resp_a, field)
|
||||
val_b = getattr(resp_b, field)
|
||||
|
|
@ -882,7 +974,7 @@ class PolicyRegistry:
|
|||
self,
|
||||
policy_name: str,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> Dict[str, str]:
|
||||
) -> Mapping[str, str]:
|
||||
"""
|
||||
Delete all versions of a policy. Also removes from in-memory registry.
|
||||
|
||||
|
|
@ -894,7 +986,7 @@ class PolicyRegistry:
|
|||
Dict with success message
|
||||
"""
|
||||
try:
|
||||
await PolicyRepository(prisma_client).table.delete_many(where={"policy_name": policy_name})
|
||||
await _policy_table(prisma_client).delete_many(where={"policy_name": policy_name})
|
||||
self.remove_policy(policy_name)
|
||||
return {"message": f"All versions of policy '{policy_name}' deleted successfully"}
|
||||
except Exception as e:
|
||||
|
|
@ -903,7 +995,7 @@ class PolicyRegistry:
|
|||
|
||||
|
||||
# Global singleton instance
|
||||
_policy_registry: Optional[PolicyRegistry] = None
|
||||
_policy_registry: PolicyRegistry | None = None
|
||||
|
||||
|
||||
def get_policy_registry() -> PolicyRegistry:
|
||||
|
|
|
|||
|
|
@ -3,18 +3,37 @@ VerificationToken repository for database operations on LiteLLM_VerificationToke
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Iterator, Mapping
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional, Type
|
||||
from typing import TYPE_CHECKING, Any, Protocol
|
||||
|
||||
from litellm.models.verification_token import (
|
||||
LiteLLM_VerificationToken,
|
||||
)
|
||||
from litellm.repositories.base_repository import BaseRepository
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma.models import (
|
||||
LiteLLM_VerificationToken as PrismaVerificationToken,
|
||||
)
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
class _DictConvertible(Protocol):
|
||||
def dict(self) -> dict[str, object]: ...
|
||||
|
||||
def __iter__(self) -> Iterator[tuple[str, object]]: ...
|
||||
|
||||
|
||||
class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
|
||||
"""Repository for verification token (API key) database operations."""
|
||||
|
||||
@property
|
||||
def prisma_client(self) -> "PrismaClient":
|
||||
prisma_client: PrismaClient = super().prisma_client
|
||||
return prisma_client
|
||||
|
||||
@property
|
||||
def table(self) -> Any:
|
||||
return self.prisma_client.db.litellm_verificationtoken
|
||||
|
|
@ -24,10 +43,10 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
|
|||
return self.prisma_client.db.litellm_deletedverificationtoken
|
||||
|
||||
@property
|
||||
def model_class(self) -> Type[LiteLLM_VerificationToken]:
|
||||
def model_class(self) -> type[LiteLLM_VerificationToken]:
|
||||
return LiteLLM_VerificationToken
|
||||
|
||||
def _to_model(self, record: Any) -> Optional[LiteLLM_VerificationToken]:
|
||||
def _to_model(self, record: _DictConvertible | None) -> LiteLLM_VerificationToken | None:
|
||||
"""Convert a database record to a VerificationToken model."""
|
||||
if record is None:
|
||||
return None
|
||||
|
|
@ -46,42 +65,43 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
|
|||
"litellm_budget_table",
|
||||
]
|
||||
for field in json_fields:
|
||||
if isinstance(data.get(field), str):
|
||||
data[field] = json.loads(data[field])
|
||||
value = data.get(field)
|
||||
if isinstance(value, str):
|
||||
data[field] = json.loads(value)
|
||||
|
||||
if data.get("org_id") is None and data.get("organization_id") is not None:
|
||||
data["org_id"] = data["organization_id"]
|
||||
|
||||
return LiteLLM_VerificationToken(**data)
|
||||
return LiteLLM_VerificationToken.model_validate(data)
|
||||
|
||||
async def find_by_id(self, token: str, id_field: str = "token") -> Optional[LiteLLM_VerificationToken]:
|
||||
async def find_by_id(self, token: str, id_field: str = "token") -> LiteLLM_VerificationToken | None:
|
||||
return await super().find_by_id(token, id_field)
|
||||
|
||||
async def find_by_alias(self, key_alias: str) -> Optional[LiteLLM_VerificationToken]:
|
||||
async def find_by_alias(self, key_alias: str) -> LiteLLM_VerificationToken | None:
|
||||
"""Find a token by key alias."""
|
||||
records = await self.table.find_many(where={"key_alias": key_alias})
|
||||
records: list[PrismaVerificationToken] = await self.table.find_many(where={"key_alias": key_alias})
|
||||
if records:
|
||||
return self._to_model(records[0])
|
||||
return None
|
||||
|
||||
async def find_by_user_id(self, user_id: str) -> List[LiteLLM_VerificationToken]:
|
||||
async def find_by_user_id(self, user_id: str) -> list[LiteLLM_VerificationToken]:
|
||||
"""Find all tokens belonging to a user."""
|
||||
records = await self.table.find_many(where={"user_id": user_id})
|
||||
records: list[PrismaVerificationToken] = await self.table.find_many(where={"user_id": user_id})
|
||||
return self._to_model_list(records)
|
||||
|
||||
async def find_by_team_id(self, team_id: str) -> List[LiteLLM_VerificationToken]:
|
||||
async def find_by_team_id(self, team_id: str) -> list[LiteLLM_VerificationToken]:
|
||||
"""Find all tokens belonging to a team."""
|
||||
records = await self.table.find_many(where={"team_id": team_id})
|
||||
records: list[PrismaVerificationToken] = await self.table.find_many(where={"team_id": team_id})
|
||||
return self._to_model_list(records)
|
||||
|
||||
async def find_by_project_id(self, project_id: str) -> List[LiteLLM_VerificationToken]:
|
||||
async def find_by_project_id(self, project_id: str) -> list[LiteLLM_VerificationToken]:
|
||||
"""Find all tokens belonging to a project."""
|
||||
records = await self.table.find_many(where={"project_id": project_id})
|
||||
records: list[PrismaVerificationToken] = await self.table.find_many(where={"project_id": project_id})
|
||||
return self._to_model_list(records)
|
||||
|
||||
async def find_active_tokens(self) -> List[LiteLLM_VerificationToken]:
|
||||
async def find_active_tokens(self) -> list[LiteLLM_VerificationToken]:
|
||||
"""Find all active (non-expired, non-blocked) tokens."""
|
||||
records = await self.table.find_many(
|
||||
records: list[PrismaVerificationToken] = await self.table.find_many(
|
||||
where={
|
||||
"blocked": {"not": True},
|
||||
"OR": [{"expires": None}, {"expires": {"gt": datetime.utcnow()}}],
|
||||
|
|
@ -92,31 +112,31 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
|
|||
def _build_token_data(
|
||||
self,
|
||||
token: str,
|
||||
key_name: Optional[str] = None,
|
||||
key_alias: Optional[str] = None,
|
||||
max_budget: Optional[float] = None,
|
||||
expires: Optional[datetime] = None,
|
||||
models: Optional[List[str]] = None,
|
||||
aliases: Optional[Dict[str, str]] = None,
|
||||
config: Optional[Dict[str, Any]] = None,
|
||||
user_id: Optional[str] = None,
|
||||
team_id: Optional[str] = None,
|
||||
agent_id: Optional[str] = None,
|
||||
project_id: Optional[str] = None,
|
||||
max_parallel_requests: Optional[int] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
tpm_limit: Optional[int] = None,
|
||||
rpm_limit: Optional[int] = None,
|
||||
budget_duration: Optional[str] = None,
|
||||
allowed_cache_controls: Optional[List[str]] = None,
|
||||
allowed_routes: Optional[List[str]] = None,
|
||||
permissions: Optional[Dict[str, Any]] = None,
|
||||
org_id: Optional[str] = None,
|
||||
created_by: Optional[str] = None,
|
||||
object_permission_id: Optional[str] = None,
|
||||
access_group_ids: Optional[List[str]] = None,
|
||||
budget_id: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
key_name: str | None = None,
|
||||
key_alias: str | None = None,
|
||||
max_budget: float | None = None,
|
||||
expires: datetime | None = None,
|
||||
models: list[str] | None = None,
|
||||
aliases: dict[str, str] | None = None,
|
||||
config: Mapping[str, object] | None = None,
|
||||
user_id: str | None = None,
|
||||
team_id: str | None = None,
|
||||
agent_id: str | None = None,
|
||||
project_id: str | None = None,
|
||||
max_parallel_requests: int | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
tpm_limit: int | None = None,
|
||||
rpm_limit: int | None = None,
|
||||
budget_duration: str | None = None,
|
||||
allowed_cache_controls: list[str] | None = None,
|
||||
allowed_routes: list[str] | None = None,
|
||||
permissions: Mapping[str, object] | None = None,
|
||||
org_id: str | None = None,
|
||||
created_by: str | None = None,
|
||||
object_permission_id: str | None = None,
|
||||
access_group_ids: list[str] | None = None,
|
||||
budget_id: str | None = None,
|
||||
) -> dict[str, object]:
|
||||
"""Build data dictionary for token creation."""
|
||||
json_fields = {
|
||||
"aliases": aliases,
|
||||
|
|
@ -145,7 +165,7 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
|
|||
"access_group_ids": access_group_ids,
|
||||
"budget_id": budget_id,
|
||||
}
|
||||
data: Dict[str, Any] = {k: v for k, v in simple_fields.items() if v is not None}
|
||||
data: dict[str, object] = {k: v for k, v in simple_fields.items() if v is not None}
|
||||
for key, val in json_fields.items():
|
||||
if val is not None:
|
||||
data[key] = json.dumps(val)
|
||||
|
|
@ -159,30 +179,30 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
|
|||
async def create_token(
|
||||
self,
|
||||
token: str,
|
||||
key_name: Optional[str] = None,
|
||||
key_alias: Optional[str] = None,
|
||||
max_budget: Optional[float] = None,
|
||||
expires: Optional[datetime] = None,
|
||||
models: Optional[List[str]] = None,
|
||||
aliases: Optional[Dict[str, str]] = None,
|
||||
config: Optional[Dict[str, Any]] = None,
|
||||
user_id: Optional[str] = None,
|
||||
team_id: Optional[str] = None,
|
||||
agent_id: Optional[str] = None,
|
||||
project_id: Optional[str] = None,
|
||||
max_parallel_requests: Optional[int] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
tpm_limit: Optional[int] = None,
|
||||
rpm_limit: Optional[int] = None,
|
||||
budget_duration: Optional[str] = None,
|
||||
allowed_cache_controls: Optional[List[str]] = None,
|
||||
allowed_routes: Optional[List[str]] = None,
|
||||
permissions: Optional[Dict[str, Any]] = None,
|
||||
org_id: Optional[str] = None,
|
||||
created_by: Optional[str] = None,
|
||||
object_permission_id: Optional[str] = None,
|
||||
access_group_ids: Optional[List[str]] = None,
|
||||
budget_id: Optional[str] = None,
|
||||
key_name: str | None = None,
|
||||
key_alias: str | None = None,
|
||||
max_budget: float | None = None,
|
||||
expires: datetime | None = None,
|
||||
models: list[str] | None = None,
|
||||
aliases: dict[str, str] | None = None,
|
||||
config: Mapping[str, object] | None = None,
|
||||
user_id: str | None = None,
|
||||
team_id: str | None = None,
|
||||
agent_id: str | None = None,
|
||||
project_id: str | None = None,
|
||||
max_parallel_requests: int | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
tpm_limit: int | None = None,
|
||||
rpm_limit: int | None = None,
|
||||
budget_duration: str | None = None,
|
||||
allowed_cache_controls: list[str] | None = None,
|
||||
allowed_routes: list[str] | None = None,
|
||||
permissions: Mapping[str, object] | None = None,
|
||||
org_id: str | None = None,
|
||||
created_by: str | None = None,
|
||||
object_permission_id: str | None = None,
|
||||
access_group_ids: list[str] | None = None,
|
||||
budget_id: str | None = None,
|
||||
) -> LiteLLM_VerificationToken:
|
||||
"""Create a new verification token."""
|
||||
data = self._build_token_data(
|
||||
|
|
@ -217,28 +237,28 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
|
|||
async def update_token(
|
||||
self,
|
||||
token: str,
|
||||
updated_by: Optional[str] = None,
|
||||
key_name: Optional[str] = None,
|
||||
key_alias: Optional[str] = None,
|
||||
max_budget: Optional[float] = None,
|
||||
expires: Optional[datetime] = None,
|
||||
models: Optional[List[str]] = None,
|
||||
aliases: Optional[Dict[str, str]] = None,
|
||||
config: Optional[Dict[str, Any]] = None,
|
||||
max_parallel_requests: Optional[int] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
tpm_limit: Optional[int] = None,
|
||||
rpm_limit: Optional[int] = None,
|
||||
budget_duration: Optional[str] = None,
|
||||
allowed_cache_controls: Optional[List[str]] = None,
|
||||
allowed_routes: Optional[List[str]] = None,
|
||||
permissions: Optional[Dict[str, Any]] = None,
|
||||
blocked: Optional[bool] = None,
|
||||
object_permission_id: Optional[str] = None,
|
||||
access_group_ids: Optional[List[str]] = None,
|
||||
) -> Optional[LiteLLM_VerificationToken]:
|
||||
updated_by: str | None = None,
|
||||
key_name: str | None = None,
|
||||
key_alias: str | None = None,
|
||||
max_budget: float | None = None,
|
||||
expires: datetime | None = None,
|
||||
models: list[str] | None = None,
|
||||
aliases: dict[str, str] | None = None,
|
||||
config: Mapping[str, object] | None = None,
|
||||
max_parallel_requests: int | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
tpm_limit: int | None = None,
|
||||
rpm_limit: int | None = None,
|
||||
budget_duration: str | None = None,
|
||||
allowed_cache_controls: list[str] | None = None,
|
||||
allowed_routes: list[str] | None = None,
|
||||
permissions: Mapping[str, object] | None = None,
|
||||
blocked: bool | None = None,
|
||||
object_permission_id: str | None = None,
|
||||
access_group_ids: list[str] | None = None,
|
||||
) -> LiteLLM_VerificationToken | None:
|
||||
"""Update a verification token."""
|
||||
data: Dict[str, Any] = {}
|
||||
data: dict[str, object] = {}
|
||||
if updated_by is not None:
|
||||
data["updated_by"] = updated_by
|
||||
if key_name is not None:
|
||||
|
|
@ -283,10 +303,10 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
|
|||
async def delete_token(
|
||||
self,
|
||||
token: str,
|
||||
deleted_by: Optional[str] = None,
|
||||
deleted_by_api_key: Optional[str] = None,
|
||||
litellm_changed_by: Optional[str] = None,
|
||||
) -> Optional[LiteLLM_VerificationToken]:
|
||||
deleted_by: str | None = None,
|
||||
deleted_by_api_key: str | None = None,
|
||||
litellm_changed_by: str | None = None,
|
||||
) -> LiteLLM_VerificationToken | None:
|
||||
"""Delete a token and archive it to the deleted tokens table.
|
||||
|
||||
Uses a transaction to ensure atomicity of the archive-then-delete operation.
|
||||
|
|
@ -307,14 +327,14 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
|
|||
|
||||
return token_record
|
||||
|
||||
def _build_archive_data(self, token: LiteLLM_VerificationToken) -> Dict[str, Any]:
|
||||
def _build_archive_data(self, token: LiteLLM_VerificationToken) -> dict[str, object]:
|
||||
"""Build archive data with only columns present in LiteLLM_DeletedVerificationToken.
|
||||
|
||||
Serializes JSON columns to strings (the archive table stores them as JSON
|
||||
columns the same way the live table does) and maps ``org_id`` onto the
|
||||
``organization_id`` column so the foreign key is preserved.
|
||||
"""
|
||||
data = token.model_dump(exclude_none=True)
|
||||
data: dict[str, object] = token.model_dump(exclude_none=True)
|
||||
for field in ("object_permission", "litellm_budget_table", "budget_limits"):
|
||||
data.pop(field, None)
|
||||
|
||||
|
|
@ -336,24 +356,24 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
|
|||
data[field] = json.dumps(data[field])
|
||||
return data
|
||||
|
||||
async def update_spend(self, token: str, spend: float) -> Optional[LiteLLM_VerificationToken]:
|
||||
async def update_spend(self, token: str, spend: float) -> LiteLLM_VerificationToken | None:
|
||||
"""Update token spend."""
|
||||
return await self.update(token, {"spend": spend}, id_field="token")
|
||||
|
||||
async def update_last_active(self, token: str) -> Optional[LiteLLM_VerificationToken]:
|
||||
async def update_last_active(self, token: str) -> LiteLLM_VerificationToken | None:
|
||||
"""Update the last_active timestamp."""
|
||||
return await self.update(token, {"last_active": datetime.utcnow()}, id_field="token")
|
||||
|
||||
async def block_token(self, token: str, updated_by: Optional[str] = None) -> Optional[LiteLLM_VerificationToken]:
|
||||
async def block_token(self, token: str, updated_by: str | None = None) -> LiteLLM_VerificationToken | None:
|
||||
"""Block a token."""
|
||||
data: Dict[str, Any] = {"blocked": True}
|
||||
data: dict[str, object] = {"blocked": True}
|
||||
if updated_by is not None:
|
||||
data["updated_by"] = updated_by
|
||||
return await self.update(token, data, id_field="token")
|
||||
|
||||
async def unblock_token(self, token: str, updated_by: Optional[str] = None) -> Optional[LiteLLM_VerificationToken]:
|
||||
async def unblock_token(self, token: str, updated_by: str | None = None) -> LiteLLM_VerificationToken | None:
|
||||
"""Unblock a token."""
|
||||
data: Dict[str, Any] = {"blocked": False}
|
||||
data: dict[str, object] = {"blocked": False}
|
||||
if updated_by is not None:
|
||||
data["updated_by"] = updated_by
|
||||
return await self.update(token, data, id_field="token")
|
||||
|
|
|
|||
|
|
@ -1,12 +1,12 @@
|
|||
{
|
||||
"ANN001": {
|
||||
"limit": 3152
|
||||
"limit": 3142
|
||||
},
|
||||
"ANN002": {
|
||||
"limit": 69
|
||||
},
|
||||
"ANN003": {
|
||||
"limit": 835
|
||||
"limit": 831
|
||||
},
|
||||
"ANN201": {
|
||||
"limit": 2138
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 130
|
||||
},
|
||||
"ANN401": {
|
||||
"limit": 2075
|
||||
"limit": 2015
|
||||
},
|
||||
"ASYNC230": {
|
||||
"limit": 14
|
||||
|
|
@ -123,7 +123,7 @@
|
|||
"limit": 52
|
||||
},
|
||||
"I001": {
|
||||
"limit": 273
|
||||
"limit": 270
|
||||
},
|
||||
"LOG015": {
|
||||
"limit": 8
|
||||
|
|
@ -135,7 +135,7 @@
|
|||
"limit": 30
|
||||
},
|
||||
"PERF401": {
|
||||
"limit": 146
|
||||
"limit": 144
|
||||
},
|
||||
"PERF402": {
|
||||
"limit": 9
|
||||
|
|
@ -222,7 +222,7 @@
|
|||
"limit": 38
|
||||
},
|
||||
"RET504": {
|
||||
"limit": 719
|
||||
"limit": 717
|
||||
},
|
||||
"RUF010": {
|
||||
"limit": 874
|
||||
|
|
@ -306,7 +306,7 @@
|
|||
"limit": 9
|
||||
},
|
||||
"TID251": {
|
||||
"limit": 2701
|
||||
"limit": 2652
|
||||
},
|
||||
"TRY002": {
|
||||
"limit": 548
|
||||
|
|
@ -324,10 +324,10 @@
|
|||
"limit": 883
|
||||
},
|
||||
"UP006": {
|
||||
"limit": 12789
|
||||
"limit": 12147
|
||||
},
|
||||
"UP007": {
|
||||
"limit": 2570
|
||||
"limit": 2526
|
||||
},
|
||||
"UP008": {
|
||||
"limit": 5
|
||||
|
|
@ -354,7 +354,7 @@
|
|||
"limit": 4
|
||||
},
|
||||
"UP035": {
|
||||
"limit": 2284
|
||||
"limit": 2232
|
||||
},
|
||||
"UP036": {
|
||||
"limit": 4
|
||||
|
|
@ -363,6 +363,6 @@
|
|||
"limit": 105
|
||||
},
|
||||
"UP045": {
|
||||
"limit": 18461
|
||||
"limit": 17824
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 23408
|
||||
"limit": 23287
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 27511
|
||||
"limit": 27473
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 292
|
||||
|
|
@ -15,7 +15,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT006": {
|
||||
"limit": 1111
|
||||
"limit": 1109
|
||||
},
|
||||
"LIT007": {
|
||||
"limit": 0
|
||||
|
|
@ -24,6 +24,6 @@
|
|||
"limit": 1004
|
||||
},
|
||||
"LIT009": {
|
||||
"limit": 2501
|
||||
"limit": 2495
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue