mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
merge: sync OCR cutover with staging
Some checks failed
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled
Some checks failed
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled
Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
commit
5846128476
67 changed files with 4337 additions and 2355 deletions
1
.github/workflows/test-linting.yml
vendored
1
.github/workflows/test-linting.yml
vendored
|
|
@ -104,6 +104,7 @@ jobs:
|
|||
- name: Check basedpyright budget (delta vs base)
|
||||
env:
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
NODE_OPTIONS: --max-old-space-size=12288
|
||||
run: |
|
||||
(uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py --base "$BASE_SHA"
|
||||
|
||||
|
|
|
|||
|
|
@ -70,6 +70,10 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/project/",
|
||||
"/memory/",
|
||||
"/mcp/",
|
||||
# Control plane (see the List Endpoints + Tables standard). Every resource
|
||||
# eventually moves under this prefix, so allowlist it once rather than
|
||||
# per-resource.
|
||||
"/management/v1/",
|
||||
# Spend / analytics
|
||||
"/spend/",
|
||||
"/analytics/",
|
||||
|
|
|
|||
|
|
@ -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}},
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from typing import (
|
|||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Sequence,
|
||||
Tuple,
|
||||
|
|
@ -1449,6 +1450,8 @@ class PrometheusLogger(CustomLogger):
|
|||
prompt_details = usage_object.get("prompt_tokens_details") or {}
|
||||
completion_details = usage_object.get("completion_tokens_details") or {}
|
||||
|
||||
cache_creation_detail_tokens = PrometheusLogger._resolve_cache_write_tokens(prompt_details)
|
||||
|
||||
detail_metrics: List[Tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]] = [
|
||||
(
|
||||
self.litellm_input_cached_tokens_metric,
|
||||
|
|
@ -1458,7 +1461,7 @@ class PrometheusLogger(CustomLogger):
|
|||
(
|
||||
self.litellm_input_cache_creation_tokens_metric,
|
||||
"litellm_input_cache_creation_tokens_metric",
|
||||
(prompt_details.get("cache_creation_tokens") if isinstance(prompt_details, dict) else None),
|
||||
cache_creation_detail_tokens,
|
||||
),
|
||||
(
|
||||
self.litellm_input_audio_tokens_metric,
|
||||
|
|
@ -1597,27 +1600,12 @@ class PrometheusLogger(CustomLogger):
|
|||
)
|
||||
|
||||
# Provider prompt caching metrics are independent of LiteLLM cache_hit.
|
||||
provider_cache_read_tokens = 0
|
||||
provider_cache_creation_tokens = 0
|
||||
usage_obj = (standard_logging_payload.get("metadata", {}) or {}).get("usage_object")
|
||||
if isinstance(usage_obj, dict):
|
||||
# Prefer explicit provider cache fields when available.
|
||||
_read = usage_obj.get("cache_read_input_tokens")
|
||||
_write = usage_obj.get("cache_creation_input_tokens")
|
||||
|
||||
if isinstance(_read, int):
|
||||
provider_cache_read_tokens = _read
|
||||
if isinstance(_write, int):
|
||||
provider_cache_creation_tokens = _write
|
||||
|
||||
# Fallback to prompt_tokens_details.cached_tokens (common normalization point).
|
||||
# Only fallback when the explicit field is genuinely absent (None).
|
||||
if _read is None:
|
||||
prompt_details = usage_obj.get("prompt_tokens_details")
|
||||
if isinstance(prompt_details, dict):
|
||||
cached_tokens = prompt_details.get("cached_tokens")
|
||||
if isinstance(cached_tokens, int):
|
||||
provider_cache_read_tokens = cached_tokens
|
||||
(
|
||||
provider_cache_read_tokens,
|
||||
provider_cache_creation_tokens,
|
||||
) = PrometheusLogger._resolve_provider_cache_tokens(usage_obj)
|
||||
|
||||
if provider_cache_read_tokens > 0:
|
||||
PrometheusLogger._inc_labeled_counter(
|
||||
|
|
@ -1639,6 +1627,40 @@ class PrometheusLogger(CustomLogger):
|
|||
amount=float(provider_cache_creation_tokens),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_provider_cache_tokens(usage_obj: Mapping[str, object]) -> tuple[int, int]:
|
||||
# Prefer explicit provider cache fields when available.
|
||||
_read = usage_obj.get("cache_read_input_tokens")
|
||||
_write = usage_obj.get("cache_creation_input_tokens")
|
||||
|
||||
provider_cache_read_tokens = _read if isinstance(_read, int) else 0
|
||||
provider_cache_creation_tokens = _write if isinstance(_write, int) else 0
|
||||
|
||||
# Fallback to prompt_tokens_details (common normalization point).
|
||||
# Only fallback when the explicit field is genuinely absent (None).
|
||||
prompt_details = usage_obj.get("prompt_tokens_details")
|
||||
if _read is None and isinstance(prompt_details, dict):
|
||||
cached_tokens = prompt_details.get("cached_tokens")
|
||||
if isinstance(cached_tokens, int):
|
||||
provider_cache_read_tokens = cached_tokens
|
||||
|
||||
if _write is None:
|
||||
write_tokens = PrometheusLogger._resolve_cache_write_tokens(prompt_details)
|
||||
if write_tokens is not None:
|
||||
provider_cache_creation_tokens = write_tokens
|
||||
|
||||
return provider_cache_read_tokens, provider_cache_creation_tokens
|
||||
|
||||
@staticmethod
|
||||
def _resolve_cache_write_tokens(prompt_details: object) -> int | None:
|
||||
if not isinstance(prompt_details, dict):
|
||||
return None
|
||||
for key in ("cache_write_tokens", "cache_creation_tokens"):
|
||||
value = prompt_details.get(key)
|
||||
if isinstance(value, int) and not isinstance(value, bool):
|
||||
return value
|
||||
return None
|
||||
|
||||
def _increment_mcp_tool_call_metrics(
|
||||
self,
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -6,8 +6,9 @@ import base64
|
|||
import mimetypes
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Coroutine, Mapping
|
||||
from io import IOBase
|
||||
from typing import Any, Coroutine, Union, cast
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
from typing_extensions import Never
|
||||
|
|
@ -109,7 +110,7 @@ def _raise_ocr_exception(
|
|||
|
||||
|
||||
def _timeout_to_seconds(
|
||||
timeout: Union[float, httpx.Timeout] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> float | None:
|
||||
"""Convert the Python OCR timeout to a single seconds value for the Rust bridge.
|
||||
|
||||
|
|
@ -126,22 +127,22 @@ def _timeout_to_seconds(
|
|||
|
||||
def _resolve_ocr_call_context(
|
||||
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],
|
||||
) -> tuple[
|
||||
str,
|
||||
dict[str, Any],
|
||||
dict[str, object],
|
||||
str | None,
|
||||
str | None,
|
||||
str,
|
||||
dict[str, object] | None,
|
||||
dict[str, object],
|
||||
Union[float, httpx.Timeout],
|
||||
float | httpx.Timeout,
|
||||
LiteLLMLoggingObj,
|
||||
]:
|
||||
litellm_logging_obj = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj"))
|
||||
|
|
@ -208,11 +209,11 @@ def _resolve_ocr_call_context(
|
|||
|
||||
return (
|
||||
model,
|
||||
document,
|
||||
cast(dict[str, object], document),
|
||||
api_key,
|
||||
api_base,
|
||||
custom_llm_provider,
|
||||
cast(dict[str, object] | None, extra_headers),
|
||||
extra_headers,
|
||||
cast(dict[str, object], optional_params),
|
||||
effective_timeout,
|
||||
litellm_logging_obj,
|
||||
|
|
@ -252,7 +253,7 @@ def _run_rust_ocr(
|
|||
custom_llm_provider: str,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
timeout: float | httpx.Timeout,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
) -> OCRResponse:
|
||||
_run_pre_call_logging(
|
||||
|
|
@ -293,7 +294,7 @@ async def _run_rust_aocr(
|
|||
custom_llm_provider: str,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
timeout: float | httpx.Timeout,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
) -> OCRResponse:
|
||||
_run_pre_call_logging(
|
||||
|
|
@ -322,13 +323,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.
|
||||
|
|
@ -580,14 +581,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
|
|
@ -632,7 +632,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
# Reads end users out of spend logs, scoped to the caller's own rows and
|
||||
# permitted teams exactly like /spend/logs/ui — it belongs to the same
|
||||
# access tier, not to customer management.
|
||||
"/customer/aliases",
|
||||
"/management/v1/spend_logs/end_users",
|
||||
"/cost/estimate",
|
||||
]
|
||||
|
||||
|
|
@ -822,12 +822,13 @@ class LiteLLMRoutes(enum.Enum):
|
|||
# Customer / end-user listing (handlers already gate on
|
||||
# PROXY_ADMIN_VIEW_ONLY — the route gate must match).
|
||||
"/customer/list",
|
||||
"/customer/aliases",
|
||||
"/customer/info",
|
||||
# UI Logs page detail drawer (single + session). The list endpoint
|
||||
# `/spend/logs/ui` is covered via spend_tracking_routes below.
|
||||
# UI Logs page detail drawer (single + session) and the end-user filter
|
||||
# facet. The list endpoint `/spend/logs/ui` is covered via
|
||||
# spend_tracking_routes below.
|
||||
"/spend/logs/ui/{logId}",
|
||||
"/spend/logs/session/ui",
|
||||
"/management/v1/spend_logs/end_users",
|
||||
# Settings / observability read endpoints exposed in admin-only
|
||||
# sidebar groups (Logging & Alerts, Admin Settings, Budgets,
|
||||
# Invitations).
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -47,6 +47,11 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.content_text import (
|
||||
content_to_text,
|
||||
is_all_text_parts,
|
||||
merge_rewritten_text_parts,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.guardrails import GuardrailEventHooks, Mode
|
||||
from litellm.types.integrations.custom_logger import (
|
||||
|
|
@ -144,48 +149,20 @@ def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isin
|
|||
return isinstance(value, list)
|
||||
|
||||
|
||||
def _content_to_text(content: object) -> str:
|
||||
"""Collapse a message ``content`` (str or list-of-parts) to plain text.
|
||||
|
||||
For the multimodal list shape, joins ``{type: "text", text: ...}`` parts
|
||||
with blank-line separators; non-text parts are ignored.
|
||||
"""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
parts: list[str] = []
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
text = part.get("text")
|
||||
if isinstance(text, str):
|
||||
parts.append(text)
|
||||
return "\n\n".join(parts)
|
||||
return ""
|
||||
|
||||
|
||||
def _replace_text_in_content(content: object, new_text: str) -> object:
|
||||
"""Write ``new_text`` back into a ``content`` value, preserving shape.
|
||||
|
||||
``str`` content is replaced directly. For list-of-parts content the first
|
||||
text part carries ``new_text``, later text parts are dropped, and
|
||||
non-text parts (images, audio, files) pass through untouched.
|
||||
``str`` content is replaced directly. An all-text part list collapses to a
|
||||
single part carrying the last declared cache_control breakpoint. Anything
|
||||
else is returned unchanged: breakpoints are positional, so one compressed
|
||||
string cannot be written back across a non-text part without moving text
|
||||
to the other side of it.
|
||||
"""
|
||||
if isinstance(content, str):
|
||||
return new_text
|
||||
if isinstance(content, list):
|
||||
out: list[object] = []
|
||||
replaced = False
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
if not replaced:
|
||||
out.append({**part, "text": new_text})
|
||||
replaced = True
|
||||
continue
|
||||
out.append(part)
|
||||
if not replaced:
|
||||
out.insert(0, {"type": "text", "text": new_text})
|
||||
return out
|
||||
return new_text
|
||||
if _is_object_list(content) and is_all_text_parts(content):
|
||||
return merge_rewritten_text_parts(content, new_text)
|
||||
return content
|
||||
|
||||
|
||||
def _render_tool_intent(fn: dict[str, object]) -> str:
|
||||
|
|
@ -422,7 +399,7 @@ def _assistant_text_from_response(response: object) -> str | None:
|
|||
if isinstance(choices, list) and choices:
|
||||
message = get_attribute_or_key(choices[0], "message", None)
|
||||
if message is not None:
|
||||
text = _content_to_text(get_attribute_or_key(message, "content", None))
|
||||
text = content_to_text(get_attribute_or_key(message, "content", None))
|
||||
if text:
|
||||
return text
|
||||
content = get_attribute_or_key(response, "content", None)
|
||||
|
|
@ -905,7 +882,10 @@ class CompresrGuardrail(CustomGuardrail):
|
|||
continue
|
||||
else:
|
||||
continue
|
||||
if len(_content_to_text(msg.get("content"))) < self.min_chars_to_compress:
|
||||
content = msg.get("content")
|
||||
if _is_object_list(content) and not is_all_text_parts(content):
|
||||
continue
|
||||
if len(content_to_text(content)) < self.min_chars_to_compress:
|
||||
continue
|
||||
targets.append(idx)
|
||||
return targets
|
||||
|
|
@ -916,7 +896,7 @@ class CompresrGuardrail(CustomGuardrail):
|
|||
) -> tuple[str, int | None]:
|
||||
for idx in range(len(messages) - 1, -1, -1):
|
||||
if messages[idx].get("role") == "user":
|
||||
return _content_to_text(messages[idx].get("content")), idx
|
||||
return content_to_text(messages[idx].get("content")), idx
|
||||
return "", None
|
||||
|
||||
def _apply_compression_results(
|
||||
|
|
@ -1034,7 +1014,7 @@ class CompresrGuardrail(CustomGuardrail):
|
|||
verbose_proxy_logger.debug("Compresr: no messages eligible for compression")
|
||||
return inputs
|
||||
|
||||
contexts = [_content_to_text(messages[idx].get("content")) for idx in targets]
|
||||
contexts = [content_to_text(messages[idx].get("content")) for idx in targets]
|
||||
|
||||
start_time = time.monotonic()
|
||||
results = await self._call_compress(contexts=contexts, queries=queries)
|
||||
|
|
|
|||
55
litellm/proxy/guardrails/guardrail_hooks/content_text.py
Normal file
55
litellm/proxy/guardrails/guardrail_hooks/content_text.py
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
"""Shared content-part helpers for compression guardrails (headroom, compresr).
|
||||
|
||||
Compression services only transform plain-string message content: every
|
||||
transform in the service pipeline gates on ``isinstance(content, str)`` and
|
||||
silently skips the OpenAI list-of-parts shape. Guardrails that send messages
|
||||
to such a service collapse text-bearing part lists to strings here, and write
|
||||
the rewritten text back through ``merge_rewritten_text_parts``.
|
||||
|
||||
Anthropic ``cache_control`` breakpoints are positional: each one caches the
|
||||
prefix ending at the part that carries it. A single compressed string can
|
||||
therefore only be written back over a run of text parts, never across a
|
||||
non-text part, which is what ``is_all_text_parts`` gates.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
|
||||
def content_to_text(content: object) -> str:
|
||||
"""Collapse a message ``content`` (str or list-of-parts) to plain text.
|
||||
|
||||
For the multimodal list shape, joins ``{type: "text", text: ...}`` parts
|
||||
with blank-line separators; non-text parts are ignored.
|
||||
"""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
parts: list[str] = []
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
text = part.get("text")
|
||||
if isinstance(text, str):
|
||||
parts.append(text)
|
||||
return "\n\n".join(parts)
|
||||
return ""
|
||||
|
||||
|
||||
def is_all_text_parts(content: object) -> bool:
|
||||
"""True when ``content`` is a non-empty part list holding only text parts."""
|
||||
if not isinstance(content, list) or not content:
|
||||
return False
|
||||
return all(isinstance(part, dict) and part.get("type") == "text" for part in content)
|
||||
|
||||
|
||||
def merge_rewritten_text_parts(parts: Sequence[object], new_text: str) -> list[object]:
|
||||
"""Collapse a rewritten all-text part list into one part carrying ``new_text``.
|
||||
|
||||
Only all-text rows are ever flattened, so the merged part IS the whole row:
|
||||
it keeps the first part's fields and the LAST declared cache_control
|
||||
breakpoint. A breakpoint caches the prefix ending at its part, so after the
|
||||
merge the last one (and its TTL) is the one that still describes the row.
|
||||
"""
|
||||
dict_parts = tuple(part for part in parts if isinstance(part, dict))
|
||||
breakpoints = tuple(part["cache_control"] for part in dict_parts if part.get("cache_control") is not None)
|
||||
base = {**dict_parts[0], "text": new_text} if dict_parts else {"type": "text", "text": new_text}
|
||||
return [{**base, "cache_control": breakpoints[-1]} if breakpoints else base]
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -10,12 +10,11 @@ All /customer management endpoints
|
|||
"""
|
||||
|
||||
#### END-USER/CUSTOMER MANAGEMENT ####
|
||||
from collections.abc import MutableSequence
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Annotated, Any, List, Optional
|
||||
from datetime import datetime, timedelta
|
||||
from typing import List, Optional
|
||||
|
||||
import fastapi
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
|
|
@ -28,7 +27,7 @@ from litellm.proxy.management_helpers.object_permission_utils import (
|
|||
_set_object_permission,
|
||||
handle_update_object_permission_common,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy
|
||||
from litellm.proxy.utils import handle_exception_on_proxy
|
||||
from litellm.repositories.budget_repository import BudgetRepository
|
||||
from litellm.repositories.table_repositories import EndUserRepository
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
||||
|
|
@ -36,7 +35,6 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
|||
)
|
||||
from litellm.types.proxy.management_endpoints.customer_endpoints import (
|
||||
BlockUsersResponse,
|
||||
CustomerAliasesResponse,
|
||||
CustomerResponse,
|
||||
DeleteCustomersResponse,
|
||||
UnblockUsersResponse,
|
||||
|
|
@ -44,11 +42,6 @@ from litellm.types.proxy.management_endpoints.customer_endpoints import (
|
|||
|
||||
router = APIRouter()
|
||||
|
||||
# Rows the end-user filter query may read out of LiteLLM_SpendLogs before DISTINCT.
|
||||
# Matches SPEND_LOGS_PAGINATION_COUNT_CAP, the equivalent bound ui_view_spend_logs
|
||||
# puts on its count query, so both reads of the same table stop at the same depth.
|
||||
SPEND_LOGS_FILTER_SCAN_CAP = 10000
|
||||
|
||||
|
||||
def _to_customer_response(record: BaseModel) -> CustomerResponse:
|
||||
"""Validate a raw end-user DB row into the typed customer response.
|
||||
|
|
@ -792,168 +785,6 @@ async def list_end_user(
|
|||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
def _parse_spend_log_window_bound(value: str, param: str) -> datetime:
|
||||
try:
|
||||
return datetime.strptime(value.strip(), "%Y-%m-%d %H:%M:%S").replace(tzinfo=timezone.utc)
|
||||
except ValueError:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": f"Invalid {param}: {value}. Expected 'YYYY-MM-DD HH:MM:SS'"},
|
||||
)
|
||||
|
||||
|
||||
async def _build_end_user_scope_condition(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient,
|
||||
query_params: MutableSequence[Any],
|
||||
) -> str | None:
|
||||
"""SQL predicate restricting end users to the logs this caller may read.
|
||||
|
||||
Returns None when the caller is a proxy admin (no restriction). Mirrors the
|
||||
scoping ``/spend/logs/ui`` applies, so the dropdown can never offer an
|
||||
end user whose rows the caller could not open.
|
||||
"""
|
||||
from litellm.proxy.spend_tracking.spend_management_endpoints import (
|
||||
_get_permitted_team_ids_for_spend_logs,
|
||||
_is_admin_view_safe,
|
||||
)
|
||||
|
||||
if _is_admin_view_safe(user_api_key_dict=user_api_key_dict):
|
||||
return None
|
||||
|
||||
try:
|
||||
permitted_team_ids = await _get_permitted_team_ids_for_spend_logs(
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
except Exception:
|
||||
permitted_team_ids = []
|
||||
|
||||
caller_user_id = user_api_key_dict.user_id
|
||||
user_clause: tuple[str, ...] = ()
|
||||
if caller_user_id is not None:
|
||||
query_params.append(caller_user_id)
|
||||
user_clause = (f'"user" = ${len(query_params)}',)
|
||||
|
||||
team_clause: tuple[str, ...] = ()
|
||||
if permitted_team_ids:
|
||||
# = ANY(::text[]) rather than an expanded IN list, matching the clause
|
||||
# ui_view_spend_logs builds: one parameter whatever the team count.
|
||||
query_params.append(permitted_team_ids)
|
||||
team_clause = (f"team_id = ANY(${len(query_params)}::text[])",)
|
||||
|
||||
scope_parts = user_clause + team_clause
|
||||
if not scope_parts:
|
||||
return "FALSE"
|
||||
return f"({' OR '.join(scope_parts)})"
|
||||
|
||||
|
||||
@router.get(
|
||||
"/customer/aliases",
|
||||
tags=["Customer Management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=CustomerAliasesResponse,
|
||||
)
|
||||
async def list_customer_aliases(
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
start_date: Annotated[str, Query(description="Window start, 'YYYY-MM-DD HH:MM:SS' (UTC)")],
|
||||
end_date: Annotated[str, Query(description="Window end, 'YYYY-MM-DD HH:MM:SS' (UTC)")],
|
||||
page: Annotated[int, Query(ge=1, description="Page number")] = 1,
|
||||
size: Annotated[int, Query(ge=1, le=100, description="Page size")] = 50,
|
||||
search: Annotated[
|
||||
str | None,
|
||||
Query(description="Case-insensitive partial match on the customer id"),
|
||||
] = None,
|
||||
) -> CustomerAliasesResponse:
|
||||
"""
|
||||
List the end users seen in spend logs over a time window, for UI filter dropdowns.
|
||||
|
||||
Scoped like `/spend/logs/ui`: a proxy admin sees every end user in the window,
|
||||
anyone else sees only end users from their own requests or from teams they
|
||||
administer (or hold the `/spend/logs` permission on).
|
||||
|
||||
Reads spend logs rather than LiteLLM_EndUserTable because only spend logs carry
|
||||
the team attribution this scoping needs. The window is required and the inner
|
||||
scan is capped at SPEND_LOGS_FILTER_SCAN_CAP rows, so the query
|
||||
cannot degrade into a full-table scan the way `/global/all_end_users` does.
|
||||
|
||||
Example curl:
|
||||
```
|
||||
curl --location 'http://0.0.0.0:4000/customer/aliases?start_date=2026-07-23%2000:00:00&end_date=2026-07-24%2000:00:00&size=50&search=acme' \
|
||||
--header 'Authorization: Bearer sk-1234'
|
||||
```
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
start_dt = _parse_spend_log_window_bound(start_date, "start_date")
|
||||
end_dt = _parse_spend_log_window_bound(end_date, "end_date")
|
||||
|
||||
query_params: List[Any] = [start_dt, end_dt]
|
||||
where_parts = [
|
||||
"\"startTime\" >= ($1::timestamptz AT TIME ZONE 'UTC')",
|
||||
"\"startTime\" <= ($2::timestamptz AT TIME ZONE 'UTC')",
|
||||
"end_user IS NOT NULL",
|
||||
"end_user != ''",
|
||||
]
|
||||
|
||||
if search:
|
||||
# Escape LIKE metacharacters so a literal '_' or '%' matches itself.
|
||||
escaped = search.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
|
||||
query_params.append(f"%{escaped}%")
|
||||
where_parts.append(f"end_user ILIKE ${len(query_params)} ESCAPE '\\'")
|
||||
|
||||
scope_condition = await _build_end_user_scope_condition(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
query_params=query_params,
|
||||
)
|
||||
if scope_condition is not None:
|
||||
where_parts.append(scope_condition)
|
||||
|
||||
# The inner LIMIT is the safety bound: it walks the startTime index newest
|
||||
# first and stops, so DISTINCT never runs over an unbounded row set.
|
||||
# request_id breaks startTime ties so the cut-off row is deterministic and
|
||||
# successive OFFSET pages agree on the set they are paging through; the
|
||||
# (startTime, request_id) index means the tiebreaker costs nothing.
|
||||
# size + 1: one row beyond the page reveals has_more without a COUNT(*).
|
||||
params = query_params + [SPEND_LOGS_FILTER_SCAN_CAP, size + 1, (page - 1) * size]
|
||||
scan_idx = len(params) - 2
|
||||
aliases_sql = (
|
||||
f"SELECT DISTINCT end_user FROM ("
|
||||
f" SELECT end_user"
|
||||
f' FROM "LiteLLM_SpendLogs"'
|
||||
f" WHERE {' AND '.join(where_parts)}"
|
||||
f' ORDER BY "startTime" DESC, request_id DESC'
|
||||
f" LIMIT ${scan_idx}"
|
||||
f") recent"
|
||||
f" ORDER BY end_user ASC"
|
||||
f" LIMIT ${scan_idx + 1} OFFSET ${scan_idx + 2}"
|
||||
)
|
||||
rows = await prisma_client.db.query_raw(aliases_sql, *params)
|
||||
aliases: List[str] = [row["end_user"] for row in rows if row.get("end_user")]
|
||||
|
||||
return CustomerAliasesResponse(
|
||||
aliases=aliases[:size],
|
||||
current_page=page,
|
||||
size=size,
|
||||
has_more=len(aliases) > size,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.customer_endpoints.list_customer_aliases(): "
|
||||
"Exception occured - {}".format(str(e))
|
||||
)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/customer/daily/activity",
|
||||
tags=["Customer Management"],
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
|
|||
12
litellm/proxy/management_endpoints/management_v1/__init__.py
Normal file
12
litellm/proxy/management_endpoints/management_v1/__init__.py
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
"""The `/management/v1` control-plane surface."""
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from litellm.proxy.management_endpoints.management_v1.spend_logs import (
|
||||
router as spend_logs_router,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
router.include_router(spend_logs_router)
|
||||
|
||||
__all__ = ["router"]
|
||||
77
litellm/proxy/management_endpoints/management_v1/common.py
Normal file
77
litellm/proxy/management_endpoints/management_v1/common.py
Normal file
|
|
@ -0,0 +1,77 @@
|
|||
"""Contract machinery shared by every `/management/v1` route."""
|
||||
|
||||
from urllib.parse import urlencode
|
||||
|
||||
from fastapi import Request
|
||||
from fastapi.dependencies.utils import get_flat_dependant
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import (
|
||||
PageLinks,
|
||||
ProblemDetail,
|
||||
)
|
||||
|
||||
MANAGEMENT_V1_PREFIX = "/management/v1"
|
||||
PROBLEM_CONTENT_TYPE = "application/problem+json"
|
||||
# A URN, not an https URL: RFC 9457 only asks that `type` identify the problem
|
||||
# type, and an https URI promises documentation at that address. Switch to an
|
||||
# https base only when pages actually exist to serve.
|
||||
PROBLEM_TYPE_BASE = "urn:litellm:error:"
|
||||
|
||||
|
||||
class ManagementProblem(Exception):
|
||||
"""Raised to return an RFC 9457 problem instead of the proxy's OpenAI error shape."""
|
||||
|
||||
def __init__(self, problem: ProblemDetail) -> None:
|
||||
self.problem = problem
|
||||
super().__init__(problem.detail)
|
||||
|
||||
|
||||
def problem_response(problem: ProblemDetail) -> JSONResponse:
|
||||
return JSONResponse(
|
||||
status_code=problem.status,
|
||||
content=problem.model_dump(exclude_none=True),
|
||||
media_type=PROBLEM_CONTENT_TYPE,
|
||||
)
|
||||
|
||||
|
||||
def _declared_query_params(request: Request) -> frozenset[str]:
|
||||
route = request.scope.get("route")
|
||||
dependant = getattr(route, "dependant", None)
|
||||
if dependant is None:
|
||||
return frozenset()
|
||||
return frozenset(field.alias for field in get_flat_dependant(dependant, skip_repeats=True).query_params)
|
||||
|
||||
|
||||
async def reject_unknown_query_params(request: Request) -> None:
|
||||
"""Reject any query param the route did not declare.
|
||||
|
||||
A silently ignored filter over-returns data, which is worse than a rejected
|
||||
request; a fresh surface is the only chance to be strict about it.
|
||||
"""
|
||||
declared = _declared_query_params(request)
|
||||
unknown: tuple[str, ...] = tuple(sorted(name for name in request.query_params if name not in declared))
|
||||
if not unknown:
|
||||
return
|
||||
raise ManagementProblem(
|
||||
ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}unknown-query-parameter",
|
||||
title="Unknown query parameter",
|
||||
status=400,
|
||||
detail=f"Unrecognized query parameter(s): {', '.join(unknown)}.",
|
||||
allowed=sorted(declared),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _page_url(request: Request, page: int) -> str:
|
||||
others = tuple((key, value) for key, value in request.query_params.multi_items() if key != "page")
|
||||
return f"{request.url.path}?{urlencode((*others, ('page', page)))}"
|
||||
|
||||
|
||||
def build_page_links(request: Request, page: int, has_more: bool) -> PageLinks:
|
||||
return PageLinks(
|
||||
self_link=_page_url(request, page),
|
||||
prev=_page_url(request, page - 1) if page > 1 else None,
|
||||
next=_page_url(request, page + 1) if has_more else None,
|
||||
)
|
||||
203
litellm/proxy/management_endpoints/management_v1/spend_logs.py
Normal file
203
litellm/proxy/management_endpoints/management_v1/spend_logs.py
Normal file
|
|
@ -0,0 +1,203 @@
|
|||
"""`/management/v1/spend_logs` facets."""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.management_endpoints.management_v1.common import (
|
||||
MANAGEMENT_V1_PREFIX,
|
||||
PROBLEM_TYPE_BASE,
|
||||
ManagementProblem,
|
||||
build_page_links,
|
||||
reject_unknown_query_params,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import (
|
||||
FacetListResponse,
|
||||
PageMeta,
|
||||
ProblemDetail,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix=MANAGEMENT_V1_PREFIX)
|
||||
|
||||
# Rows the facet query may read out of LiteLLM_SpendLogs before DISTINCT. Matches
|
||||
# SPEND_LOGS_PAGINATION_COUNT_CAP, the bound ui_view_spend_logs puts on its count
|
||||
# query, so both reads of the same table stop at the same depth.
|
||||
SPEND_LOGS_FACET_SCAN_CAP = 10000
|
||||
|
||||
|
||||
def _as_utc(value: datetime) -> datetime:
|
||||
return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _escape_like(value: str) -> str:
|
||||
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
|
||||
|
||||
|
||||
async def _end_user_scope_clause(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient,
|
||||
next_param_index: int,
|
||||
) -> tuple[str | None, tuple[Any, ...]]:
|
||||
"""SQL predicate restricting the facet to spend logs this caller may read.
|
||||
|
||||
Returns ``(None, ())`` for a proxy admin. Mirrors the scoping ``/spend/logs/ui``
|
||||
applies, so the dropdown can never offer an end user whose rows the caller
|
||||
could not open.
|
||||
"""
|
||||
from litellm.proxy.spend_tracking.spend_management_endpoints import (
|
||||
_get_permitted_team_ids_for_spend_logs,
|
||||
_is_admin_view_safe,
|
||||
)
|
||||
|
||||
if _is_admin_view_safe(user_api_key_dict=user_api_key_dict):
|
||||
return None, ()
|
||||
|
||||
try:
|
||||
permitted_team_ids = await _get_permitted_team_ids_for_spend_logs(
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
except Exception:
|
||||
permitted_team_ids = []
|
||||
|
||||
caller_user_id = user_api_key_dict.user_id
|
||||
# = ANY(::text[]) rather than an expanded IN list, matching the clause
|
||||
# ui_view_spend_logs builds: one parameter whatever the team count.
|
||||
templates = (('"user" = ${}',) if caller_user_id is not None else ()) + (
|
||||
("team_id = ANY(${}::text[])",) if permitted_team_ids else ()
|
||||
)
|
||||
params = ((caller_user_id,) if caller_user_id is not None else ()) + (
|
||||
(permitted_team_ids,) if permitted_team_ids else ()
|
||||
)
|
||||
if not templates:
|
||||
return "FALSE", ()
|
||||
clauses = tuple(template.format(next_param_index + offset) for offset, template in enumerate(templates))
|
||||
return f"({' OR '.join(clauses)})", params
|
||||
|
||||
|
||||
@router.get(
|
||||
"/spend_logs/end_users",
|
||||
tags=["Budget & Spend Tracking"],
|
||||
dependencies=[Depends(user_api_key_auth), Depends(reject_unknown_query_params)],
|
||||
response_model=FacetListResponse,
|
||||
)
|
||||
async def list_spend_log_end_users(
|
||||
request: Request,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
start_time: Annotated[
|
||||
datetime,
|
||||
Query(alias="filter[startTime][gte]", description="Window start (UTC when no offset is given)"),
|
||||
],
|
||||
end_time: Annotated[
|
||||
datetime,
|
||||
Query(alias="filter[startTime][lte]", description="Window end (UTC when no offset is given)"),
|
||||
],
|
||||
q: Annotated[str | None, Query(description="Case-insensitive partial match on the end user id")] = None,
|
||||
page: Annotated[int, Query(ge=1, description="Page number")] = 1,
|
||||
page_size: Annotated[int, Query(ge=1, le=100, description="Page size")] = 50,
|
||||
) -> FacetListResponse:
|
||||
"""
|
||||
The distinct end users appearing in spend logs over a time window, for the logs
|
||||
page filter dropdown.
|
||||
|
||||
Scoped like `/spend/logs/ui`: a proxy admin sees every end user in the window,
|
||||
anyone else sees only end users from their own requests or from teams they
|
||||
administer (or hold the `/spend/logs` permission on).
|
||||
|
||||
The window is required and the inner scan is capped at SPEND_LOGS_FACET_SCAN_CAP
|
||||
rows, so the query cannot degrade into a full-table scan the way
|
||||
`/global/all_end_users` does.
|
||||
|
||||
Example curl:
|
||||
```
|
||||
curl --location --globoff 'http://0.0.0.0:4000/management/v1/spend_logs/end_users?filter[startTime][gte]=2026-07-23T00:00:00Z&filter[startTime][lte]=2026-07-24T00:00:00Z&page_size=50&q=acme' \
|
||||
--header 'Authorization: Bearer sk-1234'
|
||||
```
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise ManagementProblem(
|
||||
ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}database-not-connected",
|
||||
title="Database not connected",
|
||||
status=503,
|
||||
detail=CommonProxyErrors.db_not_connected_error.value,
|
||||
)
|
||||
)
|
||||
|
||||
window_params: tuple[Any, ...] = (_as_utc(start_time), _as_utc(end_time))
|
||||
search_params: tuple[Any, ...] = (f"%{_escape_like(q)}%",) if q else ()
|
||||
search_clause = (f"end_user ILIKE ${len(window_params) + 1} ESCAPE '\\'",) if q else ()
|
||||
|
||||
scope_clause, scope_params = await _end_user_scope_clause(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
next_param_index=len(window_params) + len(search_params) + 1,
|
||||
)
|
||||
|
||||
where_parts = (
|
||||
(
|
||||
"\"startTime\" >= ($1::timestamptz AT TIME ZONE 'UTC')",
|
||||
"\"startTime\" <= ($2::timestamptz AT TIME ZONE 'UTC')",
|
||||
"end_user IS NOT NULL",
|
||||
"end_user != ''",
|
||||
)
|
||||
+ search_clause
|
||||
+ ((scope_clause,) if scope_clause is not None else ())
|
||||
)
|
||||
|
||||
# The inner LIMIT is the safety bound: it walks the startTime index newest
|
||||
# first and stops, so DISTINCT never runs over an unbounded row set.
|
||||
# request_id breaks startTime ties so the cut-off row is deterministic and
|
||||
# successive OFFSET pages agree on the set they are paging through.
|
||||
# page_size + 1: one row beyond the page reveals has_more without a COUNT(*).
|
||||
params = (
|
||||
window_params
|
||||
+ search_params
|
||||
+ scope_params
|
||||
+ (SPEND_LOGS_FACET_SCAN_CAP, page_size + 1, (page - 1) * page_size)
|
||||
)
|
||||
scan_idx = len(params) - 2
|
||||
facet_sql = (
|
||||
f"SELECT DISTINCT end_user FROM ("
|
||||
f" SELECT end_user"
|
||||
f' FROM "LiteLLM_SpendLogs"'
|
||||
f" WHERE {' AND '.join(where_parts)}"
|
||||
f' ORDER BY "startTime" DESC, request_id DESC'
|
||||
f" LIMIT ${scan_idx}"
|
||||
f") recent"
|
||||
f" ORDER BY end_user ASC"
|
||||
f" LIMIT ${scan_idx + 1} OFFSET ${scan_idx + 2}"
|
||||
)
|
||||
rows = await prisma_client.db.query_raw(facet_sql, *params)
|
||||
end_users: list[str] = [row["end_user"] for row in rows if row.get("end_user")]
|
||||
has_more = len(end_users) > page_size
|
||||
|
||||
return FacetListResponse(
|
||||
data=end_users[:page_size],
|
||||
meta=PageMeta(page=page, page_size=page_size, has_more=has_more),
|
||||
links=build_page_links(request=request, page=page, has_more=has_more),
|
||||
)
|
||||
|
||||
except ManagementProblem:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.management_v1.spend_logs.list_spend_log_end_users(): "
|
||||
"Exception occured - {}".format(str(e))
|
||||
)
|
||||
raise ManagementProblem(
|
||||
ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}internal-server-error",
|
||||
title="Internal server error",
|
||||
status=500,
|
||||
detail="Failed to list spend log end users.",
|
||||
)
|
||||
)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -393,6 +393,16 @@ from litellm.proxy.management_endpoints.cost_tracking_settings import (
|
|||
from litellm.proxy.management_endpoints.customer_endpoints import (
|
||||
router as customer_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.management_v1 import (
|
||||
router as management_v1_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.management_v1.common import (
|
||||
MANAGEMENT_V1_PREFIX,
|
||||
PROBLEM_TYPE_BASE,
|
||||
ManagementProblem,
|
||||
problem_response,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import ProblemDetail
|
||||
from litellm.proxy.management_endpoints.fallback_management_endpoints import (
|
||||
router as fallback_management_router,
|
||||
)
|
||||
|
|
@ -1438,8 +1448,27 @@ def _close_dangling_otel_server_span(request: Request, status_code: int, exc: Op
|
|||
request.state.parent_otel_span = None
|
||||
|
||||
|
||||
@app.exception_handler(ManagementProblem)
|
||||
async def management_problem_exception_handler(request: Request, exc: ManagementProblem):
|
||||
_close_dangling_otel_server_span(request, exc.problem.status, exc=exc)
|
||||
return problem_response(exc.problem)
|
||||
|
||||
|
||||
@app.exception_handler(RequestValidationError)
|
||||
async def otel_request_validation_exception_handler(request: Request, exc: RequestValidationError):
|
||||
if request.url.path.startswith(MANAGEMENT_V1_PREFIX):
|
||||
_close_dangling_otel_server_span(request, 400, exc=exc)
|
||||
return problem_response(
|
||||
ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}invalid-query-parameter",
|
||||
title="Invalid query parameter",
|
||||
status=400,
|
||||
detail="; ".join(
|
||||
f"{'.'.join(str(part) for part in error['loc'][1:])}: {error['msg']}" for error in exc.errors()
|
||||
)
|
||||
or "The request query parameters are invalid.",
|
||||
)
|
||||
)
|
||||
_close_dangling_otel_server_span(request, 422, exc=exc)
|
||||
return JSONResponse(
|
||||
status_code=422,
|
||||
|
|
@ -16306,6 +16335,7 @@ app.include_router(team_router)
|
|||
app.include_router(ui_sso_router)
|
||||
app.include_router(organization_router)
|
||||
app.include_router(customer_router)
|
||||
app.include_router(management_v1_router)
|
||||
app.include_router(spend_management_router)
|
||||
app.include_router(caching_router)
|
||||
app.include_router(analytics_router)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -17,25 +17,6 @@ class CustomerResponse(LiteLLM_EndUserTable):
|
|||
litellm_budget_table: Optional[LiteLLM_BudgetTableFull] = None # pyright: ignore
|
||||
|
||||
|
||||
class CustomerAliasesResponse(BaseModel):
|
||||
"""Paginated, id-only customer listing used by UI filter dropdowns.
|
||||
|
||||
Deliberately excludes budget/object-permission relations so a proxy with a
|
||||
large LiteLLM_EndUserTable can back a search-as-you-type control without
|
||||
materializing every row (see /customer/list for the full objects).
|
||||
|
||||
Reports ``has_more`` rather than a total count on purpose: a total requires
|
||||
COUNT(*) over the whole match set on every keystroke, which is the exact
|
||||
cost this endpoint exists to avoid. Fetching one row beyond the page is
|
||||
enough to drive an infinite-scroll dropdown.
|
||||
"""
|
||||
|
||||
aliases: List[str]
|
||||
current_page: int
|
||||
size: int
|
||||
has_more: bool
|
||||
|
||||
|
||||
class BlockUsersResponse(BaseModel):
|
||||
blocked_users: List[LiteLLM_EndUserTable]
|
||||
|
||||
|
|
|
|||
39
litellm/types/proxy/management_endpoints/management_v1.py
Normal file
39
litellm/types/proxy/management_endpoints/management_v1.py
Normal file
|
|
@ -0,0 +1,39 @@
|
|||
"""Shared response shapes for the `/management/v1` control-plane surface."""
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
|
||||
class ProblemDetail(BaseModel):
|
||||
"""RFC 9457 problem details, served as `application/problem+json`."""
|
||||
|
||||
type: str
|
||||
title: str
|
||||
status: int
|
||||
detail: str
|
||||
allowed: list[str] | None = None
|
||||
|
||||
|
||||
class PageLinks(BaseModel):
|
||||
"""Hypermedia for a paginated list. No `first`/`last`: without a total count the last page is unknown."""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
self_link: str = Field(alias="self")
|
||||
prev: str | None = None
|
||||
next: str | None = None
|
||||
|
||||
|
||||
class PageMeta(BaseModel):
|
||||
"""`has_more` rather than `total_count`, which would need a COUNT(*) over the whole match set per keystroke."""
|
||||
|
||||
page: int
|
||||
page_size: int
|
||||
has_more: bool
|
||||
|
||||
|
||||
class FacetListResponse(BaseModel):
|
||||
"""The distinct values one column takes over a filtered query. `data` holds bare values, not entity rows."""
|
||||
|
||||
data: list[str]
|
||||
meta: PageMeta
|
||||
links: PageLinks
|
||||
|
|
@ -556,10 +556,10 @@ def _custom_logger_class_exists_in_success_callbacks(
|
|||
e.g if `LangfusePromptManagement` is passed in, it will return True if an instance of `LangfusePromptManagement` exists in litellm.success_callback or litellm._async_success_callback
|
||||
|
||||
Prevents double adding a custom logger callback to the litellm callbacks
|
||||
|
||||
Matches on the exact class; an instance of a subclass does not count as registered
|
||||
"""
|
||||
return any(
|
||||
isinstance(cb, type(callback_class)) for cb in litellm.success_callback + litellm._async_success_callback
|
||||
)
|
||||
return any(type(cb) is type(callback_class) for cb in litellm.success_callback + litellm._async_success_callback)
|
||||
|
||||
|
||||
def _custom_logger_class_exists_in_failure_callbacks(
|
||||
|
|
@ -571,10 +571,10 @@ def _custom_logger_class_exists_in_failure_callbacks(
|
|||
e.g if `LangfusePromptManagement` is passed in, it will return True if an instance of `LangfusePromptManagement` exists in litellm.failure_callback or litellm._async_failure_callback
|
||||
|
||||
Prevents double adding a custom logger callback to the litellm callbacks
|
||||
|
||||
Matches on the exact class; an instance of a subclass does not count as registered
|
||||
"""
|
||||
return any(
|
||||
isinstance(cb, type(callback_class)) for cb in litellm.failure_callback + litellm._async_failure_callback
|
||||
)
|
||||
return any(type(cb) is type(callback_class) for cb in litellm.failure_callback + litellm._async_failure_callback)
|
||||
|
||||
|
||||
def get_request_guardrails(kwargs: Dict[str, Any]) -> List[str]:
|
||||
|
|
@ -764,7 +764,7 @@ def function_setup(
|
|||
llm_router=None, # type: ignore
|
||||
)
|
||||
if callback is None or any(
|
||||
isinstance(cb, type(callback)) for cb in litellm._async_success_callback
|
||||
type(cb) is type(callback) for cb in litellm._async_success_callback
|
||||
): # don't double add a callback
|
||||
continue
|
||||
if callback not in litellm.input_callback:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -95,9 +95,10 @@ if [ -z "$UV_BIN" ] || [ "${CURRENT_UV_VERSION:-}" != "$UV_VERSION" ]; then
|
|||
fi
|
||||
|
||||
# ── install ────────────────────────────────────────────────────────────────
|
||||
# --python-preference system: reuse a compatible system Python when present,
|
||||
# otherwise download a managed one. Either way uv honours litellm's requires-python,
|
||||
# so a too-old (3.9) or too-new (3.14+) system Python is skipped, not forced.
|
||||
# --python mirrors requires-python in pyproject.toml (keep in sync): uv selects the
|
||||
# interpreter before resolving, so an unconstrained request accepts a too-old system
|
||||
# Python (stock macOS ships 3.9) and fails resolution instead of downloading a
|
||||
# managed one. --python-preference system still reuses a compatible system Python.
|
||||
echo ""
|
||||
if [ -n "${LITELLM_CLI_REF:-}" ]; then
|
||||
header "Installing litellm[cli] from ${LITELLM_CLI_REF}…"
|
||||
|
|
@ -106,8 +107,8 @@ else
|
|||
fi
|
||||
echo ""
|
||||
|
||||
"$UV_BIN" tool install --python-preference system --force "${LITELLM_PACKAGE}" \
|
||||
|| die "uv tool install failed. Try manually: $UV_BIN tool install '${LITELLM_PACKAGE}'"
|
||||
"$UV_BIN" tool install --python '>=3.10,<3.15' --python-preference system --force "${LITELLM_PACKAGE}" \
|
||||
|| die "uv tool install failed. Try manually: $UV_BIN tool install --python '>=3.10,<3.15' '${LITELLM_PACKAGE}'"
|
||||
|
||||
# ── find the lite binary installed by uv tool ──────────────────────────────
|
||||
SCRIPTS_DIR="$("$UV_BIN" tool dir --bin)"
|
||||
|
|
|
|||
|
|
@ -100,11 +100,12 @@ else
|
|||
fi
|
||||
echo ""
|
||||
|
||||
# --python-preference system: reuse a compatible system Python when present,
|
||||
# otherwise download a managed one. Either way uv honours litellm's requires-python,
|
||||
# so a too-old (3.9) or too-new (3.14+) system Python is skipped, not forced.
|
||||
"$UV_BIN" tool install --python-preference system --force "${LITELLM_PACKAGE}" \
|
||||
|| die "uv tool install failed. Try manually: $UV_BIN tool install '${LITELLM_PACKAGE}'"
|
||||
# --python mirrors requires-python in pyproject.toml (keep in sync): uv selects the
|
||||
# interpreter before resolving, so an unconstrained request accepts a too-old system
|
||||
# Python (stock macOS ships 3.9) and fails resolution instead of downloading a
|
||||
# managed one. --python-preference system still reuses a compatible system Python.
|
||||
"$UV_BIN" tool install --python '>=3.10,<3.15' --python-preference system --force "${LITELLM_PACKAGE}" \
|
||||
|| die "uv tool install failed. Try manually: $UV_BIN tool install --python '>=3.10,<3.15' '${LITELLM_PACKAGE}'"
|
||||
|
||||
# ── find the litellm binary installed by uv tool ───────────────────────────
|
||||
SCRIPTS_DIR="$("$UV_BIN" tool dir --bin)"
|
||||
|
|
|
|||
|
|
@ -1,3 +1,9 @@
|
|||
export const UI_BASE_URL = (
|
||||
process.env.E2E_UI_BASE_URL ||
|
||||
process.env.LITELLM_PROXY_URL ||
|
||||
"http://localhost:4000"
|
||||
).replace(/\/+$/, "");
|
||||
|
||||
// Storage state paths for each role
|
||||
export const ADMIN_STORAGE_PATH = "admin.storageState.json";
|
||||
export const ADMIN_VIEWER_STORAGE_PATH = "adminViewer.storageState.json";
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import { chromium, expect, request } from "@playwright/test";
|
||||
import { users, Role, STORAGE_PATHS } from "./fixtures/users";
|
||||
import { UI_BASE_URL } from "./constants";
|
||||
import * as fs from "fs";
|
||||
|
||||
async function globalSetup() {
|
||||
|
|
@ -12,7 +13,7 @@ async function globalSetup() {
|
|||
// the admin UI toggle does; the projects migration smoke needs the link.
|
||||
const masterKey = process.env.LITELLM_MASTER_KEY || "sk-1234";
|
||||
const api = await request.newContext();
|
||||
const settingsRes = await api.patch(`http://localhost:4000${rootPath}/update/ui_settings`, {
|
||||
const settingsRes = await api.patch(`${UI_BASE_URL}${rootPath}/update/ui_settings`, {
|
||||
headers: { Authorization: `Bearer ${masterKey}` },
|
||||
data: { enable_projects_ui: true },
|
||||
});
|
||||
|
|
@ -26,7 +27,7 @@ async function globalSetup() {
|
|||
const storagePath = STORAGE_PATHS[role];
|
||||
const page = await browser.newPage();
|
||||
try {
|
||||
await page.goto(`http://localhost:4000${rootPath}/ui/login`);
|
||||
await page.goto(`${UI_BASE_URL}${rootPath}/ui/login`);
|
||||
await page.getByPlaceholder("Enter your username").fill(email);
|
||||
await page.getByPlaceholder("Enter your password").fill(password);
|
||||
await page.getByRole("button", { name: "Login", exact: true }).click();
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import { defineConfig, devices } from "@playwright/test";
|
||||
import { UI_BASE_URL } from "./constants";
|
||||
|
||||
/**
|
||||
* App Router migration smoke under a non-root mount. Boot the proxy with the same
|
||||
|
|
@ -15,7 +16,7 @@ export default defineConfig({
|
|||
workers: process.env.CI ? 1 : undefined,
|
||||
reporter: "list",
|
||||
use: {
|
||||
baseURL: "http://localhost:4000",
|
||||
baseURL: UI_BASE_URL,
|
||||
trace: "on-first-retry",
|
||||
actionTimeout: 15 * 1000,
|
||||
navigationTimeout: 30 * 1000,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import { defineConfig, devices } from "@playwright/test";
|
||||
import { UI_BASE_URL } from "./constants";
|
||||
|
||||
/**
|
||||
* See https://playwright.dev/docs/test-configuration.
|
||||
|
|
@ -20,7 +21,7 @@ export default defineConfig({
|
|||
/* Shared settings for all the projects below. See https://playwright.dev/docs/api/class-testoptions. */
|
||||
use: {
|
||||
/* Base URL to use in actions like `await page.goto('/')`. */
|
||||
baseURL: "http://localhost:4000",
|
||||
baseURL: UI_BASE_URL,
|
||||
|
||||
/* Collect trace when retrying the failed test. See https://playwright.dev/docs/trace-viewer */
|
||||
trace: "on-first-retry",
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import { test, expect } from "@playwright/test";
|
|||
|
||||
test.describe("Authentication Checks", () => {
|
||||
test("should redirect unauthenticated user from a protected page", async ({ page }) => {
|
||||
const protectedPageUrl = "http://localhost:4000/ui?page=llm-playground";
|
||||
const protectedPageUrl = "/ui?page=llm-playground";
|
||||
await page.goto(protectedPageUrl, { waitUntil: "domcontentloaded" });
|
||||
await expect(page).toHaveURL(/\/ui\/login/);
|
||||
await expect(page.getByRole("heading", { name: "Login" })).toBeVisible();
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import { users } from "../../fixtures/users";
|
|||
import { Role } from "../../fixtures/roles";
|
||||
|
||||
test("user can log in", async ({ page }) => {
|
||||
await page.goto("http://localhost:4000/ui/login");
|
||||
await page.goto("/ui/login");
|
||||
await page.getByPlaceholder("Enter your username").fill(users[Role.ProxyAdmin].email);
|
||||
await page.getByPlaceholder("Enter your password").fill(users[Role.ProxyAdmin].password);
|
||||
const loginButton = page.getByRole("button", { name: "Login", exact: true });
|
||||
|
|
|
|||
|
|
@ -132,7 +132,7 @@ test.describe("Proxy Admin - Teams", () => {
|
|||
const masterKey = process.env.LITELLM_MASTER_KEY || "sk-1234";
|
||||
const seededModels = ["fake-openai-gpt-4", "fake-anthropic-claude"];
|
||||
const restore = async () => {
|
||||
const res = await request.post("http://localhost:4000/team/update", {
|
||||
const res = await request.post("/team/update", {
|
||||
headers: { Authorization: `Bearer ${masterKey}` },
|
||||
data: { team_id: E2E_TEAM_CRUD_ID, models: seededModels },
|
||||
});
|
||||
|
|
|
|||
|
|
@ -23,7 +23,7 @@ async function clearFallbackForPrimary(request: import("@playwright/test").APIRe
|
|||
const masterKey = users[Role.ProxyAdmin].password;
|
||||
const auth = { Authorization: `Bearer ${masterKey}` };
|
||||
|
||||
const current = await request.get("http://localhost:4000/get/config/callbacks", { headers: auth });
|
||||
const current = await request.get("/get/config/callbacks", { headers: auth });
|
||||
if (!current.ok()) return;
|
||||
const body = await current.json();
|
||||
const router = body?.router_settings ?? {};
|
||||
|
|
@ -31,7 +31,7 @@ async function clearFallbackForPrimary(request: import("@playwright/test").APIRe
|
|||
const next = existing.filter((entry) => !(entry && PRIMARY in entry));
|
||||
if (next.length === existing.length) return;
|
||||
|
||||
await request.post("http://localhost:4000/config/update", {
|
||||
await request.post("/config/update", {
|
||||
headers: auth,
|
||||
data: { router_settings: { ...router, fallbacks: next } },
|
||||
});
|
||||
|
|
@ -111,7 +111,6 @@ test.describe("Router Settings - Fallbacks", () => {
|
|||
type ConfigYAML = components["schemas"]["ConfigYAML"];
|
||||
type RouterSettingsResponse = components["schemas"]["RouterSettingsResponse"];
|
||||
|
||||
const BASE_URL = "http://localhost:4000";
|
||||
const ADMIN_AUTH = { Authorization: `Bearer ${users[Role.ProxyAdmin].password}` };
|
||||
|
||||
/**
|
||||
|
|
@ -123,7 +122,7 @@ async function patchRouterSettings(
|
|||
request: import("@playwright/test").APIRequestContext,
|
||||
patch: Partial<NonNullable<ConfigYAML["router_settings"]>>,
|
||||
) {
|
||||
const res = await request.post(`${BASE_URL}/config/update`, {
|
||||
const res = await request.post(`/config/update`, {
|
||||
headers: ADMIN_AUTH,
|
||||
data: { router_settings: patch },
|
||||
});
|
||||
|
|
@ -179,7 +178,7 @@ test.describe("Router Settings - Loadbalancing", () => {
|
|||
await expect
|
||||
.poll(
|
||||
async () => {
|
||||
const res = await request.get(`${BASE_URL}/router/settings`, { headers: ADMIN_AUTH });
|
||||
const res = await request.get(`/router/settings`, { headers: ADMIN_AUTH });
|
||||
const data = (await res.json()) as RouterSettingsResponse;
|
||||
return data.current_values?.num_retries;
|
||||
},
|
||||
|
|
|
|||
|
|
@ -23,11 +23,13 @@ from litellm.integrations._types.open_inference import ErrorAttributes
|
|||
from ._helpers import assert_server_span_attrs, get_server_span
|
||||
|
||||
|
||||
def _fake_request(parent_otel_span=None):
|
||||
def _fake_request(parent_otel_span=None, path="/key/generate"):
|
||||
"""A real Request always carries a url; the validation handler reads its path to
|
||||
decide whether the caller is on a surface with its own error contract."""
|
||||
state = types.SimpleNamespace()
|
||||
if parent_otel_span is not None:
|
||||
state.parent_otel_span = parent_otel_span
|
||||
return types.SimpleNamespace(state=state)
|
||||
return types.SimpleNamespace(state=state, url=types.SimpleNamespace(path=path))
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -41,7 +43,7 @@ def wired_otel(otel_with_exporter, monkeypatch):
|
|||
def test_close_dangling_span_stamps_status(
|
||||
wired_otel, server_span_factory, status, path
|
||||
):
|
||||
request = _fake_request(parent_otel_span=server_span_factory(path))
|
||||
request = _fake_request(parent_otel_span=server_span_factory(path), path=path)
|
||||
_close_dangling_otel_server_span(request, status)
|
||||
assert_server_span_attrs(
|
||||
wired_otel,
|
||||
|
|
@ -59,7 +61,7 @@ def test_close_dangling_span_noop_when_no_span(wired_otel):
|
|||
|
||||
def test_close_dangling_span_noop_when_otel_absent(server_span_factory, monkeypatch):
|
||||
monkeypatch.setattr(proxy_server_module, "open_telemetry_logger", None)
|
||||
request = _fake_request(parent_otel_span=server_span_factory("/key/generate"))
|
||||
request = _fake_request(parent_otel_span=server_span_factory("/key/generate"), path="/key/generate")
|
||||
_close_dangling_otel_server_span(request, 500)
|
||||
|
||||
|
||||
|
|
@ -83,7 +85,7 @@ def test_close_dangling_span_noop_when_otel_absent(server_span_factory, monkeypa
|
|||
def test_exception_handler_closes_span(
|
||||
wired_otel, server_span_factory, handler, exc, status, path
|
||||
):
|
||||
request = _fake_request(parent_otel_span=server_span_factory(path))
|
||||
request = _fake_request(parent_otel_span=server_span_factory(path), path=path)
|
||||
response = asyncio.run(handler(request, exc))
|
||||
assert response.status_code == status
|
||||
assert_server_span_attrs(
|
||||
|
|
@ -94,6 +96,25 @@ def test_exception_handler_closes_span(
|
|||
)
|
||||
|
||||
|
||||
def test_validation_handler_closes_span_on_the_control_plane_too(wired_otel, server_span_factory):
|
||||
"""The control plane answers validation errors with a 400 problem document
|
||||
instead of the proxy-wide 422, and that branch returns early. It must still
|
||||
close the dangling SERVER span, or those requests leak a span apiece."""
|
||||
path = "/management/v1/spend_logs/end_users"
|
||||
request = _fake_request(parent_otel_span=server_span_factory(path), path=path)
|
||||
|
||||
response = asyncio.run(otel_request_validation_exception_handler(request, RequestValidationError(errors=[])))
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.media_type == "application/problem+json"
|
||||
assert_server_span_attrs(
|
||||
wired_otel,
|
||||
expected_status=400,
|
||||
expected_url_path=path,
|
||||
where="otel_request_validation_exception_handler (control plane)",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", ["/team/list", "/organization/list"])
|
||||
def test_openai_exception_handler_stamps_structured_error_on_span(
|
||||
wired_otel, server_span_factory, path
|
||||
|
|
@ -103,7 +124,7 @@ def test_openai_exception_handler_stamps_structured_error_on_span(
|
|||
ProxyException stringified to "" so error.message was dropped — the span
|
||||
showed an error with no message."""
|
||||
msg = "Authentication Error, Invalid proxy server token passed."
|
||||
request = _fake_request(parent_otel_span=server_span_factory(path))
|
||||
request = _fake_request(parent_otel_span=server_span_factory(path), path=path)
|
||||
exc = ProxyException(message=msg, type="auth_error", param="key", code=401)
|
||||
|
||||
response = asyncio.run(openai_exception_handler(request, exc))
|
||||
|
|
@ -123,7 +144,7 @@ def test_openai_exception_handler_stamps_structured_error_on_span(
|
|||
|
||||
def test_unhandled_handler_reraises_known_exceptions(wired_otel, server_span_factory):
|
||||
"""ProxyException / HTTPException / RequestValidationError have dedicated handlers."""
|
||||
request = _fake_request(parent_otel_span=server_span_factory("/key/generate"))
|
||||
request = _fake_request(parent_otel_span=server_span_factory("/key/generate"), path="/key/generate")
|
||||
with pytest.raises(HTTPException):
|
||||
asyncio.run(
|
||||
otel_unhandled_exception_handler(
|
||||
|
|
@ -147,7 +168,7 @@ def test_unhandled_handler_reraises_known_exceptions(wired_otel, server_span_fac
|
|||
def test_openai_exception_handler_closes_span(
|
||||
wired_otel, server_span_factory, code, path
|
||||
):
|
||||
request = _fake_request(parent_otel_span=server_span_factory(path))
|
||||
request = _fake_request(parent_otel_span=server_span_factory(path), path=path)
|
||||
exc = ProxyException(
|
||||
message="boom",
|
||||
type="invalid_request_error",
|
||||
|
|
|
|||
|
|
@ -258,6 +258,158 @@ class TestPrometheusCacheMetrics:
|
|||
# Should not emit read metric, because explicit provider value is zero.
|
||||
mock_logger.litellm_provider_cache_read_input_tokens_metric.labels.assert_not_called()
|
||||
|
||||
def test_provider_cache_creation_fallback_to_cache_write_tokens(
|
||||
self, sample_enum_values
|
||||
):
|
||||
"""OpenAI-style usage (prompt_tokens_details.cache_write_tokens, no top-level
|
||||
cache_creation_input_tokens) must populate the provider cache creation metric."""
|
||||
mock_logger = MagicMock()
|
||||
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
standard_logging_payload = {
|
||||
"cache_hit": False,
|
||||
"total_tokens": 12100,
|
||||
"prompt_tokens": 12000,
|
||||
"completion_tokens": 100,
|
||||
"model_group": "openai",
|
||||
"request_tags": [],
|
||||
"metadata": {
|
||||
"usage_object": {
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 0,
|
||||
"cache_write_tokens": 800,
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
mock_logger.litellm_cache_hits_metric = MagicMock()
|
||||
mock_logger.litellm_cache_misses_metric = MagicMock()
|
||||
mock_logger.litellm_cached_tokens_metric = MagicMock()
|
||||
mock_logger.litellm_provider_cache_read_input_tokens_metric = MagicMock()
|
||||
mock_logger.litellm_provider_cache_creation_input_tokens_metric = MagicMock()
|
||||
mock_logger.get_labels_for_metric = MagicMock(
|
||||
return_value=[
|
||||
"model",
|
||||
"hashed_api_key",
|
||||
"api_key_alias",
|
||||
"team",
|
||||
"team_alias",
|
||||
"end_user",
|
||||
"user",
|
||||
]
|
||||
)
|
||||
|
||||
PrometheusLogger._increment_cache_metrics(
|
||||
mock_logger,
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
enum_values=sample_enum_values,
|
||||
)
|
||||
|
||||
mock_logger.litellm_provider_cache_creation_input_tokens_metric.labels().inc.assert_called_once_with(
|
||||
800
|
||||
)
|
||||
|
||||
def test_provider_cache_creation_fallback_to_cache_creation_tokens(
|
||||
self, sample_enum_values
|
||||
):
|
||||
"""Normalized litellm usage dumps carry cache_creation_tokens in
|
||||
prompt_tokens_details; the fallback must read it when cache_write_tokens is absent."""
|
||||
mock_logger = MagicMock()
|
||||
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
standard_logging_payload = {
|
||||
"cache_hit": False,
|
||||
"total_tokens": 100,
|
||||
"prompt_tokens": 50,
|
||||
"completion_tokens": 50,
|
||||
"model_group": "openai",
|
||||
"request_tags": [],
|
||||
"metadata": {
|
||||
"usage_object": {
|
||||
"prompt_tokens_details": {"cache_creation_tokens": 42},
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
mock_logger.litellm_cache_hits_metric = MagicMock()
|
||||
mock_logger.litellm_cache_misses_metric = MagicMock()
|
||||
mock_logger.litellm_cached_tokens_metric = MagicMock()
|
||||
mock_logger.litellm_provider_cache_read_input_tokens_metric = MagicMock()
|
||||
mock_logger.litellm_provider_cache_creation_input_tokens_metric = MagicMock()
|
||||
mock_logger.get_labels_for_metric = MagicMock(
|
||||
return_value=[
|
||||
"model",
|
||||
"hashed_api_key",
|
||||
"api_key_alias",
|
||||
"team",
|
||||
"team_alias",
|
||||
"end_user",
|
||||
"user",
|
||||
]
|
||||
)
|
||||
|
||||
PrometheusLogger._increment_cache_metrics(
|
||||
mock_logger,
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
enum_values=sample_enum_values,
|
||||
)
|
||||
|
||||
mock_logger.litellm_provider_cache_creation_input_tokens_metric.labels().inc.assert_called_once_with(
|
||||
42
|
||||
)
|
||||
|
||||
def test_provider_cache_creation_does_not_fallback_on_explicit_zero(
|
||||
self, sample_enum_values
|
||||
):
|
||||
"""Explicit cache_creation_input_tokens=0 must not trigger fallback to
|
||||
prompt_tokens_details, mirroring the cache-read semantics."""
|
||||
mock_logger = MagicMock()
|
||||
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
standard_logging_payload = {
|
||||
"cache_hit": False,
|
||||
"total_tokens": 100,
|
||||
"prompt_tokens": 50,
|
||||
"completion_tokens": 50,
|
||||
"model_group": "openai",
|
||||
"request_tags": [],
|
||||
"metadata": {
|
||||
"usage_object": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"prompt_tokens_details": {"cache_write_tokens": 800},
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
mock_logger.litellm_cache_hits_metric = MagicMock()
|
||||
mock_logger.litellm_cache_misses_metric = MagicMock()
|
||||
mock_logger.litellm_cached_tokens_metric = MagicMock()
|
||||
mock_logger.litellm_provider_cache_read_input_tokens_metric = MagicMock()
|
||||
mock_logger.litellm_provider_cache_creation_input_tokens_metric = MagicMock()
|
||||
mock_logger.get_labels_for_metric = MagicMock(
|
||||
return_value=[
|
||||
"model",
|
||||
"hashed_api_key",
|
||||
"api_key_alias",
|
||||
"team",
|
||||
"team_alias",
|
||||
"end_user",
|
||||
"user",
|
||||
]
|
||||
)
|
||||
|
||||
PrometheusLogger._increment_cache_metrics(
|
||||
mock_logger,
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
enum_values=sample_enum_values,
|
||||
)
|
||||
|
||||
mock_logger.litellm_provider_cache_creation_input_tokens_metric.labels.assert_not_called()
|
||||
|
||||
def test_increment_cache_metrics_when_cache_hit_is_none(self, sample_enum_values):
|
||||
"""Test that no metrics are incremented when cache_hit is None"""
|
||||
# Create mock for PrometheusLogger instance
|
||||
|
|
|
|||
|
|
@ -150,6 +150,57 @@ class TestIncrementTokenDetailMetrics:
|
|||
10.0
|
||||
)
|
||||
|
||||
def test_cache_creation_falls_back_to_cache_write_tokens(self, sample_enum_values):
|
||||
logger = _make_mock_logger()
|
||||
payload = {
|
||||
"metadata": {
|
||||
"usage_object": {
|
||||
"prompt_tokens": 12000,
|
||||
"completion_tokens": 100,
|
||||
"total_tokens": 12100,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 0,
|
||||
"cache_write_tokens": 800,
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
PrometheusLogger._increment_token_detail_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=sample_enum_values,
|
||||
)
|
||||
|
||||
logger.litellm_input_cache_creation_tokens_metric.labels().inc.assert_called_once_with(
|
||||
800.0
|
||||
)
|
||||
|
||||
def test_cache_write_tokens_takes_precedence_over_cache_creation_tokens(
|
||||
self, sample_enum_values
|
||||
):
|
||||
logger = _make_mock_logger()
|
||||
payload = {
|
||||
"metadata": {
|
||||
"usage_object": {
|
||||
"prompt_tokens_details": {
|
||||
"cache_creation_tokens": 25,
|
||||
"cache_write_tokens": 800,
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
PrometheusLogger._increment_token_detail_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=sample_enum_values,
|
||||
)
|
||||
|
||||
logger.litellm_input_cache_creation_tokens_metric.labels().inc.assert_called_once_with(
|
||||
800.0
|
||||
)
|
||||
|
||||
def test_skips_metrics_when_value_is_zero(self, sample_enum_values):
|
||||
logger = _make_mock_logger()
|
||||
payload = {
|
||||
|
|
|
|||
|
|
@ -252,9 +252,7 @@ def test_transform_cancel_eval_response(config: OpenAIEvalsConfig):
|
|||
"object": "eval",
|
||||
"status": "cancelled",
|
||||
},
|
||||
request=httpx.Request(
|
||||
"POST", "https://api.openai.com/v1/evals/eval_123/cancel"
|
||||
),
|
||||
request=httpx.Request("POST", "https://api.openai.com/v1/evals/eval_123/cancel"),
|
||||
)
|
||||
|
||||
result = config.transform_cancel_eval_response(
|
||||
|
|
@ -276,8 +274,169 @@ def test_transform_run_requests_encode_eval_and_run_ids(config: OpenAIEvalsConfi
|
|||
headers={},
|
||||
)
|
||||
|
||||
assert (
|
||||
url
|
||||
== "https://api.openai.com/v1/evals/..%2F..%2Fevals%3Fx%3D1%23frag/runs/..%2Fruns%23other/cancel"
|
||||
)
|
||||
assert url == "https://api.openai.com/v1/evals/..%2F..%2Fevals%3Fx%3D1%23frag/runs/..%2Fruns%23other/cancel"
|
||||
assert request_body == {}
|
||||
|
||||
|
||||
def _eval_json_response(url: str, method: str = "GET") -> httpx.Response:
|
||||
return httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"id": "eval_123",
|
||||
"object": "eval",
|
||||
"created_at": 1234567890,
|
||||
"name": "Test Eval",
|
||||
"data_source_config": {"type": "stored_completions"},
|
||||
"testing_criteria": [],
|
||||
},
|
||||
request=httpx.Request(method, url),
|
||||
)
|
||||
|
||||
|
||||
def _run_json(run_id: str = "evalrun_123", status: str = "queued") -> dict:
|
||||
return {
|
||||
"id": run_id,
|
||||
"object": "eval.run",
|
||||
"created_at": 1234567890,
|
||||
"status": status,
|
||||
"data_source": {"type": "completions"},
|
||||
"eval_id": "eval_123",
|
||||
}
|
||||
|
||||
|
||||
def test_transform_get_eval_response(config: OpenAIEvalsConfig):
|
||||
result = config.transform_get_eval_response(
|
||||
raw_response=_eval_json_response("https://api.openai.com/v1/evals/eval_123"),
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result.id == "eval_123"
|
||||
assert result.object == "eval"
|
||||
assert result.name == "Test Eval"
|
||||
|
||||
|
||||
def test_transform_update_eval_response(config: OpenAIEvalsConfig):
|
||||
result = config.transform_update_eval_response(
|
||||
raw_response=_eval_json_response("https://api.openai.com/v1/evals/eval_123", method="POST"),
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result.id == "eval_123"
|
||||
assert result.name == "Test Eval"
|
||||
|
||||
|
||||
def test_transform_create_run_response(config: OpenAIEvalsConfig):
|
||||
response = httpx.Response(
|
||||
status_code=200,
|
||||
json=_run_json(),
|
||||
request=httpx.Request("POST", "https://api.openai.com/v1/evals/eval_123/runs"),
|
||||
)
|
||||
|
||||
result = config.transform_create_run_response(
|
||||
raw_response=response,
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result.id == "evalrun_123"
|
||||
assert result.status == "queued"
|
||||
assert result.eval_id == "eval_123"
|
||||
|
||||
|
||||
def test_transform_list_runs_request(config: OpenAIEvalsConfig):
|
||||
url, query_params = config.transform_list_runs_request(
|
||||
eval_id="eval_123",
|
||||
list_params={"limit": 5, "after": "evalrun_1", "order": "asc"},
|
||||
litellm_params=GenericLiteLLMParams(api_base="https://api.openai.com"),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert url == "https://api.openai.com/v1/evals/eval_123/runs"
|
||||
assert query_params == {"limit": 5, "after": "evalrun_1", "order": "asc"}
|
||||
|
||||
|
||||
def test_transform_list_runs_response(config: OpenAIEvalsConfig):
|
||||
response = httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [_run_json()],
|
||||
"first_id": "evalrun_123",
|
||||
"last_id": "evalrun_123",
|
||||
"has_more": False,
|
||||
},
|
||||
request=httpx.Request("GET", "https://api.openai.com/v1/evals/eval_123/runs"),
|
||||
)
|
||||
|
||||
result = config.transform_list_runs_response(
|
||||
raw_response=response,
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result.object == "list"
|
||||
assert len(result.data) == 1
|
||||
assert result.data[0].id == "evalrun_123"
|
||||
assert result.has_more is False
|
||||
|
||||
|
||||
def test_transform_get_run_response(config: OpenAIEvalsConfig):
|
||||
response = httpx.Response(
|
||||
status_code=200,
|
||||
json=_run_json(status="completed"),
|
||||
request=httpx.Request("GET", "https://api.openai.com/v1/evals/eval_123/runs/evalrun_123"),
|
||||
)
|
||||
|
||||
result = config.transform_get_run_response(
|
||||
raw_response=response,
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result.id == "evalrun_123"
|
||||
assert result.status == "completed"
|
||||
|
||||
|
||||
def test_transform_cancel_run_response(config: OpenAIEvalsConfig):
|
||||
response = httpx.Response(
|
||||
status_code=200,
|
||||
json={"id": "evalrun_123", "object": "eval.run", "status": "cancelled"},
|
||||
request=httpx.Request(
|
||||
"POST",
|
||||
"https://api.openai.com/v1/evals/eval_123/runs/evalrun_123/cancel",
|
||||
),
|
||||
)
|
||||
|
||||
result = config.transform_cancel_run_response(
|
||||
raw_response=response,
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result.id == "evalrun_123"
|
||||
assert result.status == "cancelled"
|
||||
|
||||
|
||||
def test_transform_delete_run_request(config: OpenAIEvalsConfig):
|
||||
url, headers, request_body = config.transform_delete_run_request(
|
||||
eval_id="eval_123",
|
||||
run_id="evalrun_123",
|
||||
api_base="https://api.openai.com",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert url == "https://api.openai.com/v1/evals/eval_123/runs/evalrun_123"
|
||||
assert request_body == {}
|
||||
|
||||
|
||||
def test_transform_delete_run_response(config: OpenAIEvalsConfig):
|
||||
response = httpx.Response(
|
||||
status_code=200,
|
||||
json={"run_id": "evalrun_123", "object": "eval.run.deleted", "deleted": True},
|
||||
request=httpx.Request("DELETE", "https://api.openai.com/v1/evals/eval_123/runs/evalrun_123"),
|
||||
)
|
||||
|
||||
result = config.transform_delete_run_response(
|
||||
raw_response=response,
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result.run_id == "evalrun_123"
|
||||
assert result.deleted is True
|
||||
|
|
|
|||
|
|
@ -4,9 +4,11 @@ Tests for Volcengine Responses API transformation.
|
|||
|
||||
import os
|
||||
import sys
|
||||
from typing import List, Literal, Optional, Union
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
|
||||
|
|
@ -32,12 +34,10 @@ class TestVolcengineResponsesAPITransformation:
|
|||
)
|
||||
|
||||
assert config is not None, "Config should not be None for Volcengine provider"
|
||||
assert isinstance(
|
||||
config, VolcEngineResponsesAPIConfig
|
||||
), f"Expected VolcEngineResponsesAPIConfig, got {type(config)}"
|
||||
assert (
|
||||
config.custom_llm_provider == LlmProviders.VOLCENGINE
|
||||
), "custom_llm_provider should be VOLCENGINE"
|
||||
assert isinstance(config, VolcEngineResponsesAPIConfig), (
|
||||
f"Expected VolcEngineResponsesAPIConfig, got {type(config)}"
|
||||
)
|
||||
assert config.custom_llm_provider == LlmProviders.VOLCENGINE, "custom_llm_provider should be VOLCENGINE"
|
||||
|
||||
def test_parallel_tool_calls_dropped(self):
|
||||
"""Volcengine does not list parallel_tool_calls; ensure it is removed."""
|
||||
|
|
@ -54,9 +54,7 @@ class TestVolcengineResponsesAPITransformation:
|
|||
drop_params=False,
|
||||
)
|
||||
|
||||
assert (
|
||||
"parallel_tool_calls" not in mapped
|
||||
), "parallel_tool_calls must be dropped"
|
||||
assert "parallel_tool_calls" not in mapped, "parallel_tool_calls must be dropped"
|
||||
assert mapped.get("temperature") == 0.5
|
||||
assert "metadata" not in mapped, "Undocumented params should not be included"
|
||||
|
||||
|
|
@ -91,14 +89,10 @@ class TestVolcengineResponsesAPITransformation:
|
|||
default_url = config.get_complete_url(api_base=None, litellm_params={})
|
||||
assert default_url == "https://ark.cn-beijing.volces.com/api/v3/responses"
|
||||
|
||||
api_base_with_api = config.get_complete_url(
|
||||
api_base="https://custom.volc.com/api/v3", litellm_params={}
|
||||
)
|
||||
api_base_with_api = config.get_complete_url(api_base="https://custom.volc.com/api/v3", litellm_params={})
|
||||
assert api_base_with_api == "https://custom.volc.com/api/v3/responses"
|
||||
|
||||
api_base_full = config.get_complete_url(
|
||||
api_base="https://custom.volc.com/api/v3/responses", litellm_params={}
|
||||
)
|
||||
api_base_full = config.get_complete_url(api_base="https://custom.volc.com/api/v3/responses", litellm_params={})
|
||||
assert api_base_full == "https://custom.volc.com/api/v3/responses"
|
||||
|
||||
def test_response_id_path_requests_encode_response_id(self):
|
||||
|
|
@ -112,10 +106,7 @@ class TestVolcengineResponsesAPITransformation:
|
|||
headers={},
|
||||
)
|
||||
|
||||
assert (
|
||||
url
|
||||
== "https://custom.volc.com/api/v3/responses/..%2F..%2Fresponses%2Fother%3Fx%3D1%23frag/cancel"
|
||||
)
|
||||
assert url == "https://custom.volc.com/api/v3/responses/..%2F..%2Fresponses%2Fother%3Fx%3D1%23frag/cancel"
|
||||
assert params == {}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -125,9 +116,7 @@ class TestVolcengineResponsesAPITransformation:
|
|||
(GenericLiteLLMParams(api_key="attr-key"), "attr-key"),
|
||||
],
|
||||
)
|
||||
def test_validate_environment_uses_api_key(
|
||||
self, monkeypatch, litellm_params, expected_key
|
||||
):
|
||||
def test_validate_environment_uses_api_key(self, monkeypatch, litellm_params, expected_key):
|
||||
"""validate_environment should pull api key from params/env and attach headers."""
|
||||
config = VolcEngineResponsesAPIConfig()
|
||||
|
||||
|
|
@ -135,9 +124,7 @@ class TestVolcengineResponsesAPITransformation:
|
|||
monkeypatch.delenv("ARK_API_KEY", raising=False)
|
||||
monkeypatch.delenv("VOLCENGINE_API_KEY", raising=False)
|
||||
|
||||
headers = config.validate_environment(
|
||||
headers={}, model="volcengine/demo-model", litellm_params=litellm_params
|
||||
)
|
||||
headers = config.validate_environment(headers={}, model="volcengine/demo-model", litellm_params=litellm_params)
|
||||
|
||||
assert headers.get("Authorization") == f"Bearer {expected_key}"
|
||||
assert headers.get("Content-Type") == "application/json"
|
||||
|
|
@ -151,9 +138,7 @@ class TestVolcengineResponsesAPITransformation:
|
|||
monkeypatch.delenv("VOLCENGINE_API_KEY", raising=False)
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
config.validate_environment(
|
||||
headers={}, model="volcengine/demo", litellm_params={}
|
||||
)
|
||||
config.validate_environment(headers={}, model="volcengine/demo", litellm_params={})
|
||||
|
||||
def test_unsupported_params_are_dropped_with_extra_body(self):
|
||||
"""Unknown fields (including extra_body) should be dropped before send."""
|
||||
|
|
@ -240,9 +225,7 @@ class TestVolcengineResponsesAPITransformation:
|
|||
|
||||
# Use class name comparison instead of isinstance to avoid issues with
|
||||
# module reloading during parallel test execution (conftest reloads litellm)
|
||||
assert (
|
||||
type(error).__name__ == "VolcEngineError"
|
||||
), f"Expected VolcEngineError, got {type(error).__name__}"
|
||||
assert type(error).__name__ == "VolcEngineError", f"Expected VolcEngineError, got {type(error).__name__}"
|
||||
assert error.status_code == 400
|
||||
assert error.message == "bad request"
|
||||
assert error.headers.get("x") == "y"
|
||||
|
|
@ -296,3 +279,206 @@ class TestVolcengineResponsesAPITransformation:
|
|||
|
||||
assert isinstance(result, DeleteResponseResult)
|
||||
assert result.deleted is True
|
||||
|
||||
def test_transform_streaming_response_fills_missing_required_fields(self):
|
||||
config = VolcEngineResponsesAPIConfig()
|
||||
|
||||
event = config.transform_streaming_response(
|
||||
model="volcengine/demo-model",
|
||||
parsed_chunk={"type": "response.completed", "response": {"id": "resp_1"}},
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert type(event).__name__ == "ResponseCompletedEvent"
|
||||
assert event.type == "response.completed"
|
||||
assert event.response.id == "resp_1"
|
||||
assert event.response.output == []
|
||||
assert event.response.created_at == 0
|
||||
|
||||
def test_transform_response_api_response_falls_back_to_model_construct(self):
|
||||
config = VolcEngineResponsesAPIConfig()
|
||||
http_response = httpx.Response(
|
||||
status_code=200,
|
||||
json={"id": "resp_fallback", "created_at": 123, "output": "not-a-list"},
|
||||
request=httpx.Request("POST", "https://example.com/responses"),
|
||||
headers={"x-test": "1"},
|
||||
)
|
||||
|
||||
result = config.transform_response_api_response(
|
||||
model="volcengine/demo-model",
|
||||
raw_response=http_response,
|
||||
logging_obj=type(
|
||||
"Logger",
|
||||
(),
|
||||
{"post_call": staticmethod(lambda **kwargs: None)},
|
||||
),
|
||||
)
|
||||
|
||||
assert result.id == "resp_fallback"
|
||||
assert result.output == "not-a-list"
|
||||
assert result._hidden_params["headers"].get("x-test") == "1"
|
||||
|
||||
def test_transform_delete_response_api_request_builds_url(self):
|
||||
config = VolcEngineResponsesAPIConfig()
|
||||
|
||||
url, data = config.transform_delete_response_api_request(
|
||||
response_id="resp_123",
|
||||
api_base="https://custom.volc.com/api/v3/responses",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert url == "https://custom.volc.com/api/v3/responses/resp_123"
|
||||
assert data == {}
|
||||
|
||||
def test_transform_get_response_api_request_and_response(self):
|
||||
config = VolcEngineResponsesAPIConfig()
|
||||
|
||||
url, data = config.transform_get_response_api_request(
|
||||
response_id="resp 123",
|
||||
api_base="https://custom.volc.com/api/v3/responses",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert url == "https://custom.volc.com/api/v3/responses/resp%20123"
|
||||
assert data == {}
|
||||
|
||||
http_response = httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"id": "resp_123",
|
||||
"object": "response",
|
||||
"created_at": 123,
|
||||
"status": "completed",
|
||||
"output": [],
|
||||
"model": "demo-model",
|
||||
},
|
||||
request=httpx.Request("GET", url),
|
||||
headers={"x-test": "1"},
|
||||
)
|
||||
|
||||
result = config.transform_get_response_api_response(
|
||||
raw_response=http_response,
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result.id == "resp_123"
|
||||
assert result._hidden_params["headers"].get("x-test") == "1"
|
||||
|
||||
def test_transform_cancel_response_api_response_parses_json(self):
|
||||
config = VolcEngineResponsesAPIConfig()
|
||||
http_response = httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"id": "resp_123",
|
||||
"object": "response",
|
||||
"created_at": 123,
|
||||
"status": "cancelled",
|
||||
"output": [],
|
||||
"model": "demo-model",
|
||||
},
|
||||
request=httpx.Request("POST", "https://example.com/responses/resp_123/cancel"),
|
||||
headers={"x-test": "1"},
|
||||
)
|
||||
|
||||
result = config.transform_cancel_response_api_response(
|
||||
raw_response=http_response,
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result.id == "resp_123"
|
||||
assert result.status == "cancelled"
|
||||
assert result._hidden_params["headers"].get("x-test") == "1"
|
||||
|
||||
def test_transform_list_input_items_request_builds_query_params(self):
|
||||
config = VolcEngineResponsesAPIConfig()
|
||||
|
||||
url, params = config.transform_list_input_items_request(
|
||||
response_id="resp_123",
|
||||
api_base="https://custom.volc.com/api/v3/responses",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
after="item_a",
|
||||
before="item_b",
|
||||
include=["metadata", "usage"],
|
||||
limit=5,
|
||||
order="asc",
|
||||
)
|
||||
|
||||
assert url == "https://custom.volc.com/api/v3/responses/resp_123/input_items"
|
||||
assert params == {
|
||||
"after": "item_a",
|
||||
"before": "item_b",
|
||||
"include": "metadata,usage",
|
||||
"limit": 5,
|
||||
"order": "asc",
|
||||
}
|
||||
|
||||
def test_transform_list_input_items_response_returns_parsed_body(self):
|
||||
config = VolcEngineResponsesAPIConfig()
|
||||
payload = {"object": "list", "data": [{"id": "item_1"}]}
|
||||
http_response = httpx.Response(
|
||||
status_code=200,
|
||||
json=payload,
|
||||
request=httpx.Request("GET", "https://example.com/responses/resp_123/input_items"),
|
||||
)
|
||||
|
||||
result = config.transform_list_input_items_response(
|
||||
raw_response=http_response,
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result == payload
|
||||
|
||||
|
||||
class _FillWidget(BaseModel):
|
||||
type: Literal["widget"]
|
||||
count: int
|
||||
parts: List[str]
|
||||
label: Optional[str]
|
||||
|
||||
|
||||
class _FillGadget(BaseModel):
|
||||
type: Literal["gadget"]
|
||||
name: str
|
||||
|
||||
|
||||
class _FillEnvelope(BaseModel):
|
||||
kind: str = "envelope"
|
||||
tags: List[str] = Field(default_factory=lambda: ["default-tag"])
|
||||
payload: Union[_FillWidget, _FillGadget]
|
||||
entries: List[_FillWidget]
|
||||
note: Optional[str]
|
||||
values: Union[List[str], str]
|
||||
|
||||
|
||||
class TestVolcengineStreamingFieldFill:
|
||||
def test_fill_uses_defaults_factories_and_heuristics(self):
|
||||
filled = VolcEngineResponsesAPIConfig._fill_missing_fields(
|
||||
{"payload": {"type": "gadget", "name": "g"}, "entries": [{"type": "widget"}]},
|
||||
_FillEnvelope,
|
||||
)
|
||||
|
||||
assert filled["kind"] == "envelope"
|
||||
assert filled["tags"] == ["default-tag"]
|
||||
assert filled["note"] is None
|
||||
assert filled["values"] == []
|
||||
|
||||
validated = _FillEnvelope.model_validate(filled)
|
||||
assert isinstance(validated.payload, _FillGadget)
|
||||
assert validated.entries[0].count == 0
|
||||
assert validated.entries[0].parts == []
|
||||
assert validated.entries[0].label is None
|
||||
|
||||
def test_fill_selects_union_member_by_type_literal(self):
|
||||
filled = VolcEngineResponsesAPIConfig._fill_missing_fields(
|
||||
{"payload": {"type": "widget"}, "entries": []},
|
||||
_FillEnvelope,
|
||||
)
|
||||
|
||||
validated = _FillEnvelope.model_validate(filled)
|
||||
assert isinstance(validated.payload, _FillWidget)
|
||||
assert validated.payload.count == 0
|
||||
assert validated.payload.parts == []
|
||||
assert validated.payload.label is None
|
||||
|
|
|
|||
|
|
@ -6,7 +6,8 @@ Tests cover:
|
|||
resolved via tool_call_id, falling back to the last user message)
|
||||
- target selection: tool outputs by default, system/history opt-in, min-chars
|
||||
threshold, targets without a derivable query are left uncompressed
|
||||
- multimodal content: text parts replaced, non-text parts preserved
|
||||
- multimodal content: all-text rows merge into one part carrying the last
|
||||
cache_control breakpoint, rows holding a non-text part are left uncompressed
|
||||
- recovery: hash marker appended, compresr_retrieve tool injected, originals
|
||||
stored per litellm_call_id, agentic loop returns the original content and
|
||||
rejects hashes not issued for the current request
|
||||
|
|
@ -560,7 +561,7 @@ async def test_short_messages_skipped(guardrail: CompresrGuardrail):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multimodal_text_replaced_non_text_preserved(
|
||||
async def test_multimodal_row_with_non_text_part_is_not_compressed(
|
||||
guardrail: CompresrGuardrail,
|
||||
):
|
||||
image_part = {"type": "image_url", "image_url": {"url": "https://example.com/x.png"}}
|
||||
|
|
@ -572,6 +573,35 @@ async def test_multimodal_text_replaced_non_text_preserved(
|
|||
"content": [{"type": "text", "text": TOOL_OUTPUT}, image_part],
|
||||
},
|
||||
]
|
||||
expected = json.loads(json.dumps(messages[1]["content"]))
|
||||
mock_post = AsyncMock(return_value=_make_single_compress_response())
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", mock_post):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=_apply_inputs(messages),
|
||||
request_data={"model": "gpt-4o"},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
mock_post.assert_not_called()
|
||||
assert result["structured_messages"][1]["content"] == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_all_text_row_merges_and_keeps_last_cache_control(
|
||||
guardrail: CompresrGuardrail,
|
||||
):
|
||||
messages = [
|
||||
{"role": "user", "content": USER_QUESTION},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "c1",
|
||||
"content": [
|
||||
{"type": "text", "text": TOOL_OUTPUT, "cache_control": {"type": "ephemeral"}},
|
||||
{"type": "text", "text": TOOL_OUTPUT, "cache_control": {"type": "ephemeral", "ttl": "1h"}},
|
||||
],
|
||||
},
|
||||
]
|
||||
mock_post = AsyncMock(return_value=_make_single_compress_response())
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", mock_post):
|
||||
|
|
@ -583,9 +613,41 @@ async def test_multimodal_text_replaced_non_text_preserved(
|
|||
|
||||
content = result["structured_messages"][1]["content"]
|
||||
assert isinstance(content, list)
|
||||
assert len(content) == 1
|
||||
assert content[0]["type"] == "text"
|
||||
assert content[0]["text"].startswith("compressed summary")
|
||||
assert content[1] == image_part
|
||||
assert content[0]["cache_control"] == {"type": "ephemeral", "ttl": "1h"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_around_non_text_part_is_never_relocated(
|
||||
guardrail: CompresrGuardrail,
|
||||
):
|
||||
image_part = {"type": "image_url", "image_url": {"url": "https://example.com/x.png"}}
|
||||
messages = [
|
||||
{"role": "user", "content": USER_QUESTION},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "c1",
|
||||
"content": [
|
||||
{"type": "text", "text": TOOL_OUTPUT},
|
||||
image_part,
|
||||
{"type": "text", "text": TOOL_OUTPUT, "cache_control": {"type": "ephemeral"}},
|
||||
],
|
||||
},
|
||||
]
|
||||
expected = json.loads(json.dumps(messages[1]["content"]))
|
||||
mock_post = AsyncMock(return_value=_make_single_compress_response())
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", mock_post):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=_apply_inputs(messages),
|
||||
request_data={"model": "gpt-4o"},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
mock_post.assert_not_called()
|
||||
assert result["structured_messages"][1]["content"] == expected
|
||||
|
||||
|
||||
# ── passthrough / bypass ─────────────────────────────────────────────
|
||||
|
|
|
|||
|
|
@ -0,0 +1,433 @@
|
|||
from datetime import datetime, timezone
|
||||
from typing import List
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.exceptions import RequestValidationError
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.proxy._types import LiteLLMRoutes, LitellmUserRoles
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
|
||||
from litellm.proxy.management_endpoints.management_v1 import router
|
||||
from litellm.proxy.management_endpoints.management_v1.common import (
|
||||
MANAGEMENT_V1_PREFIX,
|
||||
PROBLEM_TYPE_BASE,
|
||||
ManagementProblem,
|
||||
problem_response,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import ProblemDetail
|
||||
|
||||
app = FastAPI()
|
||||
|
||||
|
||||
@app.exception_handler(ManagementProblem)
|
||||
async def management_problem_exception_handler(request: Request, exc: ManagementProblem):
|
||||
return problem_response(exc.problem)
|
||||
|
||||
|
||||
@app.exception_handler(RequestValidationError)
|
||||
async def validation_exception_handler(request: Request, exc: RequestValidationError):
|
||||
return problem_response(
|
||||
ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}invalid-query-parameter",
|
||||
title="Invalid query parameter",
|
||||
status=400,
|
||||
detail="; ".join(
|
||||
f"{'.'.join(str(part) for part in error['loc'][1:])}: {error['msg']}" for error in exc.errors()
|
||||
)
|
||||
or "The request query parameters are invalid.",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
app.include_router(router)
|
||||
client = TestClient(app)
|
||||
|
||||
END_USERS_PATH = f"{MANAGEMENT_V1_PREFIX}/spend_logs/end_users"
|
||||
WINDOW = "filter[startTime][gte]=2026-07-23T00:00:00Z&filter[startTime][lte]=2026-07-24T00:00:00Z"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_prisma_client(monkeypatch):
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.query_raw = AsyncMock(return_value=[])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
|
||||
return prisma_client
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def as_proxy_admin():
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
yield
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
def _mock_rows(mock_prisma_client, end_users: List[str]) -> AsyncMock:
|
||||
query_raw = AsyncMock(return_value=[{"end_user": eu} for eu in end_users])
|
||||
mock_prisma_client.db.query_raw = query_raw
|
||||
return query_raw
|
||||
|
||||
|
||||
def _as_role(role: LitellmUserRoles, user_id):
|
||||
original = app.dependency_overrides.copy()
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_id=user_id, user_role=role)
|
||||
return original
|
||||
|
||||
|
||||
def _get(query: str = WINDOW):
|
||||
suffix = f"?{query}" if query else ""
|
||||
return client.get(f"{END_USERS_PATH}{suffix}", headers={"Authorization": "Bearer k"})
|
||||
|
||||
|
||||
def test_returns_the_control_plane_envelope(mock_prisma_client, as_proxy_admin):
|
||||
"""`{data, meta, links}` is the contract; a bare list or a legacy `aliases` key is not."""
|
||||
_mock_rows(mock_prisma_client, ["a", "b"])
|
||||
|
||||
response = _get()
|
||||
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["data"] == ["a", "b"]
|
||||
assert body["meta"] == {"page": 1, "page_size": 50, "has_more": False}
|
||||
assert set(body) == {"data", "meta", "links"}
|
||||
assert "aliases" not in body
|
||||
assert "total_count" not in body["meta"]
|
||||
|
||||
|
||||
def test_links_let_a_client_page_without_building_urls(mock_prisma_client, as_proxy_admin):
|
||||
"""The UI follows links.next; if it is absent the client has to recompute page params,
|
||||
which is what makes a later switch to cursor pagination a breaking change."""
|
||||
_mock_rows(mock_prisma_client, [f"u{i}" for i in range(4)])
|
||||
|
||||
links = _get(f"{WINDOW}&page=2&page_size=3").json()["links"]
|
||||
|
||||
assert links["self"].startswith(f"{END_USERS_PATH}?")
|
||||
assert "page=2" in links["self"]
|
||||
assert "page=1" in links["prev"] and "page_size=3" in links["prev"]
|
||||
assert "page=3" in links["next"] and "page_size=3" in links["next"]
|
||||
|
||||
|
||||
def test_next_link_is_absent_on_the_last_page(mock_prisma_client, as_proxy_admin):
|
||||
_mock_rows(mock_prisma_client, ["u0", "u1"])
|
||||
|
||||
body = _get(f"{WINDOW}&page_size=3").json()
|
||||
|
||||
assert body["meta"]["has_more"] is False
|
||||
assert body["links"]["next"] is None
|
||||
assert body["links"]["prev"] is None
|
||||
|
||||
|
||||
def test_reads_spend_logs_not_the_end_user_table(mock_prisma_client, as_proxy_admin):
|
||||
"""Team scoping only exists in spend logs, so that is the source of truth."""
|
||||
query_raw = _mock_rows(mock_prisma_client, ["a"])
|
||||
|
||||
_get()
|
||||
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert '"LiteLLM_SpendLogs"' in sql
|
||||
assert "LiteLLM_EndUserTable" not in sql
|
||||
|
||||
|
||||
def test_caps_the_rows_it_scans(mock_prisma_client, as_proxy_admin):
|
||||
"""The inner LIMIT is the crash guard: DISTINCT must never see an unbounded set."""
|
||||
from litellm.proxy.management_endpoints.management_v1.spend_logs import (
|
||||
SPEND_LOGS_FACET_SCAN_CAP,
|
||||
)
|
||||
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
|
||||
_get()
|
||||
|
||||
sql = query_raw.call_args.args[0]
|
||||
inner = sql[sql.index("FROM (") : sql.index(") recent")]
|
||||
assert "LIMIT $3" in inner
|
||||
assert query_raw.call_args.args[3] == SPEND_LOGS_FACET_SCAN_CAP
|
||||
assert 'ORDER BY "startTime" DESC' in inner
|
||||
|
||||
|
||||
def test_scan_cap_matches_the_logs_page_bound():
|
||||
"""Pin the cap's value, not just that it is passed through.
|
||||
|
||||
Asserting the param equals the constant is tautological: raising the constant
|
||||
to a billion keeps that assertion green while removing the bound entirely.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.management_v1.spend_logs import (
|
||||
SPEND_LOGS_FACET_SCAN_CAP,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.spend_management_endpoints import (
|
||||
SPEND_LOGS_PAGINATION_COUNT_CAP,
|
||||
)
|
||||
|
||||
assert SPEND_LOGS_FACET_SCAN_CAP == SPEND_LOGS_PAGINATION_COUNT_CAP
|
||||
|
||||
|
||||
def test_breaks_start_time_ties_deterministically(mock_prisma_client, as_proxy_admin):
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
|
||||
_get()
|
||||
|
||||
assert 'ORDER BY "startTime" DESC, request_id DESC' in query_raw.call_args.args[0]
|
||||
|
||||
|
||||
def test_bounds_the_window_on_the_indexed_start_time(mock_prisma_client, as_proxy_admin):
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
|
||||
_get()
|
||||
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert "\"startTime\" >= ($1::timestamptz AT TIME ZONE 'UTC')" in sql
|
||||
assert "\"startTime\" <= ($2::timestamptz AT TIME ZONE 'UTC')" in sql
|
||||
assert query_raw.call_args.args[1] == datetime(2026, 7, 23, tzinfo=timezone.utc)
|
||||
assert query_raw.call_args.args[2] == datetime(2026, 7, 24, tzinfo=timezone.utc)
|
||||
|
||||
|
||||
def test_a_naive_window_bound_is_read_as_utc(mock_prisma_client, as_proxy_admin):
|
||||
"""The dashboard sends 'YYYY-MM-DD HH:MM:SS' with no offset; reading it as
|
||||
server-local time would shift the window off what the logs table is showing."""
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
|
||||
_get("filter[startTime][gte]=2026-07-23 00:00:00&filter[startTime][lte]=2026-07-24 00:00:00")
|
||||
|
||||
assert query_raw.call_args.args[1] == datetime(2026, 7, 23, tzinfo=timezone.utc)
|
||||
assert query_raw.call_args.args[2] == datetime(2026, 7, 24, tzinfo=timezone.utc)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"query",
|
||||
["", "filter[startTime][gte]=2026-07-23T00:00:00Z"],
|
||||
ids=["no-window", "half-window"],
|
||||
)
|
||||
def test_requires_a_time_window(mock_prisma_client, as_proxy_admin, query):
|
||||
"""No window means no index bound, which is the unbounded scan we must not allow."""
|
||||
_mock_rows(mock_prisma_client, [])
|
||||
|
||||
response = _get(query)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.headers["content-type"].startswith("application/problem+json")
|
||||
|
||||
|
||||
def test_rejects_a_malformed_window_as_a_problem_document(mock_prisma_client, as_proxy_admin):
|
||||
_mock_rows(mock_prisma_client, [])
|
||||
|
||||
response = _get(f"filter[startTime][gte]=yesterday&filter[startTime][lte]=2026-07-24T00:00:00Z")
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.headers["content-type"].startswith("application/problem+json")
|
||||
body = response.json()
|
||||
assert body["type"].startswith(PROBLEM_TYPE_BASE)
|
||||
assert body["status"] == 400
|
||||
assert body["title"] and body["detail"]
|
||||
assert "error" not in body
|
||||
|
||||
|
||||
def test_problem_type_is_an_identifier_not_a_dead_docs_link(mock_prisma_client, as_proxy_admin):
|
||||
"""RFC 9457 only asks that `type` identify the problem type. An https URI promises
|
||||
human-readable documentation at that address, and https://docs.litellm.ai/errors/
|
||||
is a 404, so emitting one would ship a broken link in every error body."""
|
||||
_mock_rows(mock_prisma_client, [])
|
||||
|
||||
problem_type = _get("filter[startTime][gte]=yesterday&filter[startTime][lte]=2026-07-24T00:00:00Z").json()["type"]
|
||||
|
||||
assert problem_type.startswith("urn:")
|
||||
assert "docs.litellm.ai" not in problem_type
|
||||
assert not problem_type.startswith("http")
|
||||
|
||||
|
||||
def test_rejects_an_unknown_query_parameter(mock_prisma_client, as_proxy_admin):
|
||||
"""A silently ignored filter over-returns data, which is worse than a rejected request."""
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
|
||||
response = _get(f"{WINDOW}&q_typo=acme")
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.headers["content-type"].startswith("application/problem+json")
|
||||
body = response.json()
|
||||
assert "q_typo" in body["detail"]
|
||||
assert "q" in body["allowed"]
|
||||
query_raw.assert_not_called()
|
||||
|
||||
|
||||
def test_accepts_every_declared_parameter(mock_prisma_client, as_proxy_admin):
|
||||
"""Guards the unknown-param check against rejecting the endpoint's own contract."""
|
||||
_mock_rows(mock_prisma_client, [])
|
||||
|
||||
assert _get(f"{WINDOW}&q=acme&page=2&page_size=10").status_code == 200
|
||||
|
||||
|
||||
def test_caps_page_size(mock_prisma_client, as_proxy_admin):
|
||||
_mock_rows(mock_prisma_client, [])
|
||||
|
||||
assert _get(f"{WINDOW}&page_size=100000").status_code == 400
|
||||
|
||||
|
||||
def test_applies_no_scope_for_a_proxy_admin(mock_prisma_client, as_proxy_admin):
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
|
||||
_get()
|
||||
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert '"user" =' not in sql
|
||||
assert "team_id" not in sql
|
||||
|
||||
|
||||
@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
|
||||
def test_scopes_a_team_admin_to_their_own_rows_and_teams(mock_prisma_client, role):
|
||||
"""A team admin must not see end users belonging to teams they cannot read."""
|
||||
query_raw = _mock_rows(mock_prisma_client, ["cust-a"])
|
||||
original = _as_role(role, user_id="team-admin-1")
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs",
|
||||
new=AsyncMock(return_value=["team-a", "team-b"]),
|
||||
):
|
||||
response = _get()
|
||||
finally:
|
||||
app.dependency_overrides = original
|
||||
|
||||
assert response.status_code == 200
|
||||
# Same clause shape ui_view_spend_logs builds, so the two cannot diverge.
|
||||
assert '("user" = $3 OR team_id = ANY($4::text[]))' in query_raw.call_args.args[0]
|
||||
assert query_raw.call_args.args[3] == "team-admin-1"
|
||||
assert query_raw.call_args.args[4] == ["team-a", "team-b"]
|
||||
|
||||
|
||||
def test_scopes_a_teamless_user_to_their_own_rows(mock_prisma_client):
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id="solo")
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs",
|
||||
new=AsyncMock(return_value=[]),
|
||||
):
|
||||
response = _get()
|
||||
finally:
|
||||
app.dependency_overrides = original
|
||||
|
||||
assert response.status_code == 200
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert '("user" = $3)' in sql
|
||||
assert "team_id" not in sql
|
||||
assert query_raw.call_args.args[3] == "solo"
|
||||
|
||||
|
||||
def test_returns_nothing_when_the_caller_owns_no_scope(mock_prisma_client):
|
||||
"""Unidentifiable caller must match no rows, never fall through to unscoped."""
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id=None)
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs",
|
||||
new=AsyncMock(return_value=[]),
|
||||
):
|
||||
response = _get()
|
||||
finally:
|
||||
app.dependency_overrides = original
|
||||
|
||||
assert response.status_code == 200
|
||||
assert "FALSE" in query_raw.call_args.args[0]
|
||||
|
||||
|
||||
def test_scopes_when_the_permitted_team_lookup_fails(mock_prisma_client):
|
||||
"""A failed team lookup must degrade to own-rows-only, never to unscoped."""
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id="solo")
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs",
|
||||
new=AsyncMock(side_effect=RuntimeError("db down")),
|
||||
):
|
||||
response = _get()
|
||||
finally:
|
||||
app.dependency_overrides = original
|
||||
|
||||
assert response.status_code == 200
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert '("user" = $3)' in sql
|
||||
assert "team_id" not in sql
|
||||
|
||||
|
||||
def test_fetches_one_extra_row_and_trims_it(mock_prisma_client, as_proxy_admin):
|
||||
query_raw = _mock_rows(mock_prisma_client, [f"u{i}" for i in range(4)])
|
||||
|
||||
body = _get(f"{WINDOW}&page_size=3").json()
|
||||
|
||||
assert body["data"] == ["u0", "u1", "u2"]
|
||||
assert body["meta"]["has_more"] is True
|
||||
assert query_raw.call_args.args[4:] == (4, 0)
|
||||
|
||||
|
||||
def test_reports_no_more_pages_on_an_exactly_full_page(mock_prisma_client, as_proxy_admin):
|
||||
_mock_rows(mock_prisma_client, ["u0", "u1", "u2"])
|
||||
|
||||
body = _get(f"{WINDOW}&page_size=3").json()
|
||||
|
||||
assert body["data"] == ["u0", "u1", "u2"]
|
||||
assert body["meta"]["has_more"] is False
|
||||
|
||||
|
||||
def test_offsets_by_page(mock_prisma_client, as_proxy_admin):
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
|
||||
body = _get(f"{WINDOW}&page=3&page_size=25").json()
|
||||
|
||||
assert body["meta"]["page"] == 3
|
||||
assert query_raw.call_args.args[4:] == (26, 50)
|
||||
|
||||
|
||||
def test_q_escapes_like_metacharacters(mock_prisma_client, as_proxy_admin):
|
||||
"""End-user ids routinely contain '_'; unescaped it is a wildcard."""
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
|
||||
_get(f"{WINDOW}&q=device_id%25")
|
||||
|
||||
assert "end_user ILIKE $3 ESCAPE" in query_raw.call_args.args[0]
|
||||
assert query_raw.call_args.args[3] == r"%device\_id\%%"
|
||||
|
||||
|
||||
def test_q_placeholder_precedes_the_scan_limit_and_offset(mock_prisma_client, as_proxy_admin):
|
||||
query_raw = _mock_rows(mock_prisma_client, [])
|
||||
|
||||
_get(f"{WINDOW}&q=acme&page_size=10")
|
||||
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert "LIMIT $4" in sql
|
||||
assert "LIMIT $5 OFFSET $6" in sql
|
||||
assert query_raw.call_args.args[3] == "%acme%"
|
||||
assert query_raw.call_args.args[5:] == (11, 0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"role",
|
||||
[
|
||||
LitellmUserRoles.PROXY_ADMIN,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||||
LitellmUserRoles.INTERNAL_USER,
|
||||
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
|
||||
],
|
||||
)
|
||||
def test_is_reachable_by_every_role_that_can_open_the_logs_page(role):
|
||||
"""Route-level auth gate, which the dependency_overrides in the other tests bypass.
|
||||
|
||||
Handler-side team scoping is dead code if RouteChecks rejects the role first.
|
||||
"""
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
|
||||
for allowed in (
|
||||
LiteLLMRoutes.internal_user_routes.value,
|
||||
LiteLLMRoutes.internal_user_view_only_routes.value,
|
||||
):
|
||||
assert ("/spend/logs/ui" in allowed) == (END_USERS_PATH in allowed)
|
||||
|
||||
if role in (LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY):
|
||||
allowed_routes = (
|
||||
LiteLLMRoutes.internal_user_routes.value
|
||||
if role == LitellmUserRoles.INTERNAL_USER
|
||||
else LiteLLMRoutes.internal_user_view_only_routes.value
|
||||
)
|
||||
assert RouteChecks.check_route_access(route=END_USERS_PATH, allowed_routes=allowed_routes)
|
||||
else:
|
||||
assert END_USERS_PATH in LiteLLMRoutes.admin_viewer_routes.value
|
||||
|
|
@ -1,4 +1,3 @@
|
|||
from datetime import datetime, timezone
|
||||
from typing import List
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -10,7 +9,6 @@ from fastapi.testclient import TestClient
|
|||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_EndUserTable,
|
||||
LiteLLMRoutes,
|
||||
LitellmUserRoles,
|
||||
ProxyException,
|
||||
)
|
||||
|
|
@ -784,291 +782,3 @@ def test_char_delete_body(mock_prisma_client, mock_user_api_key_auth):
|
|||
"deleted_customers": 2,
|
||||
"message": "Successfully deleted customers with ids: ['c1', 'c2']",
|
||||
}
|
||||
|
||||
|
||||
WINDOW = "start_date=2026-07-23+00%3A00%3A00&end_date=2026-07-24+00%3A00%3A00"
|
||||
|
||||
|
||||
def _mock_alias_rows(mock_prisma_client, end_users: List[str]) -> AsyncMock:
|
||||
query_raw = AsyncMock(return_value=[{"end_user": eu} for eu in end_users])
|
||||
mock_prisma_client.db.query_raw = query_raw
|
||||
return query_raw
|
||||
|
||||
|
||||
def _as_role(role: LitellmUserRoles, user_id: str = "u1"):
|
||||
"""Override auth for one request; returns a context-manager-free setter/teardown pair."""
|
||||
original = app.dependency_overrides.copy()
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_id=user_id, user_role=role)
|
||||
return original
|
||||
|
||||
|
||||
def test_customer_aliases_reads_spend_logs_not_the_end_user_table(mock_prisma_client, mock_user_api_key_auth):
|
||||
"""Team scoping only exists in spend logs, so that is the source of truth."""
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, ["a", "b"])
|
||||
|
||||
response = client.get(f"/customer/aliases?{WINDOW}", headers={"Authorization": "Bearer k"})
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"aliases": ["a", "b"], "current_page": 1, "size": 50, "has_more": False}
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert '"LiteLLM_SpendLogs"' in sql
|
||||
assert "LiteLLM_EndUserTable" not in sql
|
||||
mock_prisma_client.db.litellm_endusertable.find_many.assert_not_called()
|
||||
|
||||
|
||||
def test_customer_aliases_caps_the_rows_it_scans(mock_prisma_client, mock_user_api_key_auth):
|
||||
"""The inner LIMIT is the crash guard: DISTINCT must never see an unbounded set."""
|
||||
from litellm.proxy.management_endpoints.customer_endpoints import SPEND_LOGS_FILTER_SCAN_CAP
|
||||
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, [])
|
||||
|
||||
client.get(f"/customer/aliases?{WINDOW}", headers={"Authorization": "Bearer k"})
|
||||
|
||||
sql = query_raw.call_args.args[0]
|
||||
inner = sql[sql.index("FROM (") : sql.index(") recent")]
|
||||
assert "LIMIT $3" in inner
|
||||
assert query_raw.call_args.args[3] == SPEND_LOGS_FILTER_SCAN_CAP
|
||||
assert 'ORDER BY "startTime" DESC' in inner
|
||||
|
||||
|
||||
def test_spend_logs_filter_scan_cap_matches_the_logs_page_bound():
|
||||
"""Pin the cap's value, not just that it is passed through.
|
||||
|
||||
Asserting the param equals the constant is tautological: raising the constant
|
||||
to a billion keeps that assertion green while removing the bound entirely.
|
||||
The documented rationale is that both reads of LiteLLM_SpendLogs stop at the
|
||||
same depth, so tie it to the count cap ui_view_spend_logs already uses.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.customer_endpoints import SPEND_LOGS_FILTER_SCAN_CAP
|
||||
from litellm.proxy.spend_tracking.spend_management_endpoints import (
|
||||
SPEND_LOGS_PAGINATION_COUNT_CAP,
|
||||
)
|
||||
|
||||
assert SPEND_LOGS_FILTER_SCAN_CAP == SPEND_LOGS_PAGINATION_COUNT_CAP
|
||||
|
||||
|
||||
def test_customer_aliases_breaks_start_time_ties_deterministically(mock_prisma_client, mock_user_api_key_auth):
|
||||
"""Without a unique tiebreaker the capped scan can cut differently per request,
|
||||
so OFFSET page 2 would page through a different set than page 1 did."""
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, [])
|
||||
|
||||
client.get(f"/customer/aliases?{WINDOW}", headers={"Authorization": "Bearer k"})
|
||||
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert 'ORDER BY "startTime" DESC, request_id DESC' in sql
|
||||
|
||||
|
||||
def test_customer_aliases_requires_a_time_window(mock_prisma_client, mock_user_api_key_auth):
|
||||
"""No window means no index bound, which is the unbounded scan we must not allow."""
|
||||
_mock_alias_rows(mock_prisma_client, [])
|
||||
|
||||
assert client.get("/customer/aliases", headers={"Authorization": "Bearer k"}).status_code == 422
|
||||
assert (
|
||||
client.get(
|
||||
"/customer/aliases?start_date=2026-07-23+00%3A00%3A00", headers={"Authorization": "Bearer k"}
|
||||
).status_code
|
||||
== 422
|
||||
)
|
||||
|
||||
|
||||
def test_customer_aliases_bounds_the_window_on_the_indexed_start_time(mock_prisma_client, mock_user_api_key_auth):
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, [])
|
||||
|
||||
client.get(f"/customer/aliases?{WINDOW}", headers={"Authorization": "Bearer k"})
|
||||
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert "\"startTime\" >= ($1::timestamptz AT TIME ZONE 'UTC')" in sql
|
||||
assert "\"startTime\" <= ($2::timestamptz AT TIME ZONE 'UTC')" in sql
|
||||
assert query_raw.call_args.args[1] == datetime(2026, 7, 23, tzinfo=timezone.utc)
|
||||
assert query_raw.call_args.args[2] == datetime(2026, 7, 24, tzinfo=timezone.utc)
|
||||
|
||||
|
||||
def test_customer_aliases_rejects_a_malformed_window(mock_prisma_client, mock_user_api_key_auth):
|
||||
_mock_alias_rows(mock_prisma_client, [])
|
||||
|
||||
response = client.get(
|
||||
f"/customer/aliases?start_date=yesterday&end_date=2026-07-24+00%3A00%3A00",
|
||||
headers={"Authorization": "Bearer k"},
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
|
||||
|
||||
def test_customer_aliases_applies_no_scope_for_a_proxy_admin(mock_prisma_client, mock_user_api_key_auth):
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, [])
|
||||
|
||||
client.get(f"/customer/aliases?{WINDOW}", headers={"Authorization": "Bearer k"})
|
||||
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert '"user" =' not in sql
|
||||
assert "team_id" not in sql
|
||||
|
||||
|
||||
@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
|
||||
def test_customer_aliases_scopes_a_team_admin_to_their_own_rows_and_teams(mock_prisma_client, role):
|
||||
"""A team admin must not see end users belonging to teams they cannot read."""
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, ["cust-a"])
|
||||
original = _as_role(role, user_id="team-admin-1")
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs",
|
||||
new=AsyncMock(return_value=["team-a", "team-b"]),
|
||||
):
|
||||
response = client.get(f"/customer/aliases?{WINDOW}", headers={"Authorization": "Bearer k"})
|
||||
finally:
|
||||
app.dependency_overrides = original
|
||||
|
||||
assert response.status_code == 200
|
||||
sql = query_raw.call_args.args[0]
|
||||
# Same clause shape ui_view_spend_logs builds, so the two cannot diverge.
|
||||
assert '("user" = $3 OR team_id = ANY($4::text[]))' in sql
|
||||
assert query_raw.call_args.args[3] == "team-admin-1"
|
||||
assert query_raw.call_args.args[4] == ["team-a", "team-b"]
|
||||
|
||||
|
||||
def test_customer_aliases_scopes_a_teamless_user_to_their_own_rows(mock_prisma_client):
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, [])
|
||||
original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id="solo")
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs",
|
||||
new=AsyncMock(return_value=[]),
|
||||
):
|
||||
response = client.get(f"/customer/aliases?{WINDOW}", headers={"Authorization": "Bearer k"})
|
||||
finally:
|
||||
app.dependency_overrides = original
|
||||
|
||||
assert response.status_code == 200
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert '("user" = $3)' in sql
|
||||
assert "team_id" not in sql
|
||||
assert query_raw.call_args.args[3] == "solo"
|
||||
|
||||
|
||||
def test_customer_aliases_returns_nothing_when_the_caller_owns_no_scope(mock_prisma_client):
|
||||
"""Unidentifiable caller must match no rows, never fall through to unscoped."""
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, [])
|
||||
original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id=None)
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs",
|
||||
new=AsyncMock(return_value=[]),
|
||||
):
|
||||
response = client.get(f"/customer/aliases?{WINDOW}", headers={"Authorization": "Bearer k"})
|
||||
finally:
|
||||
app.dependency_overrides = original
|
||||
|
||||
assert response.status_code == 200
|
||||
assert "FALSE" in query_raw.call_args.args[0]
|
||||
|
||||
|
||||
def test_customer_aliases_scopes_when_permitted_team_lookup_fails(mock_prisma_client):
|
||||
"""A failed team lookup must degrade to own-rows-only, never to unscoped."""
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, [])
|
||||
original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id="solo")
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs",
|
||||
new=AsyncMock(side_effect=RuntimeError("db down")),
|
||||
):
|
||||
response = client.get(f"/customer/aliases?{WINDOW}", headers={"Authorization": "Bearer k"})
|
||||
finally:
|
||||
app.dependency_overrides = original
|
||||
|
||||
assert response.status_code == 200
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert '("user" = $3)' in sql
|
||||
assert "team_id" not in sql
|
||||
|
||||
|
||||
def test_customer_aliases_fetches_one_extra_row_and_trims_it(mock_prisma_client, mock_user_api_key_auth):
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, [f"u{i}" for i in range(4)])
|
||||
|
||||
response = client.get(f"/customer/aliases?{WINDOW}&size=3", headers={"Authorization": "Bearer k"})
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["aliases"] == ["u0", "u1", "u2"]
|
||||
assert response.json()["has_more"] is True
|
||||
assert query_raw.call_args.args[4:] == (4, 0)
|
||||
|
||||
|
||||
def test_customer_aliases_reports_no_more_pages_on_an_exactly_full_page(mock_prisma_client, mock_user_api_key_auth):
|
||||
_mock_alias_rows(mock_prisma_client, ["u0", "u1", "u2"])
|
||||
|
||||
response = client.get(f"/customer/aliases?{WINDOW}&size=3", headers={"Authorization": "Bearer k"})
|
||||
|
||||
assert response.json()["aliases"] == ["u0", "u1", "u2"]
|
||||
assert response.json()["has_more"] is False
|
||||
|
||||
|
||||
def test_customer_aliases_offsets_by_page(mock_prisma_client, mock_user_api_key_auth):
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, [])
|
||||
|
||||
response = client.get(f"/customer/aliases?{WINDOW}&page=3&size=25", headers={"Authorization": "Bearer k"})
|
||||
|
||||
assert response.json()["current_page"] == 3
|
||||
assert query_raw.call_args.args[4:] == (26, 50)
|
||||
|
||||
|
||||
def test_customer_aliases_search_escapes_like_metacharacters(mock_prisma_client, mock_user_api_key_auth):
|
||||
"""End-user ids routinely contain '_'; unescaped it is a wildcard."""
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, [])
|
||||
|
||||
client.get(f"/customer/aliases?{WINDOW}&search=device_id%25", headers={"Authorization": "Bearer k"})
|
||||
|
||||
assert "end_user ILIKE $3 ESCAPE" in query_raw.call_args.args[0]
|
||||
assert query_raw.call_args.args[3] == r"%device\_id\%%"
|
||||
|
||||
|
||||
def test_customer_aliases_search_placeholder_precedes_scan_limit_and_offset(mock_prisma_client, mock_user_api_key_auth):
|
||||
query_raw = _mock_alias_rows(mock_prisma_client, [])
|
||||
|
||||
client.get(f"/customer/aliases?{WINDOW}&search=acme&size=10", headers={"Authorization": "Bearer k"})
|
||||
|
||||
sql = query_raw.call_args.args[0]
|
||||
assert "LIMIT $4" in sql
|
||||
assert "LIMIT $5 OFFSET $6" in sql
|
||||
assert query_raw.call_args.args[3] == "%acme%"
|
||||
assert query_raw.call_args.args[5:] == (11, 0)
|
||||
|
||||
|
||||
def test_customer_aliases_caps_page_size(mock_prisma_client, mock_user_api_key_auth):
|
||||
_mock_alias_rows(mock_prisma_client, [])
|
||||
|
||||
response = client.get(f"/customer/aliases?{WINDOW}&size=100000", headers={"Authorization": "Bearer k"})
|
||||
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"role",
|
||||
[
|
||||
LitellmUserRoles.PROXY_ADMIN,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||||
LitellmUserRoles.INTERNAL_USER,
|
||||
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
|
||||
],
|
||||
)
|
||||
def test_customer_aliases_is_reachable_by_every_role_that_can_open_the_logs_page(role):
|
||||
"""Route-level auth gate, which the dependency_overrides in the other tests bypass.
|
||||
|
||||
Handler-side team scoping is dead code if RouteChecks rejects the role first,
|
||||
so pin that /customer/aliases travels in the same access tier as /spend/logs/ui.
|
||||
"""
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
|
||||
for allowed in (
|
||||
LiteLLMRoutes.internal_user_routes.value,
|
||||
LiteLLMRoutes.internal_user_view_only_routes.value,
|
||||
):
|
||||
assert ("/spend/logs/ui" in allowed) == ("/customer/aliases" in allowed)
|
||||
|
||||
if role in (LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY):
|
||||
allowed_routes = (
|
||||
LiteLLMRoutes.internal_user_routes.value
|
||||
if role == LitellmUserRoles.INTERNAL_USER
|
||||
else LiteLLMRoutes.internal_user_view_only_routes.value
|
||||
)
|
||||
assert RouteChecks.check_route_access(route="/customer/aliases", allowed_routes=allowed_routes)
|
||||
else:
|
||||
assert "/customer/aliases" in LiteLLMRoutes.admin_viewer_routes.value
|
||||
|
|
|
|||
|
|
@ -28,9 +28,11 @@ from litellm.proxy.proxy_server import (
|
|||
from .conftest import normalize
|
||||
|
||||
|
||||
def _make_request(parent_otel_span=None):
|
||||
def _make_request(parent_otel_span=None, path="/chat/completions"):
|
||||
"""A real Request always carries a url; the validation handler reads its path to
|
||||
decide whether the caller is on a surface with its own error contract."""
|
||||
state = SimpleNamespace(parent_otel_span=parent_otel_span)
|
||||
return SimpleNamespace(state=state)
|
||||
return SimpleNamespace(state=state, url=SimpleNamespace(path=path))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -221,6 +223,40 @@ async def test_otel_request_validation_exception_handler_empty_errors_invalid_pa
|
|||
assert body == {"detail": []}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_otel_request_validation_exception_handler_returns_a_problem_on_the_control_plane():
|
||||
"""`/management/v1` answers validation errors as RFC 9457, so a caller there gets a
|
||||
400 problem document rather than the proxy-wide 422 `{"detail": [...]}` shape."""
|
||||
errors = [{"loc": ["query", "page_size"], "msg": "Input should be less than or equal to 100", "type": "less_than_equal"}]
|
||||
exc = RequestValidationError(errors)
|
||||
request = _make_request(path="/management/v1/spend_logs/end_users")
|
||||
|
||||
response = await otel_request_validation_exception_handler(request=request, exc=exc)
|
||||
body = json.loads(response.body)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.media_type == "application/problem+json"
|
||||
assert body["type"].startswith("urn:")
|
||||
assert body["status"] == 400
|
||||
assert "page_size" in body["detail"]
|
||||
assert "detail" in body and not isinstance(body["detail"], list)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_otel_request_validation_exception_handler_leaves_other_routes_on_422():
|
||||
"""The problem+json branch is scoped by path prefix. A route that merely contains
|
||||
the word management, or sits above the prefix, keeps the shape its callers parse."""
|
||||
exc = RequestValidationError([])
|
||||
|
||||
for path in ("/management", "/v1/management/foo", "/customer/list"):
|
||||
response = await otel_request_validation_exception_handler(
|
||||
request=_make_request(path=path), exc=exc
|
||||
)
|
||||
|
||||
assert response.status_code == 422, path
|
||||
assert json.loads(response.body) == {"detail": []}, path
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# otel_unhandled_exception_handler
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -4916,3 +4916,94 @@ def test_is_prompt_caching_valid_prompt_explicit_min_token_count_overrides_model
|
|||
is_prompt_caching_valid_prompt(model="claude-opus-4-8", messages=PROMPT_CACHE_MESSAGES, min_token_count=8192)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
def test_custom_logger_guards_ignore_subclass_instances(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Regression LIT-4392: the success/failure existence guards used isinstance, so a user
|
||||
subclass of a built-in logger already promoted into the callback lists made the guard
|
||||
report the built-in itself as registered and the configured logger was silently skipped.
|
||||
The exact-class assertions must hold alongside the subclass assertions: the guards still
|
||||
have to dedup a second instance of the same class, only a subclass must stop matching."""
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.utils import (
|
||||
_custom_logger_class_exists_in_failure_callbacks,
|
||||
_custom_logger_class_exists_in_success_callbacks,
|
||||
)
|
||||
|
||||
class BuiltinLogger(CustomLogger):
|
||||
pass
|
||||
|
||||
class UserSubclassLogger(BuiltinLogger):
|
||||
pass
|
||||
|
||||
builtin_instance = BuiltinLogger()
|
||||
|
||||
monkeypatch.setattr(litellm, "success_callback", [UserSubclassLogger()])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [UserSubclassLogger()])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
assert _custom_logger_class_exists_in_success_callbacks(builtin_instance) is False
|
||||
assert _custom_logger_class_exists_in_failure_callbacks(builtin_instance) is False
|
||||
|
||||
monkeypatch.setattr(litellm, "success_callback", [BuiltinLogger()])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [BuiltinLogger()])
|
||||
assert _custom_logger_class_exists_in_success_callbacks(builtin_instance) is True
|
||||
assert _custom_logger_class_exists_in_failure_callbacks(builtin_instance) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_s3_v2_success_callback_registers_alongside_user_subclass(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Regression LIT-4392: with a user S3Logger subclass registered via litellm_settings.callbacks
|
||||
and success_callback ["s3_v2"], the built-in s3_v2 logger was never added and S3 logs were
|
||||
silently dropped while requests kept returning 200."""
|
||||
from litellm.integrations.s3_v2 import S3Logger
|
||||
from litellm.utils import _add_custom_logger_callback_to_specific_event
|
||||
|
||||
class UserS3Logger(S3Logger):
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
pass
|
||||
|
||||
user_logger = UserS3Logger()
|
||||
monkeypatch.setattr(litellm, "success_callback", [user_logger, "s3_v2"])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [user_logger])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
|
||||
_add_custom_logger_callback_to_specific_event("s3_v2", "success")
|
||||
|
||||
assert any(type(cb) is S3Logger for cb in litellm.success_callback)
|
||||
assert any(type(cb) is S3Logger for cb in litellm._async_success_callback)
|
||||
assert "s3_v2" not in litellm.success_callback
|
||||
assert user_logger in litellm.success_callback
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_builtin_string_callback_registers_when_subclass_already_active(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Regression LIT-4392, litellm.callbacks path: the inline dedup in function_setup also
|
||||
matched subclass instances, so a built-in name in litellm.callbacks was dropped whenever a
|
||||
user subclass was already promoted into _async_success_callback."""
|
||||
from litellm.integrations.s3_v2 import S3Logger
|
||||
|
||||
class UserS3Logger(S3Logger):
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
pass
|
||||
|
||||
user_logger = UserS3Logger()
|
||||
monkeypatch.setattr(litellm, "callbacks", ["s3_v2"])
|
||||
monkeypatch.setattr(litellm, "input_callback", [])
|
||||
monkeypatch.setattr(litellm, "success_callback", [user_logger])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [user_logger])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
|
||||
await litellm.acompletion(
|
||||
model="gpt-5.6",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="ok",
|
||||
)
|
||||
|
||||
assert any(type(cb) is S3Logger for cb in litellm._async_success_callback)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
8
ui/litellm-dashboard/package-lock.json
generated
8
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -5529,16 +5529,16 @@
|
|||
}
|
||||
},
|
||||
"node_modules/brace-expansion": {
|
||||
"version": "5.0.7",
|
||||
"resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-5.0.7.tgz",
|
||||
"integrity": "sha512-7oFy703dxfY3/NLxC1fh2SUCQ0H9rmAY+5EpDVfXjUTTs+HEwR2nYaqLv+GWcTsumwxPfiz6CzCNkwXwBUwqCA==",
|
||||
"version": "5.0.8",
|
||||
"resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-5.0.8.tgz",
|
||||
"integrity": "sha512-JZyDyq3D4AUifKTPOB7DELf6XsB3WdPuNxCtob1vFXPsSXhdAiHBWJ/tJ8HAc9aH84BK+5JFZLNkJKx3G9kzQg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"balanced-match": "^4.0.2"
|
||||
},
|
||||
"engines": {
|
||||
"node": "18 || 20 || >=22"
|
||||
"node": "20 || >=22"
|
||||
}
|
||||
},
|
||||
"node_modules/braces": {
|
||||
|
|
|
|||
|
|
@ -90,7 +90,7 @@
|
|||
"overrides": {
|
||||
"prismjs": "1.30.0",
|
||||
"js-yaml": "4.3.0",
|
||||
"brace-expansion": "5.0.7",
|
||||
"brace-expansion": "5.0.8",
|
||||
"glob": "13.0.0",
|
||||
"minimatch": "10.2.4",
|
||||
"ws": "8.21.0",
|
||||
|
|
|
|||
|
|
@ -1,22 +0,0 @@
|
|||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { $api } from "@/lib/http/api";
|
||||
import type { components } from "@/lib/http/schema";
|
||||
|
||||
type EndUserAliasesPage = components["schemas"]["CustomerAliasesResponse"];
|
||||
|
||||
export interface EndUserAliasesWindow {
|
||||
start_date: string;
|
||||
end_date: string;
|
||||
}
|
||||
|
||||
export const useInfiniteEndUserAliases = (window: EndUserAliasesWindow, size: number = 50, search?: string) => {
|
||||
const { accessToken } = useAuthorized();
|
||||
const query = { ...window, size, ...(search !== undefined && search !== "" ? { search } : {}) };
|
||||
const options = {
|
||||
pageParamName: "page",
|
||||
initialPageParam: 1,
|
||||
getNextPageParam: (lastPage: EndUserAliasesPage) => (lastPage.has_more ? lastPage.current_page + 1 : undefined),
|
||||
enabled: Boolean(accessToken),
|
||||
};
|
||||
return $api.useInfiniteQuery("get", "/customer/aliases", { params: { query } }, options);
|
||||
};
|
||||
|
|
@ -0,0 +1,80 @@
|
|||
import { renderHook } from "@testing-library/react";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
|
||||
const useInfiniteQuery = vi.fn();
|
||||
vi.mock("@/lib/http/api", () => ({ $api: { useInfiniteQuery: (...args: unknown[]) => useInfiniteQuery(...args) } }));
|
||||
|
||||
const mockUseAuthorized = vi.fn();
|
||||
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
|
||||
default: () => mockUseAuthorized(),
|
||||
}));
|
||||
|
||||
import { nextPageFromLinks, useInfiniteSpendLogEndUsers } from "./useSpendLogEndUsers";
|
||||
|
||||
const WINDOW = { start_date: "2026-07-23 00:00:00", end_date: "2026-07-24 00:00:00" };
|
||||
|
||||
const page = (next: string | null) => ({
|
||||
data: ["cust-a"],
|
||||
meta: { page: 1, page_size: 50, has_more: next !== null },
|
||||
links: { self: "/management/v1/spend_logs/end_users?page=1", prev: null, next },
|
||||
});
|
||||
|
||||
describe("useInfiniteSpendLogEndUsers", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
mockUseAuthorized.mockReturnValue({ accessToken: "test-token" });
|
||||
});
|
||||
|
||||
it("calls the control plane path", () => {
|
||||
renderHook(() => useInfiniteSpendLogEndUsers(WINDOW, 50));
|
||||
|
||||
expect(useInfiniteQuery.mock.calls[0][1]).toBe("/management/v1/spend_logs/end_users");
|
||||
});
|
||||
|
||||
it("sends the window as filter params and the page size as page_size", () => {
|
||||
renderHook(() => useInfiniteSpendLogEndUsers(WINDOW, 25));
|
||||
|
||||
const query = useInfiniteQuery.mock.calls[0][2].params.query;
|
||||
expect(query).toEqual({
|
||||
"filter[startTime][gte]": "2026-07-23 00:00:00",
|
||||
"filter[startTime][lte]": "2026-07-24 00:00:00",
|
||||
page_size: 25,
|
||||
});
|
||||
expect(query).not.toHaveProperty("start_date");
|
||||
expect(query).not.toHaveProperty("end_date");
|
||||
expect(query).not.toHaveProperty("size");
|
||||
});
|
||||
|
||||
it("sends free text as q, not search", () => {
|
||||
renderHook(() => useInfiniteSpendLogEndUsers(WINDOW, 50, "acme"));
|
||||
|
||||
const query = useInfiniteQuery.mock.calls[0][2].params.query;
|
||||
expect(query.q).toBe("acme");
|
||||
expect(query).not.toHaveProperty("search");
|
||||
});
|
||||
|
||||
it("omits q entirely when the search box is empty", () => {
|
||||
renderHook(() => useInfiniteSpendLogEndUsers(WINDOW, 50, ""));
|
||||
|
||||
expect(useInfiniteQuery.mock.calls[0][2].params.query).not.toHaveProperty("q");
|
||||
});
|
||||
|
||||
it("derives the next page from the server's links.next", () => {
|
||||
renderHook(() => useInfiniteSpendLogEndUsers(WINDOW, 50));
|
||||
|
||||
const { getNextPageParam } = useInfiniteQuery.mock.calls[0][3];
|
||||
expect(getNextPageParam(page("/management/v1/spend_logs/end_users?page_size=50&page=7"))).toBe(7);
|
||||
});
|
||||
});
|
||||
|
||||
describe("nextPageFromLinks", () => {
|
||||
it("reads the page the server pointed at rather than incrementing", () => {
|
||||
/* An endpoint that later switches to cursor pagination changes links.next and
|
||||
nothing else; a client that computed page+1 would silently break. */
|
||||
expect(nextPageFromLinks(page("/management/v1/spend_logs/end_users?page_size=50&page=7"))).toBe(7);
|
||||
});
|
||||
|
||||
it("stops paging when the server omits links.next", () => {
|
||||
expect(nextPageFromLinks(page(null))).toBeUndefined();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,36 @@
|
|||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { $api } from "@/lib/http/api";
|
||||
import type { components } from "@/lib/http/schema";
|
||||
|
||||
type EndUsersPage = components["schemas"]["FacetListResponse"];
|
||||
|
||||
export interface SpendLogsWindow {
|
||||
start_date: string;
|
||||
end_date: string;
|
||||
}
|
||||
|
||||
/** Reads the server's `links.next` instead of computing the next page, so the
|
||||
* endpoint can move to cursor pagination without touching this hook. */
|
||||
export const nextPageFromLinks = (lastPage: EndUsersPage): number | undefined => {
|
||||
const next = lastPage.links.next;
|
||||
if (!next) return undefined;
|
||||
const page = new URLSearchParams(next.slice(next.indexOf("?") + 1)).get("page");
|
||||
return page === null ? undefined : Number(page);
|
||||
};
|
||||
|
||||
export const useInfiniteSpendLogEndUsers = (window: SpendLogsWindow, pageSize: number = 50, q?: string) => {
|
||||
const { accessToken } = useAuthorized();
|
||||
const query = {
|
||||
"filter[startTime][gte]": window.start_date,
|
||||
"filter[startTime][lte]": window.end_date,
|
||||
page_size: pageSize,
|
||||
...(q !== undefined && q !== "" ? { q } : {}),
|
||||
};
|
||||
const options = {
|
||||
pageParamName: "page",
|
||||
initialPageParam: 1,
|
||||
getNextPageParam: nextPageFromLinks,
|
||||
enabled: Boolean(accessToken),
|
||||
};
|
||||
return $api.useInfiniteQuery("get", "/management/v1/spend_logs/end_users", { params: { query } }, options);
|
||||
};
|
||||
|
|
@ -14,11 +14,11 @@ vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({
|
|||
useInfiniteModelInfo: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/customers/useEndUserAliases", () => ({
|
||||
useInfiniteEndUserAliases: vi.fn(),
|
||||
vi.mock("@/app/(dashboard)/hooks/spendLogs/useSpendLogEndUsers", () => ({
|
||||
useInfiniteSpendLogEndUsers: vi.fn(),
|
||||
}));
|
||||
|
||||
import { useInfiniteEndUserAliases } from "@/app/(dashboard)/hooks/customers/useEndUserAliases";
|
||||
import { useInfiniteSpendLogEndUsers } from "@/app/(dashboard)/hooks/spendLogs/useSpendLogEndUsers";
|
||||
import { useInfiniteKeyAliases } from "@/app/(dashboard)/hooks/keys/useKeyAliases";
|
||||
import { useInfiniteModelInfo } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
|
||||
|
|
@ -50,8 +50,8 @@ describe("RequestLogsFilters", () => {
|
|||
vi.mocked(useInfiniteModelInfo).mockReturnValue(
|
||||
emptyInfiniteQuery as unknown as ReturnType<typeof useInfiniteModelInfo>,
|
||||
);
|
||||
vi.mocked(useInfiniteEndUserAliases).mockReturnValue(
|
||||
emptyInfiniteQuery as unknown as ReturnType<typeof useInfiniteEndUserAliases>,
|
||||
vi.mocked(useInfiniteSpendLogEndUsers).mockReturnValue(
|
||||
emptyInfiniteQuery as unknown as ReturnType<typeof useInfiniteSpendLogEndUsers>,
|
||||
);
|
||||
});
|
||||
|
||||
|
|
@ -98,8 +98,8 @@ describe("RequestLogsFilters", () => {
|
|||
it("asks the server for a bounded page of end users scoped to the visible time window", async () => {
|
||||
renderFilters();
|
||||
|
||||
await waitFor(() => expect(useInfiniteEndUserAliases).toHaveBeenCalled());
|
||||
expect(useInfiniteEndUserAliases).toHaveBeenCalledWith(LOGS_WINDOW, 50, undefined);
|
||||
await waitFor(() => expect(useInfiniteSpendLogEndUsers).toHaveBeenCalled());
|
||||
expect(useInfiniteSpendLogEndUsers).toHaveBeenCalledWith(LOGS_WINDOW, 50, undefined);
|
||||
});
|
||||
|
||||
it("pushes the End User query to the server rather than filtering a preloaded list", async () => {
|
||||
|
|
@ -110,14 +110,23 @@ describe("RequestLogsFilters", () => {
|
|||
await user.click(input);
|
||||
await user.type(input, "acme");
|
||||
|
||||
await waitFor(() => expect(useInfiniteEndUserAliases).toHaveBeenCalledWith(LOGS_WINDOW, 50, "acme"));
|
||||
await waitFor(() => expect(useInfiniteSpendLogEndUsers).toHaveBeenCalledWith(LOGS_WINDOW, 50, "acme"));
|
||||
});
|
||||
|
||||
it("renders only the end users the current page returned", async () => {
|
||||
vi.mocked(useInfiniteEndUserAliases).mockReturnValue({
|
||||
vi.mocked(useInfiniteSpendLogEndUsers).mockReturnValue({
|
||||
...emptyInfiniteQuery,
|
||||
data: { pages: [{ aliases: ["cust-a", "cust-b"], current_page: 1, size: 50, has_more: true }], pageParams: [1] },
|
||||
} as unknown as ReturnType<typeof useInfiniteEndUserAliases>);
|
||||
data: {
|
||||
pages: [
|
||||
{
|
||||
data: ["cust-a", "cust-b"],
|
||||
meta: { page: 1, page_size: 50, has_more: true },
|
||||
links: { self: "", next: "?page=2" },
|
||||
},
|
||||
],
|
||||
pageParams: [1],
|
||||
},
|
||||
} as unknown as ReturnType<typeof useInfiniteSpendLogEndUsers>);
|
||||
const user = userEvent.setup();
|
||||
renderFilters();
|
||||
|
||||
|
|
@ -129,12 +138,17 @@ describe("RequestLogsFilters", () => {
|
|||
|
||||
it("loads the next page when the End User list is scrolled near the end", async () => {
|
||||
const fetchNextPage = vi.fn();
|
||||
vi.mocked(useInfiniteEndUserAliases).mockReturnValue({
|
||||
vi.mocked(useInfiniteSpendLogEndUsers).mockReturnValue({
|
||||
...emptyInfiniteQuery,
|
||||
fetchNextPage,
|
||||
hasNextPage: true,
|
||||
data: { pages: [{ aliases: ["cust-a"], current_page: 1, size: 50, has_more: true }], pageParams: [1] },
|
||||
} as unknown as ReturnType<typeof useInfiniteEndUserAliases>);
|
||||
data: {
|
||||
pages: [
|
||||
{ data: ["cust-a"], meta: { page: 1, page_size: 50, has_more: true }, links: { self: "", next: "?page=2" } },
|
||||
],
|
||||
pageParams: [1],
|
||||
},
|
||||
} as unknown as ReturnType<typeof useInfiniteSpendLogEndUsers>);
|
||||
const user = userEvent.setup();
|
||||
renderFilters();
|
||||
|
||||
|
|
@ -152,6 +166,6 @@ describe("RequestLogsFilters", () => {
|
|||
const otherWindow = { start_date: "2026-01-01 00:00:00", end_date: "2026-01-02 00:00:00" };
|
||||
renderWithProviders(<RequestLogsFilters get={() => undefined} set={vi.fn()} teams={[]} logsWindow={otherWindow} />);
|
||||
|
||||
await waitFor(() => expect(useInfiniteEndUserAliases).toHaveBeenCalledWith(otherWindow, 50, undefined));
|
||||
await waitFor(() => expect(useInfiniteSpendLogEndUsers).toHaveBeenCalledWith(otherWindow, 50, undefined));
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
import { useMemo, useState } from "react";
|
||||
|
||||
import { useInfiniteEndUserAliases } from "@/app/(dashboard)/hooks/customers/useEndUserAliases";
|
||||
import { useInfiniteSpendLogEndUsers } from "@/app/(dashboard)/hooks/spendLogs/useSpendLogEndUsers";
|
||||
import { useInfiniteKeyAliases } from "@/app/(dashboard)/hooks/keys/useKeyAliases";
|
||||
import { useInfiniteModelInfo } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
import { DataTableFilterField } from "@/components/shared/DataTable";
|
||||
|
|
@ -154,7 +154,7 @@ function EndUserFilterField({
|
|||
logsWindow: LogsWindow;
|
||||
}) {
|
||||
const [search, setSearch] = useState("");
|
||||
const { data, fetchNextPage, hasNextPage, isFetchingNextPage, isLoading } = useInfiniteEndUserAliases(
|
||||
const { data, fetchNextPage, hasNextPage, isFetchingNextPage, isLoading } = useInfiniteSpendLogEndUsers(
|
||||
logsWindow,
|
||||
PAGE_SIZE,
|
||||
emptyToUndefined(search),
|
||||
|
|
@ -163,10 +163,10 @@ function EndUserFilterField({
|
|||
const options = useMemo<SearchSelectOption[]>(() => {
|
||||
const seen = new Set<string>();
|
||||
return (data?.pages ?? []).flatMap((page) =>
|
||||
page.aliases.flatMap((alias) => {
|
||||
if (!alias || seen.has(alias)) return [];
|
||||
seen.add(alias);
|
||||
return [{ label: alias, value: alias }];
|
||||
page.data.flatMap((endUser) => {
|
||||
if (!endUser || seen.has(endUser)) return [];
|
||||
seen.add(endUser);
|
||||
return [{ label: endUser, value: endUser }];
|
||||
}),
|
||||
);
|
||||
}, [data]);
|
||||
|
|
|
|||
205
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
205
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -2772,40 +2772,6 @@ export interface paths {
|
|||
patch: operations["cursor_proxy_route_cursor__endpoint__patch"];
|
||||
trace?: never;
|
||||
};
|
||||
"/customer/aliases": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
/**
|
||||
* List Customer Aliases
|
||||
* @description List the end users seen in spend logs over a time window, for UI filter dropdowns.
|
||||
*
|
||||
* Scoped like `/spend/logs/ui`: a proxy admin sees every end user in the window,
|
||||
* anyone else sees only end users from their own requests or from teams they
|
||||
* administer (or hold the `/spend/logs` permission on).
|
||||
*
|
||||
* Reads spend logs rather than LiteLLM_EndUserTable because only spend logs carry
|
||||
* the team attribution this scoping needs. The window is required and the inner
|
||||
* scan is capped at SPEND_LOGS_FILTER_SCAN_CAP rows, so the query
|
||||
* cannot degrade into a full-table scan the way `/global/all_end_users` does.
|
||||
*
|
||||
* Example curl:
|
||||
* ```
|
||||
* curl --location 'http://0.0.0.0:4000/customer/aliases?start_date=2026-07-23%2000:00:00&end_date=2026-07-24%2000:00:00&size=50&search=acme' --header 'Authorization: Bearer sk-1234'
|
||||
* ```
|
||||
*/
|
||||
get: operations["list_customer_aliases_customer_aliases_get"];
|
||||
put?: never;
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/customer/block": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -7219,6 +7185,40 @@ export interface paths {
|
|||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/management/v1/spend_logs/end_users": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
/**
|
||||
* List Spend Log End Users
|
||||
* @description The distinct end users appearing in spend logs over a time window, for the logs
|
||||
* page filter dropdown.
|
||||
*
|
||||
* Scoped like `/spend/logs/ui`: a proxy admin sees every end user in the window,
|
||||
* anyone else sees only end users from their own requests or from teams they
|
||||
* administer (or hold the `/spend/logs` permission on).
|
||||
*
|
||||
* The window is required and the inner scan is capped at SPEND_LOGS_FACET_SCAN_CAP
|
||||
* rows, so the query cannot degrade into a full-table scan the way
|
||||
* `/global/all_end_users` does.
|
||||
*
|
||||
* Example curl:
|
||||
* ```
|
||||
* curl --location --globoff 'http://0.0.0.0:4000/management/v1/spend_logs/end_users?filter[startTime][gte]=2026-07-23T00:00:00Z&filter[startTime][lte]=2026-07-24T00:00:00Z&page_size=50&q=acme' --header 'Authorization: Bearer sk-1234'
|
||||
* ```
|
||||
*/
|
||||
get: operations["list_spend_log_end_users_management_v1_spend_logs_end_users_get"];
|
||||
put?: never;
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/mcp-rest/test/connection": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -23329,29 +23329,6 @@ export interface components {
|
|||
[key: string]: unknown;
|
||||
};
|
||||
};
|
||||
/**
|
||||
* CustomerAliasesResponse
|
||||
* @description Paginated, id-only customer listing used by UI filter dropdowns.
|
||||
*
|
||||
* Deliberately excludes budget/object-permission relations so a proxy with a
|
||||
* large LiteLLM_EndUserTable can back a search-as-you-type control without
|
||||
* materializing every row (see /customer/list for the full objects).
|
||||
*
|
||||
* Reports ``has_more`` rather than a total count on purpose: a total requires
|
||||
* COUNT(*) over the whole match set on every keystroke, which is the exact
|
||||
* cost this endpoint exists to avoid. Fetching one row beyond the page is
|
||||
* enough to drive an infinite-scroll dropdown.
|
||||
*/
|
||||
CustomerAliasesResponse: {
|
||||
/** Aliases */
|
||||
aliases: string[];
|
||||
/** Current Page */
|
||||
current_page: number;
|
||||
/** Has More */
|
||||
has_more: boolean;
|
||||
/** Size */
|
||||
size: number;
|
||||
};
|
||||
/**
|
||||
* CustomerResponse
|
||||
* @description Customer object returned by the /customer read+write endpoints.
|
||||
|
|
@ -23893,6 +23870,16 @@ export interface components {
|
|||
/** Updated At */
|
||||
updated_at?: number | null;
|
||||
};
|
||||
/**
|
||||
* FacetListResponse
|
||||
* @description The distinct values one column takes over a filtered query. `data` holds bare values, not entity rows.
|
||||
*/
|
||||
FacetListResponse: {
|
||||
/** Data */
|
||||
data: string[];
|
||||
links: components["schemas"]["PageLinks"];
|
||||
meta: components["schemas"]["PageMeta"];
|
||||
};
|
||||
/**
|
||||
* FailedKeyUpdate
|
||||
* @description Failed key update with reason
|
||||
|
|
@ -28852,6 +28839,30 @@ export interface components {
|
|||
/** Tpm Limit */
|
||||
tpm_limit?: number | null;
|
||||
};
|
||||
/**
|
||||
* PageLinks
|
||||
* @description Hypermedia for a paginated list. No `first`/`last`: without a total count the last page is unknown.
|
||||
*/
|
||||
PageLinks: {
|
||||
/** Next */
|
||||
next?: string | null;
|
||||
/** Prev */
|
||||
prev?: string | null;
|
||||
/** Self */
|
||||
self: string;
|
||||
};
|
||||
/**
|
||||
* PageMeta
|
||||
* @description `has_more` rather than `total_count`, which would need a COUNT(*) over the whole match set per keystroke.
|
||||
*/
|
||||
PageMeta: {
|
||||
/** Has More */
|
||||
has_more: boolean;
|
||||
/** Page */
|
||||
page: number;
|
||||
/** Page Size */
|
||||
page_size: number;
|
||||
};
|
||||
/**
|
||||
* PaginatedAuditLogResponse
|
||||
* @description Response model for paginated audit logs
|
||||
|
|
@ -38457,46 +38468,6 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
list_customer_aliases_customer_aliases_get: {
|
||||
parameters: {
|
||||
query: {
|
||||
/** @description Window start, 'YYYY-MM-DD HH:MM:SS' (UTC) */
|
||||
start_date: string;
|
||||
/** @description Window end, 'YYYY-MM-DD HH:MM:SS' (UTC) */
|
||||
end_date: string;
|
||||
/** @description Page number */
|
||||
page?: number;
|
||||
/** @description Page size */
|
||||
size?: number;
|
||||
/** @description Case-insensitive partial match on the customer id */
|
||||
search?: string | null;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["CustomerAliasesResponse"];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
block_user_customer_block_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -43424,6 +43395,46 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
list_spend_log_end_users_management_v1_spend_logs_end_users_get: {
|
||||
parameters: {
|
||||
query: {
|
||||
/** @description Window start (UTC when no offset is given) */
|
||||
"filter[startTime][gte]": string;
|
||||
/** @description Window end (UTC when no offset is given) */
|
||||
"filter[startTime][lte]": string;
|
||||
/** @description Case-insensitive partial match on the end user id */
|
||||
q?: string | null;
|
||||
/** @description Page number */
|
||||
page?: number;
|
||||
/** @description Page size */
|
||||
page_size?: number;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["FacetListResponse"];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
test_connection_mcp_rest_test_connection_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
|
|||
8
uv.lock
generated
8
uv.lock
generated
|
|
@ -10,7 +10,7 @@ resolution-markers = [
|
|||
]
|
||||
|
||||
[options]
|
||||
exclude-newer = "2026-07-22T23:28:30.575519Z"
|
||||
exclude-newer = "2026-07-24T16:43:28.506903Z"
|
||||
exclude-newer-span = "P3D"
|
||||
|
||||
[manifest]
|
||||
|
|
@ -2378,14 +2378,14 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "gitpython"
|
||||
version = "3.1.54"
|
||||
version = "3.1.55"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "gitdb" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/5e/d5/3da0b92033887033f4c27f2dd109a303c4ca62813c7b3bb2511edb4777de/gitpython-3.1.54.tar.gz", hash = "sha256:53f2085e24a2cda300eed7c3fc5f1559ae289634b725e98acaf4791940247aa0", size = 225076, upload-time = "2026-07-22T04:08:51.403Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/b2/ab/ba0d29f2fa2277ed6256b2ac09003494045355f3a10bf32f351761287870/gitpython-3.1.55.tar.gz", hash = "sha256:781e3b1624dad81b24e9524bf0297b69786a0706db2cbceec1e2b05c38e5152f", size = 225071, upload-time = "2026-07-23T02:52:43.246Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/d1/b9/876f442a28df5c068ca69b0122d5c35e65fd2d2fa9992ea5cb5944ea00a6/gitpython-3.1.54-py3-none-any.whl", hash = "sha256:b90d7b3d9bc0238681d24369130826f0dcdb0ceaa45db67cf1d4ffa4c302dedf", size = 216575, upload-time = "2026-07-22T04:08:50.05Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/20/6a/d3b8208d2f8aac66abe8ccc1c23fa2c89464ec42cc71a601e95d05902428/gitpython-3.1.55-py3-none-any.whl", hash = "sha256:7c9ec1e69c158c081632ab35c41471e302c96db2ae42165036a5d2403378812e", size = 216590, upload-time = "2026-07-23T02:52:41.932Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue