mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
commit
e68c60a66e
18 changed files with 963 additions and 172 deletions
|
|
@ -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
|
||||
|
|
|
|||
37
.github/workflows/_test-unit-services-base.yml
vendored
37
.github/workflows/_test-unit-services-base.yml
vendored
|
|
@ -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 \
|
||||
|
|
|
|||
219
.github/workflows/test-unit-proxy-db.yml
vendored
219
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
8
.github/workflows/test-unit-security.yml
vendored
8
.github/workflows/test-unit-security.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
};
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
8
ui/litellm-dashboard/src/utils/budgetUtils.ts
Normal file
8
ui/litellm-dashboard/src/utils/budgetUtils.ts
Normal 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
4
uv.lock
generated
|
|
@ -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" },
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue