Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_proxy_test_master_key_leak

# Conflicts:
#	tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py
This commit is contained in:
Yuneng Jiang 2026-04-23 18:20:01 -07:00
commit e68c60a66e
No known key found for this signature in database
18 changed files with 963 additions and 172 deletions

View file

@ -98,6 +98,19 @@ commands:
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
start_redis:
description: "Start a redis container on port 6379 and wait until it accepts connections. Use this to isolate a job from the shared remote Redis so concurrent CI pipelines don't contend for pod locks or buffer keys."
steps:
- run:
name: Start Redis
command: |
docker run -d \
--name redis-cache \
-p 6379:6379 \
redis:7-alpine@sha256:7aec734b2bb298a1d769fd8729f13b8514a41bf90fcdd1f38ec52267fbaa8ee6
- wait_for_service:
url: tcp://localhost:6379
timeout: "60"
setup_litellm_enterprise_pip:
steps:
- run:
@ -563,39 +576,6 @@ jobs:
paths:
- realtime_translation_coverage.xml
- realtime_translation_coverage
mcp_testing:
docker:
- *python312_image
working_directory: ~/project
steps:
- checkout
- setup_google_dns
- install_uv
- run:
name: Install Dependencies
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
# Run pytest and generate JUnit XML report
- run:
name: Run tests
command: |
uv run --no-sync python -m pytest -vv tests/mcp_tests --cov=litellm --cov-report=xml -x -s -v --junitxml=test-results/junit.xml --durations=5 -n 2
no_output_timeout: 15m
- run:
name: Rename the coverage files
command: |
mv coverage.xml mcp_coverage.xml
mv .coverage mcp_coverage
# Store test results
- store_test_results:
path: test-results
- persist_to_workspace:
root: .
paths:
- mcp_coverage.xml
- mcp_coverage
agent_testing:
docker:
- *python312_image
@ -794,39 +774,6 @@ jobs:
paths:
- search_coverage.xml
- search_coverage
# Split litellm_mapped_tests into parallel jobs
litellm_mapped_tests_proxy_part1:
docker:
- *python312_image
working_directory: ~/project
resource_class: large
steps:
- setup_litellm_test_deps
- run:
name: Run proxy tests part 1 (high-volume directories)
command: |
uv run --no-sync python -m prisma generate
export PYTHONUNBUFFERED=1
uv run --no-sync python -m pytest tests/test_litellm/proxy/guardrails tests/test_litellm/proxy/management_endpoints tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/client tests/test_litellm/proxy/auth --junitxml=test-results/junit-proxy-part1.xml --durations=10 -n 4 --maxfail=5 --timeout=60 -vv --log-cli-level=WARNING -r A
no_output_timeout: 15m
- store_test_results:
path: test-results
litellm_mapped_tests_proxy_part2:
docker:
- *python312_image
working_directory: ~/project
resource_class: large
steps:
- setup_litellm_test_deps
- run:
name: Run proxy tests part 2 (all other tests)
command: |
uv run --no-sync python -m prisma generate
export PYTHONUNBUFFERED=1
uv run --no-sync python -m pytest tests/test_litellm/proxy --ignore=tests/test_litellm/proxy/guardrails --ignore=tests/test_litellm/proxy/management_endpoints --ignore=tests/test_litellm/proxy/_experimental --ignore=tests/test_litellm/proxy/client --ignore=tests/test_litellm/proxy/auth --junitxml=test-results/junit-proxy-part2.xml --durations=10 -n 4 --maxfail=5 --timeout=120 -vv --log-cli-level=WARNING -r A
no_output_timeout: 15m
- store_test_results:
path: test-results
litellm_mapped_enterprise_tests:
docker:
- *python312_image
@ -1591,6 +1538,7 @@ jobs:
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
- start_postgres
- start_redis
- attach_workspace:
at: ~/project
- run:
@ -1600,15 +1548,18 @@ jobs:
docker images | grep litellm-docker-database
- run:
name: Run Docker container
# intentionally give bad redis credentials here
# the OTEL test - should get this as a trace
# Point the proxy at the job-local Redis (start_redis) instead of the
# shared remote Redis. The Redis transaction buffer uses a single
# global pod-lock key (cronjob_lock:db_spend_update_job) and a single
# global buffer list (litellm_spend_update_buffer); sharing those
# across concurrent CI pipelines causes spend flushes to stall or
# land in the wrong DB, which is what makes this test flaky.
command: |
docker run -d \
-p 4000:4000 \
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
-e REDIS_HOST=$REDIS_HOST \
-e REDIS_PASSWORD=$REDIS_PASSWORD \
-e REDIS_PORT=$REDIS_PORT \
-e REDIS_HOST=host.docker.internal \
-e REDIS_PORT=6379 \
-e LITELLM_MASTER_KEY="sk-1234" \
-e OPENAI_API_KEY=$OPENAI_API_KEY \
-e LITELLM_LICENSE=$LITELLM_LICENSE \
@ -1638,12 +1589,14 @@ jobs:
command: |
uv run --no-sync python -m pytest -vv tests/spend_tracking_tests -x --junitxml=test-results/junit.xml --durations=5
no_output_timeout: 15m
# Clean up first container
- run:
name: Stop and remove first container
when: always
command: |
docker stop my-app
docker rm my-app
docker stop redis-cache
docker rm redis-cache
proxy_multi_instance_tests:
machine:
@ -2072,7 +2025,7 @@ jobs:
- run:
name: Combine Coverage
command: |
uv tool run --from 'coverage[toml]==7.10.6' coverage combine realtime_translation_coverage ocr_coverage search_coverage mcp_coverage logging_coverage audio_coverage local_testing_part1_coverage local_testing_part2_coverage pass_through_unit_tests_coverage batches_coverage guardrails_coverage redis_caching_coverage
uv tool run --from 'coverage[toml]==7.10.6' coverage combine realtime_translation_coverage ocr_coverage search_coverage logging_coverage audio_coverage local_testing_part1_coverage local_testing_part2_coverage pass_through_unit_tests_coverage batches_coverage guardrails_coverage redis_caching_coverage
uv tool run --from 'coverage[toml]==7.10.6' coverage xml
- codecov/upload:
file: ./coverage.xml
@ -2407,8 +2360,6 @@ workflows:
filters: *main_branches
- realtime_translation_testing:
filters: *main_branches
- mcp_testing:
filters: *main_branches
- agent_testing:
filters: *main_branches
- guardrails_testing:
@ -2423,10 +2374,6 @@ workflows:
filters: *main_branches
- litellm_mapped_enterprise_tests:
filters: *main_branches
- litellm_mapped_tests_proxy_part1:
filters: *main_branches
- litellm_mapped_tests_proxy_part2:
filters: *main_branches
- batches_testing:
filters: *main_branches
- litellm_utils_testing:
@ -2444,14 +2391,11 @@ workflows:
- upload-coverage:
requires:
- realtime_translation_testing
- mcp_testing
- agent_testing
- google_generate_content_endpoint_testing
- guardrails_testing
- ocr_testing
- search_testing
- litellm_mapped_tests_proxy_part1
- litellm_mapped_tests_proxy_part2
- litellm_mapped_enterprise_tests
- batches_testing
- litellm_utils_testing

View file

@ -32,41 +32,39 @@ on:
required: false
type: boolean
default: false
dist:
description: "pytest-xdist distribution mode (loadscope|load|worksteal|loadfile|no)"
required: false
type: string
default: "loadscope"
artifact-name:
description: "Unique name for the coverage artifact (must be unique per run)"
required: false
type: string
default: "run"
secrets:
DATABASE_URL:
required: false
POSTGRES_USER:
required: false
POSTGRES_PASSWORD:
required: false
permissions:
contents: read
# The postgres service container below is spawned per-job on localhost and
# destroyed with the job. Nothing outside the runner can reach it. The
# user/password/database here are not secrets — they're bootstrap values
# for a throwaway container — so we hardcode them instead of attaching
# every matrix shard to a GHA environment just to read three "secrets"
# (which also produces a "temporarily deployed to …" notification on the
# PR timeline per shard per push).
jobs:
run:
name: Run tests
runs-on: ubuntu-latest
timeout-minutes: ${{ inputs.timeout-minutes }}
# Environment is derived from the enable-* flags, not caller-controllable.
# This prevents callers from passing arbitrary environment names to bypass secret scoping.
environment: >-
${{
inputs.enable-postgres && 'integration-postgres' ||
''
}}
services:
postgres:
image: postgres@sha256:705a5d5b5836f3fcba0d02c4d281e6a7dd9ed2dd4078640f08a1e1e9896e097d # postgres:14
env:
POSTGRES_USER: ${{ secrets.POSTGRES_USER }}
POSTGRES_PASSWORD: ${{ secrets.POSTGRES_PASSWORD }}
POSTGRES_USER: litellm
POSTGRES_PASSWORD: litellm
POSTGRES_DB: litellm_test
ports:
- 5432:5432
@ -114,7 +112,7 @@ jobs:
- name: Run Prisma migrations
if: ${{ inputs.enable-postgres }}
env:
DATABASE_URL: ${{ secrets.DATABASE_URL }}
DATABASE_URL: "postgresql://litellm:litellm@localhost:5432/litellm_test"
run: |
uv run --no-sync prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
@ -124,7 +122,8 @@ jobs:
MAX_FAILURES: ${{ inputs.max-failures }}
WORKERS: ${{ inputs.workers }}
RERUNS: ${{ inputs.reruns }}
DATABASE_URL: ${{ inputs.enable-postgres && secrets.DATABASE_URL || '' }}
DIST: ${{ inputs.dist }}
DATABASE_URL: ${{ inputs.enable-postgres && 'postgresql://litellm:litellm@localhost:5432/litellm_test' || '' }}
run: |
if [ "${WORKERS}" = "0" ]; then
uv run --no-sync pytest ${TEST_PATH:?} \
@ -143,7 +142,7 @@ jobs:
-n "${WORKERS}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--dist=loadscope \
--dist="${DIST}" \
--durations=20 \
--cov=litellm \
--cov-report=xml:coverage.xml \

View file

@ -12,8 +12,74 @@ concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true
# Semantic matrix: each shard groups tests by concern (auth, server, logging, …)
# rather than alphabetical letter ranges. Adding a new test file means adding it
# to whichever group it belongs to, not reshuffling slices.
#
# Design targets:
# * Every shard runs in <= 7 minutes of wall-clock on the default runner.
# Most of a shard's time is pytest plugin load + xdist worker imports +
# pytest-cov instrumentation, not the tests themselves. Keeping per-shard
# work low and matching worker count to runner cores is what controls it.
# * workers: 4 matches the 4-core ubuntu-latest runner. -n 8 on 4 cores
# oversubscribes 2x and workers fight for CPU during their cold-start
# imports (measured ~441% CPU for -n 8 locally, i.e. ~55% effective).
# * test_key_generate_prisma.py stays serial (workers=0) — it has event-loop
# conflicts with the logging worker when run in parallel.
# * test_proxy_utils.py runs as a single shard with --dist=worksteal so
# xdist balances its 188 parametrized cases across workers instead of
# pinning the whole file to one worker (the default --dist=loadscope
# behavior for single-file targets).
# * test_db_schema_migration.py is isolated because one test in it
# (test_aaaasschema_migration_check) takes ~170s — by itself it
# determines the shard's wall-clock floor.
jobs:
# Fast guard — fails the workflow if a test_*.py file under
# tests/proxy_unit_tests/ is not referenced by any matrix entry below.
# The semantic-shard design (no catch-all "remaining" bucket) relies on
# every test file being explicitly assigned; this guard prevents a new
# file from silently dropping out of CI.
assert-shard-coverage:
runs-on: ubuntu-latest
timeout-minutes: 2
permissions:
contents: read
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Assert every test_*.py is in a matrix shard
run: |
python3 - <<'PY'
import pathlib, sys, yaml
wf = yaml.safe_load(open(".github/workflows/test-unit-proxy-db.yml"))
matrix = wf["jobs"]["proxy-db"]["strategy"]["matrix"]["include"]
referenced = set()
for entry in matrix:
for token in entry["test-path"].split():
if token.startswith("tests/proxy_unit_tests/"):
referenced.add(pathlib.PurePosixPath(token).name)
actual = {p.name for p in pathlib.Path("tests/proxy_unit_tests").iterdir()
if p.name.startswith("test_") and (p.suffix == ".py" or p.is_dir())
and p.name != "test_configs"}
orphans = sorted(actual - referenced)
if orphans:
print("ERROR: the following files/dirs under tests/proxy_unit_tests/")
print(" are not assigned to any shard in test-unit-proxy-db.yml:")
for o in orphans:
print(f" - {o}")
print()
print("Add each to whichever semantic shard it belongs to.")
sys.exit(1)
print(f"OK: all {len(actual)} files assigned to a shard.")
PY
proxy-db:
needs: assert-shard-coverage
# Display only the semantic shard name in the checks UI instead of GHA's
# default "proxy-db (key-generation, tests/proxy_unit_tests/…, 0, loadscope, 20)"
# which includes every matrix field and gets truncated past the test-path.
name: ${{ matrix.test-group }}
permissions:
contents: read
id-token: write
@ -22,26 +88,146 @@ jobs:
fail-fast: false
matrix:
include:
# Key generation tests must NOT run in parallel (event loop conflicts with logging worker)
# Must run serially — event-loop conflict with the logging worker.
- test-group: key-generation
test-path: "tests/proxy_unit_tests/test_key_generate_prisma.py"
workers: 0
timeout: 30
- test-group: auth-checks
test-path: "tests/proxy_unit_tests/test_auth_checks.py tests/proxy_unit_tests/test_user_api_key_auth.py"
workers: 8
dist: loadscope
timeout: 20
# test_proxy_utils.py is large (168+ parametrized tests) — run it on its
# own matrix so --dist=loadscope doesn't pin all of it to a single xdist
# worker and push the "remaining" group past the job timeout.
# ---- auth: split into 2 shards ----
- test-group: auth-checks
test-path: >-
tests/proxy_unit_tests/test_auth_checks.py
tests/proxy_unit_tests/test_user_api_key_auth.py
workers: 4
dist: loadscope
timeout: 15
- test-group: jwt-and-keys
test-path: >-
tests/proxy_unit_tests/test_jwt.py
tests/proxy_unit_tests/test_jwt_key_mapping.py
tests/proxy_unit_tests/test_proxy_custom_auth.py
tests/proxy_unit_tests/test_key_generate_dynamodb.py
tests/proxy_unit_tests/test_deployed_proxy_keygen.py
workers: 4
dist: loadscope
timeout: 15
# ---- test_proxy_utils.py, single shard, worksteal distribution ----
- test-group: proxy-utils
test-path: "tests/proxy_unit_tests/test_proxy_utils.py"
workers: 8
timeout: 20
- test-group: remaining
test-path: "tests/proxy_unit_tests --ignore=tests/proxy_unit_tests/test_key_generate_prisma.py --ignore=tests/proxy_unit_tests/test_auth_checks.py --ignore=tests/proxy_unit_tests/test_user_api_key_auth.py --ignore=tests/proxy_unit_tests/test_proxy_utils.py"
workers: 8
timeout: 30
workers: 4
dist: worksteal
timeout: 15
# ---- proxy server: split into 2 shards ----
- test-group: proxy-server-core
test-path: >-
tests/proxy_unit_tests/test_proxy_server.py
tests/proxy_unit_tests/test_proxy_server_keys.py
tests/proxy_unit_tests/test_proxy_server_caching.py
tests/proxy_unit_tests/test_proxy_server_langfuse.py
tests/proxy_unit_tests/test_proxy_server_spend.py
tests/proxy_unit_tests/test_aproxy_startup.py
workers: 4
dist: loadscope
timeout: 15
- test-group: proxy-runtime
test-path: >-
tests/proxy_unit_tests/test_proxy_config_unit_test.py
tests/proxy_unit_tests/test_proxy_routes.py
tests/proxy_unit_tests/test_proxy_gunicorn.py
tests/proxy_unit_tests/test_server_root_path.py
tests/proxy_unit_tests/test_proxy_pass_user_config.py
tests/proxy_unit_tests/test_proxy_token_counter.py
workers: 4
dist: loadscope
timeout: 15
# ---- logging: split into 2 shards ----
- test-group: custom-logging
test-path: >-
tests/proxy_unit_tests/test_custom_callback_input.py
tests/proxy_unit_tests/test_custom_logger_s3_gcs.py
tests/proxy_unit_tests/test_proxy_custom_logger.py
workers: 4
dist: loadscope
timeout: 15
- test-group: logging-misc
test-path: >-
tests/proxy_unit_tests/test_proxy_reject_logging.py
tests/proxy_unit_tests/test_audit_logs_proxy.py
tests/proxy_unit_tests/test_search_api_logging.py
workers: 4
dist: loadscope
timeout: 15
# ---- db-and-spend: isolate the 170s schema-migration test ----
# test_db_schema_migration.py has exactly one test, and that test
# is mostly waiting on `prisma migrate deploy` / `prisma migrate
# diff` subprocesses (~170s). It does no CPU-bound Python work
# inside the test. Running with workers=0 (serial, no xdist)
# skips the 4-worker cold-start cost we'd otherwise pay for a
# single test, saving ~4 minutes of wall-clock.
- test-group: schema-migration
test-path: "tests/proxy_unit_tests/test_db_schema_migration.py"
workers: 0
dist: loadscope
timeout: 15
- test-group: db-and-spend
test-path: >-
tests/proxy_unit_tests/test_prisma_client_backoff_retry.py
tests/proxy_unit_tests/test_db_schema_changes.py
tests/proxy_unit_tests/test_e2e_pod_lock_manager.py
tests/proxy_unit_tests/test_skills_db.py
tests/proxy_unit_tests/test_update_daily_tag_spend.py
tests/proxy_unit_tests/test_update_spend.py
tests/proxy_unit_tests/test_project_endpoints_prisma.py
tests/proxy_unit_tests/test_proxy_encrypt_decrypt.py
workers: 4
dist: loadscope
timeout: 15
# ---- guardrails + budget + hooks: split into 2 ----
- test-group: guardrails-hooks
test-path: >-
tests/proxy_unit_tests/test_proxy_setting_guardrails.py
tests/proxy_unit_tests/test_banned_keyword_list.py
tests/proxy_unit_tests/test_unit_test_proxy_hooks.py
workers: 4
dist: loadscope
timeout: 15
- test-group: budgets
test-path: >-
tests/proxy_unit_tests/test_default_end_user_budget_simple.py
tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py
tests/proxy_unit_tests/test_zero_cost_model_budget_bypass.py
workers: 4
dist: loadscope
timeout: 15
- test-group: endpoints-and-responses
test-path: >-
tests/proxy_unit_tests/test_blog_posts_endpoint.py
tests/proxy_unit_tests/test_models_fallback_endpoint.py
tests/proxy_unit_tests/test_google_endpoint_routing.py
tests/proxy_unit_tests/test_google_gemini_proxy_request.py
tests/proxy_unit_tests/test_get_favicon.py
tests/proxy_unit_tests/test_get_image.py
tests/proxy_unit_tests/test_ui_path_detection.py
tests/proxy_unit_tests/test_prompt_test_endpoint.py
tests/proxy_unit_tests/test_check_batch_cost.py
tests/proxy_unit_tests/test_check_responses_cost.py
tests/proxy_unit_tests/test_response_polling_handler.py
tests/proxy_unit_tests/test_response_polling_pre_call_checks.py
tests/proxy_unit_tests/test_realtime_cache.py
tests/proxy_unit_tests/test_proxy_exception_mapping.py
tests/proxy_unit_tests/test_custom_tokenizer_bug.py
tests/proxy_unit_tests/test_model_response_typing
workers: 4
dist: loadscope
timeout: 15
uses: ./.github/workflows/_test-unit-services-base.yml
with:
test-path: ${{ matrix.test-path }}
@ -49,8 +235,5 @@ jobs:
reruns: 2
timeout-minutes: ${{ matrix.timeout }}
enable-postgres: true
dist: ${{ matrix.dist }}
artifact-name: proxy-db-${{ matrix.test-group }}
secrets:
DATABASE_URL: ${{ secrets.DATABASE_URL }}
POSTGRES_USER: ${{ secrets.POSTGRES_USER }}
POSTGRES_PASSWORD: ${{ secrets.POSTGRES_PASSWORD }}

View file

@ -36,6 +36,8 @@ jobs:
tests/test_litellm/proxy/health_endpoints
tests/test_litellm/proxy/public_endpoints
tests/test_litellm/proxy/prompts
tests/test_litellm/proxy/rag_endpoints
tests/test_litellm/proxy/realtime_endpoints
tests/test_litellm/proxy/ui_crud_endpoints
workers: 2
reruns: 2

View file

@ -1,6 +1,8 @@
name: "Unit Tests: Security"
# Uses DATABASE_URL secret — only runs on trusted branches, not PRs.
# Kept push-only (was previously required by DATABASE_URL secret scoping;
# now the postgres credentials are ephemeral localhost values but the
# push-trigger stays to match the proxy-db workflow cadence).
on:
push:
branches: [main, "litellm_**"]
@ -24,7 +26,3 @@ jobs:
timeout-minutes: 20
enable-postgres: true
artifact-name: security
secrets:
DATABASE_URL: ${{ secrets.DATABASE_URL }}
POSTGRES_USER: ${{ secrets.POSTGRES_USER }}
POSTGRES_PASSWORD: ${{ secrets.POSTGRES_PASSWORD }}

View file

@ -3907,7 +3907,7 @@ class OrganizationMemberUpdateResponse(MemberUpdateResponse):
class TeamInfoResponseObjectTeamTable(LiteLLM_TeamTable):
team_member_budget_table: Optional[LiteLLM_BudgetTable] = None
team_member_budget_table: Optional[LiteLLM_BudgetTableFull] = None
# Resources inherited from access groups (separate from direct assignments)
access_group_models: Optional[List[str]] = None
access_group_mcp_server_ids: Optional[List[str]] = None

View file

@ -632,20 +632,27 @@ class ResetBudgetJob:
now = datetime.utcnow()
# Note on raw SQL: prisma-client-python does not support null-filtering
# on `Json?` columns (no DbNull/JsonNull sentinel — see
# RobertCraigie/prisma-client-py#714). We use `query_raw` with
# `IS NOT NULL` so we don't materialize every key/team row on each
# tick of the reset job. Writes still go through the ORM.
# --- Keys ---
try:
all_keys = await self.prisma_client.db.litellm_verificationtoken.find_many(
where={"budget_limits": {"not": None}} # type: ignore[arg-type]
key_rows = await self.prisma_client.db.query_raw(
'SELECT token, budget_limits FROM "LiteLLM_VerificationToken" '
"WHERE budget_limits IS NOT NULL"
)
for key in all_keys:
raw = key.budget_limits # type: ignore[attr-defined]
for row in key_rows:
raw = row["budget_limits"]
if not raw:
continue
windows: list = raw if isinstance(raw, list) else json.loads(raw)
changed = False
for window in windows:
counter_key = (
f"spend:key:{key.token}:window:{window['budget_duration']}"
f"spend:key:{row['token']}:window:{window['budget_duration']}"
)
if await ResetBudgetJob._reset_expired_window(
window, counter_key, spend_counter_cache, now
@ -653,7 +660,7 @@ class ResetBudgetJob:
changed = True
if changed:
await self.prisma_client.db.litellm_verificationtoken.update(
where={"token": key.token},
where={"token": row["token"]},
data={"budget_limits": json.dumps(windows)}, # type: ignore[arg-type]
)
except Exception as e:
@ -663,26 +670,25 @@ class ResetBudgetJob:
# --- Teams ---
try:
all_teams = await self.prisma_client.db.litellm_teamtable.find_many(
where={"budget_limits": {"not": None}} # type: ignore[arg-type]
team_rows = await self.prisma_client.db.query_raw(
'SELECT team_id, budget_limits FROM "LiteLLM_TeamTable" '
"WHERE budget_limits IS NOT NULL"
)
for team in all_teams:
raw = team.budget_limits # type: ignore[attr-defined]
for row in team_rows:
raw = row["budget_limits"]
if not raw:
continue
windows = raw if isinstance(raw, list) else json.loads(raw)
changed = False
for window in windows:
counter_key = (
f"spend:team:{team.team_id}:window:{window['budget_duration']}"
)
counter_key = f"spend:team:{row['team_id']}:window:{window['budget_duration']}"
if await ResetBudgetJob._reset_expired_window(
window, counter_key, spend_counter_cache, now
):
changed = True
if changed:
await self.prisma_client.db.litellm_teamtable.update(
where={"team_id": team.team_id},
where={"team_id": row["team_id"]},
data={"budget_limits": json.dumps(windows)}, # type: ignore[arg-type]
)
except Exception as e:

View file

@ -52,12 +52,17 @@ from litellm.proxy._experimental.mcp_server.utils import (
from litellm.proxy._experimental.mcp_server.utils import (
validate_and_normalize_mcp_server_payload as _base_validate_and_normalize_mcp_server_payload,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
router = APIRouter(prefix="/v1/mcp", tags=["mcp"])
MCP_AVAILABLE: bool = True
TEMPORARY_MCP_SERVER_TTL_SECONDS = 300
TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX = "litellm:mcp:temporary_server"
def does_mcp_server_exist(
@ -329,13 +334,115 @@ if MCP_AVAILABLE:
)
return server
def get_cached_temporary_mcp_server(
async def _cache_temporary_mcp_server_in_redis(
server: MCPServer, ttl_seconds: int
) -> None:
"""
Best-effort write-through to Redis so temporary MCP OAuth sessions are
shared across proxy instances. Keep local in-memory cache as fallback.
"""
if litellm.cache is None or not hasattr(litellm.cache, "cache"):
return
cache_backend = getattr(litellm.cache, "cache", None)
if cache_backend is None or not hasattr(cache_backend, "async_set_cache"):
return
payload: Dict[str, Any] = server.model_dump(mode="json")
payload_json = json.dumps(payload)
try:
encrypted_payload = encrypt_value_helper(payload_json)
except Exception as e:
verbose_proxy_logger.debug(
f"Failed to encrypt temporary MCP server payload for Redis cache: {str(e)}"
)
return
if not isinstance(encrypted_payload, str):
verbose_proxy_logger.debug(
"Encrypted temporary MCP payload is not a string; skipping Redis cache write"
)
return
try:
await cache_backend.async_set_cache(
key=f"{TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX}:{server.server_id}",
value=encrypted_payload,
ttl=max(1, ttl_seconds),
)
except Exception as e:
verbose_proxy_logger.debug(
f"Failed to write temporary MCP server to Redis cache: {str(e)}"
)
async def _get_temporary_mcp_server_from_redis(
server_id: str,
) -> Optional[MCPServer]:
"""
Best-effort read from Redis shared cache. Returns None on miss/errors.
Values must be encrypted strings (same contract as _cache_temporary_mcp_server_in_redis);
legacy plaintext dict payloads are rejected.
"""
if litellm.cache is None or not hasattr(litellm.cache, "cache"):
return None
cache_backend = getattr(litellm.cache, "cache", None)
if cache_backend is None or not hasattr(cache_backend, "async_get_cache"):
return None
try:
cached_server = await cache_backend.async_get_cache(
key=f"{TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX}:{server_id}"
)
except Exception as e:
verbose_proxy_logger.debug(
f"Failed reading temporary MCP server from Redis cache: {str(e)}"
)
return None
if not isinstance(cached_server, str):
verbose_proxy_logger.debug(
"Temporary MCP Redis cache value must be an encrypted string; rejecting non-string payload"
)
return None
decrypted_json = decrypt_value_helper(
value=cached_server,
key="temporary_mcp_server",
exception_type="debug",
)
if decrypted_json is None:
return None
try:
loaded = json.loads(decrypted_json)
except Exception as e:
verbose_proxy_logger.debug(
f"Invalid decrypted temporary MCP payload in Redis cache: {str(e)}"
)
return None
if not isinstance(loaded, dict):
return None
payload_dict: Dict[str, Any] = loaded
try:
return MCPServer(**payload_dict)
except Exception as e:
verbose_proxy_logger.debug(
f"Invalid temporary MCP server payload in Redis cache: {str(e)}"
)
return None
async def get_cached_temporary_mcp_server(
server_id: str,
) -> Optional[MCPServer]:
_prune_expired_temporary_mcp_servers()
entry = _temporary_mcp_servers.get(server_id)
if entry is None:
return None
redis_server = await _get_temporary_mcp_server_from_redis(server_id)
if redis_server is None:
return None
# Intentionally avoid repopulating local cache from Redis to prevent
# extending effective lifetime beyond the remaining Redis TTL.
return redis_server
return entry.server
def _redact_mcp_credentials(
@ -1325,6 +1432,10 @@ if MCP_AVAILABLE:
temporary_server,
ttl_seconds=TEMPORARY_MCP_SERVER_TTL_SECONDS,
)
await _cache_temporary_mcp_server_in_redis(
temporary_server,
ttl_seconds=TEMPORARY_MCP_SERVER_TTL_SECONDS,
)
except Exception as e:
verbose_proxy_logger.exception(
f"Error caching temporary mcp server: {str(e)}"
@ -1336,10 +1447,10 @@ if MCP_AVAILABLE:
return _redact_mcp_credentials(temp_record)
def _get_cached_temporary_mcp_server_or_404(
async def _get_cached_temporary_mcp_server_or_404(
server_id: str, request: Optional[Request] = None
) -> MCPServer:
server = get_cached_temporary_mcp_server(server_id)
server = await get_cached_temporary_mcp_server(server_id)
if server is None:
# Fall back to real DB/config server (e.g. for the user-side OAuth flow
# which calls these endpoints with a real server_id, not a temp session id).
@ -1378,7 +1489,9 @@ if MCP_AVAILABLE:
response_type: Optional[str] = None,
scope: Optional[str] = None,
):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
mcp_server = await _get_cached_temporary_mcp_server_or_404(
server_id, request=request
)
# Use the server's stored client_id when the caller doesn't supply one
resolved_client_id = mcp_server.client_id or client_id or ""
if not resolved_client_id:
@ -1422,7 +1535,9 @@ if MCP_AVAILABLE:
refresh_token: Optional[str] = Form(None),
scope: Optional[str] = Form(None),
):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
mcp_server = await _get_cached_temporary_mcp_server_or_404(
server_id, request=request
)
resolved_client_id = mcp_server.client_id or client_id or ""
if not resolved_client_id:
raise HTTPException(
@ -1458,7 +1573,9 @@ if MCP_AVAILABLE:
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
mcp_server = await _get_cached_temporary_mcp_server_or_404(
server_id, request=request
)
request_data = await _read_request_body(request=request)
data: dict = {**request_data}

View file

@ -9,6 +9,7 @@ from fastapi import HTTPException, Request
import litellm
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy._types import ( # key request types; user request types; team request types; customer request types
BudgetNewRequest,
DeleteCustomerRequest,
@ -192,6 +193,13 @@ async def _clone_team_default_budget_for_member(
continue
cloned_data[field] = value
# Start the member's budget window at clone time, not the pool's reset
# timestamp — otherwise a member joining mid-cycle inherits a stale reset.
if cloned_data.get("budget_duration"):
cloned_data["budget_reset_at"] = get_budget_reset_time(
cloned_data["budget_duration"]
)
new_budget = await prisma_client.db.litellm_budgettable.create(data=cloned_data)
return new_budget.budget_id

View file

@ -8598,7 +8598,9 @@ class Router:
# No match found
return None
def map_team_model(self, team_model_name: str, team_id: str) -> Optional[str]:
def map_team_model(
self, team_model_name: Optional[str], team_id: str
) -> Optional[str]:
"""
Check if team_model_name resolves to team-specific deployments.
@ -8606,6 +8608,11 @@ class Router:
sibling deployments via team_id filtering, instead of collapsing to a
single internal model_name.
When team_model_name is None (e.g. vector store / file endpoints that
don't include a model in their request), returns the first matching
team deployment's team_public_model_name so the router can inject BYOK
credentials from the team-scoped deployment.
Returns:
- str: the team_model_name if team deployments exist for this team
- None: if no team-specific model is found
@ -8615,6 +8622,13 @@ class Router:
return None
for model in models:
if model.get("model_info", {}).get("team_id") == team_id:
if team_model_name is None:
# No model was specified (e.g. vector store endpoints).
# Return the deployment's public model name so the router
# can route to it and inject the BYOK API key.
return model.get("model_info", {}).get(
"team_public_model_name"
) or model.get("model_name")
return team_model_name
# No team-scoped deployment found; wildcard/pattern routes are

View file

@ -1,6 +1,6 @@
[project]
name = "litellm"
version = "1.83.12"
version = "1.83.13"
description = "Library to easily interface with LLM API providers"
readme = "README.md"
requires-python = ">=3.10, <3.14"
@ -236,7 +236,7 @@ source-exclude = [
profile = "black"
[tool.commitizen]
version = "1.83.12"
version = "1.83.13"
version_files = [
"pyproject.toml:^version",
]

View file

@ -1,7 +1,9 @@
import asyncio
import json
import os
import sys
import time
import types
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, List
from unittest.mock import AsyncMock, MagicMock
@ -696,9 +698,9 @@ def test_reset_budget_resets_endusers_with_null_budget_id(
# Both end users should have been reset
updated = mock_prisma_client.updated_data["enduser"]
assert len(updated) == 2, (
f"Expected 2 endusers reset (1 explicit + 1 implicit), got {len(updated)}"
)
assert (
len(updated) == 2
), f"Expected 2 endusers reset (1 explicit + 1 implicit), got {len(updated)}"
user_ids = {u.user_id for u in updated}
assert "enduser-explicit" in user_ids
@ -819,3 +821,231 @@ def test_reset_budget_for_team_members_preserves_total_spend():
assert call_kwargs["where"]["budget_id"]["in"] == ["budget-1"]
assert call_kwargs["data"] == {"spend": 0}
assert "total_spend" not in call_kwargs["data"]
# ---------------------------------------------------------------------------
# reset_budget_windows (per-key / per-team concurrent window resets)
# ---------------------------------------------------------------------------
def _make_reset_budget_windows_job(
monkeypatch,
key_rows: List[Dict[str, Any]],
team_rows: List[Dict[str, Any]],
):
"""Build a ResetBudgetJob with a fully-mocked prisma client and a fake
`litellm.proxy.proxy_server` module exposing a stub `spend_counter_cache`.
Returns (job, prisma_client_mock, spend_counter_cache_mock).
"""
prisma_client = MagicMock()
async def fake_query_raw(query: str, *args, **kwargs):
# Dispatch by table name in the SQL so a single stub covers both calls.
if '"LiteLLM_VerificationToken"' in query:
return key_rows
if '"LiteLLM_TeamTable"' in query:
return team_rows
raise AssertionError(f"Unexpected query_raw call: {query}")
prisma_client.db.query_raw = AsyncMock(side_effect=fake_query_raw)
prisma_client.db.litellm_verificationtoken.update = AsyncMock(return_value=None)
prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=None)
# Stub out litellm.proxy.proxy_server so the in-function
# `from litellm.proxy.proxy_server import spend_counter_cache` resolves
# without importing the real (heavy) module.
spend_counter_cache = MagicMock()
spend_counter_cache.in_memory_cache.set_cache = MagicMock()
spend_counter_cache.redis_cache = None # skip the async redis branch
fake_module = types.ModuleType("litellm.proxy.proxy_server")
fake_module.spend_counter_cache = spend_counter_cache
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module)
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
return job, prisma_client, spend_counter_cache
def test_reset_budget_windows_uses_is_not_null_filter(monkeypatch):
"""Regression guard for the Prisma client limitation documented in
RobertCraigie/prisma-client-py#714: `{"not": None}` on a `Json?` column
raises `MissingRequiredValueError`. We work around it by using `query_raw`
with `IS NOT NULL`. If someone reverts to the ORM filter, this test fails.
"""
job, prisma_client, _ = _make_reset_budget_windows_job(
monkeypatch, key_rows=[], team_rows=[]
)
asyncio.run(job.reset_budget_windows())
queries = [call.args[0] for call in prisma_client.db.query_raw.await_args_list]
assert len(queries) == 2, queries
key_query, team_query = queries
assert '"LiteLLM_VerificationToken"' in key_query
assert "budget_limits IS NOT NULL" in key_query
assert '"LiteLLM_TeamTable"' in team_query
assert "budget_limits IS NOT NULL" in team_query
def test_reset_budget_windows_resets_expired_key_window(monkeypatch):
"""A key whose window's `reset_at` has passed gets an update with a new
`reset_at` in the future, and the in-memory spend counter is cleared."""
now = datetime.utcnow()
expired = (now - timedelta(minutes=5)).isoformat() + "Z"
key_rows = [
{
"token": "sk-expired",
"budget_limits": [{"budget_duration": "1d", "reset_at": expired}],
}
]
job, prisma_client, spend_counter_cache = _make_reset_budget_windows_job(
monkeypatch, key_rows=key_rows, team_rows=[]
)
asyncio.run(job.reset_budget_windows())
# Update should have been called exactly once with the expired token.
prisma_client.db.litellm_verificationtoken.update.assert_awaited_once()
call_kwargs = prisma_client.db.litellm_verificationtoken.update.await_args.kwargs
assert call_kwargs["where"] == {"token": "sk-expired"}
# The `budget_limits` payload is re-serialized JSON with a bumped reset_at.
written_windows = json.loads(call_kwargs["data"]["budget_limits"])
assert len(written_windows) == 1
new_reset_at = datetime.fromisoformat(
written_windows[0]["reset_at"].replace("Z", "+00:00")
).replace(tzinfo=None)
assert new_reset_at > now
# The spend counter for this key+window was cleared.
spend_counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:key:sk-expired:window:1d", value=0.0
)
def test_reset_budget_windows_skips_unexpired_key_window(monkeypatch):
"""If `reset_at` is in the future, no write should happen for that key."""
now = datetime.utcnow()
future = (now + timedelta(hours=1)).isoformat() + "Z"
key_rows = [
{
"token": "sk-future",
"budget_limits": [{"budget_duration": "1d", "reset_at": future}],
}
]
job, prisma_client, _ = _make_reset_budget_windows_job(
monkeypatch, key_rows=key_rows, team_rows=[]
)
asyncio.run(job.reset_budget_windows())
prisma_client.db.litellm_verificationtoken.update.assert_not_awaited()
def test_reset_budget_windows_resets_expired_team_window(monkeypatch):
"""Same as the key test, but for teams."""
now = datetime.utcnow()
expired = (now - timedelta(minutes=1)).isoformat() + "Z"
team_rows = [
{
"team_id": "team-expired",
"budget_limits": [{"budget_duration": "30d", "reset_at": expired}],
}
]
job, prisma_client, spend_counter_cache = _make_reset_budget_windows_job(
monkeypatch, key_rows=[], team_rows=team_rows
)
asyncio.run(job.reset_budget_windows())
prisma_client.db.litellm_teamtable.update.assert_awaited_once()
call_kwargs = prisma_client.db.litellm_teamtable.update.await_args.kwargs
assert call_kwargs["where"] == {"team_id": "team-expired"}
assert "budget_limits" in call_kwargs["data"]
spend_counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:team:team-expired:window:30d", value=0.0
)
def test_reset_budget_windows_handles_string_budget_limits(monkeypatch):
"""Defensive: if `query_raw` returns `budget_limits` as a JSON-encoded
string (driver-dependent), the code still parses and resets it.
"""
now = datetime.utcnow()
expired = (now - timedelta(minutes=1)).isoformat() + "Z"
key_rows = [
{
"token": "sk-string-limits",
"budget_limits": json.dumps(
[{"budget_duration": "1d", "reset_at": expired}]
),
}
]
job, prisma_client, _ = _make_reset_budget_windows_job(
monkeypatch, key_rows=key_rows, team_rows=[]
)
asyncio.run(job.reset_budget_windows())
prisma_client.db.litellm_verificationtoken.update.assert_awaited_once()
def test_reset_budget_windows_skips_row_with_empty_budget_limits(monkeypatch):
"""A row whose `budget_limits` comes back as an empty/falsy payload
(shouldn't happen given the WHERE filter, but we guard anyway) must not
trigger an update or crash the loop."""
key_rows = [
{"token": "sk-empty-list", "budget_limits": []},
{"token": "sk-empty-str", "budget_limits": ""},
]
job, prisma_client, _ = _make_reset_budget_windows_job(
monkeypatch, key_rows=key_rows, team_rows=[]
)
asyncio.run(job.reset_budget_windows())
prisma_client.db.litellm_verificationtoken.update.assert_not_awaited()
def test_reset_budget_windows_query_error_does_not_break_team_path(monkeypatch):
"""If the key query raises, the teams path still runs (and vice-versa).
Each side has its own try/except; this locks that in."""
now = datetime.utcnow()
expired = (now - timedelta(minutes=1)).isoformat() + "Z"
prisma_client = MagicMock()
async def fake_query_raw(query: str, *args, **kwargs):
if '"LiteLLM_VerificationToken"' in query:
raise RuntimeError("boom")
if '"LiteLLM_TeamTable"' in query:
return [
{
"team_id": "team-ok",
"budget_limits": [{"budget_duration": "1d", "reset_at": expired}],
}
]
raise AssertionError(query)
prisma_client.db.query_raw = AsyncMock(side_effect=fake_query_raw)
prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=None)
spend_counter_cache = MagicMock()
spend_counter_cache.in_memory_cache.set_cache = MagicMock()
spend_counter_cache.redis_cache = None
fake_module = types.ModuleType("litellm.proxy.proxy_server")
fake_module.spend_counter_cache = spend_counter_cache
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module)
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_windows()) # must not raise
prisma_client.db.litellm_teamtable.update.assert_awaited_once()

View file

@ -1,6 +1,7 @@
import os
import sys
import types
import json
from datetime import datetime, timedelta
from types import SimpleNamespace
from typing import List, Optional
@ -1311,7 +1312,8 @@ class TestTemporaryMCPSessionEndpoints:
assert cache["temp-cache"].server is server
assert cache["temp-cache"].expires_at > datetime.utcnow()
def test_get_cached_temporary_mcp_server_prunes_expired_entries(self):
@pytest.mark.asyncio
async def test_get_cached_temporary_mcp_server_prunes_expired_entries(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_TemporaryMCPServerEntry,
get_cached_temporary_mcp_server,
@ -1327,12 +1329,13 @@ class TestTemporaryMCPSessionEndpoints:
"litellm.proxy.management_endpoints.mcp_management_endpoints._temporary_mcp_servers",
cache,
):
result = get_cached_temporary_mcp_server("expired")
result = await get_cached_temporary_mcp_server("expired")
assert result is None
assert "expired" not in cache
def test_get_cached_temporary_mcp_server_or_404(self):
@pytest.mark.asyncio
async def test_get_cached_temporary_mcp_server_or_404(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_get_cached_temporary_mcp_server_or_404,
)
@ -1343,17 +1346,17 @@ class TestTemporaryMCPSessionEndpoints:
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
return_value=server,
) as get_cached:
result = _get_cached_temporary_mcp_server_or_404("cached")
result = await _get_cached_temporary_mcp_server_or_404("cached")
assert result is server
get_cached.assert_called_once_with("cached")
get_cached.assert_awaited_once_with("cached")
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
return_value=None,
):
with pytest.raises(HTTPException) as exc_info:
_get_cached_temporary_mcp_server_or_404("missing")
await _get_cached_temporary_mcp_server_or_404("missing")
assert exc_info.value.status_code == 404
@ -1403,6 +1406,10 @@ class TestTemporaryMCPSessionEndpoints:
"litellm.proxy.management_endpoints.mcp_management_endpoints._cache_temporary_mcp_server",
MagicMock(),
) as cache_mock,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._cache_temporary_mcp_server_in_redis",
AsyncMock(),
) as redis_cache_mock,
):
response = await add_session_mcp_server(
payload=payload,
@ -1414,6 +1421,9 @@ class TestTemporaryMCPSessionEndpoints:
cache_mock.assert_called_once_with(
built_server, ttl_seconds=TEMPORARY_MCP_SERVER_TTL_SECONDS
)
redis_cache_mock.assert_awaited_once_with(
built_server, ttl_seconds=TEMPORARY_MCP_SERVER_TTL_SECONDS
)
args, _ = mock_manager.build_mcp_server_from_table.call_args
temp_record = args[0]
@ -1486,7 +1496,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is authorize_response
get_server.assert_called_once_with("server-1", request=request)
get_server.assert_awaited_once_with("server-1", request=request)
authorize_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@ -1533,7 +1543,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is exchange_response
get_server.assert_called_once_with("server-1", request=request)
get_server.assert_awaited_once_with("server-1", request=request)
exchange_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@ -1581,7 +1591,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is exchange_response
get_server.assert_called_once_with("server-1", request=request)
get_server.assert_awaited_once_with("server-1", request=request)
exchange_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@ -1628,7 +1638,7 @@ class TestTemporaryMCPSessionEndpoints:
result = await mcp_register(request=request, server_id="server-1")
assert result is register_response
get_server.assert_called_once_with("server-1", request=request)
get_server.assert_awaited_once_with("server-1", request=request)
read_body.assert_awaited_once_with(request=request)
register_mock.assert_awaited_once_with(
request=request,
@ -1640,6 +1650,218 @@ class TestTemporaryMCPSessionEndpoints:
fallback_client_id="server-1",
)
@pytest.mark.asyncio
async def test_get_cached_temporary_mcp_server_falls_back_to_redis(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
get_cached_temporary_mcp_server,
)
server = generate_mock_mcp_server_config_record(server_id="from-redis")
serialized = json.dumps(server.model_dump(mode="json"))
mock_cache_backend = SimpleNamespace(
async_get_cache=AsyncMock(return_value="encrypted-payload")
)
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._temporary_mcp_servers",
{},
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.decrypt_value_helper",
return_value=serialized,
):
result = await get_cached_temporary_mcp_server("from-redis")
finally:
mgmt_endpoints.litellm.cache = original_cache
assert result is not None
assert result.server_id == "from-redis"
mock_cache_backend.async_get_cache.assert_awaited_once_with(
key="litellm:mcp:temporary_server:from-redis"
)
@pytest.mark.asyncio
async def test_cache_temporary_mcp_server_in_redis_uses_ttl_and_key(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_cache_temporary_mcp_server_in_redis,
)
server = generate_mock_mcp_server_config_record(server_id="to-redis")
mock_cache_backend = SimpleNamespace(async_set_cache=AsyncMock())
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.encrypt_value_helper",
return_value="encrypted-payload",
):
await _cache_temporary_mcp_server_in_redis(server, ttl_seconds=123)
finally:
mgmt_endpoints.litellm.cache = original_cache
mock_cache_backend.async_set_cache.assert_awaited_once()
call_kwargs = mock_cache_backend.async_set_cache.await_args.kwargs
assert call_kwargs["key"] == "litellm:mcp:temporary_server:to-redis"
assert call_kwargs["ttl"] == 123
@pytest.mark.asyncio
async def test_cache_temporary_mcp_server_in_redis_encrypts_payload(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_cache_temporary_mcp_server_in_redis,
)
server = generate_mock_mcp_server_config_record(server_id="to-redis-encrypted")
mock_cache_backend = SimpleNamespace(async_set_cache=AsyncMock())
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.encrypt_value_helper",
return_value="encrypted-payload",
) as encrypt_mock:
await _cache_temporary_mcp_server_in_redis(server, ttl_seconds=60)
finally:
mgmt_endpoints.litellm.cache = original_cache
encrypt_mock.assert_called_once()
call_kwargs = mock_cache_backend.async_set_cache.await_args.kwargs
assert call_kwargs["value"] == "encrypted-payload"
@pytest.mark.asyncio
async def test_get_temporary_mcp_server_from_redis_decrypts_payload(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_get_temporary_mcp_server_from_redis,
)
server = generate_mock_mcp_server_config_record(server_id="from-redis-encrypted")
serialized = json.dumps(server.model_dump(mode="json"))
mock_cache_backend = SimpleNamespace(
async_get_cache=AsyncMock(return_value="encrypted-payload")
)
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.decrypt_value_helper",
return_value=serialized,
) as decrypt_mock:
result = await _get_temporary_mcp_server_from_redis(
"from-redis-encrypted"
)
finally:
mgmt_endpoints.litellm.cache = original_cache
assert result is not None
assert result.server_id == "from-redis-encrypted"
decrypt_mock.assert_called_once()
@pytest.mark.asyncio
async def test_cache_temporary_mcp_server_in_redis_skips_on_encrypt_failure(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_cache_temporary_mcp_server_in_redis,
)
server = generate_mock_mcp_server_config_record(server_id="encrypt-fail")
mock_cache_backend = SimpleNamespace(async_set_cache=AsyncMock())
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.encrypt_value_helper",
side_effect=Exception("boom"),
):
await _cache_temporary_mcp_server_in_redis(server, ttl_seconds=60)
finally:
mgmt_endpoints.litellm.cache = original_cache
mock_cache_backend.async_set_cache.assert_not_awaited()
@pytest.mark.asyncio
async def test_cache_temporary_mcp_server_in_redis_skips_non_string_encryption_result(
self,
):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_cache_temporary_mcp_server_in_redis,
)
server = generate_mock_mcp_server_config_record(server_id="encrypt-non-string")
mock_cache_backend = SimpleNamespace(async_set_cache=AsyncMock())
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.encrypt_value_helper",
return_value={"not": "a-string"},
):
await _cache_temporary_mcp_server_in_redis(server, ttl_seconds=60)
finally:
mgmt_endpoints.litellm.cache = original_cache
mock_cache_backend.async_set_cache.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_temporary_mcp_server_from_redis_returns_none_on_invalid_decrypt_json(
self,
):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_get_temporary_mcp_server_from_redis,
)
mock_cache_backend = SimpleNamespace(async_get_cache=AsyncMock(return_value="enc"))
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.decrypt_value_helper",
return_value="{not json}",
):
result = await _get_temporary_mcp_server_from_redis("bad-json")
finally:
mgmt_endpoints.litellm.cache = original_cache
assert result is None
@pytest.mark.asyncio
async def test_get_temporary_mcp_server_from_redis_returns_none_on_decrypt_none(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_get_temporary_mcp_server_from_redis,
)
mock_cache_backend = SimpleNamespace(async_get_cache=AsyncMock(return_value="enc"))
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.decrypt_value_helper",
return_value=None,
):
result = await _get_temporary_mcp_server_from_redis("decrypt-none")
finally:
mgmt_endpoints.litellm.cache = original_cache
assert result is None
@pytest.mark.asyncio
async def test_get_temporary_mcp_server_from_redis_rejects_plain_dict_payload(self):
"""Plain dict values in Redis are not accepted (write path is encrypted-only)."""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_get_temporary_mcp_server_from_redis,
)
server = generate_mock_mcp_server_config_record(server_id="legacy-dict")
mock_cache_backend = SimpleNamespace(
async_get_cache=AsyncMock(return_value=server.model_dump(mode="json"))
)
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
result = await _get_temporary_mcp_server_from_redis("legacy-dict")
finally:
mgmt_endpoints.litellm.cache = original_cache
assert result is None
class TestUpdateMCPServer:
"""Test suite for update MCP server functionality"""

View file

@ -422,6 +422,26 @@ def reset_router_callbacks():
litellm.logging_callback_manager._reset_all_callbacks()
@pytest.fixture(autouse=True)
def reset_proxy_auth_globals(monkeypatch):
"""
Pin proxy auth-related globals to a known baseline so tests don't inherit
leaked state (master_key, prisma_client, custom auth, cached tokens) from
earlier tests. Individual tests can still override via their own
monkeypatch calls — those run after this fixture and revert first.
"""
import litellm.proxy.proxy_server as ps
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(ps, "master_key", None)
monkeypatch.setattr(ps, "user_custom_auth", None)
monkeypatch.setattr(ps, "general_settings", {})
try:
ps.user_api_key_cache.in_memory_cache.cache_dict.clear()
except AttributeError:
pass
@pytest.mark.asyncio
async def test_ui_view_spend_logs_with_user_id(client, monkeypatch):
mock_spend_logs = [
@ -1150,14 +1170,14 @@ async def test_ui_view_spend_logs_date_range_filter(client, monkeypatch):
async def test_ui_view_spend_logs_unauthorized(client):
# Test without authorization header
response = client.get("/spend/logs/ui")
assert response.status_code == 401 or response.status_code == 403
assert response.status_code in (401, 403), response.text
# Test with invalid authorization
response = client.get(
"/spend/logs/ui",
headers={"Authorization": "Bearer invalid-token"},
)
assert response.status_code == 401 or response.status_code == 403
assert response.status_code in (401, 403), response.text
@pytest.mark.asyncio

View file

@ -60,6 +60,7 @@ export interface TeamMembership {
team_id: string;
budget_id: string;
spend: number;
total_spend: number | null;
litellm_budget_table: {
budget_id: string;
soft_budget: number | null;
@ -69,6 +70,7 @@ export interface TeamMembership {
rpm_limit: number | null;
model_max_budget: Record<string, number> | null;
budget_duration: string | null;
budget_reset_at: string | null;
allowed_models?: string[] | null;
};
}

View file

@ -1,6 +1,7 @@
import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { Member } from "@/components/networking";
import { formatBudgetReset } from "@/utils/budgetUtils";
import { formatNumberWithCommas } from "@/utils/dataUtils";
import { isProxyAdminRole, isUserTeamAdminForSingleTeam } from "@/utils/roles";
import { InfoCircleOutlined } from "@ant-design/icons";
@ -45,11 +46,16 @@ export default function TeamMemberTab({
return "0";
};
// Helper function to get spend for a user
const getUserSpend = (userId: string | null): number | null => {
const getUserCurrentCycleSpend = (userId: string | null): number => {
if (!userId) return 0;
const membership = teamData.team_memberships.find((tm) => tm.user_id === userId);
return membership?.spend || 0;
return membership?.spend ?? 0;
};
const getUserTotalSpend = (userId: string | null): number => {
if (!userId) return 0;
const membership = teamData.team_memberships.find((tm) => tm.user_id === userId);
return membership?.total_spend ?? 0;
};
const getUserBudget = (userId: string | null): string | null => {
@ -89,6 +95,12 @@ export default function TeamMemberTab({
return models && models.length > 0 ? models : null;
};
const getUserBudgetReset = (userId: string | null): string | null => {
if (!userId) return null;
const membership = teamData.team_memberships.find((tm) => tm.user_id === userId);
return formatBudgetReset(membership?.litellm_budget_table?.budget_reset_at);
};
const extraColumns: ColumnsType<Member> = [
{
title: (
@ -124,15 +136,29 @@ export default function TeamMemberTab({
{
title: (
<Space direction="horizontal">
Team Member Spend (USD)
<Tooltip title="This is the amount spent by a user in the team.">
Current Cycle Spend (USD)
<Tooltip title="Spend for the current budget cycle. Resets to $0 when the member's budget window rolls over. This is the value checked against the member's budget.">
<InfoCircleOutlined />
</Tooltip>
</Space>
),
key: "spend",
render: (_: unknown, record: Member) => (
<Typography.Text>${formatNumberWithCommas(getUserSpend(record.user_id), 4)}</Typography.Text>
<Typography.Text>${formatNumberWithCommas(getUserCurrentCycleSpend(record.user_id), 4)}</Typography.Text>
),
},
{
title: (
<Space direction="horizontal">
Total Spend (USD)
<Tooltip title="Cumulative spend by this member within this team, across all budget cycles. Tracking began 2026-04-21; spend from before that date is not included.">
<InfoCircleOutlined />
</Tooltip>
</Space>
),
key: "total_spend",
render: (_: unknown, record: Member) => (
<Typography.Text>${formatNumberWithCommas(getUserTotalSpend(record.user_id), 4)}</Typography.Text>
),
},
{
@ -147,6 +173,18 @@ export default function TeamMemberTab({
);
},
},
{
title: "Budget Reset",
key: "budget_reset",
render: (_: unknown, record: Member) => {
const reset = getUserBudgetReset(record.user_id);
return reset ? (
<Typography.Text>{reset}</Typography.Text>
) : (
<Typography.Text type="secondary">—</Typography.Text>
);
},
},
{
title: (
<Space direction="horizontal">

View file

@ -0,0 +1,8 @@
import dayjs from "dayjs";
export function formatBudgetReset(iso: string | null | undefined): string | null {
if (!iso) return null;
const resetDate = dayjs(iso);
if (!resetDate.isValid()) return null;
return resetDate.format("MMM D, YYYY");
}

4
uv.lock generated
View file

@ -9,7 +9,7 @@ resolution-markers = [
]
[options]
exclude-newer = "2026-04-20T01:21:50.985363Z"
exclude-newer = "2026-04-21T00:00:09.504288Z"
exclude-newer-span = "P3D"
[manifest]
@ -3085,7 +3085,7 @@ wheels = [
[[package]]
name = "litellm"
version = "1.83.12"
version = "1.83.13"
source = { editable = "." }
dependencies = [
{ name = "aiohttp" },