mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge branch 'main' into fix/dashscope-logo
This commit is contained in:
commit
9d5afd9cd9
114 changed files with 6092 additions and 5921 deletions
|
|
@ -1330,6 +1330,57 @@ jobs:
|
|||
paths:
|
||||
- audio_coverage.xml
|
||||
- audio_coverage
|
||||
redis_caching_unit_tests:
|
||||
docker:
|
||||
- image: cimg/python:3.11
|
||||
auth:
|
||||
username: ${DOCKERHUB_USERNAME}
|
||||
password: ${DOCKERHUB_PASSWORD}
|
||||
working_directory: ~/project
|
||||
|
||||
steps:
|
||||
- checkout
|
||||
- setup_google_dns
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
command: |
|
||||
python -m pip install --upgrade pip uv
|
||||
uv pip install --system -r requirements.txt
|
||||
pip install "pytest==7.3.1"
|
||||
pip install "pytest-retry==1.6.3"
|
||||
pip install "pytest-cov==5.0.0"
|
||||
pip install "pytest-asyncio==0.21.1"
|
||||
pip install "pytest-xdist==3.6.1"
|
||||
pip install "pytest-rerunfailures==14.0"
|
||||
# Run pytest and generate JUnit XML report
|
||||
- run:
|
||||
name: Run tests
|
||||
command: |
|
||||
pwd
|
||||
ls
|
||||
python -m pytest -vv \
|
||||
tests/local_testing/test_dual_cache.py \
|
||||
tests/local_testing/test_redis_batch_optimizations.py \
|
||||
tests/local_testing/test_router_utils.py \
|
||||
--cov=litellm --cov-report=xml \
|
||||
-x -s -v --junitxml=test-results/junit.xml \
|
||||
--durations=5 -n 2 \
|
||||
--reruns 2 --reruns-delay 1
|
||||
no_output_timeout: 20m
|
||||
- run:
|
||||
name: Rename the coverage files
|
||||
command: |
|
||||
mv coverage.xml redis_caching_coverage.xml
|
||||
mv .coverage redis_caching_coverage
|
||||
|
||||
# Store test results
|
||||
- store_test_results:
|
||||
path: test-results
|
||||
- persist_to_workspace:
|
||||
root: .
|
||||
paths:
|
||||
- redis_caching_coverage.xml
|
||||
- redis_caching_coverage
|
||||
installing_litellm_on_python:
|
||||
docker:
|
||||
- image: cimg/python:3.11
|
||||
|
|
@ -2868,114 +2919,6 @@ jobs:
|
|||
- store_test_results:
|
||||
path: test-results
|
||||
|
||||
proxy_e2e_azure_batches_tests:
|
||||
machine:
|
||||
image: ubuntu-2204:2023.10.1
|
||||
resource_class: large
|
||||
working_directory: ~/project
|
||||
steps:
|
||||
- checkout
|
||||
- setup_google_dns
|
||||
- run:
|
||||
name: Install Docker CLI
|
||||
command: |
|
||||
curl -fsSL https://get.docker.com | sh
|
||||
sudo usermod -aG docker $USER
|
||||
docker version
|
||||
- run:
|
||||
name: Install Python 3.12
|
||||
command: |
|
||||
curl https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh --output miniconda.sh
|
||||
bash miniconda.sh -b -p $HOME/miniconda
|
||||
export PATH="$HOME/miniconda/bin:$PATH"
|
||||
conda init bash
|
||||
source ~/.bashrc
|
||||
conda create -n myenv python=3.12 -y
|
||||
conda activate myenv
|
||||
python --version
|
||||
- run:
|
||||
name: Install Poetry
|
||||
command: |
|
||||
export PATH="$HOME/miniconda/bin:$PATH"
|
||||
source $HOME/miniconda/etc/profile.d/conda.sh
|
||||
conda activate myenv
|
||||
pip install poetry
|
||||
- run:
|
||||
name: Install dockerize
|
||||
command: |
|
||||
wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start PostgreSQL Database
|
||||
command: |
|
||||
docker run -d \
|
||||
--name postgres-db \
|
||||
-e POSTGRES_USER=llmproxy \
|
||||
-e POSTGRES_PASSWORD=dbpassword9090 \
|
||||
-e POSTGRES_DB=litellm \
|
||||
-p 5432:5432 \
|
||||
postgres:15
|
||||
- run:
|
||||
name: Wait for PostgreSQL to be ready
|
||||
command: dockerize -wait tcp://localhost:5432 -timeout 1m
|
||||
- run:
|
||||
name: Install system dependencies
|
||||
command: |
|
||||
sudo apt-get update -y
|
||||
sudo apt-get install -y libpq-dev
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
command: |
|
||||
export PATH="$HOME/miniconda/bin:$PATH"
|
||||
source $HOME/miniconda/etc/profile.d/conda.sh
|
||||
conda activate myenv
|
||||
poetry config virtualenvs.in-project true
|
||||
poetry install --with dev,proxy-dev --extras "proxy"
|
||||
poetry run pip install psycopg2-binary uvicorn fastapi httpx tenacity
|
||||
- run:
|
||||
name: Setup litellm-enterprise
|
||||
command: |
|
||||
export PATH="$HOME/miniconda/bin:$PATH"
|
||||
source $HOME/miniconda/etc/profile.d/conda.sh
|
||||
conda activate myenv
|
||||
poetry run pip install --force-reinstall --no-deps -e enterprise/
|
||||
- run:
|
||||
name: Generate Prisma client
|
||||
command: |
|
||||
export PATH="$HOME/miniconda/bin:$PATH"
|
||||
source $HOME/miniconda/etc/profile.d/conda.sh
|
||||
conda activate myenv
|
||||
poetry run prisma generate --schema litellm/proxy/schema.prisma
|
||||
- run:
|
||||
name: Run Prisma migrations
|
||||
command: |
|
||||
export PATH="$HOME/miniconda/bin:$PATH"
|
||||
source $HOME/miniconda/etc/profile.d/conda.sh
|
||||
conda activate myenv
|
||||
export DATABASE_URL=postgresql://llmproxy:dbpassword9090@localhost:5432/litellm
|
||||
cd litellm/proxy
|
||||
poetry run prisma migrate deploy --schema schema.prisma
|
||||
cd ../..
|
||||
- run:
|
||||
name: Run Azure Batch E2E Tests
|
||||
command: |
|
||||
export PATH="$HOME/miniconda/bin:$PATH"
|
||||
source $HOME/miniconda/etc/profile.d/conda.sh
|
||||
conda activate myenv
|
||||
export DATABASE_URL=postgresql://llmproxy:dbpassword9090@localhost:5432/litellm
|
||||
export USE_LOCAL_LITELLM=true
|
||||
export USE_MOCK_MODELS=true
|
||||
export USE_STATE_TRACKER=true
|
||||
export LITELLM_LOG=DEBUG
|
||||
poetry run pytest tests/proxy_e2e_azure_batches_tests/test_proxy_e2e_azure_batches.py \
|
||||
-vv -s -k "test_e2e_managed_batch" \
|
||||
--tb=short \
|
||||
--maxfail=3 \
|
||||
--durations=10 \
|
||||
--junitxml=test-results/junit.xml
|
||||
no_output_timeout: 15m
|
||||
|
||||
upload-coverage:
|
||||
docker:
|
||||
- image: cimg/python:3.9
|
||||
|
|
@ -2997,7 +2940,7 @@ jobs:
|
|||
python -m venv venv
|
||||
. venv/bin/activate
|
||||
pip install coverage
|
||||
coverage combine realtime_translation_coverage ocr_coverage search_coverage mcp_coverage litellm_mcps_tests_coverage logging_coverage audio_coverage local_testing_part1_coverage local_testing_part2_coverage pass_through_unit_tests_coverage batches_coverage guardrails_coverage
|
||||
coverage combine realtime_translation_coverage ocr_coverage search_coverage mcp_coverage litellm_mcps_tests_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
|
||||
coverage xml
|
||||
- codecov/upload:
|
||||
file: ./coverage.xml
|
||||
|
|
@ -3182,6 +3125,117 @@ jobs:
|
|||
CI=true npm run test -- --run \
|
||||
--pool forks --poolOptions.forks.maxForks=8
|
||||
|
||||
e2e_ui_testing:
|
||||
docker:
|
||||
- image: cimg/python:3.12-browsers
|
||||
auth:
|
||||
username: ${DOCKERHUB_USERNAME}
|
||||
password: ${DOCKERHUB_PASSWORD}
|
||||
- image: cimg/postgres:16.0
|
||||
environment:
|
||||
POSTGRES_USER: e2euser
|
||||
POSTGRES_PASSWORD: e2epassword
|
||||
POSTGRES_DB: litellm_e2e
|
||||
resource_class: large
|
||||
working_directory: ~/project
|
||||
environment:
|
||||
DATABASE_URL: "postgresql://e2euser:e2epassword@localhost:5432/litellm_e2e"
|
||||
CI: "true"
|
||||
steps:
|
||||
- checkout
|
||||
- setup_google_dns
|
||||
- restore_cache:
|
||||
keys:
|
||||
- ui-e2e-py-deps-v1-{{ checksum "requirements.txt" }}
|
||||
- run:
|
||||
name: Install Python dependencies
|
||||
command: |
|
||||
python -m pip install --upgrade pip uv
|
||||
uv pip install --system -r requirements.txt
|
||||
pip install "prisma==0.11.0"
|
||||
prisma generate --schema litellm/proxy/schema.prisma
|
||||
- save_cache:
|
||||
key: ui-e2e-py-deps-v1-{{ checksum "requirements.txt" }}
|
||||
paths:
|
||||
- ~/.local/lib
|
||||
- ~/.local/bin
|
||||
- restore_cache:
|
||||
keys:
|
||||
- ui-e2e-node-deps-v1-{{ checksum "ui/litellm-dashboard/package-lock.json" }}
|
||||
- run:
|
||||
name: Install Node dependencies and Playwright
|
||||
command: |
|
||||
cd ui/litellm-dashboard
|
||||
npm ci
|
||||
npx playwright install chromium --with-deps
|
||||
- save_cache:
|
||||
key: ui-e2e-node-deps-v1-{{ checksum "ui/litellm-dashboard/package-lock.json" }}
|
||||
paths:
|
||||
- ui/litellm-dashboard/node_modules
|
||||
- run:
|
||||
name: Build UI from source
|
||||
command: |
|
||||
cd ui/litellm-dashboard
|
||||
npm run build
|
||||
cp -r out/ ../../litellm/proxy/_experimental/out/
|
||||
# Restructure HTML so extensionless routes work (login.html -> login/index.html)
|
||||
find ../../litellm/proxy/_experimental/out -name '*.html' ! -name 'index.html' | while read -r f; do
|
||||
d="${f%.html}"; mkdir -p "$d"; mv "$f" "$d/index.html"
|
||||
done
|
||||
- run:
|
||||
name: Wait for PostgreSQL
|
||||
command: dockerize -wait tcp://localhost:5432 -timeout 30s
|
||||
- run:
|
||||
name: Push Prisma schema
|
||||
command: prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
|
||||
- run:
|
||||
name: Seed database
|
||||
command: |
|
||||
PGPASSWORD=e2epassword psql -h localhost -p 5432 -U e2euser -d litellm_e2e \
|
||||
-f ui/litellm-dashboard/e2e_tests/fixtures/seed.sql
|
||||
- run:
|
||||
name: Start mock LLM server
|
||||
command: python ui/litellm-dashboard/e2e_tests/fixtures/mock_llm_server/server.py
|
||||
background: true
|
||||
- run:
|
||||
name: Start LiteLLM proxy
|
||||
environment:
|
||||
LITELLM_MASTER_KEY: "sk-1234"
|
||||
MOCK_LLM_URL: "http://127.0.0.1:8090/v1"
|
||||
DISABLE_SCHEMA_UPDATE: "true"
|
||||
SERVER_ROOT_PATH: ""
|
||||
PROXY_LOGOUT_URL: ""
|
||||
command: |
|
||||
python -m litellm.proxy.proxy_cli \
|
||||
--config ui/litellm-dashboard/e2e_tests/fixtures/config.yml \
|
||||
--port 4000
|
||||
background: true
|
||||
- run:
|
||||
name: Wait for proxy to be ready
|
||||
command: |
|
||||
for i in $(seq 1 60); do
|
||||
HTTP_CODE=$(curl -s -o /dev/null -w "%{http_code}" http://127.0.0.1:4000/health -H "Authorization: Bearer sk-1234" 2>/dev/null || true)
|
||||
if [ "$HTTP_CODE" = "200" ]; then
|
||||
echo "Proxy is ready"
|
||||
exit 0
|
||||
fi
|
||||
sleep 2
|
||||
done
|
||||
echo "Proxy failed to start"
|
||||
exit 1
|
||||
- run:
|
||||
name: Run Playwright E2E tests
|
||||
command: |
|
||||
cd ui/litellm-dashboard
|
||||
npx playwright test --config e2e_tests/playwright.config.ts
|
||||
no_output_timeout: 10m
|
||||
- store_artifacts:
|
||||
path: ui/litellm-dashboard/test-results
|
||||
destination: e2e-test-results
|
||||
- store_artifacts:
|
||||
path: ui/litellm-dashboard/playwright-report
|
||||
destination: e2e-playwright-report
|
||||
|
||||
build_docker_database_image:
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
|
|
@ -3207,80 +3261,6 @@ jobs:
|
|||
paths:
|
||||
- litellm-docker-database.tar.zst
|
||||
|
||||
e2e_ui_testing:
|
||||
machine:
|
||||
image: ubuntu-2204:2023.10.1
|
||||
resource_class: large
|
||||
working_directory: ~/project
|
||||
parameters:
|
||||
browser:
|
||||
type: string
|
||||
steps:
|
||||
- checkout
|
||||
- setup_google_dns
|
||||
- attach_workspace:
|
||||
at: ~/project
|
||||
- run:
|
||||
name: Load Docker Database Image
|
||||
command: |
|
||||
zstd -d litellm-docker-database.tar.zst --stdout | docker load
|
||||
docker images | grep litellm-docker-database
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
command: |
|
||||
npm install -D @playwright/test
|
||||
- run:
|
||||
name: Install Playwright Browsers
|
||||
command: |
|
||||
npx playwright install
|
||||
- run:
|
||||
name: Run Docker container
|
||||
command: |
|
||||
docker run -d \
|
||||
-p 4000:4000 \
|
||||
-e DATABASE_URL=$E2E_UI_TEST_DATABASE_URL \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
-e OPENAI_API_KEY=$OPENAI_API_KEY \
|
||||
-e UI_USERNAME="admin" \
|
||||
-e UI_PASSWORD="gm" \
|
||||
-e LITELLM_LICENSE=$LITELLM_LICENSE \
|
||||
--name litellm-docker-database-<< parameters.browser >> \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/simple_config.yaml:/app/config.yaml \
|
||||
litellm-docker-database:ci \
|
||||
--config /app/config.yaml \
|
||||
--port 4000 \
|
||||
--detailed_debug
|
||||
- run:
|
||||
name: Install curl and dockerize
|
||||
command: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y curl
|
||||
sudo wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start outputting logs
|
||||
command: docker logs -f litellm-docker-database-<< parameters.browser >>
|
||||
background: true
|
||||
- run:
|
||||
name: Wait for app to be ready
|
||||
command: dockerize -wait http://localhost:4000 -timeout 5m
|
||||
- run:
|
||||
name: Run Playwright Tests
|
||||
command: |
|
||||
npx playwright test \
|
||||
--project << parameters.browser >> \
|
||||
--config ui/litellm-dashboard/e2e_tests/playwright.config.ts \
|
||||
--reporter=html \
|
||||
--output=test-results
|
||||
no_output_timeout: 15m
|
||||
- store_artifacts:
|
||||
path: test-results
|
||||
destination: playwright-results
|
||||
|
||||
- store_artifacts:
|
||||
path: playwright-report
|
||||
destination: playwright-report
|
||||
|
||||
prisma_schema_sync:
|
||||
machine:
|
||||
|
|
@ -3509,32 +3489,12 @@ workflows:
|
|||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
# - e2e_ui_testing:
|
||||
# name: e2e_ui_testing_chromium
|
||||
# browser: chromium
|
||||
# context: e2e_ui_tests
|
||||
# requires:
|
||||
# - ui_build
|
||||
# - build_docker_database_image
|
||||
# - prisma_schema_sync
|
||||
# filters:
|
||||
# branches:
|
||||
# only:
|
||||
# - main
|
||||
# - /litellm_.*/
|
||||
# - e2e_ui_testing:
|
||||
# name: e2e_ui_testing_firefox
|
||||
# browser: firefox
|
||||
# context: e2e_ui_tests
|
||||
# requires:
|
||||
# - ui_build
|
||||
# - build_docker_database_image
|
||||
# - prisma_schema_sync
|
||||
# filters:
|
||||
# branches:
|
||||
# only:
|
||||
# - main
|
||||
# - /litellm_.*/
|
||||
- e2e_ui_testing:
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- build_and_test:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
|
|
@ -3605,12 +3565,6 @@ workflows:
|
|||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- proxy_e2e_azure_batches_tests:
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- llm_translation_testing:
|
||||
filters:
|
||||
branches:
|
||||
|
|
@ -3729,6 +3683,12 @@ workflows:
|
|||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- redis_caching_unit_tests:
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- upload-coverage:
|
||||
requires:
|
||||
- realtime_translation_testing
|
||||
|
|
@ -3747,6 +3707,7 @@ workflows:
|
|||
- image_gen_testing
|
||||
- logging_testing
|
||||
- audio_testing
|
||||
- redis_caching_unit_tests
|
||||
- langfuse_logging_unit_tests
|
||||
- local_testing_part1
|
||||
- local_testing_part2
|
||||
|
|
|
|||
48
.github/workflows/_test-unit-base.yml
vendored
48
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -27,6 +27,10 @@ on:
|
|||
required: false
|
||||
type: number
|
||||
default: 10
|
||||
artifact-name:
|
||||
description: "Unique name for the coverage artifact (must be unique per run)"
|
||||
required: true
|
||||
type: string
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
@ -93,4 +97,46 @@ jobs:
|
|||
--reruns "${RERUNS}" \
|
||||
--reruns-delay 1 \
|
||||
--dist=loadscope \
|
||||
--durations=20
|
||||
--durations=20 \
|
||||
--cov=litellm \
|
||||
--cov-report=xml:coverage.xml \
|
||||
--cov-config=pyproject.toml
|
||||
|
||||
- name: Save coverage report
|
||||
if: always()
|
||||
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
|
||||
with:
|
||||
name: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
|
||||
path: coverage.xml
|
||||
retention-days: 1
|
||||
|
||||
upload-coverage:
|
||||
name: Upload coverage to Codecov
|
||||
needs: run
|
||||
if: always()
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Download coverage report
|
||||
uses: actions/download-artifact@95815c38cf2ff2164869cbab79da8d1f422bc89e # v4.2.1
|
||||
with:
|
||||
pattern: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
|
||||
path: coverage-reports
|
||||
merge-multiple: true
|
||||
|
||||
- name: Upload to Codecov
|
||||
uses: codecov/codecov-action@75cd11691c0faa626561e295848008c8a7dddffe # v5.5.4
|
||||
with:
|
||||
use_oidc: true
|
||||
directory: coverage-reports
|
||||
root_dir: ${{ github.workspace }}
|
||||
fail_ci_if_error: false
|
||||
|
|
|
|||
71
.github/workflows/_test-unit-services-base.yml
vendored
71
.github/workflows/_test-unit-services-base.yml
vendored
|
|
@ -27,23 +27,17 @@ on:
|
|||
required: false
|
||||
type: number
|
||||
default: 10
|
||||
enable-redis:
|
||||
description: "Pass Redis Cloud credentials to tests via REDIS_HOST/PORT/PASSWORD env vars"
|
||||
required: false
|
||||
type: boolean
|
||||
default: false
|
||||
enable-postgres:
|
||||
description: "Start a local Postgres service container and run Prisma migrations"
|
||||
required: false
|
||||
type: boolean
|
||||
default: false
|
||||
artifact-name:
|
||||
description: "Unique name for the coverage artifact (must be unique per run)"
|
||||
required: false
|
||||
type: string
|
||||
default: "run"
|
||||
secrets:
|
||||
REDIS_HOST:
|
||||
required: false
|
||||
REDIS_PORT:
|
||||
required: false
|
||||
REDIS_PASSWORD:
|
||||
required: false
|
||||
DATABASE_URL:
|
||||
required: false
|
||||
POSTGRES_USER:
|
||||
|
|
@ -61,11 +55,8 @@ jobs:
|
|||
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.
|
||||
# Note: Postgres service container always starts (GHA limitation), so any Redis job
|
||||
# also needs Postgres secrets → uses integration-redis-postgres, not integration-redis.
|
||||
environment: >-
|
||||
${{
|
||||
inputs.enable-redis && 'integration-redis-postgres' ||
|
||||
inputs.enable-postgres && 'integration-postgres' ||
|
||||
''
|
||||
}}
|
||||
|
|
@ -141,9 +132,6 @@ jobs:
|
|||
WORKERS: ${{ inputs.workers }}
|
||||
RERUNS: ${{ inputs.reruns }}
|
||||
DATABASE_URL: ${{ inputs.enable-postgres && secrets.DATABASE_URL || '' }}
|
||||
REDIS_HOST: ${{ inputs.enable-redis && secrets.REDIS_HOST || '' }}
|
||||
REDIS_PORT: ${{ inputs.enable-redis && secrets.REDIS_PORT || '' }}
|
||||
REDIS_PASSWORD: ${{ inputs.enable-redis && secrets.REDIS_PASSWORD || '' }}
|
||||
run: |
|
||||
if [ "${WORKERS}" = "0" ]; then
|
||||
poetry run pytest ${TEST_PATH:?} \
|
||||
|
|
@ -151,7 +139,10 @@ jobs:
|
|||
--maxfail="${MAX_FAILURES}" \
|
||||
--reruns "${RERUNS}" \
|
||||
--reruns-delay 1 \
|
||||
--durations=20
|
||||
--durations=20 \
|
||||
--cov=litellm \
|
||||
--cov-report=xml:coverage.xml \
|
||||
--cov-config=pyproject.toml
|
||||
else
|
||||
poetry run pytest ${TEST_PATH:?} \
|
||||
--tb=short -vv \
|
||||
|
|
@ -160,5 +151,47 @@ jobs:
|
|||
--reruns "${RERUNS}" \
|
||||
--reruns-delay 1 \
|
||||
--dist=loadscope \
|
||||
--durations=20
|
||||
--durations=20 \
|
||||
--cov=litellm \
|
||||
--cov-report=xml:coverage.xml \
|
||||
--cov-config=pyproject.toml
|
||||
fi
|
||||
|
||||
- name: Save coverage report
|
||||
if: always()
|
||||
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
|
||||
with:
|
||||
name: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
|
||||
path: coverage.xml
|
||||
retention-days: 1
|
||||
|
||||
upload-coverage:
|
||||
name: Upload coverage to Codecov
|
||||
needs: run
|
||||
if: always()
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Download coverage report
|
||||
uses: actions/download-artifact@95815c38cf2ff2164869cbab79da8d1f422bc89e # v4.2.1
|
||||
with:
|
||||
pattern: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
|
||||
path: coverage-reports
|
||||
merge-multiple: true
|
||||
|
||||
- name: Upload to Codecov
|
||||
uses: codecov/codecov-action@75cd11691c0faa626561e295848008c8a7dddffe # v5.5.4
|
||||
with:
|
||||
use_oidc: true
|
||||
directory: coverage-reports
|
||||
root_dir: ${{ github.workspace }}
|
||||
fail_ci_if_error: false
|
||||
|
|
|
|||
16
.github/workflows/create-release.yml
vendored
16
.github/workflows/create-release.yml
vendored
|
|
@ -48,7 +48,21 @@ jobs:
|
|||
const cosignSection = [
|
||||
`## Verify Docker Image Signature`,
|
||||
``,
|
||||
`All LiteLLM Docker images are signed with [cosign](https://docs.sigstore.dev/cosign/overview/). To verify the integrity of an image before deploying:`,
|
||||
`All LiteLLM Docker images are signed with [cosign](https://docs.sigstore.dev/cosign/overview/). Every release is signed with the same key introduced in [commit \`0112e53\`](https://github.com/BerriAI/litellm/commit/0112e53046018d726492c814b3644b7d376029d0).`,
|
||||
``,
|
||||
`**Verify using the pinned commit hash (recommended):**`,
|
||||
``,
|
||||
`A commit hash is cryptographically immutable, so this is the strongest way to ensure you are using the original signing key:`,
|
||||
``,
|
||||
'```bash',
|
||||
`cosign verify \\`,
|
||||
` --key https://raw.githubusercontent.com/BerriAI/litellm/0112e53046018d726492c814b3644b7d376029d0/cosign.pub \\`,
|
||||
` ghcr.io/berriai/litellm:${tag}`,
|
||||
'```',
|
||||
``,
|
||||
`**Verify using the release tag (convenience):**`,
|
||||
``,
|
||||
`Tags are protected in this repository and resolve to the same key. This option is easier to read but relies on tag protection rules:`,
|
||||
``,
|
||||
'```bash',
|
||||
`cosign verify \\`,
|
||||
|
|
|
|||
0
.github/workflows/run_llm_translation_tests.py
vendored
Executable file → Normal file
0
.github/workflows/run_llm_translation_tests.py
vendored
Executable file → Normal file
214
.github/workflows/test-litellm-matrix.yml
vendored
214
.github/workflows/test-litellm-matrix.yml
vendored
|
|
@ -1,214 +0,0 @@
|
|||
name: LiteLLM Unit Tests (Matrix)
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches: [main]
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
# Cancel in-progress runs for the same PR
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20 # Increased from 15 to 20
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
test-group:
|
||||
# tests/test_litellm split by subdirectory (~560 files total)
|
||||
# Vertex AI tests separated for better isolation (prevent auth/env pollution)
|
||||
- name: "llms-vertex"
|
||||
path: "tests/test_litellm/llms/vertex_ai"
|
||||
workers: 1
|
||||
reruns: 2
|
||||
- name: "llms-other"
|
||||
path: "tests/test_litellm/llms --ignore=tests/test_litellm/llms/vertex_ai"
|
||||
workers: 2
|
||||
reruns: 2
|
||||
# tests/test_litellm/proxy split by subdirectory (~180 files total)
|
||||
- name: "proxy-guardrails"
|
||||
path: "tests/test_litellm/proxy/guardrails tests/test_litellm/proxy/management_endpoints tests/test_litellm/proxy/management_helpers"
|
||||
workers: 2
|
||||
reruns: 2
|
||||
- name: "proxy-core"
|
||||
path: "tests/test_litellm/proxy/auth tests/test_litellm/proxy/client tests/test_litellm/proxy/db tests/test_litellm/proxy/hooks tests/test_litellm/proxy/policy_engine"
|
||||
workers: 2
|
||||
reruns: 2
|
||||
- name: "proxy-misc"
|
||||
path: "tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py"
|
||||
workers: 2
|
||||
reruns: 2
|
||||
- name: "integrations"
|
||||
path: "tests/test_litellm/integrations"
|
||||
workers: 2
|
||||
reruns: 3 # Integration tests tend to be flakier
|
||||
- name: "core-utils"
|
||||
path: "tests/test_litellm/litellm_core_utils"
|
||||
workers: 2
|
||||
reruns: 1
|
||||
- name: "other-1"
|
||||
# responses (5942) + caching (1723) + types (819) ≈ 8.5k lines
|
||||
path: "tests/test_litellm/responses tests/test_litellm/caching tests/test_litellm/types"
|
||||
workers: 2
|
||||
reruns: 2
|
||||
- name: "other-2"
|
||||
# enterprise (3062) + google_genai (2511) + router_utils (1982) ≈ 7.6k lines
|
||||
path: "tests/test_litellm/enterprise tests/test_litellm/google_genai tests/test_litellm/router_utils"
|
||||
workers: 2
|
||||
reruns: 2
|
||||
- name: "other-3"
|
||||
# remaining dirs ≈ 8.0k lines
|
||||
path: "tests/test_litellm/router_strategy tests/test_litellm/secret_managers tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/test_litellm/experimental_mcp_client tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/vector_stores"
|
||||
workers: 2
|
||||
reruns: 2
|
||||
- name: "root"
|
||||
path: "tests/test_litellm/test_*.py"
|
||||
workers: 2
|
||||
reruns: 2
|
||||
# tests/proxy_unit_tests split alphabetically (~48 files total)
|
||||
- name: "proxy-unit-a1"
|
||||
# test_[a-j]*.py: jwt (1564) + auth_checks (978) + google_gemini (478) + e2e_pod_lock (437) + rest
|
||||
path: "tests/proxy_unit_tests/test_[a-j]*.py"
|
||||
workers: 2
|
||||
reruns: 1
|
||||
- name: "proxy-unit-a2"
|
||||
# test_[k-o]*.py: key_generate_prisma (4346) + key_generate_dynamodb + models_fallback
|
||||
path: "tests/proxy_unit_tests/test_[k-o]*.py"
|
||||
workers: 2
|
||||
reruns: 1
|
||||
- name: "proxy-unit-b1"
|
||||
# lighter config/utility proxy tests (prisma, project, prompt, proxy_[c-r]*)
|
||||
path: "tests/proxy_unit_tests/test_prisma*.py tests/proxy_unit_tests/test_project*.py tests/proxy_unit_tests/test_prompt*.py tests/proxy_unit_tests/test_proxy_[c-r]*.py"
|
||||
workers: 2
|
||||
reruns: 1
|
||||
- name: "proxy-unit-b2"
|
||||
# proxy_server.py alone (2750 lines) - isolated to avoid blocking smaller tests
|
||||
path: "tests/proxy_unit_tests/test_proxy_server.py"
|
||||
workers: 2
|
||||
reruns: 1
|
||||
- name: "proxy-unit-b3"
|
||||
# proxy_server_* (618) + proxy_setting_guardrails (71) - smaller server-related tests
|
||||
path: "tests/proxy_unit_tests/test_proxy_server_*.py tests/proxy_unit_tests/test_proxy_setting_guardrails.py"
|
||||
workers: 2
|
||||
reruns: 1
|
||||
- name: "proxy-unit-b4"
|
||||
# proxy_utils.py alone (2339 lines) - isolated to avoid blocking token counter
|
||||
path: "tests/proxy_unit_tests/test_proxy_utils.py"
|
||||
workers: 2
|
||||
reruns: 1
|
||||
- name: "proxy-unit-b5"
|
||||
# proxy_token_counter (1279) - runs independently from utils
|
||||
path: "tests/proxy_unit_tests/test_proxy_token_counter.py"
|
||||
workers: 2
|
||||
reruns: 1
|
||||
- name: "proxy-unit-b6"
|
||||
# test_[r-t]*.py: response_polling (1399) + search_api_logging (202) + server_root (64) + skills_db (261) + realtime_cache (62)
|
||||
path: "tests/proxy_unit_tests/test_[r-t]*.py"
|
||||
workers: 2
|
||||
reruns: 1
|
||||
- name: "proxy-unit-b7"
|
||||
# test_[u-z]*.py: user_api_key_auth (1136) + zero_cost (590) + update_spend (305) + unit_test_* (206) + ui_path (157)
|
||||
path: "tests/proxy_unit_tests/test_[u-z]*.py"
|
||||
workers: 2
|
||||
reruns: 1
|
||||
|
||||
name: test (${{ matrix.test-group.name }})
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Install Poetry
|
||||
run: pip install 'poetry==2.3.2'
|
||||
|
||||
- name: Cache Poetry dependencies
|
||||
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.0.0
|
||||
with:
|
||||
path: |
|
||||
~/.cache/pypoetry
|
||||
~/.cache/pip
|
||||
.venv
|
||||
key: ${{ runner.os }}-poetry-${{ hashFiles('poetry.lock') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-poetry-
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
poetry config virtualenvs.in-project true
|
||||
poetry install --with dev,proxy-dev --extras "proxy semantic-router"
|
||||
# pytest-rerunfailures and pytest-xdist are in pyproject.toml dev dependencies
|
||||
poetry run pip install google-genai==1.22.0 \
|
||||
google-cloud-aiplatform==1.115.0 fastapi-offline==1.7.3 python-multipart==0.0.22 openapi-core==0.23.0
|
||||
|
||||
- name: Setup litellm-enterprise
|
||||
run: |
|
||||
poetry run pip install --force-reinstall --no-deps -e enterprise/
|
||||
|
||||
- name: Generate Prisma client
|
||||
env:
|
||||
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
|
||||
run: |
|
||||
poetry run pip install nodejs-wheel-binaries==24.13.1
|
||||
poetry run prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Run tests - ${{ matrix.test-group.name }}
|
||||
run: |
|
||||
poetry run pytest ${{ matrix.test-group.path }} \
|
||||
--tb=short -vv \
|
||||
--maxfail=10 \
|
||||
-n ${{ matrix.test-group.workers }} \
|
||||
--reruns ${{ matrix.test-group.reruns }} \
|
||||
--reruns-delay 1 \
|
||||
--dist=loadscope \
|
||||
--durations=20 \
|
||||
--cov=litellm \
|
||||
--cov-report=xml:coverage-${{ matrix.test-group.name }}.xml \
|
||||
--cov-config=pyproject.toml
|
||||
|
||||
- name: Save coverage report
|
||||
if: always()
|
||||
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
|
||||
with:
|
||||
name: coverage-${{ matrix.test-group.name }}
|
||||
path: coverage-${{ matrix.test-group.name }}.xml
|
||||
retention-days: 1
|
||||
|
||||
upload-coverage:
|
||||
name: Upload coverage to Codecov
|
||||
needs: test
|
||||
if: always()
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write # Required for OIDC tokenless upload
|
||||
pull-requests: write # Required for Codecov PR comments
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
|
||||
- name: Download all coverage reports
|
||||
uses: actions/download-artifact@95815c38cf2ff2164869cbab79da8d1f422bc89e # v4.2.1
|
||||
with:
|
||||
pattern: coverage-*
|
||||
path: coverage-reports
|
||||
merge-multiple: true
|
||||
|
||||
- name: Upload to Codecov
|
||||
uses: codecov/codecov-action@aa56896cf108bd10b5eb883cd1d24196da57f695 # v5.5.4
|
||||
with:
|
||||
use_oidc: true
|
||||
directory: coverage-reports
|
||||
root_dir: ${{ github.workspace }}
|
||||
fail_ci_if_error: false
|
||||
|
|
@ -1,97 +0,0 @@
|
|||
name: Proxy E2E Azure Batches Tests
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches: [main]
|
||||
workflow_dispatch:
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
proxy_e2e_azure_batches_tests:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
services:
|
||||
postgres:
|
||||
image: postgres:15
|
||||
env:
|
||||
POSTGRES_USER: llmproxy
|
||||
POSTGRES_PASSWORD: dbpassword9090
|
||||
POSTGRES_DB: litellm
|
||||
ports:
|
||||
- 5432:5432
|
||||
options: >-
|
||||
--health-cmd pg_isready
|
||||
--health-interval 10s
|
||||
--health-timeout 5s
|
||||
--health-retries 5
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Install Poetry
|
||||
run: pip install 'poetry==2.3.2'
|
||||
|
||||
- name: Cache Poetry dependencies
|
||||
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.0.0
|
||||
with:
|
||||
path: |
|
||||
~/.cache/pypoetry
|
||||
~/.cache/pip
|
||||
.venv
|
||||
key: ${{ runner.os }}-poetry-e2e-batches-${{ hashFiles('poetry.lock') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-poetry-e2e-batches-
|
||||
${{ runner.os }}-poetry-
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
poetry config virtualenvs.in-project true
|
||||
poetry install --with dev,proxy-dev --extras "proxy"
|
||||
poetry run pip install psycopg2-binary==2.9.11 uvicorn==0.42.0 fastapi==0.135.2 httpx==0.28.1 tenacity==9.1.4
|
||||
|
||||
- name: Setup litellm-enterprise
|
||||
run: |
|
||||
poetry run pip install --force-reinstall --no-deps -e enterprise/
|
||||
|
||||
- name: Generate Prisma client
|
||||
env:
|
||||
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
|
||||
run: |
|
||||
poetry run pip install nodejs-wheel-binaries==24.13.1
|
||||
poetry run prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Run Prisma migrations
|
||||
env:
|
||||
DATABASE_URL: postgresql://llmproxy:dbpassword9090@localhost:5432/litellm
|
||||
run: |
|
||||
cd litellm/proxy
|
||||
poetry run prisma migrate deploy --schema schema.prisma
|
||||
cd ../..
|
||||
|
||||
- name: Run Azure Batch E2E Tests
|
||||
env:
|
||||
DATABASE_URL: postgresql://llmproxy:dbpassword9090@localhost:5432/litellm
|
||||
USE_LOCAL_LITELLM: "true"
|
||||
USE_MOCK_MODELS: "true"
|
||||
USE_STATE_TRACKER: "true"
|
||||
LITELLM_LOG: DEBUG
|
||||
run: |
|
||||
poetry run pytest tests/proxy_e2e_azure_batches_tests/test_proxy_e2e_azure_batches.py \
|
||||
-vv -s -k "test_e2e_managed_batch" \
|
||||
--tb=short \
|
||||
--maxfail=3 \
|
||||
--durations=10
|
||||
38
.github/workflows/test-unit-caching-redis.yml
vendored
38
.github/workflows/test-unit-caching-redis.yml
vendored
|
|
@ -1,38 +0,0 @@
|
|||
name: "Unit Tests: Caching (Redis)"
|
||||
|
||||
# Uses cloud Redis credentials — only runs on trusted branches, not PRs.
|
||||
# This prevents external PRs from accessing Redis credentials.
|
||||
on:
|
||||
push:
|
||||
branches: [main, "litellm_*"]
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
caching-redis:
|
||||
uses: ./.github/workflows/_test-unit-services-base.yml
|
||||
with:
|
||||
# Redis-only tests that do NOT require provider API keys.
|
||||
# Tests needing API keys (test_caching.py, test_caching_ssl.py, test_prometheus_service.py,
|
||||
# test_router_caching.py) are in Phase 3 integration workflows.
|
||||
test-path: >-
|
||||
tests/local_testing/test_dual_cache.py
|
||||
tests/local_testing/test_redis_batch_optimizations.py
|
||||
tests/local_testing/test_router_utils.py
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
enable-redis: true
|
||||
enable-postgres: false
|
||||
secrets:
|
||||
REDIS_HOST: ${{ secrets.REDIS_HOST }}
|
||||
REDIS_PORT: ${{ secrets.REDIS_PORT }}
|
||||
REDIS_PASSWORD: ${{ secrets.REDIS_PASSWORD }}
|
||||
DATABASE_URL: ${{ secrets.DATABASE_URL }}
|
||||
POSTGRES_USER: ${{ secrets.POSTGRES_USER }}
|
||||
POSTGRES_PASSWORD: ${{ secrets.POSTGRES_PASSWORD }}
|
||||
3
.github/workflows/test-unit-core-utils.yml
vendored
3
.github/workflows/test-unit-core-utils.yml
vendored
|
|
@ -6,6 +6,8 @@ on:
|
|||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -18,3 +20,4 @@ jobs:
|
|||
test-path: "tests/test_litellm/litellm_core_utils"
|
||||
workers: 2
|
||||
reruns: 1
|
||||
artifact-name: core-utils
|
||||
|
|
|
|||
|
|
@ -6,6 +6,8 @@ on:
|
|||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -22,3 +24,4 @@ jobs:
|
|||
tests/test_litellm/router_strategy
|
||||
workers: 2
|
||||
reruns: 2
|
||||
artifact-name: enterprise-routing
|
||||
|
|
|
|||
3
.github/workflows/test-unit-integrations.yml
vendored
3
.github/workflows/test-unit-integrations.yml
vendored
|
|
@ -6,6 +6,8 @@ on:
|
|||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -18,3 +20,4 @@ jobs:
|
|||
test-path: "tests/test_litellm/integrations"
|
||||
workers: 2
|
||||
reruns: 3
|
||||
artifact-name: integrations
|
||||
|
|
|
|||
10
.github/workflows/test-unit-llm-providers.yml
vendored
10
.github/workflows/test-unit-llm-providers.yml
vendored
|
|
@ -14,16 +14,26 @@ concurrency:
|
|||
jobs:
|
||||
vertex-ai:
|
||||
name: Vertex AI
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: "tests/test_litellm/llms/vertex_ai"
|
||||
workers: 1
|
||||
reruns: 2
|
||||
artifact-name: llm-vertex-ai
|
||||
|
||||
other-providers:
|
||||
name: All Other Providers
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: "tests/test_litellm/llms --ignore=tests/test_litellm/llms/vertex_ai"
|
||||
workers: 2
|
||||
reruns: 2
|
||||
artifact-name: llm-other-providers
|
||||
|
|
|
|||
3
.github/workflows/test-unit-misc.yml
vendored
3
.github/workflows/test-unit-misc.yml
vendored
|
|
@ -6,6 +6,8 @@ on:
|
|||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -29,3 +31,4 @@ jobs:
|
|||
tests/test_litellm/test_*.py
|
||||
workers: 2
|
||||
reruns: 2
|
||||
artifact-name: misc
|
||||
|
|
|
|||
3
.github/workflows/test-unit-proxy-auth.yml
vendored
3
.github/workflows/test-unit-proxy-auth.yml
vendored
|
|
@ -6,6 +6,8 @@ on:
|
|||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -18,3 +20,4 @@ jobs:
|
|||
test-path: "tests/test_litellm/proxy/auth tests/test_litellm/proxy/hooks tests/test_litellm/proxy/policy_engine tests/test_litellm/proxy/client"
|
||||
workers: 2
|
||||
reruns: 2
|
||||
artifact-name: proxy-auth
|
||||
|
|
|
|||
6
.github/workflows/test-unit-proxy-db.yml
vendored
6
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -14,6 +14,10 @@ concurrency:
|
|||
|
||||
jobs:
|
||||
proxy-db:
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
|
|
@ -37,8 +41,8 @@ jobs:
|
|||
workers: ${{ matrix.workers }}
|
||||
reruns: 2
|
||||
timeout-minutes: ${{ matrix.timeout }}
|
||||
enable-redis: false
|
||||
enable-postgres: true
|
||||
artifact-name: proxy-db-${{ matrix.test-group }}
|
||||
secrets:
|
||||
DATABASE_URL: ${{ secrets.DATABASE_URL }}
|
||||
POSTGRES_USER: ${{ secrets.POSTGRES_USER }}
|
||||
|
|
|
|||
|
|
@ -6,6 +6,8 @@ on:
|
|||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -33,3 +35,4 @@ jobs:
|
|||
tests/test_litellm/proxy/ui_crud_endpoints
|
||||
workers: 2
|
||||
reruns: 2
|
||||
artifact-name: proxy-endpoints
|
||||
|
|
|
|||
3
.github/workflows/test-unit-proxy-infra.yml
vendored
3
.github/workflows/test-unit-proxy-infra.yml
vendored
|
|
@ -6,6 +6,8 @@ on:
|
|||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -26,3 +28,4 @@ jobs:
|
|||
tests/test_litellm/proxy/test_*.py
|
||||
workers: 2
|
||||
reruns: 2
|
||||
artifact-name: proxy-infra
|
||||
|
|
|
|||
|
|
@ -6,6 +6,8 @@ on:
|
|||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -18,3 +20,4 @@ jobs:
|
|||
test-path: "tests/test_litellm/responses tests/test_litellm/caching tests/test_litellm/types"
|
||||
workers: 2
|
||||
reruns: 2
|
||||
artifact-name: responses-caching-types
|
||||
|
|
|
|||
4
.github/workflows/test-unit-security.yml
vendored
4
.github/workflows/test-unit-security.yml
vendored
|
|
@ -7,6 +7,8 @@ on:
|
|||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.ref }}
|
||||
|
|
@ -20,8 +22,8 @@ jobs:
|
|||
workers: 1
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
enable-redis: false
|
||||
enable-postgres: true
|
||||
artifact-name: security
|
||||
secrets:
|
||||
DATABASE_URL: ${{ secrets.DATABASE_URL }}
|
||||
POSTGRES_USER: ${{ secrets.POSTGRES_USER }}
|
||||
|
|
|
|||
12
.trivyignore
12
.trivyignore
|
|
@ -1,12 +0,0 @@
|
|||
# LiteLLM Trivy Ignore File
|
||||
# CVEs listed here are temporarily allowlisted pending fixes
|
||||
|
||||
# Next.js vulnerabilities in UI dashboard (next@14.2.35)
|
||||
# Allowlisted: 2026-01-31, 7-day fix timeline
|
||||
# Fix: Upgrade to Next.js 15.5.10+ or 16.1.5+
|
||||
|
||||
# HIGH: DoS via request deserialization
|
||||
GHSA-h25m-26qc-wcjf
|
||||
|
||||
# MEDIUM: Image Optimizer DoS
|
||||
CVE-2025-59471
|
||||
|
|
@ -254,7 +254,7 @@ See `CLAUDE.md` and the `Makefile` for standard commands. Key notes:
|
|||
- `openapi-core` must be installed (`poetry run pip install openapi-core`) for the OpenAPI compliance tests in `tests/test_litellm/interactions/`.
|
||||
- The `--timeout` pytest flag is NOT available; don't pass it.
|
||||
- Unit tests: `poetry run pytest tests/test_litellm/ -x -vv -n 4`
|
||||
- Black `--check` may report pre-existing formatting issues; this does not block test runs.
|
||||
- **Before committing, always run `poetry run black .` to format your code.** Black formatting is enforced in CI.
|
||||
- If `poetry install` fails with "pyproject.toml changed significantly since poetry.lock was last generated", run `poetry lock` first to regenerate the lock file.
|
||||
|
||||
### Lint
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ This file provides guidance to Claude Code (claude.ai/code) when working with co
|
|||
- `make format` - Apply Black code formatting
|
||||
- `make lint-ruff` - Run Ruff linting only
|
||||
- `make lint-mypy` - Run MyPy type checking only
|
||||
- **Before committing, always run `poetry run black .` to format your code.** Black formatting is enforced in CI.
|
||||
|
||||
### Single Test Files
|
||||
- `poetry run pytest tests/path/to/test_file.py -v` - Run specific test file
|
||||
|
|
|
|||
|
|
@ -149,6 +149,19 @@ Apply formatting (auto-fixes issues):
|
|||
make format
|
||||
```
|
||||
|
||||
> **Black formatting is enforced in CI.** All PRs must pass the Black formatting check.
|
||||
>
|
||||
> - **AI coding agents** (Claude Code, Copilot, Cursor, etc.): `AGENTS.md` and `CLAUDE.md` instruct agents to run `poetry run black .` before committing.
|
||||
> - **VS Code users**: Install the [Black Formatter extension](https://marketplace.visualstudio.com/items?itemName=ms-python.black-formatter) and enable format-on-save:
|
||||
> ```json
|
||||
> {
|
||||
> "[python]": {
|
||||
> "editor.defaultFormatter": "ms-python.black-formatter",
|
||||
> "editor.formatOnSave": true
|
||||
> }
|
||||
> }
|
||||
> ```
|
||||
|
||||
### CI Compatibility
|
||||
|
||||
To ensure your changes will pass CI, run the exact same checks locally:
|
||||
|
|
|
|||
26
README.md
26
README.md
|
|
@ -404,6 +404,32 @@ Support for more providers. Missing a provider or LLM Platform, raise a [feature
|
|||
2. Install dependencies `npm install`
|
||||
3. Run `npm run dev` to start the dashboard
|
||||
|
||||
# Verify Docker Image Signatures
|
||||
|
||||
All LiteLLM Docker images published to GHCR are signed with [cosign](https://docs.sigstore.dev/cosign/overview/). Every release is signed with the same key introduced in [commit `0112e53`](https://github.com/BerriAI/litellm/commit/0112e53046018d726492c814b3644b7d376029d0).
|
||||
|
||||
**Verify using the pinned commit hash (recommended):**
|
||||
|
||||
A commit hash is cryptographically immutable, so this is the strongest way to ensure you are using the original signing key:
|
||||
|
||||
```bash
|
||||
cosign verify \
|
||||
--key https://raw.githubusercontent.com/BerriAI/litellm/0112e53046018d726492c814b3644b7d376029d0/cosign.pub \
|
||||
ghcr.io/berriai/litellm:<release-tag>
|
||||
```
|
||||
|
||||
**Verify using a release tag (convenience):**
|
||||
|
||||
Tags are protected in this repository and resolve to the same key. This option is easier to read but relies on tag protection rules:
|
||||
|
||||
```bash
|
||||
cosign verify \
|
||||
--key https://raw.githubusercontent.com/BerriAI/litellm/<release-tag>/cosign.pub \
|
||||
ghcr.io/berriai/litellm:<release-tag>
|
||||
```
|
||||
|
||||
Replace `<release-tag>` with the version you are deploying (e.g. `v1.83.0-stable`).
|
||||
|
||||
# Enterprise
|
||||
For companies that need better security, user management and professional support
|
||||
|
||||
|
|
|
|||
|
|
@ -1,36 +0,0 @@
|
|||
ignore:
|
||||
- vulnerability: CVE-2026-22184
|
||||
reason: no fixed zlib package is available yet in the Wolfi repositories, so this is ignored temporarily until an upstream release exists
|
||||
# Wolfi base image: Python 3.13 and Node from apk have no fixed builds in Wolfi yet / not applicable
|
||||
- vulnerability: CVE-2025-55130
|
||||
reason: Node in Wolfi apk; only used for Admin UI build/prisma
|
||||
- vulnerability: CVE-2025-59465
|
||||
reason: Node in Wolfi apk; only used for Admin UI build/prisma
|
||||
- vulnerability: CVE-2025-55131
|
||||
reason: Node in Wolfi apk; only used for Admin UI build/prisma
|
||||
- vulnerability: CVE-2025-59466
|
||||
reason: Node in Wolfi apk; only used for Admin UI build/prisma
|
||||
- vulnerability: CVE-2026-21637
|
||||
reason: Node in Wolfi apk; only used for Admin UI build/prisma
|
||||
- vulnerability: CVE-2025-55132
|
||||
reason: Node in Wolfi apk; only used for Admin UI build/prisma
|
||||
- vulnerability: GHSA-hx9q-6w63-j58v
|
||||
reason: orjson dumps recursion; allowlisted
|
||||
- vulnerability: GHSA-73rr-hh4g-fpgx
|
||||
reason: diff npm transitive dep; override in package.json, allowlisted
|
||||
- vulnerability: CVE-2026-0865
|
||||
reason: Python 3.13 in Wolfi base; no fixed apk build yet
|
||||
- vulnerability: CVE-2025-15282
|
||||
reason: Python 3.13 in Wolfi base; no fixed apk build yet
|
||||
- vulnerability: CVE-2026-0672
|
||||
reason: Python 3.13 in Wolfi base; no fixed apk build yet
|
||||
- vulnerability: CVE-2025-15366
|
||||
reason: Python 3.13 in Wolfi base; no fixed apk build yet
|
||||
- vulnerability: CVE-2025-15367
|
||||
reason: Python 3.13 in Wolfi base; no fixed apk build yet
|
||||
- vulnerability: CVE-2025-11468
|
||||
reason: Python 3.13 in Wolfi base; no fixed apk build yet
|
||||
- vulnerability: CVE-2025-12781
|
||||
reason: Python 3.13 in Wolfi base; no fixed apk build yet
|
||||
- vulnerability: CVE-2026-1299
|
||||
reason: Python 3.13 in Wolfi base; no fixed apk build yet
|
||||
|
|
@ -1,261 +0,0 @@
|
|||
#!/bin/bash
|
||||
|
||||
# Security Scans Script for LiteLLM
|
||||
# This script runs comprehensive security scans including Trivy and Grype
|
||||
|
||||
set -e
|
||||
|
||||
echo "Starting security scans for LiteLLM..."
|
||||
|
||||
# Function to install Trivy and required tools
|
||||
install_trivy() {
|
||||
echo "Installing Trivy and required tools..."
|
||||
TRIVY_VERSION="0.35.0"
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y wget jq curl bsdmainutils
|
||||
wget -qO trivy.deb "https://github.com/aquasecurity/trivy/releases/download/v${TRIVY_VERSION}/trivy_${TRIVY_VERSION}_Linux-64bit.deb"
|
||||
sudo dpkg -i trivy.deb
|
||||
rm trivy.deb
|
||||
echo "Trivy ${TRIVY_VERSION} installed successfully"
|
||||
}
|
||||
|
||||
# Function to install Grype
|
||||
install_grype() {
|
||||
echo "Installing Grype..."
|
||||
curl -sSfL https://raw.githubusercontent.com/anchore/grype/main/install.sh | sudo sh -s -- -b /usr/local/bin
|
||||
echo "Grype installed successfully"
|
||||
}
|
||||
|
||||
# Function to install ggshield
|
||||
install_ggshield() {
|
||||
echo "Installing ggshield..."
|
||||
pip3 install --upgrade pip
|
||||
pip3 install ggshield
|
||||
echo "ggshield installed successfully"
|
||||
}
|
||||
|
||||
# # Function to run secret detection scans
|
||||
# run_secret_detection() {
|
||||
# echo "Running secret detection scans..."
|
||||
|
||||
# if ! command -v ggshield &> /dev/null; then
|
||||
# install_ggshield
|
||||
# fi
|
||||
|
||||
# # Check if GITGUARDIAN_API_KEY is set (required for CI/CD)
|
||||
# if [ -z "$GITGUARDIAN_API_KEY" ]; then
|
||||
# echo "Warning: GITGUARDIAN_API_KEY environment variable is not set."
|
||||
# echo "ggshield requires a GitGuardian API key to scan for secrets."
|
||||
# echo "Please set GITGUARDIAN_API_KEY in your CI/CD environment variables."
|
||||
# exit 1
|
||||
# fi
|
||||
|
||||
# echo "Scanning codebase for secrets..."
|
||||
# echo "Note: Large codebases may take several minutes due to API rate limits (50 requests/minute on free plan)"
|
||||
# echo "ggshield will automatically handle rate limits and retry as needed."
|
||||
# echo "Binary files, cache files, and build artifacts are excluded via .gitguardian.yaml"
|
||||
|
||||
# # Use --recursive for directory scanning and auto-confirm if prompted
|
||||
# # .gitguardian.yaml will automatically exclude binary files, wheel files, etc.
|
||||
# # GITGUARDIAN_API_KEY environment variable will be used for authentication
|
||||
# echo y | ggshield secret scan path . --recursive || {
|
||||
# echo ""
|
||||
# echo "=========================================="
|
||||
# echo "ERROR: Secret Detection Failed"
|
||||
# echo "=========================================="
|
||||
# echo "ggshield has detected secrets in the codebase."
|
||||
# echo "Please review discovered secrets above, revoke any actively used secrets"
|
||||
# echo "from underlying systems and make changes to inject secrets dynamically at runtime."
|
||||
# echo ""
|
||||
# echo "For more information, see: https://docs.gitguardian.com/secrets-detection/"
|
||||
# echo "=========================================="
|
||||
# echo ""
|
||||
# exit 1
|
||||
# }
|
||||
|
||||
# echo "Secret detection scans completed successfully"
|
||||
# }
|
||||
|
||||
# Function to run Trivy scans
|
||||
run_trivy_scans() {
|
||||
echo "Running Trivy scans..."
|
||||
|
||||
echo "Scanning LiteLLM Docs..."
|
||||
trivy fs --ignorefile .trivyignore --scanners vuln --dependency-tree --exit-code 1 --severity HIGH,CRITICAL,MEDIUM ./docs/
|
||||
|
||||
echo "Scanning LiteLLM UI..."
|
||||
trivy fs --ignorefile .trivyignore --scanners vuln --dependency-tree --exit-code 1 --severity HIGH,CRITICAL,MEDIUM ./ui/
|
||||
|
||||
echo "Trivy scans completed successfully"
|
||||
}
|
||||
|
||||
# Function to build and scan Docker images with Grype
|
||||
run_grype_scans() {
|
||||
echo "Running Grype scans..."
|
||||
|
||||
# Temporarily add wheel files to .dockerignore for security scans
|
||||
echo "Temporarily modifying .dockerignore to exclude problematic wheel files..."
|
||||
cp .dockerignore .dockerignore.backup 2>/dev/null || touch .dockerignore.backup
|
||||
echo "/*.whl" >> .dockerignore
|
||||
|
||||
# Build and scan Dockerfile.database
|
||||
echo "Building and scanning Dockerfile.database..."
|
||||
docker build --no-cache -t litellm-database:latest -f ./docker/Dockerfile.database .
|
||||
grype litellm-database:latest --config ci_cd/.grype.yaml --fail-on critical
|
||||
|
||||
# Build and scan main Dockerfile
|
||||
echo "Building and scanning main Dockerfile..."
|
||||
docker build --no-cache -t litellm:latest .
|
||||
grype litellm:latest --config ci_cd/.grype.yaml --fail-on critical
|
||||
|
||||
# Restore original .dockerignore
|
||||
echo "Restoring original .dockerignore..."
|
||||
mv .dockerignore.backup .dockerignore
|
||||
|
||||
# Scan the locally built LiteLLM image for vulnerabilities with CVSS >= 4.0
|
||||
echo "Scanning locally built LiteLLM image for high-severity vulnerabilities..."
|
||||
echo "Using locally built image: litellm:latest"
|
||||
|
||||
# Allowlist of CVEs to be ignored in failure threshold/reporting
|
||||
# - CVE-2025-8869: Not applicable on Python >=3.13 (PEP 706 implemented); pip fallback unused; no OS-level fix
|
||||
# - GHSA-4xh5-x5gv-qwph: GitHub Security Advisory alias for CVE-2025-8869
|
||||
# - GHSA-5j98-mcp5-4vw2: glob CLI command injection via -c/--cmd; glob CLI is not used in the litellm runtime image,
|
||||
# and the vulnerable versions are pulled in only via OS-level/node tooling outside of our application code
|
||||
ALLOWED_CVES=(
|
||||
"CVE-2025-8869"
|
||||
"GHSA-4xh5-x5gv-qwph"
|
||||
"CVE-2025-8291" # no fix available as of Oct 11, 2025
|
||||
"GHSA-5j98-mcp5-4vw2"
|
||||
"CVE-2025-13836" # Python 3.13 HTTP response reading OOM/DoS - no fix available in base image
|
||||
"CVE-2025-12084" # Python 3.13 xml.dom.minidom quadratic algorithm - no fix available in base image
|
||||
"CVE-2025-60876" # BusyBox wget HTTP request splitting - no fix available in Chainguard Wolfi base image
|
||||
"CVE-2026-0861" # Wolfi glibc still flagged even on 2.42-r5; upstream patched build unavailable yet
|
||||
"CVE-2010-4756" # glibc glob DoS - awaiting patched Wolfi glibc build
|
||||
"CVE-2019-1010022" # glibc stack guard bypass - awaiting patched Wolfi glibc build
|
||||
"CVE-2019-1010023" # glibc ldd remap issue - awaiting patched Wolfi glibc build
|
||||
"CVE-2019-1010024" # glibc ASLR mitigation bypass - awaiting patched Wolfi glibc build
|
||||
"CVE-2019-1010025" # glibc pthread heap address leak - awaiting patched Wolfi glibc build
|
||||
"CVE-2026-22184" # zlib untgz buffer overflow - untgz unused + no fixed Wolfi build yet
|
||||
"GHSA-58pv-8j8x-9vj2" # jaraco.context path traversal - setuptools vendored only (v5.3.0), not used in application code (using v6.1.0+)
|
||||
"GHSA-34x7-hfp2-rc4v" # node-tar hardlink path traversal - not applicable, tar CLI not exposed in application code
|
||||
"GHSA-r6q2-hw4h-h46w" # node-tar not used by application runtime, Linux-only container, not affect by macOS APFS-specific exploit
|
||||
"GHSA-8rrh-rw8j-w5fx" # wheel is from chainguard and will be handled by then TODO: Remove this after Chainguard updates the wheel
|
||||
"CVE-2025-59465" # Node only used for Admin UI build/prisma
|
||||
"CVE-2025-55131" # Node only used for Admin UI build/prisma
|
||||
"CVE-2025-59466" # Node only used for Admin UI build/prisma
|
||||
"CVE-2025-55130" # Node only used for Admin UI build/prisma
|
||||
"CVE-2025-59467" # Node only used for Admin UI build/prisma
|
||||
"CVE-2026-21637" # Node only used for Admin UI build/prisma
|
||||
"CVE-2025-55132" # Node only used for Admin UI build/prisma
|
||||
"GHSA-hx9q-6w63-j58v" # orjson dumps recursion; allowlisted
|
||||
"CVE-2025-15281" # No fix available yet
|
||||
"CVE-2026-0865" # No fix available yet
|
||||
"CVE-2025-15282" # No fix available yet
|
||||
"CVE-2026-0672" # No fix available yet
|
||||
"CVE-2025-15366" # No fix available yet
|
||||
"CVE-2025-15367" # No fix available yet
|
||||
"CVE-2025-12781" # No fix available yet
|
||||
"CVE-2025-11468" # No fix available yet
|
||||
"CVE-2026-1299" # Python 3.13 email module header injection - not applicable, LiteLLM doesn't use BytesGenerator for email serialization
|
||||
"CVE-2026-0775" # npm cli incorrect permission assignment - no fix available yet, npm is only used at build/prisma-generate time
|
||||
"GHSA-3ppc-4f35-3m26" # minimatch ReDoS via repeated wildcards - from nodejs_wheel bundled npm, not used in application runtime code
|
||||
"GHSA-83g3-92jg-28cx" # tar arbitrary file read/write via hardlink - from nodejs_wheel bundled npm, not used in application runtime code
|
||||
"CVE-2026-2297" # Python 3.13 SourcelessFileLoader audit hook bypass - no fix available in base image
|
||||
"GHSA-qffp-2rhf-9h96" # tar hardlink path traversal - from nodejs_wheel bundled npm, not used in application runtime code
|
||||
"CVE-2026-2673" # OpenSSL 3.6.1 TLS 1.3 key exchange group negotiation issue - no fix available yet
|
||||
"CVE-2026-3644" # Python 3.13 vulnerability - no fix available in base image
|
||||
"CVE-2026-4224" # Python 3.13 Expat parser stack overflow in ElementDeclHandler - no fix available in base image
|
||||
)
|
||||
|
||||
# Build JSON array of allowlisted CVE IDs for jq
|
||||
ALLOWED_IDS_JSON=$(printf '%s\n' "${ALLOWED_CVES[@]}" | jq -R . | jq -s .)
|
||||
|
||||
echo "Checking for vulnerabilities with CVSS score >= 4.0..."
|
||||
echo "Allowlisted CVEs (ignored in threshold): ${ALLOWED_CVES[*]}"
|
||||
echo ""
|
||||
|
||||
# Show all high-severity vulnerabilities for transparency
|
||||
TOTAL_HIGH_SEVERITY=$(grype litellm:latest -o json | jq -r '
|
||||
.matches[]
|
||||
| select(.vulnerability.cvss[]?.metrics.baseScore >= 4.0)
|
||||
| .vulnerability.id' | wc -l)
|
||||
|
||||
if [ "$TOTAL_HIGH_SEVERITY" -gt 0 ]; then
|
||||
echo "Total vulnerabilities found with CVSS >= 4.0: $TOTAL_HIGH_SEVERITY"
|
||||
echo ""
|
||||
echo "All high-severity vulnerabilities (including allowlisted):"
|
||||
grype litellm:latest -o json | jq --argjson allow "$ALLOWED_IDS_JSON" -r '
|
||||
["Package", "Version", "Vulnerability ID", "CVSS Score", "Allowlisted"],
|
||||
(.matches[]
|
||||
| select(.vulnerability.cvss[]?.metrics.baseScore >= 4.0)
|
||||
| [.artifact.name, .artifact.version, .vulnerability.id, .vulnerability.cvss[0].metrics.baseScore, (if (.vulnerability.id as $id | $allow | index($id)) then "YES" else "NO" end)])
|
||||
| @tsv' | column -t -s $'\t'
|
||||
echo ""
|
||||
fi
|
||||
|
||||
HIGH_SEVERITY_COUNT=$(grype litellm:latest -o json | jq --argjson allow "$ALLOWED_IDS_JSON" -r '
|
||||
.matches[]
|
||||
| select(.vulnerability.cvss[]?.metrics.baseScore >= 4.0)
|
||||
| select((.vulnerability.id as $id | $allow | index($id) | not))
|
||||
| .vulnerability.id' | wc -l)
|
||||
|
||||
if [ "$HIGH_SEVERITY_COUNT" -gt 0 ]; then
|
||||
echo ""
|
||||
echo "=========================================="
|
||||
echo "ERROR: Security Scan Failed"
|
||||
echo "=========================================="
|
||||
echo "Found $HIGH_SEVERITY_COUNT non-allowlisted vulnerabilities with CVSS score >= 4.0 in litellm:latest"
|
||||
echo ""
|
||||
echo "These vulnerabilities are NOT in the allowlist and must be addressed."
|
||||
echo "Current allowlisted CVEs: ${ALLOWED_CVES[*]}"
|
||||
echo ""
|
||||
echo "Detailed vulnerability report:"
|
||||
echo ""
|
||||
grype litellm:latest -o json | jq --argjson allow "$ALLOWED_IDS_JSON" -r '
|
||||
["Package", "Version", "Vulnerability ID", "CVSS Score", "Severity", "Fix Version", "Description"],
|
||||
(.matches[]
|
||||
| select(.vulnerability.cvss[]?.metrics.baseScore >= 4.0)
|
||||
| select((.vulnerability.id as $id | $allow | index($id) | not))
|
||||
| [.artifact.name, .artifact.version, .vulnerability.id, .vulnerability.cvss[0].metrics.baseScore, .vulnerability.severity, (.vulnerability.fix.versions[0] // "No fix available"), .vulnerability.description])
|
||||
| @tsv' | column -t -s $'\t'
|
||||
echo ""
|
||||
echo "=========================================="
|
||||
echo "Action Required:"
|
||||
echo "=========================================="
|
||||
echo "1. If a fix is available, update the package to the fixed version"
|
||||
echo "2. If the vulnerability is not applicable or has no fix:"
|
||||
echo " - Add the CVE/GHSA ID to ALLOWED_CVES array in ci_cd/security_scans.sh"
|
||||
echo " - Add a comment explaining why it's safe to ignore"
|
||||
echo ""
|
||||
echo "Note: Some vulnerabilities may have multiple IDs (CVE-XXXX and GHSA-XXXX)."
|
||||
echo "Add all relevant IDs to the allowlist if they refer to the same issue."
|
||||
echo "=========================================="
|
||||
echo ""
|
||||
exit 1
|
||||
else
|
||||
echo "No high-severity vulnerabilities (CVSS >= 4.0) found in litellm:latest"
|
||||
fi
|
||||
|
||||
echo "Grype scans completed successfully"
|
||||
}
|
||||
|
||||
# Main execution
|
||||
main() {
|
||||
echo "Installing security scanning tools..."
|
||||
install_trivy
|
||||
install_grype
|
||||
|
||||
# echo "Running secret detection scans..."
|
||||
# run_secret_detection
|
||||
|
||||
echo "Running filesystem vulnerability scans..."
|
||||
run_trivy_scans
|
||||
|
||||
echo "Running Docker image vulnerability scans..."
|
||||
run_grype_scans
|
||||
|
||||
echo "All security scans completed successfully!"
|
||||
}
|
||||
|
||||
# Execute main function
|
||||
main "$@"
|
||||
|
|
@ -41,22 +41,24 @@ COPY . .
|
|||
ENV LITELLM_NON_ROOT=true
|
||||
|
||||
# Build Admin UI using the upstream command order while keeping a single RUN layer
|
||||
# NOTE: .npmrc (which has ignore-scripts=true and min-release-age=3d) is temporarily
|
||||
# renamed during npm install/ci. This is safe because npm ci installs from
|
||||
# NOTE: .npmrc files (which may set ignore-scripts=true and min-release-age=3d)
|
||||
# are temporarily renamed during npm install/ci so they don't block lifecycle
|
||||
# scripts needed by the build. This is safe because npm ci installs from
|
||||
# package-lock.json with pinned versions + integrity hashes.
|
||||
RUN mkdir -p /var/lib/litellm/ui && \
|
||||
mv /app/.npmrc /app/.npmrc.bak && \
|
||||
([ -f /app/.npmrc ] && mv /app/.npmrc /app/.npmrc.bak || true) && \
|
||||
npm install -g npm@11.12.1 && \
|
||||
npm install -g node-gyp@12.2.0 && \
|
||||
ln -sf /usr/local/lib/node_modules/node-gyp /usr/lib/node_modules/npm/node_modules/node-gyp && \
|
||||
ln -sf "$(npm root -g)/node-gyp" "$(npm root -g)/npm/node_modules/node-gyp" && \
|
||||
npm cache clean --force && \
|
||||
cd /app/ui/litellm-dashboard && \
|
||||
if [ -f "/app/enterprise/enterprise_ui/enterprise_colors.json" ]; then \
|
||||
cp /app/enterprise/enterprise_ui/enterprise_colors.json ./ui_colors.json; \
|
||||
fi && \
|
||||
mv .npmrc .npmrc.bak && \
|
||||
([ -f .npmrc ] && mv .npmrc .npmrc.bak || true) && \
|
||||
npm ci && \
|
||||
mv .npmrc.bak .npmrc && mv /app/.npmrc.bak /app/.npmrc && \
|
||||
([ -f .npmrc.bak ] && mv .npmrc.bak .npmrc || true) && \
|
||||
([ -f /app/.npmrc.bak ] && mv /app/.npmrc.bak /app/.npmrc || true) && \
|
||||
npm run build && \
|
||||
cp -r /app/ui/litellm-dashboard/out/* /var/lib/litellm/ui/ && \
|
||||
mkdir -p /var/lib/litellm/assets && \
|
||||
|
|
|
|||
|
|
@ -13,19 +13,19 @@ To build and run the application, you will use the `docker-compose.yml` file loc
|
|||
|
||||
### 1. Set the Master Key
|
||||
|
||||
The application requires a `MASTER_KEY` for signing and validating tokens. You must set this key as an environment variable before running the application.
|
||||
The application requires a `LITELLM_MASTER_KEY` for signing and validating tokens. You must set this key as an environment variable before running the application.
|
||||
|
||||
Create a `.env` file in the root of the project and add the following line:
|
||||
|
||||
```
|
||||
MASTER_KEY=your-secret-key
|
||||
LITELLM_MASTER_KEY=your-secret-key
|
||||
```
|
||||
|
||||
Replace `your-secret-key` with a strong, randomly generated secret.
|
||||
|
||||
### 2. Build and Run the Containers
|
||||
|
||||
Once you have set the `MASTER_KEY`, you can build and run the containers using the following command:
|
||||
Once you have set the `LITELLM_MASTER_KEY`, you can build and run the containers using the following command:
|
||||
|
||||
```bash
|
||||
docker compose up -d --build
|
||||
|
|
@ -89,4 +89,4 @@ This command should succeed (showing engine versions) even with `--network none`
|
|||
## Troubleshooting
|
||||
|
||||
- **`build_admin_ui.sh: not found`**: This error can occur if the Docker build context is not set correctly. Ensure that you are running the `docker-compose` command from the root of the project.
|
||||
- **`Master key is not initialized`**: This error means the `MASTER_key` environment variable is not set. Make sure you have created a `.env` file in the project root with the `MASTER_KEY` defined.
|
||||
- **`Master key is not initialized`**: This error means the `LITELLM_MASTER_KEY` environment variable is not set. Make sure you have created a `.env` file in the project root with the `LITELLM_MASTER_KEY` defined.
|
||||
|
|
|
|||
|
|
@ -1,7 +0,0 @@
|
|||
# js-yaml CVE-2025-64718
|
||||
# This vulnerability is not applicable because we've forced js-yaml to version 4.1.1
|
||||
# via npm overrides in package.json. Trivy incorrectly reports this based on
|
||||
# dependency requirements in the lockfile, but the actual installed version is 4.1.1.
|
||||
# Verified with: npm list js-yaml
|
||||
CVE-2025-64718
|
||||
|
||||
|
|
@ -27,6 +27,41 @@ Building on the roadmap from our [security incident](https://docs.litellm.ai/blo
|
|||
- Validation and release are separated into different repositories, making it harder for an attacker to reach release credentials.
|
||||
- Trusted Publishing for PyPI releases - this means no long-lived credentials are used to publish releases.
|
||||
- Immutable Docker release tags - this means no tampering of Docker release tags after they are published [Learn more](https://docs.docker.com/docker-hub/repos/manage/hub-images/immutable-tags/). Note: work for GHCR docker releases is planned as well.
|
||||
- Docker image signing with [Cosign](https://github.com/sigstore/cosign) - all release images are signed so users can independently verify they came from us.
|
||||
|
||||
## Verify Docker image signatures
|
||||
|
||||
Starting from `v1.83.0-nightly`, all LiteLLM Docker images published to GHCR are signed with [cosign](https://docs.sigstore.dev/cosign/overview/). Every release is signed with the same key introduced in [commit `0112e53`](https://github.com/BerriAI/litellm/commit/0112e53046018d726492c814b3644b7d376029d0).
|
||||
|
||||
**Verify using the pinned commit hash (recommended):**
|
||||
|
||||
A commit hash is cryptographically immutable, so this is the strongest way to ensure you are using the original signing key:
|
||||
|
||||
```bash
|
||||
cosign verify \
|
||||
--key https://raw.githubusercontent.com/BerriAI/litellm/0112e53046018d726492c814b3644b7d376029d0/cosign.pub \
|
||||
ghcr.io/berriai/litellm:<release-tag>
|
||||
```
|
||||
|
||||
**Verify using a release tag (convenience):**
|
||||
|
||||
Tags are protected in this repository and resolve to the same key. This option is easier to read but relies on tag protection rules:
|
||||
|
||||
```bash
|
||||
cosign verify \
|
||||
--key https://raw.githubusercontent.com/BerriAI/litellm/<release-tag>/cosign.pub \
|
||||
ghcr.io/berriai/litellm:<release-tag>
|
||||
```
|
||||
|
||||
Replace `<release-tag>` with the version you are deploying (e.g. `v1.83.0-stable`).
|
||||
|
||||
Expected output:
|
||||
|
||||
```
|
||||
The following checks were performed on each of these signatures:
|
||||
- The cosign claims were validated
|
||||
- The signatures were verified against the specified public key
|
||||
```
|
||||
|
||||
## What's next
|
||||
|
||||
|
|
|
|||
|
|
@ -143,8 +143,41 @@ This will ensure, your releases are safe, even when:
|
|||
- Tampered registry artifacts are published
|
||||
- Tag mutations are made after the release is published
|
||||
|
||||
We believe that [Cosign](https://github.com/sigstore/cosign) is a good fit for this, and have already begun working on it [PR](https://github.com/BerriAI/litellm/pull/24683).
|
||||
We believe that [Cosign](https://github.com/sigstore/cosign) is a good fit for this, and have shipped it in [PR #24683](https://github.com/BerriAI/litellm/pull/24683).
|
||||
|
||||
#### How to verify a Docker image with Cosign
|
||||
|
||||
Starting from `v1.83.0-nightly`, all LiteLLM Docker images published to GHCR are signed with [cosign](https://docs.sigstore.dev/cosign/overview/). Every release is signed with the same key that was introduced in [commit `0112e53`](https://github.com/BerriAI/litellm/commit/0112e53046018d726492c814b3644b7d376029d0).
|
||||
|
||||
**Verify using the pinned commit hash (recommended):**
|
||||
|
||||
A commit hash is cryptographically immutable, so this is the strongest way to ensure you are using the original signing key:
|
||||
|
||||
```bash
|
||||
cosign verify \
|
||||
--key https://raw.githubusercontent.com/BerriAI/litellm/0112e53046018d726492c814b3644b7d376029d0/cosign.pub \
|
||||
ghcr.io/berriai/litellm:<release-tag>
|
||||
```
|
||||
|
||||
**Verify using a release tag (convenience):**
|
||||
|
||||
Tags are protected in this repository and resolve to the same key. This option is easier to read but relies on tag protection rules:
|
||||
|
||||
```bash
|
||||
cosign verify \
|
||||
--key https://raw.githubusercontent.com/BerriAI/litellm/<release-tag>/cosign.pub \
|
||||
ghcr.io/berriai/litellm:<release-tag>
|
||||
```
|
||||
|
||||
Replace `<release-tag>` with the version you are deploying (e.g. `v1.83.0-stable`).
|
||||
|
||||
Expected output:
|
||||
|
||||
```
|
||||
The following checks were performed on each of these signatures:
|
||||
- The cosign claims were validated
|
||||
- The signatures were verified against the specified public key
|
||||
```
|
||||
|
||||
### Avoid Compromised Packages
|
||||
|
||||
|
|
|
|||
|
|
@ -708,6 +708,40 @@ The LiteLLM AI Gateway team has already taken the following steps:
|
|||
- Engaged Google's Mandiant security team to assist with forensic analysis of the build and publishing chain
|
||||
|
||||
|
||||
## Verify Docker image signatures
|
||||
|
||||
Starting from `v1.83.0-nightly`, all LiteLLM Docker images published to GHCR are signed with [cosign](https://docs.sigstore.dev/cosign/overview/). Every release is signed with the same key introduced in [commit `0112e53`](https://github.com/BerriAI/litellm/commit/0112e53046018d726492c814b3644b7d376029d0).
|
||||
|
||||
**Verify using the pinned commit hash (recommended):**
|
||||
|
||||
A commit hash is cryptographically immutable, so this is the strongest way to ensure you are using the original signing key:
|
||||
|
||||
```bash
|
||||
cosign verify \
|
||||
--key https://raw.githubusercontent.com/BerriAI/litellm/0112e53046018d726492c814b3644b7d376029d0/cosign.pub \
|
||||
ghcr.io/berriai/litellm:<release-tag>
|
||||
```
|
||||
|
||||
**Verify using a release tag (convenience):**
|
||||
|
||||
Tags are protected in this repository and resolve to the same key. This option is easier to read but relies on tag protection rules:
|
||||
|
||||
```bash
|
||||
cosign verify \
|
||||
--key https://raw.githubusercontent.com/BerriAI/litellm/<release-tag>/cosign.pub \
|
||||
ghcr.io/berriai/litellm:<release-tag>
|
||||
```
|
||||
|
||||
Replace `<release-tag>` with the version you are deploying (e.g. `v1.83.0-stable`).
|
||||
|
||||
Expected output:
|
||||
|
||||
```
|
||||
The following checks were performed on each of these signatures:
|
||||
- The cosign claims were validated
|
||||
- The signatures were verified against the specified public key
|
||||
```
|
||||
|
||||
## Verified safe versions
|
||||
|
||||
We have audited every LiteLLM release published between v1.78.0 and v1.82.6 across both PyPI and Docker. Each artifact was verified by:
|
||||
|
|
|
|||
|
|
@ -238,7 +238,7 @@ router_settings:
|
|||
| public_routes | List[str] | (Enterprise Feature) Control list of public routes |
|
||||
| alert_types | List[str] | Control list of alert types to send to slack (Doc on alert types)[./alerting.md] |
|
||||
| enforced_params | List[str] | (Enterprise Feature) List of params that must be included in all requests to the proxy |
|
||||
| enable_oauth2_auth | boolean | (Enterprise Feature) If true, enables oauth2.0 authentication |
|
||||
| enable_oauth2_auth | boolean | (Enterprise Feature) If true, enables oauth2.0 authentication on LLM + info routes |
|
||||
| use_x_forwarded_for | str | If true, uses the X-Forwarded-For header to get the client IP address |
|
||||
| service_account_settings | List[Dict[str, Any]] | Set `service_account_settings` if you want to create settings that only apply to service account keys (Doc on service accounts)[./service_accounts.md] |
|
||||
| image_generation_model | str | The default model to use for image generation - ignores model set in request |
|
||||
|
|
@ -597,6 +597,7 @@ router_settings:
|
|||
| LITELLM_MCP_TOOL_LISTING_TIMEOUT | Timeout in seconds for listing tools from an MCP server. Default is 30
|
||||
| LITELLM_MCP_METADATA_TIMEOUT | HTTP client timeout in seconds for OAuth metadata fetching. Default is 10
|
||||
| LITELLM_MCP_HEALTH_CHECK_TIMEOUT | Health check timeout in seconds for MCP servers. Default is 10
|
||||
| LITELLM_MCP_STDIO_EXTRA_COMMANDS | Comma-separated extra command basenames allowed for MCP stdio transport beyond the built-in allowlist. Example: `my-mcp-bin`. Empty by default
|
||||
| MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL | Default TTL in seconds for MCP OAuth2 token cache. Default is 3600
|
||||
| MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE | Maximum number of entries in MCP OAuth2 token cache. Default is 200
|
||||
| MCP_OAUTH2_TOKEN_CACHE_MIN_TTL | Minimum TTL in seconds for MCP OAuth2 token cache. Default is 10
|
||||
|
|
@ -1032,6 +1033,7 @@ router_settings:
|
|||
| SENDGRID_SENDER_EMAIL | Email address used as the sender in SendGrid email transactions
|
||||
| SPEND_LOGS_URL | URL for retrieving spend logs
|
||||
| SPEND_LOG_CLEANUP_BATCH_SIZE | Number of logs deleted per batch during cleanup. Default is 1000
|
||||
| STALE_OBJECT_CLEANUP_BATCH_SIZE | Max number of stale managed objects updated per cleanup cycle. Default is 1000
|
||||
| SSL_CERTIFICATE | Path to the SSL certificate file
|
||||
| SSL_ECDH_CURVE | ECDH curve for SSL/TLS key exchange (e.g., 'X25519' to disable PQC).
|
||||
| SSL_SECURITY_LEVEL | [BETA] Security level for SSL/TLS connections. E.g. `DEFAULT@SECLEVEL=1`
|
||||
|
|
|
|||
|
|
@ -65,7 +65,43 @@ docker compose up
|
|||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Docker Run
|
||||
### Verify Docker image signatures
|
||||
|
||||
All LiteLLM Docker images are signed with [cosign](https://docs.sigstore.dev/cosign/overview/). Every release is signed with the same key introduced in [commit `0112e53`](https://github.com/BerriAI/litellm/commit/0112e53046018d726492c814b3644b7d376029d0).
|
||||
|
||||
**Verify using the pinned commit hash (recommended):**
|
||||
|
||||
A commit hash is cryptographically immutable, so this is the strongest way to ensure you are using the original signing key:
|
||||
|
||||
```bash
|
||||
cosign verify \
|
||||
--key https://raw.githubusercontent.com/BerriAI/litellm/0112e53046018d726492c814b3644b7d376029d0/cosign.pub \
|
||||
ghcr.io/berriai/litellm:<release-tag>
|
||||
```
|
||||
|
||||
**Verify using a release tag (convenience):**
|
||||
|
||||
Tags are protected in this repository and resolve to the same key. This option is easier to read but relies on tag protection rules:
|
||||
|
||||
```bash
|
||||
cosign verify \
|
||||
--key https://raw.githubusercontent.com/BerriAI/litellm/<release-tag>/cosign.pub \
|
||||
ghcr.io/berriai/litellm:<release-tag>
|
||||
```
|
||||
|
||||
Replace `<release-tag>` with the version you are deploying (e.g. `v1.83.0-stable`).
|
||||
|
||||
Expected output:
|
||||
|
||||
```
|
||||
The following checks were performed on each of these signatures:
|
||||
- The cosign claims were validated
|
||||
- The signatures were verified against the specified public key
|
||||
```
|
||||
|
||||
Learn more about LiteLLM's release signing in the [CI/CD v2 announcement](https://docs.litellm.ai/blog/ci-cd-v2-improvements#verify-docker-image-signatures).
|
||||
|
||||
### Docker Run
|
||||
|
||||
#### Step 1. CREATE config.yaml
|
||||
|
||||
|
|
|
|||
|
|
@ -63,16 +63,19 @@ Start the LiteLLM Proxy with [`--detailed_debug` mode and you should see more ve
|
|||
|
||||
## Using OAuth2 + JWT Together
|
||||
|
||||
If both `enable_oauth2_auth` and `enable_jwt_auth` are enabled, LiteLLM can split auth paths:
|
||||
- JWT validation for user tokens
|
||||
- OAuth2 introspection for machine tokens
|
||||
LiteLLM supports two OAuth2 + JWT modes:
|
||||
|
||||
For JWT-shaped machine tokens, configure `litellm_jwtauth.routing_overrides`:
|
||||
1. **Global OAuth2 mode** (`enable_oauth2_auth: true`)
|
||||
OAuth2 auth is enabled on LLM + info routes.
|
||||
2. **Selective JWT override mode** (`enable_oauth2_auth: false`)
|
||||
Only JWT-shaped tokens that match `litellm_jwtauth.routing_overrides` are routed to OAuth2 on LLM + info routes.
|
||||
|
||||
For selective routing (OAuth2 only for specific JWTs), configure:
|
||||
|
||||
```yaml title="config.yaml"
|
||||
general_settings:
|
||||
enable_jwt_auth: true
|
||||
enable_oauth2_auth: true
|
||||
enable_oauth2_auth: false
|
||||
litellm_jwtauth:
|
||||
routing_overrides:
|
||||
- iss: "machine-issuer.example.com"
|
||||
|
|
|
|||
|
|
@ -792,16 +792,18 @@ litellm_jwtauth:
|
|||
|
||||
## Route JWT-Shaped Machine Tokens to OAuth2
|
||||
|
||||
Use this when both are enabled:
|
||||
Use this when:
|
||||
- `enable_jwt_auth: true` for standard JWT validation
|
||||
- `enable_oauth2_auth: true` for OAuth2 introspection
|
||||
- machine tokens are JWT-shaped and should be routed to OAuth2 based on claims
|
||||
|
||||
If some machine tokens are also JWT-shaped, configure `routing_overrides` to route matching tokens to OAuth2.
|
||||
`routing_overrides` supports two operating modes:
|
||||
- **Selective mode**: set `enable_oauth2_auth: false` to send only matching JWTs to OAuth2 on LLM + info routes
|
||||
- **Global mode**: set `enable_oauth2_auth: true` to also enable OAuth2 on LLM + info routes
|
||||
|
||||
```yaml title="config.yaml"
|
||||
general_settings:
|
||||
enable_jwt_auth: true
|
||||
enable_oauth2_auth: true
|
||||
enable_oauth2_auth: false
|
||||
litellm_jwtauth:
|
||||
user_id_jwt_field: "sub"
|
||||
routing_overrides:
|
||||
|
|
@ -822,7 +824,7 @@ general_settings:
|
|||
```yaml title="config.yaml"
|
||||
general_settings:
|
||||
enable_jwt_auth: true
|
||||
enable_oauth2_auth: true
|
||||
enable_oauth2_auth: false
|
||||
litellm_jwtauth:
|
||||
routing_overrides:
|
||||
- iss: ["machine-issuer.example.com", "backup-issuer.example.com"]
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.constants import (
|
||||
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS,
|
||||
MAX_OBJECTS_PER_POLL_CYCLE,
|
||||
STALE_OBJECT_CLEANUP_BATCH_SIZE,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -32,21 +33,49 @@ class CheckResponsesCost:
|
|||
self.prisma_client: PrismaClient = prisma_client
|
||||
self.llm_router: Router = llm_router
|
||||
|
||||
async def _expire_stale_rows(
|
||||
self, cutoff: datetime, batch_size: int
|
||||
) -> int:
|
||||
"""Execute the bounded UPDATE that marks stale rows as 'stale_expired'.
|
||||
|
||||
Isolated so it can be swapped / mocked in tests without touching the
|
||||
orchestration logic in ``_cleanup_stale_managed_objects``.
|
||||
|
||||
Uses PostgreSQL syntax (``$1::timestamptz``, ``LIMIT``, double-quoted
|
||||
identifiers) which is the only dialect the proxy supports — every
|
||||
``schema.prisma`` in the repo sets ``provider = "postgresql"``.
|
||||
Same pattern as ``spend_log_cleanup.py``.
|
||||
"""
|
||||
return await self.prisma_client.db.execute_raw(
|
||||
"""
|
||||
UPDATE "LiteLLM_ManagedObjectTable"
|
||||
SET "status" = 'stale_expired'
|
||||
WHERE "id" IN (
|
||||
SELECT "id" FROM "LiteLLM_ManagedObjectTable"
|
||||
WHERE "file_purpose" = 'response'
|
||||
AND "status" NOT IN ('completed', 'complete', 'failed', 'expired', 'cancelled', 'stale_expired')
|
||||
AND "created_at" < $1::timestamptz
|
||||
ORDER BY "created_at" ASC
|
||||
LIMIT $2
|
||||
)
|
||||
""",
|
||||
cutoff,
|
||||
batch_size,
|
||||
)
|
||||
|
||||
async def _cleanup_stale_managed_objects(self) -> None:
|
||||
"""
|
||||
Mark managed objects older than MANAGED_OBJECT_STALENESS_CUTOFF_DAYS days
|
||||
in non-terminal states as 'stale_expired'. These will never complete and
|
||||
should not be polled.
|
||||
|
||||
Runs as a single DB query with a subquery LIMIT so no rows are loaded
|
||||
into Python memory. Processes at most STALE_OBJECT_CLEANUP_BATCH_SIZE
|
||||
rows per invocation to avoid overwhelming the DB when there is a large
|
||||
backlog.
|
||||
"""
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
|
||||
result = await self.prisma_client.db.litellm_managedobjecttable.update_many(
|
||||
where={
|
||||
"file_purpose": "response",
|
||||
"status": {"not_in": ["completed", "complete", "failed", "expired", "cancelled", "stale_expired"]},
|
||||
"created_at": {"lt": cutoff},
|
||||
},
|
||||
data={"status": "stale_expired"},
|
||||
)
|
||||
result = await self._expire_stale_rows(cutoff, STALE_OBJECT_CLEANUP_BATCH_SIZE)
|
||||
if result > 0:
|
||||
verbose_proxy_logger.warning(
|
||||
f"CheckResponsesCost: marked {result} stale managed objects "
|
||||
|
|
|
|||
|
|
@ -141,6 +141,17 @@ MCP_TOOL_LISTING_TIMEOUT = float(os.getenv("LITELLM_MCP_TOOL_LISTING_TIMEOUT", "
|
|||
MCP_METADATA_TIMEOUT = float(os.getenv("LITELLM_MCP_METADATA_TIMEOUT", "10.0"))
|
||||
MCP_HEALTH_CHECK_TIMEOUT = float(os.getenv("LITELLM_MCP_HEALTH_CHECK_TIMEOUT", "10.0"))
|
||||
|
||||
# Allowlist of commands permitted for MCP stdio transport.
|
||||
# Prevents arbitrary command execution via /mcp-rest/test/* endpoints or server creation.
|
||||
# Note: allowlisted runtimes can still execute code via args (e.g. python -c "...").
|
||||
# This is an accepted residual risk since these endpoints require PROXY_ADMIN.
|
||||
# Extend via LITELLM_MCP_STDIO_EXTRA_COMMANDS env var (comma-separated).
|
||||
_MCP_STDIO_EXTRA_COMMANDS = os.getenv("LITELLM_MCP_STDIO_EXTRA_COMMANDS", "")
|
||||
MCP_STDIO_ALLOWED_COMMANDS: frozenset = frozenset(
|
||||
{"npx", "uvx", "python", "python3", "node", "docker", "deno"}
|
||||
| (set(_MCP_STDIO_EXTRA_COMMANDS.split(",")) - {""})
|
||||
)
|
||||
|
||||
LITELLM_UI_ALLOW_HEADERS = [
|
||||
"x-litellm-semantic-filter",
|
||||
"x-litellm-semantic-filter-tools",
|
||||
|
|
@ -1367,6 +1378,9 @@ MAX_OBJECTS_PER_POLL_CYCLE = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE",
|
|||
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS = max(
|
||||
1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7))
|
||||
)
|
||||
STALE_OBJECT_CLEANUP_BATCH_SIZE = max(
|
||||
1, int(os.getenv("STALE_OBJECT_CLEANUP_BATCH_SIZE", 1000))
|
||||
)
|
||||
# Set PROXY_BATCH_POLLING_ENABLED=false to disable the CheckBatchCost and
|
||||
# CheckResponsesCost background polling jobs entirely (e.g. to avoid DB load on
|
||||
# installations with large numbers of stale managed objects).
|
||||
|
|
|
|||
|
|
@ -7822,8 +7822,8 @@
|
|||
"input_cost_per_token": 3.6e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.8e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
|
|
@ -7838,6 +7838,26 @@
|
|||
"cache_read_input_token_cost": 3.6e-07,
|
||||
"cache_creation_input_token_cost": 4.5e-06
|
||||
},
|
||||
"bedrock/us-gov-east-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"bedrock/us-gov-east-1/meta.llama3-70b-instruct-v1:0": {
|
||||
"input_cost_per_token": 2.65e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -7973,8 +7993,8 @@
|
|||
"input_cost_per_token": 3.6e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.8e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
|
|
@ -7989,6 +8009,26 @@
|
|||
"cache_read_input_token_cost": 3.6e-07,
|
||||
"cache_creation_input_token_cost": 4.5e-06
|
||||
},
|
||||
"bedrock/us-gov-west-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"bedrock/us-gov-west-1/meta.llama3-70b-instruct-v1:0": {
|
||||
"input_cost_per_token": 2.65e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -28945,6 +28985,32 @@
|
|||
"tool_use_system_prompt_tokens": 346,
|
||||
"supports_native_structured_output": true
|
||||
},
|
||||
"us-gov.anthropic.claude-sonnet-4-5-20250929-v1:0": {
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 6.6e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 2.475e-05,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 64000,
|
||||
"max_tokens": 64000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"tool_use_system_prompt_tokens": 346,
|
||||
"supports_native_structured_output": true
|
||||
},
|
||||
"au.anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.375e-06,
|
||||
"cache_read_input_token_cost": 1.1e-07,
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import asyncio
|
|||
import datetime
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from typing import Any, Callable, Dict, List, Literal, Optional, Set, Tuple, Union, cast
|
||||
from urllib.parse import urlparse
|
||||
|
|
@ -35,6 +36,8 @@ from litellm.constants import (
|
|||
MCP_CLIENT_TIMEOUT,
|
||||
MCP_HEALTH_CHECK_TIMEOUT,
|
||||
MCP_METADATA_TIMEOUT,
|
||||
MCP_NPM_CACHE_DIR,
|
||||
MCP_STDIO_ALLOWED_COMMANDS,
|
||||
MCP_TOOL_LISTING_TIMEOUT,
|
||||
)
|
||||
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
|
||||
|
|
@ -1119,9 +1122,19 @@ class MCPServerManager:
|
|||
# In containers the default (~/.npm or /app/.npm) may not exist
|
||||
# or be read-only, causing npx to fail with ENOENT.
|
||||
if "NPM_CONFIG_CACHE" not in resolved_env:
|
||||
from litellm.constants import MCP_NPM_CACHE_DIR
|
||||
|
||||
resolved_env["NPM_CONFIG_CACHE"] = MCP_NPM_CACHE_DIR
|
||||
# Defense-in-depth: block commands not in the allowlist.
|
||||
# The Pydantic validator blocks new servers; this catches legacy
|
||||
# config/DB records predating the allowlist.
|
||||
if server.command:
|
||||
base_command = os.path.basename(server.command)
|
||||
if base_command not in MCP_STDIO_ALLOWED_COMMANDS:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"MCP stdio command '{server.command}' is not in the allowlist ({sorted(MCP_STDIO_ALLOWED_COMMANDS)}). "
|
||||
f"Add it to LITELLM_MCP_STDIO_EXTRA_COMMANDS to allow this command.",
|
||||
)
|
||||
|
||||
stdio_config: Optional[MCPStdioConfig] = None
|
||||
if server.command and server.args is not None:
|
||||
stdio_config = MCPStdioConfig(
|
||||
|
|
|
|||
|
|
@ -2,14 +2,14 @@ import importlib
|
|||
from datetime import datetime
|
||||
from typing import Any, Awaitable, Callable, Dict, List, Literal, Optional, Set, Union
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
|
||||
build_effective_auth_contexts,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.utils import merge_mcp_headers
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
||||
|
|
@ -1027,6 +1027,13 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
Test if we can connect to the provided MCP server before adding it
|
||||
"""
|
||||
if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={
|
||||
"error": "User does not have permission to test MCP server connections. Only PROXY_ADMIN users can perform this action."
|
||||
},
|
||||
)
|
||||
|
||||
async def _test_connection_operation(client):
|
||||
async def _noop(session):
|
||||
|
|
@ -1041,7 +1048,7 @@ if MCP_AVAILABLE:
|
|||
raw_headers=_safe_get_request_headers(request),
|
||||
)
|
||||
|
||||
@router.post("/test/tools/list")
|
||||
@router.post("/test/tools/list", dependencies=[Depends(user_api_key_auth)])
|
||||
async def test_tools_list(
|
||||
request: Request,
|
||||
new_mcp_server_request: NewMCPServerRequest,
|
||||
|
|
@ -1050,6 +1057,14 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
Preview tools available from MCP server before adding it
|
||||
"""
|
||||
if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={
|
||||
"error": "User does not have permission to test MCP server tools. Only PROXY_ADMIN users can perform this action."
|
||||
},
|
||||
)
|
||||
|
||||
# For OpenAPI spec servers, generate tools from the spec directly
|
||||
if new_mcp_server_request.spec_path:
|
||||
return await _preview_openapi_tools(new_mcp_server_request.spec_path)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import enum
|
||||
import json
|
||||
import os
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union
|
||||
|
||||
|
|
@ -15,6 +16,7 @@ from pydantic import (
|
|||
from typing_extensions import Required, TypedDict
|
||||
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import MCP_STDIO_ALLOWED_COMMANDS
|
||||
from litellm.types.integrations.slack_alerting import AlertType
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
|
|
@ -1162,6 +1164,13 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
raise ValueError("command is required for stdio transport")
|
||||
if not values.get("args"):
|
||||
raise ValueError("args is required for stdio transport")
|
||||
# Validate command against allowlist to prevent arbitrary execution
|
||||
base_command = os.path.basename(values["command"])
|
||||
if base_command not in MCP_STDIO_ALLOWED_COMMANDS:
|
||||
raise ValueError(
|
||||
f"Command '{values['command']}' is not in the allowed commands list "
|
||||
f"for stdio transport. Allowed commands: {sorted(MCP_STDIO_ALLOWED_COMMANDS)}"
|
||||
)
|
||||
elif transport in [MCPTransport.http, MCPTransport.sse]:
|
||||
if not values.get("url") and not values.get("spec_path"):
|
||||
raise ValueError(
|
||||
|
|
@ -1222,6 +1231,13 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
raise ValueError("command is required for stdio transport")
|
||||
if not values.get("args"):
|
||||
raise ValueError("args is required for stdio transport")
|
||||
# Validate command against allowlist to prevent arbitrary execution
|
||||
base_command = os.path.basename(values["command"])
|
||||
if base_command not in MCP_STDIO_ALLOWED_COMMANDS:
|
||||
raise ValueError(
|
||||
f"Command '{values['command']}' is not in the allowed commands list "
|
||||
f"for stdio transport. Allowed commands: {sorted(MCP_STDIO_ALLOWED_COMMANDS)}"
|
||||
)
|
||||
elif transport in [MCPTransport.http, MCPTransport.sse]:
|
||||
if not values.get("url") and not values.get("spec_path"):
|
||||
raise ValueError(
|
||||
|
|
|
|||
|
|
@ -690,42 +690,39 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
|
||||
########## End of Route Checks Before Reading DB / Cache for "token" ########
|
||||
|
||||
if general_settings.get("enable_oauth2_auth", False) is True:
|
||||
# Only apply OAuth2 M2M authentication to LLM API routes and info routes, not UI/management routes
|
||||
# This allows UI SSO to work separately from API M2M authentication
|
||||
# Note: Info routes are already scoped to the user
|
||||
if RouteChecks.is_llm_api_route(route=route) or RouteChecks.is_info_route(
|
||||
route=route
|
||||
):
|
||||
# When both OAuth2 and JWT auth are enabled, use token format to decide:
|
||||
# - JWT tokens (3 dot-separated parts) -> skip OAuth2, fall through to JWT handler
|
||||
# - Opaque tokens -> use OAuth2 handler
|
||||
# This allows JWT for users and OAuth2 for M2M on the same instance
|
||||
is_jwt = (
|
||||
jwt_handler.is_jwt(token=api_key)
|
||||
if general_settings.get("enable_jwt_auth", False) is True
|
||||
else False
|
||||
)
|
||||
# Routing uses unverified JWT claims only to choose auth path.
|
||||
# Final authentication is enforced by the selected validator.
|
||||
route_jwt_to_oauth2 = (
|
||||
is_jwt
|
||||
and _should_route_jwt_to_oauth2_override(
|
||||
token=api_key, jwt_handler=jwt_handler
|
||||
)
|
||||
)
|
||||
if not is_jwt or route_jwt_to_oauth2:
|
||||
# return UserAPIKeyAuth object
|
||||
# helper to check if the api_key is a valid oauth2 token
|
||||
from litellm.proxy.proxy_server import premium_user
|
||||
enable_oauth2_auth = general_settings.get("enable_oauth2_auth", False) is True
|
||||
enable_jwt_auth = general_settings.get("enable_jwt_auth", False) is True
|
||||
is_jwt = jwt_handler.is_jwt(token=api_key) if enable_jwt_auth else False
|
||||
|
||||
if premium_user is not True:
|
||||
raise ValueError(
|
||||
"Oauth2 token validation is only available for premium users"
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
)
|
||||
# Routing uses unverified JWT claims only to choose auth path.
|
||||
# Final authentication is enforced by the selected validator.
|
||||
route_jwt_to_oauth2 = (
|
||||
is_jwt
|
||||
and _should_route_jwt_to_oauth2_override(
|
||||
token=api_key, jwt_handler=jwt_handler
|
||||
)
|
||||
)
|
||||
|
||||
return await Oauth2Handler.check_oauth2_token(token=api_key)
|
||||
# OAuth2 applies for:
|
||||
# 1) when global OAuth2 auth is enabled on LLM + info routes
|
||||
# 2) JWT tokens that explicitly match routing_overrides on LLM + info routes
|
||||
should_apply_override_oauth2 = route_jwt_to_oauth2 and (
|
||||
RouteChecks.is_llm_api_route(route=route)
|
||||
or RouteChecks.is_info_route(route=route)
|
||||
)
|
||||
should_apply_global_oauth2 = enable_oauth2_auth and (
|
||||
RouteChecks.is_llm_api_route(route=route)
|
||||
or RouteChecks.is_info_route(route=route)
|
||||
)
|
||||
if (should_apply_global_oauth2 and not is_jwt) or should_apply_override_oauth2:
|
||||
from litellm.proxy.proxy_server import premium_user
|
||||
if premium_user is not True:
|
||||
raise ValueError(
|
||||
"Oauth2 token validation is only available for premium users"
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
)
|
||||
|
||||
return await Oauth2Handler.check_oauth2_token(token=api_key)
|
||||
|
||||
if general_settings.get("enable_oauth2_proxy_auth", False) is True:
|
||||
return await handle_oauth2_proxy_request(request=request)
|
||||
|
|
|
|||
|
|
@ -502,9 +502,7 @@ def _enforce_upperbound_key_params(
|
|||
|
||||
for elem in data:
|
||||
key, value = elem
|
||||
upperbound_value = getattr(
|
||||
litellm.upperbound_key_generate_params, key, None
|
||||
)
|
||||
upperbound_value = getattr(litellm.upperbound_key_generate_params, key, None)
|
||||
if upperbound_value is not None:
|
||||
if value is None:
|
||||
if fill_defaults:
|
||||
|
|
@ -524,9 +522,7 @@ def _enforce_upperbound_key_params(
|
|||
},
|
||||
)
|
||||
elif key in ["budget_duration", "duration"]:
|
||||
upperbound_duration = duration_in_seconds(
|
||||
duration=upperbound_value
|
||||
)
|
||||
upperbound_duration = duration_in_seconds(duration=upperbound_value)
|
||||
if value == "-1":
|
||||
user_duration = float("inf")
|
||||
else:
|
||||
|
|
@ -1759,9 +1755,7 @@ async def _process_single_key_update(
|
|||
decision = result.get("decision", True)
|
||||
message = result.get("message", "Authentication Failed - Custom Auth Rule")
|
||||
if not decision:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN, detail=message
|
||||
)
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=message)
|
||||
|
||||
# Enforce upperbound key params on update (don't fill defaults)
|
||||
_enforce_upperbound_key_params(update_key_request, fill_defaults=False)
|
||||
|
|
@ -2638,22 +2632,39 @@ async def info_key_fn_v2(
|
|||
detail={"message": "Malformed request. No keys passed in."},
|
||||
)
|
||||
|
||||
key_info = await prisma_client.get_data(
|
||||
token=data.keys, table_name="key", query_type="find_all"
|
||||
)
|
||||
if key_info is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail={"message": "No keys found"},
|
||||
# Resolve key_aliases to tokens so we never pass token=None (unbounded query)
|
||||
tokens_to_query = list(data.keys) if data.keys else []
|
||||
if data.key_aliases:
|
||||
alias_rows = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={"key_alias": {"in": data.key_aliases}},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
alias_tokens = [row.token for row in alias_rows if row.token]
|
||||
tokens_to_query.extend(alias_tokens)
|
||||
|
||||
if not tokens_to_query:
|
||||
return {"key": data.keys, "info": []}
|
||||
|
||||
key_info = await prisma_client.get_data(
|
||||
token=tokens_to_query, table_name="key", query_type="find_all"
|
||||
)
|
||||
if not key_info:
|
||||
return {"key": data.keys, "info": []}
|
||||
|
||||
filtered_key_info = []
|
||||
for k in key_info:
|
||||
if not await _can_user_query_key_info(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
key=k.token,
|
||||
key_info=k,
|
||||
):
|
||||
continue
|
||||
try:
|
||||
k = k.model_dump() # noqa
|
||||
k_dict = k.model_dump()
|
||||
except Exception:
|
||||
# if using pydantic v1
|
||||
k = k.dict()
|
||||
filtered_key_info.append(k)
|
||||
k_dict = k.dict()
|
||||
k_dict.pop("token", None)
|
||||
filtered_key_info.append(k_dict)
|
||||
return {"key": data.keys, "info": filtered_key_info}
|
||||
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -100,6 +100,8 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
|||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkTeamMemberAddRequest,
|
||||
BulkTeamMemberAddResponse,
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
BulkUpdateTeamMemberPermissionsResponse,
|
||||
GetTeamMemberPermissionsResponse,
|
||||
TeamListItem,
|
||||
TeamListResponse,
|
||||
|
|
@ -4274,6 +4276,151 @@ async def update_team_member_permissions(
|
|||
return updated_team
|
||||
|
||||
|
||||
@router.post(
|
||||
"/team/permissions_bulk_update",
|
||||
tags=["team management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=BulkUpdateTeamMemberPermissionsResponse,
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def bulk_update_team_member_permissions(
|
||||
data: BulkUpdateTeamMemberPermissionsRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Append permissions to existing teams.
|
||||
|
||||
Either pass team_ids to target specific teams, or set
|
||||
apply_to_all_teams=True to update every team. For each team,
|
||||
the provided permissions are merged with the team's existing
|
||||
permissions (duplicates are skipped).
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail={"error": "No db connected"})
|
||||
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": "Only proxy admins can bulk-update team permissions"},
|
||||
)
|
||||
|
||||
if not data.permissions:
|
||||
return {
|
||||
"message": "No permissions provided",
|
||||
"teams_updated": 0,
|
||||
}
|
||||
|
||||
if not data.apply_to_all_teams and not data.team_ids:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Must provide team_ids or set apply_to_all_teams=true"
|
||||
},
|
||||
)
|
||||
|
||||
if data.apply_to_all_teams and data.team_ids:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Cannot set both apply_to_all_teams=true and team_ids"
|
||||
},
|
||||
)
|
||||
|
||||
permissions_to_add = set(data.permissions)
|
||||
|
||||
if data.team_ids:
|
||||
teams_updated = await _append_permissions_to_specific_teams(
|
||||
prisma_client, data.team_ids, permissions_to_add
|
||||
)
|
||||
else:
|
||||
teams_updated = await _append_permissions_to_all_teams(
|
||||
prisma_client, permissions_to_add
|
||||
)
|
||||
|
||||
return {
|
||||
"message": "Team permissions updated successfully",
|
||||
"teams_updated": teams_updated,
|
||||
"permissions_appended": data.permissions,
|
||||
}
|
||||
|
||||
|
||||
async def _compute_and_batch_updates(prisma_client, teams, permissions_to_add: set) -> int:
|
||||
"""Compute merged permissions and batch-write updates. Returns count of teams updated."""
|
||||
updates = []
|
||||
for team in teams:
|
||||
existing = set(team.team_member_permissions or [])
|
||||
if permissions_to_add <= existing:
|
||||
continue
|
||||
merged = sorted(existing | permissions_to_add) # normalise to alphabetical order
|
||||
updates.append((team.team_id, merged))
|
||||
|
||||
if updates:
|
||||
batcher = prisma_client.db.batch_()
|
||||
for team_id, merged_perms in updates:
|
||||
batcher.litellm_teamtable.update(
|
||||
where={"team_id": team_id},
|
||||
data={"team_member_permissions": merged_perms},
|
||||
)
|
||||
await batcher.commit()
|
||||
|
||||
return len(updates)
|
||||
|
||||
|
||||
async def _append_permissions_to_specific_teams(
|
||||
prisma_client, team_ids: List[str], permissions_to_add: set
|
||||
) -> int:
|
||||
"""Fetch specific teams by ID and append permissions."""
|
||||
teams = await prisma_client.db.litellm_teamtable.find_many(
|
||||
where={"team_id": {"in": team_ids}},
|
||||
)
|
||||
|
||||
found_ids = {team.team_id for team in teams}
|
||||
missing_ids = set(team_ids) - found_ids
|
||||
if missing_ids:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"Team(s) not found: {sorted(missing_ids)}"},
|
||||
)
|
||||
|
||||
return await _compute_and_batch_updates(prisma_client, teams, permissions_to_add)
|
||||
|
||||
|
||||
async def _append_permissions_to_all_teams(
|
||||
prisma_client, permissions_to_add: set
|
||||
) -> int:
|
||||
"""Paginated read + batched write across all teams."""
|
||||
teams_updated = 0
|
||||
cursor = None
|
||||
BATCH_SIZE = 500
|
||||
|
||||
while True:
|
||||
find_args: dict = {
|
||||
"take": BATCH_SIZE,
|
||||
"order": {"team_id": "asc"},
|
||||
}
|
||||
if cursor is not None:
|
||||
find_args["cursor"] = {"team_id": cursor}
|
||||
find_args["skip"] = 1
|
||||
|
||||
teams = await prisma_client.db.litellm_teamtable.find_many(**find_args)
|
||||
|
||||
if not teams:
|
||||
break
|
||||
|
||||
teams_updated += await _compute_and_batch_updates(
|
||||
prisma_client, teams, permissions_to_add
|
||||
)
|
||||
|
||||
cursor = teams[-1].team_id
|
||||
|
||||
if len(teams) < BATCH_SIZE:
|
||||
break
|
||||
|
||||
return teams_updated
|
||||
|
||||
|
||||
@router.get(
|
||||
"/team/daily/activity",
|
||||
response_model=SpendAnalyticsPaginatedResponse,
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from typing import Any, Dict, List, Optional, Union
|
|||
from pydantic import BaseModel
|
||||
|
||||
from litellm.proxy._types import (
|
||||
KeyManagementRoutes,
|
||||
LiteLLM_DeletedTeamTable,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
|
|
@ -43,6 +44,27 @@ class UpdateTeamMemberPermissionsRequest(BaseModel):
|
|||
team_member_permissions: List[str]
|
||||
|
||||
|
||||
class BulkUpdateTeamMemberPermissionsRequest(BaseModel):
|
||||
"""Request to bulk-update team member permissions across teams."""
|
||||
|
||||
permissions: List[KeyManagementRoutes]
|
||||
"""Permissions to append to the target teams (duplicates are skipped)."""
|
||||
|
||||
team_ids: Optional[List[str]] = None
|
||||
"""Specific team IDs to update. Required unless apply_to_all_teams is True."""
|
||||
|
||||
apply_to_all_teams: bool = False
|
||||
"""When True, update all teams. Mutually exclusive with team_ids."""
|
||||
|
||||
|
||||
class BulkUpdateTeamMemberPermissionsResponse(BaseModel):
|
||||
"""Response for bulk team member permissions update."""
|
||||
|
||||
message: str
|
||||
teams_updated: int
|
||||
permissions_appended: Optional[List[str]] = None
|
||||
|
||||
|
||||
class TeamListItem(LiteLLM_TeamTable):
|
||||
"""A team item in the paginated list response, enriched with computed fields."""
|
||||
|
||||
|
|
|
|||
|
|
@ -7822,8 +7822,8 @@
|
|||
"input_cost_per_token": 3.6e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.8e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
|
|
@ -7838,6 +7838,26 @@
|
|||
"cache_read_input_token_cost": 3.6e-07,
|
||||
"cache_creation_input_token_cost": 4.5e-06
|
||||
},
|
||||
"bedrock/us-gov-east-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"bedrock/us-gov-east-1/meta.llama3-70b-instruct-v1:0": {
|
||||
"input_cost_per_token": 2.65e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -7973,8 +7993,8 @@
|
|||
"input_cost_per_token": 3.6e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.8e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
|
|
@ -7989,6 +8009,26 @@
|
|||
"cache_read_input_token_cost": 3.6e-07,
|
||||
"cache_creation_input_token_cost": 4.5e-06
|
||||
},
|
||||
"bedrock/us-gov-west-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"bedrock/us-gov-west-1/meta.llama3-70b-instruct-v1:0": {
|
||||
"input_cost_per_token": 2.65e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
|
|
@ -28930,6 +28970,32 @@
|
|||
"tool_use_system_prompt_tokens": 346,
|
||||
"supports_native_structured_output": true
|
||||
},
|
||||
"us-gov.anthropic.claude-sonnet-4-5-20250929-v1:0": {
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 6.6e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 2.475e-05,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 64000,
|
||||
"max_tokens": 64000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"tool_use_system_prompt_tokens": 346,
|
||||
"supports_native_structured_output": true
|
||||
},
|
||||
"au.anthropic.claude-haiku-4-5-20251001-v1:0": {
|
||||
"cache_creation_input_token_cost": 1.375e-06,
|
||||
"cache_read_input_token_cost": 1.1e-07,
|
||||
|
|
|
|||
3900
poetry.lock
generated
3900
poetry.lock
generated
File diff suppressed because it is too large
Load diff
|
|
@ -1,6 +1,6 @@
|
|||
[tool.poetry]
|
||||
name = "litellm"
|
||||
version = "1.83.3"
|
||||
version = "1.83.5"
|
||||
description = "Library to easily interface with LLM API providers"
|
||||
authors = ["BerriAI"]
|
||||
license = "MIT"
|
||||
|
|
@ -65,7 +65,7 @@ mcp = {version = "1.26.0", optional = true, python = ">=3.10"}
|
|||
a2a-sdk = {version = "0.3.25", optional = true, python = ">=3.10"}
|
||||
litellm-proxy-extras = {version = "0.4.65", optional = true}
|
||||
rich = {version = "13.9.4", optional = true}
|
||||
litellm-enterprise = {version = "0.1.36", optional = true}
|
||||
litellm-enterprise = {version = "0.1.37", optional = true}
|
||||
diskcache = {version = "5.6.3", optional = true}
|
||||
polars = {version = "1.39.3", optional = true, python = ">=3.10"}
|
||||
semantic-router = {version = "0.1.12", optional = true, python = ">=3.9,<3.14"}
|
||||
|
|
@ -181,7 +181,7 @@ requires = ["poetry-core", "wheel"]
|
|||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "1.83.3"
|
||||
version = "1.83.5"
|
||||
version_files = [
|
||||
"pyproject.toml:^version"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -85,4 +85,4 @@ requests-toolbelt==1.0.0 # transitive dep (langfuse)
|
|||
########################
|
||||
# LITELLM ENTERPRISE DEPENDENCIES
|
||||
########################
|
||||
litellm-enterprise==0.1.36
|
||||
litellm-enterprise==0.1.37
|
||||
|
|
|
|||
|
|
@ -373,35 +373,6 @@ def test_openai_azure_embedding_optional_arg():
|
|||
# test_openai_embedding()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model, api_base",
|
||||
[
|
||||
("embed-english-v2.0", None),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_cohere_embedding(sync_mode, model, api_base):
|
||||
try:
|
||||
# litellm.set_verbose=True
|
||||
data = {
|
||||
"model": model,
|
||||
"input": ["good morning from litellm", "this is another item"],
|
||||
"input_type": "search_query",
|
||||
"api_base": api_base,
|
||||
}
|
||||
if sync_mode:
|
||||
response = embedding(**data)
|
||||
else:
|
||||
response = await litellm.aembedding(**data)
|
||||
|
||||
print(f"response:", response)
|
||||
|
||||
assert isinstance(response.usage, litellm.Usage)
|
||||
except Exception as e:
|
||||
pytest.fail(f"Error occurred: {e}")
|
||||
|
||||
|
||||
# test_cohere_embedding()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,494 +0,0 @@
|
|||
"""Base class for LiteLLM integration tests.
|
||||
|
||||
Supports both local (mock) and remote testing modes via environment variables:
|
||||
- USE_LOCAL_LITELLM: When "true", uses local LiteLLM at localhost:4000 (default: false)
|
||||
- USE_MOCK_MODELS: When "true", uses mock model names (default: false)
|
||||
- LITELLM_API_KEY: API key for remote LiteLLM (required when USE_LOCAL_LITELLM=false)
|
||||
- LITELLM_BASE_URL: Base URL for remote LiteLLM (required when USE_LOCAL_LITELLM=false)
|
||||
"""
|
||||
|
||||
import enum
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from abc import ABC
|
||||
from collections import defaultdict
|
||||
from typing import Any, Callable, Dict, List, Tuple, Union
|
||||
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
import requests
|
||||
from urllib3.exceptions import InsecureRequestWarning
|
||||
|
||||
requests.packages.urllib3.disable_warnings(category=InsecureRequestWarning)
|
||||
|
||||
LOCAL_LITELLM_BASE_URL = "http://localhost:4000"
|
||||
LOCAL_MOCK_SERVER_URL = "http://localhost:8090"
|
||||
|
||||
if "USE_LOCAL_LITELLM" not in os.environ:
|
||||
os.environ["USE_LOCAL_LITELLM"] = "true"
|
||||
if "USE_MOCK_MODELS" not in os.environ:
|
||||
os.environ["USE_MOCK_MODELS"] = "true"
|
||||
if "USE_STATE_TRACKER" not in os.environ:
|
||||
os.environ["USE_STATE_TRACKER"] = "true"
|
||||
if "DATABASE_URL" not in os.environ:
|
||||
os.environ["DATABASE_URL"] = "postgresql://llmproxy:dbpassword9090@localhost:5432/litellm"
|
||||
|
||||
|
||||
def use_local_litellm() -> bool:
|
||||
return os.environ.get("USE_LOCAL_LITELLM", "false").lower() == "true"
|
||||
|
||||
|
||||
def use_remote_litellm() -> bool:
|
||||
return not use_local_litellm()
|
||||
|
||||
|
||||
def use_mock_models() -> bool:
|
||||
return os.environ.get("USE_MOCK_MODELS", "false").lower() == "true"
|
||||
|
||||
|
||||
def get_local_litellm_base_url() -> str:
|
||||
return LOCAL_LITELLM_BASE_URL
|
||||
|
||||
|
||||
def get_remote_litellm_base_url() -> str:
|
||||
return os.environ.get("LITELLM_BASE_URL", "").rstrip("/")
|
||||
|
||||
|
||||
def get_litellm_base_url() -> str:
|
||||
if use_local_litellm():
|
||||
return get_local_litellm_base_url()
|
||||
return get_remote_litellm_base_url()
|
||||
|
||||
|
||||
def get_litellm_api_key() -> str:
|
||||
if use_local_litellm():
|
||||
return "sk-1234"
|
||||
return os.environ.get("LITELLM_API_KEY", "")
|
||||
|
||||
|
||||
def get_mock_server_base_url() -> str:
|
||||
return LOCAL_MOCK_SERVER_URL
|
||||
|
||||
|
||||
def get_responses_model_name() -> str:
|
||||
if use_mock_models():
|
||||
return "openai-fake-gpt-4o"
|
||||
return "gpt-4o-mini-2024-07-18"
|
||||
|
||||
|
||||
def model_id(param) -> str:
|
||||
"""Generate a test ID from a model name or tuple containing model name.
|
||||
|
||||
Handles both:
|
||||
- String: "gpt-4o-mini" -> "gpt_4o_mini"
|
||||
- Tuple: ("gpt-4o", "openai/gpt-4o") -> "gpt_4o"
|
||||
"""
|
||||
if isinstance(param, tuple):
|
||||
name = param[0]
|
||||
else:
|
||||
name = param
|
||||
return name.replace("-", "_").replace(".", "_")
|
||||
|
||||
|
||||
def generate_test_id(
|
||||
params: Tuple[str, ...],
|
||||
test_name: str = "test",
|
||||
) -> str:
|
||||
"""Generate test ID from model parameters tuple.
|
||||
|
||||
Handles two tuple formats:
|
||||
- 6 elements: (provider, deployment, model_name, api_version, action, reason)
|
||||
- 7 elements: (provider, deployment, model_name, api_version, model_id, action, reason)
|
||||
|
||||
Uses model_id (position 4) if 7 elements, otherwise model_name (position 2).
|
||||
"""
|
||||
provider = params[0]
|
||||
deployment = params[1]
|
||||
api_version = params[3]
|
||||
|
||||
if len(params) == 7:
|
||||
identifier = params[4] # model_id
|
||||
else:
|
||||
identifier = params[2] # model_name
|
||||
|
||||
test_id = "/".join([provider, deployment, api_version, identifier, test_name])
|
||||
return test_id.replace("-", "_").replace(".", "_")
|
||||
|
||||
|
||||
class ModelTestAction(enum.Enum):
|
||||
NOT_APPLICABLE = 1
|
||||
SKIP = 2
|
||||
RUN = 3
|
||||
WARN_ON_FAIL = 4
|
||||
|
||||
def applicable(self) -> bool:
|
||||
return self.value != ModelTestAction.NOT_APPLICABLE.value
|
||||
|
||||
|
||||
class BaseLiteLLMIntegrationTest(ABC):
|
||||
"""Base class for all LiteLLM integration tests.
|
||||
|
||||
Supports both local/mock and remote testing based on environment variables.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def get_api_key() -> str:
|
||||
return get_litellm_api_key()
|
||||
|
||||
@staticmethod
|
||||
def get_base_url() -> str:
|
||||
return get_litellm_base_url()
|
||||
|
||||
@staticmethod
|
||||
def get_ca_bundle_path() -> str:
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
# change if needed
|
||||
|
||||
@classmethod
|
||||
def _get_ssl_verify_setting(cls) -> Union[bool, str]:
|
||||
"""Get the appropriate SSL verification setting based on mode.
|
||||
|
||||
Returns path string (not SSLContext) for compatibility with both
|
||||
requests and httpx libraries.
|
||||
"""
|
||||
if use_local_litellm():
|
||||
return False
|
||||
ca_bundle_path = cls.get_ca_bundle_path()
|
||||
if os.path.exists(ca_bundle_path):
|
||||
return ca_bundle_path
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def setup_class(cls):
|
||||
cls.api_key = cls.get_api_key()
|
||||
cls.base_url = cls.get_base_url()
|
||||
|
||||
if not cls.api_key:
|
||||
pytest.fail(
|
||||
"API key is not available. Set LITELLM_API_KEY or USE_LOCAL_LITELLM=true",
|
||||
)
|
||||
if not cls.base_url:
|
||||
pytest.fail(
|
||||
"Base URL is not available. Set LITELLM_BASE_URL or USE_LOCAL_LITELLM=true",
|
||||
)
|
||||
|
||||
verify_setting = cls._get_ssl_verify_setting()
|
||||
|
||||
if use_remote_litellm() and isinstance(verify_setting, str):
|
||||
os.environ["REQUESTS_CA_BUNDLE"] = verify_setting
|
||||
os.environ["CURL_CA_BUNDLE"] = verify_setting
|
||||
print(f"Using CA bundle: {verify_setting}")
|
||||
|
||||
cls.openai_client = openai.OpenAI(
|
||||
base_url=cls.base_url,
|
||||
api_key=cls.api_key,
|
||||
http_client=httpx.Client(verify=verify_setting),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def make_request(
|
||||
cls,
|
||||
method: str,
|
||||
endpoint: str,
|
||||
timeout_secs: int,
|
||||
**kwargs,
|
||||
) -> requests.Response:
|
||||
headers = kwargs.get("headers", {})
|
||||
headers["Authorization"] = f"Bearer {cls.api_key}"
|
||||
kwargs["headers"] = headers
|
||||
kwargs.setdefault("timeout", timeout_secs)
|
||||
kwargs.setdefault("verify", cls._get_ssl_verify_setting())
|
||||
|
||||
url = f"{cls.base_url}{endpoint}"
|
||||
return requests.request(method, url, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
def generate_request_id() -> str:
|
||||
return f"req-{uuid.uuid4().hex[:8]}"
|
||||
|
||||
@staticmethod
|
||||
def get_timeout_secs(model_name: str) -> int:
|
||||
model_lower = model_name.lower()
|
||||
slow_models = ["gpt-5", "gpt_5", "o1", "claude-opus", "claude_opus", "o3", "o4"]
|
||||
|
||||
if any(slow_model in model_lower for slow_model in slow_models):
|
||||
return 300
|
||||
return 60
|
||||
|
||||
@staticmethod
|
||||
def generate_unique_filename(extension: str = "txt") -> str:
|
||||
return f"test_{time.time()}.{extension}"
|
||||
|
||||
@staticmethod
|
||||
def extract_model_params(model_data: Dict[str, Any]) -> Tuple[str, str, str, str]:
|
||||
"""Extract standardized parameters from model data."""
|
||||
model_name = model_data.get("model_name", "")
|
||||
model_info = model_data.get("model_info", {})
|
||||
provider = model_info.get("litellm_provider", "unknown")
|
||||
litellm_params = model_data.get("litellm_params", {})
|
||||
|
||||
if provider == "azure":
|
||||
api_base = litellm_params.get("api_base", "unknown")
|
||||
if api_base != "unknown" and "//" in api_base:
|
||||
domain_name = api_base.split("//")[1]
|
||||
deployment = domain_name.split(".")[0]
|
||||
else:
|
||||
deployment = "unknown"
|
||||
api_version = litellm_params.get("api_version", "unknown")
|
||||
elif provider in ["bedrock", "bedrock_converse"]:
|
||||
deployment = litellm_params.get("aws_region_name", "unknown")
|
||||
api_version = "unknown"
|
||||
else:
|
||||
deployment = "unknown"
|
||||
api_version = "unknown"
|
||||
|
||||
return provider, deployment, model_name, api_version
|
||||
|
||||
@classmethod
|
||||
def _fetch_all_models_from_litellm(cls) -> List[Dict[str, Any]]:
|
||||
base_url = cls.get_base_url()
|
||||
api_key = cls.get_api_key()
|
||||
|
||||
if not api_key or not base_url:
|
||||
return []
|
||||
|
||||
verify_setting = cls._get_ssl_verify_setting()
|
||||
|
||||
response = requests.get(
|
||||
f"{base_url}/model/info",
|
||||
headers={"Authorization": f"Bearer {api_key}"},
|
||||
verify=verify_setting,
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise RuntimeError(
|
||||
f"Failed to fetch all models from {base_url}. Response code: {response.status_code}",
|
||||
)
|
||||
|
||||
data = response.json()
|
||||
return data.get("data", [])
|
||||
|
||||
@classmethod
|
||||
def _fetch_all_approved_models(cls) -> List[Dict[str, Any]]:
|
||||
return cls._fetch_all_models_from_litellm()
|
||||
|
||||
@classmethod
|
||||
def build_model_test_params(
|
||||
cls,
|
||||
should_skip_model: Callable[
|
||||
[str, str, str, str, Dict[str, Any]],
|
||||
Tuple["ModelTestAction", str],
|
||||
],
|
||||
include_model_id: bool = False,
|
||||
include_load_balanced: bool = False,
|
||||
) -> List[Tuple[str, ...]]:
|
||||
"""Build test parameters from all approved models.
|
||||
|
||||
Args:
|
||||
should_skip_model: Callback that determines if a model should be skipped.
|
||||
Signature: (provider, deployment, model_name, api_version, model_info) -> (action, reason)
|
||||
include_model_id: If True, includes model_id in tuple (7 elements), else 6 elements.
|
||||
include_load_balanced: If True, adds extra tests for load-balanced model groups.
|
||||
|
||||
Returns:
|
||||
List of tuples with model test parameters.
|
||||
- 6-element: (provider, deployment, model_name, api_version, action, reason)
|
||||
- 7-element: (provider, deployment, model_name, api_version, model_id, action, reason)
|
||||
"""
|
||||
models = cls._fetch_all_approved_models()
|
||||
test_params: List[Tuple[str, ...]] = []
|
||||
models_by_model_name: Dict[str, List[Tuple[str, ...]]] = defaultdict(list)
|
||||
|
||||
for model_data in models:
|
||||
model_info = model_data.get("model_info", {}) or {}
|
||||
|
||||
provider, deployment, model_name, api_version = cls.extract_model_params(
|
||||
model_data,
|
||||
)
|
||||
|
||||
model_test_action, model_test_action_reason = should_skip_model(
|
||||
provider,
|
||||
deployment,
|
||||
model_name,
|
||||
api_version,
|
||||
model_info,
|
||||
)
|
||||
|
||||
if model_test_action.applicable():
|
||||
if include_model_id:
|
||||
model_id = str(model_info.get("id"))
|
||||
params_tuple: Tuple[str, ...] = (
|
||||
provider,
|
||||
deployment,
|
||||
model_name,
|
||||
api_version,
|
||||
model_id,
|
||||
model_test_action,
|
||||
model_test_action_reason,
|
||||
)
|
||||
else:
|
||||
params_tuple = (
|
||||
provider,
|
||||
deployment,
|
||||
model_name,
|
||||
api_version,
|
||||
model_test_action,
|
||||
model_test_action_reason,
|
||||
)
|
||||
|
||||
test_params.append(params_tuple)
|
||||
|
||||
if include_load_balanced:
|
||||
models_by_model_name[model_name].append(params_tuple)
|
||||
|
||||
if include_load_balanced and include_model_id:
|
||||
for load_balanced_model_name, deployments in models_by_model_name.items():
|
||||
if len(deployments) <= 1:
|
||||
continue
|
||||
|
||||
first_deployment = deployments[0]
|
||||
test_params.append(
|
||||
(
|
||||
first_deployment[0], # provider
|
||||
"load_balanced",
|
||||
load_balanced_model_name,
|
||||
"load_balanced",
|
||||
load_balanced_model_name, # model_id = model_name for LB
|
||||
first_deployment[5], # model_test_action
|
||||
first_deployment[6], # model_test_action_reason
|
||||
),
|
||||
)
|
||||
|
||||
return test_params
|
||||
|
||||
|
||||
class UserKeyTestMixin:
|
||||
"""Mixin for tests that need to create users and API keys."""
|
||||
|
||||
allowed_routes: list[str] = []
|
||||
|
||||
_base_url: str = None
|
||||
_master_api_key: str = None
|
||||
admin_client: httpx.Client = None
|
||||
|
||||
@classmethod
|
||||
def setup_admin_client(cls):
|
||||
cls._base_url = get_litellm_base_url()
|
||||
cls._master_api_key = get_litellm_api_key()
|
||||
verify_setting = (
|
||||
False
|
||||
if use_local_litellm()
|
||||
else BaseLiteLLMIntegrationTest._get_ssl_verify_setting()
|
||||
)
|
||||
cls.admin_client = httpx.Client(base_url=cls._base_url, verify=verify_setting)
|
||||
|
||||
@classmethod
|
||||
def teardown_admin_client(cls):
|
||||
if cls.admin_client:
|
||||
cls.admin_client.close()
|
||||
|
||||
@staticmethod
|
||||
def unique_suffix() -> str:
|
||||
return f"{time.strftime('%Y%m%d%H%M%S')}{int(time.time() * 1000) % 1000:03d}"
|
||||
|
||||
@classmethod
|
||||
def create_user_and_key(cls, user_suffix: str) -> tuple[str, str, str]:
|
||||
user_email = f"test-user-{user_suffix}-{cls.unique_suffix()}@test.com"
|
||||
user_response = cls.admin_client.post(
|
||||
"/user/new",
|
||||
json={
|
||||
"user_email": user_email,
|
||||
"user_alias": user_email,
|
||||
"user_role": "internal_user",
|
||||
"auto_create_key": "false",
|
||||
},
|
||||
headers={
|
||||
"Authorization": f"Bearer {cls._master_api_key}",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
timeout=30,
|
||||
)
|
||||
assert user_response.status_code == 200, (
|
||||
f"Failed to create user: {user_response.status_code} - {user_response.text}"
|
||||
)
|
||||
user_id = user_response.json().get("user_id")
|
||||
|
||||
key_alias = user_email.replace("@", "-at-").replace(".", "-")
|
||||
key_response = cls.admin_client.post(
|
||||
"/key/generate",
|
||||
json={
|
||||
"user_id": user_id,
|
||||
"key_alias": key_alias,
|
||||
"allowed_routes": cls.allowed_routes,
|
||||
},
|
||||
headers={
|
||||
"Authorization": f"Bearer {cls._master_api_key}",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
timeout=30,
|
||||
)
|
||||
assert key_response.status_code == 200, (
|
||||
f"Failed to create key: {key_response.status_code} - {key_response.text}"
|
||||
)
|
||||
api_key = key_response.json().get("key")
|
||||
|
||||
print(f"Created user {user_email}")
|
||||
return user_id, api_key, user_email
|
||||
|
||||
@classmethod
|
||||
def create_user_key_and_client(
|
||||
cls,
|
||||
user_suffix: str,
|
||||
) -> tuple[str, str, str, openai.OpenAI]:
|
||||
user_id, api_key, user_email = cls.create_user_and_key(user_suffix)
|
||||
verify_setting = (
|
||||
False
|
||||
if use_local_litellm()
|
||||
else BaseLiteLLMIntegrationTest._get_ssl_verify_setting()
|
||||
)
|
||||
client = openai.OpenAI(
|
||||
base_url=cls._base_url,
|
||||
api_key=api_key,
|
||||
http_client=httpx.Client(verify=verify_setting),
|
||||
)
|
||||
return user_id, api_key, user_email, client
|
||||
|
||||
@classmethod
|
||||
def create_key_and_client(
|
||||
cls,
|
||||
user_id: str,
|
||||
key_suffix: str,
|
||||
) -> tuple[str, openai.OpenAI]:
|
||||
key_alias = f"additional-key-{key_suffix}-{cls.unique_suffix()}"
|
||||
key_response = cls.admin_client.post(
|
||||
"/key/generate",
|
||||
json={
|
||||
"user_id": user_id,
|
||||
"key_alias": key_alias,
|
||||
"allowed_routes": cls.allowed_routes,
|
||||
},
|
||||
headers={
|
||||
"Authorization": f"Bearer {cls._master_api_key}",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
timeout=30,
|
||||
)
|
||||
assert key_response.status_code == 200, (
|
||||
f"Failed to create additional key: {key_response.status_code} - {key_response.text}"
|
||||
)
|
||||
api_key = key_response.json().get("key")
|
||||
verify_setting = (
|
||||
False
|
||||
if use_local_litellm()
|
||||
else BaseLiteLLMIntegrationTest._get_ssl_verify_setting()
|
||||
)
|
||||
client = openai.OpenAI(
|
||||
base_url=cls._base_url,
|
||||
api_key=api_key,
|
||||
http_client=httpx.Client(verify=verify_setting),
|
||||
)
|
||||
print(f"Created additional key for user {user_id}")
|
||||
return api_key, client
|
||||
|
|
@ -1,311 +0,0 @@
|
|||
"""
|
||||
Pytest configuration for Azure Batch E2E Tests.
|
||||
|
||||
This conftest manages:
|
||||
1. Mock Azure Batch server (FastAPI on port 8090)
|
||||
2. LiteLLM proxy server (port 4000)
|
||||
3. PostgreSQL database setup
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Generator
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
_test_dir = Path(__file__).parent
|
||||
sys.path.insert(0, str(_test_dir.parent.parent)) # litellm root
|
||||
sys.path.insert(0, str(_test_dir)) # test directory for local imports
|
||||
|
||||
LOG_DIR = _test_dir
|
||||
|
||||
|
||||
def pytest_configure(config):
|
||||
"""Ensure test directory is in Python path before collection."""
|
||||
test_dir = Path(__file__).parent
|
||||
if str(test_dir) not in sys.path:
|
||||
sys.path.insert(0, str(test_dir))
|
||||
|
||||
|
||||
MOCK_SERVER_PORT = 8090
|
||||
MOCK_SERVER_URL = f"http://localhost:{MOCK_SERVER_PORT}"
|
||||
LITELLM_PROXY_PORT = 4000
|
||||
LITELLM_PROXY_URL = f"http://localhost:{LITELLM_PROXY_PORT}"
|
||||
DATABASE_URL = "postgresql://llmproxy:dbpassword9090@localhost:5432/litellm"
|
||||
|
||||
|
||||
def kill_process_on_port(port: int) -> None:
|
||||
"""Kill any process using the specified port."""
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["lsof", "-ti", f":{port}"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=5,
|
||||
)
|
||||
if result.stdout.strip():
|
||||
pids = result.stdout.strip().split("\n")
|
||||
for pid in pids:
|
||||
try:
|
||||
subprocess.run(["kill", "-9", pid.strip()], timeout=5)
|
||||
except Exception:
|
||||
pass
|
||||
time.sleep(1)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def wait_for_server(url: str, max_attempts: int = 30, delay: float = 1.0) -> bool:
|
||||
"""Wait for a server to become available at url/health.
|
||||
|
||||
Any HTTP response (including 401) means the server is up.
|
||||
Only connection errors count as "not ready yet".
|
||||
"""
|
||||
for attempt in range(max_attempts):
|
||||
try:
|
||||
response = httpx.get(f"{url}/health", timeout=2.0)
|
||||
return True
|
||||
except (httpx.ConnectError, httpx.TimeoutException, httpx.NetworkError):
|
||||
pass
|
||||
except Exception:
|
||||
pass
|
||||
if attempt < max_attempts - 1:
|
||||
time.sleep(delay)
|
||||
return False
|
||||
|
||||
|
||||
def _read_log_tail(log_path: Path, max_lines: int = 80) -> str:
|
||||
"""Read the last N lines of a log file, returning empty string if not found."""
|
||||
if not log_path.exists():
|
||||
return "(log file not found)"
|
||||
try:
|
||||
text = log_path.read_text()
|
||||
lines = text.strip().splitlines()
|
||||
if len(lines) > max_lines:
|
||||
return f"... ({len(lines) - max_lines} lines truncated) ...\n" + "\n".join(
|
||||
lines[-max_lines:]
|
||||
)
|
||||
return text
|
||||
except Exception as e:
|
||||
return f"(error reading log: {e})"
|
||||
|
||||
|
||||
def _check_process_alive(process: subprocess.Popen, label: str, log_path: Path):
|
||||
"""Check if a subprocess crashed immediately after starting.
|
||||
Raises pytest.fail with log output if the process has already exited.
|
||||
"""
|
||||
time.sleep(1)
|
||||
exit_code = process.poll()
|
||||
if exit_code is not None:
|
||||
log_output = _read_log_tail(log_path)
|
||||
pytest.fail(
|
||||
f"{label} exited immediately with code {exit_code}.\n"
|
||||
f"--- {label} log ({log_path}) ---\n{log_output}\n"
|
||||
f"--- end log ---"
|
||||
)
|
||||
|
||||
|
||||
def setup_database() -> bool:
|
||||
"""Ensure PostgreSQL database exists and is accessible."""
|
||||
try:
|
||||
import psycopg2
|
||||
|
||||
conn = psycopg2.connect(
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="litellm",
|
||||
user="llmproxy",
|
||||
password="dbpassword9090",
|
||||
connect_timeout=5,
|
||||
)
|
||||
conn.close()
|
||||
return True
|
||||
except ImportError:
|
||||
print("WARNING: psycopg2 not installed — cannot verify database")
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def mock_azure_server() -> Generator[str, None, None]:
|
||||
"""Start mock Azure batch server as a subprocess."""
|
||||
print(f"\n{'=' * 60}")
|
||||
print("Setting up Mock Azure Batch Server")
|
||||
print(f"{'=' * 60}")
|
||||
|
||||
kill_process_on_port(MOCK_SERVER_PORT)
|
||||
|
||||
runner_script = Path(__file__).parent / "fixtures" / "run_mock_server.py"
|
||||
runner_script.write_text(
|
||||
"""
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
from fixtures.mock_azure_batch_server import create_mock_azure_batch_server
|
||||
import uvicorn
|
||||
|
||||
if __name__ == "__main__":
|
||||
app = create_mock_azure_batch_server()
|
||||
uvicorn.run(app, host="0.0.0.0", port=8090, log_level="info", access_log=False)
|
||||
"""
|
||||
)
|
||||
|
||||
mock_log = LOG_DIR / "mock_server.log"
|
||||
log_file = open(mock_log, "w")
|
||||
|
||||
print(f"Starting mock server on port {MOCK_SERVER_PORT}...")
|
||||
print(f"Log file: {mock_log}")
|
||||
process = subprocess.Popen(
|
||||
[sys.executable, str(runner_script)],
|
||||
stdout=log_file,
|
||||
stderr=subprocess.STDOUT,
|
||||
cwd=Path(__file__).parent,
|
||||
)
|
||||
|
||||
_check_process_alive(process, "Mock server", mock_log)
|
||||
|
||||
if not wait_for_server(MOCK_SERVER_URL, max_attempts=30, delay=1.0):
|
||||
log_output = _read_log_tail(mock_log)
|
||||
exit_code = process.poll()
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
process.wait()
|
||||
log_file.close()
|
||||
pytest.fail(
|
||||
f"Mock server failed to start on port {MOCK_SERVER_PORT} "
|
||||
f"(process exit_code={exit_code}).\n"
|
||||
f"--- mock server log ---\n{log_output}\n--- end log ---\n"
|
||||
f"Hint: ensure 'uvicorn' and 'fastapi' are installed."
|
||||
)
|
||||
|
||||
print(f"Mock Azure server ready at {MOCK_SERVER_URL}")
|
||||
yield MOCK_SERVER_URL
|
||||
|
||||
print("\nShutting down mock server...")
|
||||
try:
|
||||
process.terminate()
|
||||
process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
process.wait()
|
||||
log_file.close()
|
||||
print("Mock server stopped")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def litellm_proxy_server(mock_azure_server: str) -> Generator[str, None, None]:
|
||||
"""Start LiteLLM proxy server for the test session."""
|
||||
print(f"\n{'=' * 60}")
|
||||
print("Setting up LiteLLM Proxy Server")
|
||||
print(f"{'=' * 60}")
|
||||
|
||||
if not setup_database():
|
||||
pytest.skip(
|
||||
"PostgreSQL database not available at localhost:5432. "
|
||||
"Start PostgreSQL and create a 'litellm' database:\n"
|
||||
" docker run -d --name litellm-db -p 5432:5432 "
|
||||
'-e POSTGRES_USER=llmproxy -e POSTGRES_PASSWORD=dbpassword9090 '
|
||||
"-e POSTGRES_DB=litellm postgres:15\n"
|
||||
"Then run: prisma db push --schema=litellm/proxy/schema.prisma"
|
||||
)
|
||||
print("Database connection verified")
|
||||
|
||||
config_path = Path(__file__).parent / "fixtures" / "config.yml"
|
||||
if not config_path.exists():
|
||||
pytest.fail(f"Config file not found: {config_path}")
|
||||
print("Config file found")
|
||||
|
||||
kill_process_on_port(LITELLM_PROXY_PORT)
|
||||
|
||||
os.environ["MOCK_SERVER_URL_V1"] = f"{mock_azure_server}/v1"
|
||||
os.environ["MOCK_SERVER_URL_OPENAI_V1"] = f"{mock_azure_server}/openai/v1"
|
||||
os.environ["DATABASE_URL"] = DATABASE_URL
|
||||
os.environ["USE_LOCAL_LITELLM"] = "true"
|
||||
os.environ["USE_MOCK_MODELS"] = "true"
|
||||
os.environ["USE_STATE_TRACKER"] = "true"
|
||||
os.environ["PROXY_BATCH_POLLING_INTERVAL"] = "10"
|
||||
|
||||
print("Environment configured")
|
||||
|
||||
print(f"Starting LiteLLM proxy on port {LITELLM_PROXY_PORT}...")
|
||||
litellm_root = Path(__file__).parent.parent.parent
|
||||
|
||||
cmd = [
|
||||
sys.executable,
|
||||
"-m",
|
||||
"litellm.proxy.proxy_cli",
|
||||
"--config",
|
||||
str(config_path),
|
||||
"--port",
|
||||
str(LITELLM_PROXY_PORT),
|
||||
"--detailed_debug",
|
||||
]
|
||||
|
||||
proxy_log = LOG_DIR / "proxy_server.log"
|
||||
log_file = open(proxy_log, "w")
|
||||
print(f"Log file: {proxy_log}")
|
||||
|
||||
process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=log_file,
|
||||
stderr=subprocess.STDOUT,
|
||||
env=os.environ.copy(),
|
||||
cwd=litellm_root,
|
||||
)
|
||||
|
||||
_check_process_alive(process, "LiteLLM proxy", proxy_log)
|
||||
|
||||
if not wait_for_server(LITELLM_PROXY_URL, max_attempts=60, delay=1.0):
|
||||
log_output = _read_log_tail(proxy_log)
|
||||
exit_code = process.poll()
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
process.wait()
|
||||
log_file.close()
|
||||
pytest.fail(
|
||||
f"LiteLLM proxy failed to start on port {LITELLM_PROXY_PORT} "
|
||||
f"(process exit_code={exit_code}).\n"
|
||||
f"--- proxy log (last 80 lines) ---\n{log_output}\n--- end log ---\n"
|
||||
f"Hints:\n"
|
||||
f" 1. Ensure Prisma client is generated: "
|
||||
f"cd {litellm_root} && prisma generate --schema=litellm/proxy/schema.prisma\n"
|
||||
f" 2. Ensure DB migrations are applied: "
|
||||
f"prisma db push --schema=litellm/proxy/schema.prisma\n"
|
||||
f" 3. Check the full log at: {proxy_log}"
|
||||
)
|
||||
|
||||
print(f"LiteLLM proxy ready at {LITELLM_PROXY_URL}")
|
||||
yield LITELLM_PROXY_URL
|
||||
|
||||
print("\nShutting down LiteLLM proxy...")
|
||||
try:
|
||||
process.terminate()
|
||||
process.wait(timeout=10)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
process.wait()
|
||||
log_file.close()
|
||||
print("LiteLLM proxy stopped")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def event_loop():
|
||||
"""Provide an event loop for async tests."""
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
yield loop
|
||||
loop.close()
|
||||
|
|
@ -1,56 +0,0 @@
|
|||
model_list:
|
||||
- model_name: openai-fake-gpt-3.5-turbo
|
||||
litellm_params:
|
||||
model: openai/openai-fake-gpt-3.5-turbo
|
||||
api_base: os.environ/MOCK_SERVER_URL_V1
|
||||
api_key: fake-key
|
||||
- model_name: openai-fake-gpt-4
|
||||
litellm_params:
|
||||
model: openai/openai-fake-gpt-4
|
||||
api_base: os.environ/MOCK_SERVER_URL_V1
|
||||
api_key: fake-key
|
||||
- model_name: openai-fake-gpt-4o
|
||||
litellm_params:
|
||||
model: openai/openai-fake-gpt-4o
|
||||
api_base: os.environ/MOCK_SERVER_URL_V1
|
||||
api_key: fake-key
|
||||
- model_name: fake-text-embedding-3-small
|
||||
litellm_params:
|
||||
model: openai/fake-text-embedding-3-small
|
||||
api_base: os.environ/MOCK_SERVER_URL_V1
|
||||
api_key: fake-key
|
||||
- model_name: o3-mini-batch-2025-01-31
|
||||
litellm_params:
|
||||
model: openai/o3-mini-batch-2025-01-31
|
||||
api_base: os.environ/MOCK_SERVER_URL_OPENAI_V1
|
||||
api_key: fake-key
|
||||
model_info:
|
||||
mode: batch
|
||||
- model_name: azure-fake-gpt-5-batch-2025-08-07
|
||||
litellm_params:
|
||||
api_base: http://0.0.0.0:8090
|
||||
api_key: fake-key
|
||||
api_version: 2025-03-01-preview
|
||||
base_model: azure/gpt-5
|
||||
model: azure/gpt-5-mini
|
||||
custom_llm_provider: azure
|
||||
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
database_url: os.environ/DATABASE_URL
|
||||
proxy_batch_polling_interval: 10
|
||||
|
||||
litellm_settings:
|
||||
drop_params: true
|
||||
set_verbose: true
|
||||
json_logs: true
|
||||
# S3 callback for batch completion logging (points to mock server)
|
||||
callbacks: ["s3_v2"]
|
||||
s3_callback_params:
|
||||
s3_bucket_name: litellm-test-bucket
|
||||
s3_region_name: us-east-1
|
||||
s3_endpoint_url: http://0.0.0.0:8090
|
||||
s3_aws_access_key_id: fake-key
|
||||
s3_aws_secret_access_key: fake-secret
|
||||
s3_use_ssl: false
|
||||
s3_verify: false
|
||||
|
|
@ -1,3 +0,0 @@
|
|||
from .server import create_mock_azure_batch_server
|
||||
|
||||
__all__ = ["create_mock_azure_batch_server"]
|
||||
|
|
@ -1,517 +0,0 @@
|
|||
import asyncio
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Query, Request, UploadFile
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class FileObject(BaseModel):
|
||||
id: str
|
||||
object: str = "file"
|
||||
bytes: int
|
||||
created_at: int
|
||||
filename: str
|
||||
purpose: str
|
||||
status: str = "processed"
|
||||
status_details: Optional[str] = None
|
||||
expires_at: Optional[int] = None
|
||||
|
||||
|
||||
class BatchObject(BaseModel):
|
||||
id: str
|
||||
object: str = "batch"
|
||||
endpoint: str
|
||||
errors: Optional[Dict] = None
|
||||
input_file_id: str
|
||||
completion_window: str
|
||||
status: str
|
||||
output_file_id: Optional[str] = None
|
||||
error_file_id: Optional[str] = None
|
||||
created_at: int
|
||||
in_progress_at: Optional[int] = None
|
||||
expires_at: Optional[int] = None
|
||||
finalizing_at: Optional[int] = None
|
||||
completed_at: Optional[int] = None
|
||||
failed_at: Optional[int] = None
|
||||
expired_at: Optional[int] = None
|
||||
cancelling_at: Optional[int] = None
|
||||
cancelled_at: Optional[int] = None
|
||||
request_counts: Optional[Dict[str, int]] = None
|
||||
metadata: Optional[Dict] = None
|
||||
|
||||
|
||||
class BatchListResponse(BaseModel):
|
||||
object: str = "list"
|
||||
data: List[Dict]
|
||||
first_id: Optional[str] = None
|
||||
last_id: Optional[str] = None
|
||||
has_more: bool = False
|
||||
|
||||
|
||||
file_storage: Dict[str, Dict] = {}
|
||||
batch_storage: Dict[str, BatchObject] = {}
|
||||
batch_results: Dict[str, List[Dict]] = {}
|
||||
|
||||
PROCESSING_DELAY_SECONDS = float(1)
|
||||
VALIDATING_DELAY_SECONDS = float(3)
|
||||
|
||||
|
||||
async def process_batch(batch_id: str):
|
||||
logger.info(f"Starting batch processing for {batch_id}")
|
||||
try:
|
||||
batch = batch_storage[batch_id]
|
||||
|
||||
await asyncio.sleep(VALIDATING_DELAY_SECONDS)
|
||||
batch.status = "in_progress"
|
||||
batch.in_progress_at = int(time.time())
|
||||
logger.info(f"Batch {batch_id} status: in_progress")
|
||||
|
||||
await process_batch_requests(batch_id)
|
||||
await asyncio.sleep(PROCESSING_DELAY_SECONDS)
|
||||
|
||||
batch.status = "finalizing"
|
||||
batch.finalizing_at = int(time.time())
|
||||
logger.info(f"Batch {batch_id} status: finalizing")
|
||||
await asyncio.sleep(PROCESSING_DELAY_SECONDS)
|
||||
|
||||
await create_output_file(batch_id)
|
||||
|
||||
batch.status = "completed"
|
||||
batch.completed_at = int(time.time())
|
||||
logger.info(f"Batch {batch_id} status: completed")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Batch {batch_id} failed: {e}")
|
||||
batch = batch_storage[batch_id]
|
||||
batch.status = "failed"
|
||||
batch.failed_at = int(time.time())
|
||||
batch.errors = {
|
||||
"object": "list",
|
||||
"data": [{"code": "processing_error", "message": str(e)}],
|
||||
}
|
||||
|
||||
|
||||
async def process_batch_requests(batch_id: str):
|
||||
batch = batch_storage[batch_id]
|
||||
input_file = file_storage[batch.input_file_id]
|
||||
|
||||
requests = []
|
||||
for line in input_file["content"].split("\n"):
|
||||
if line.strip():
|
||||
try:
|
||||
requests.append(json.loads(line))
|
||||
except json.JSONDecodeError as e:
|
||||
logger.warning(f"Invalid JSON line in batch {batch_id}: {e}")
|
||||
|
||||
logger.info(f"Batch {batch_id} has {len(requests)} requests")
|
||||
|
||||
results = []
|
||||
failed_count = 0
|
||||
for req in requests:
|
||||
result = await process_single_request(req)
|
||||
if result.get("error"):
|
||||
failed_count += 1
|
||||
results.append(result)
|
||||
|
||||
batch_results[batch_id] = results
|
||||
batch.request_counts = {
|
||||
"total": len(requests),
|
||||
"completed": len(results) - failed_count,
|
||||
"failed": failed_count,
|
||||
}
|
||||
|
||||
|
||||
async def process_single_request(request_data: Dict) -> Dict:
|
||||
custom_id = request_data.get("custom_id")
|
||||
url = request_data.get("url", "/v1/chat/completions")
|
||||
body = request_data.get("body", {})
|
||||
|
||||
if "/chat/completions" in url:
|
||||
response_body = {
|
||||
"id": f"chatcmpl-{uuid.uuid4().hex}",
|
||||
"object": "chat.completion",
|
||||
"created": int(time.time()),
|
||||
"model": body.get("model", "gpt-4o"),
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "Mock batch response."},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
|
||||
}
|
||||
status_code = 200
|
||||
else:
|
||||
response_body = {"error": {"message": f"Unsupported endpoint: {url}"}}
|
||||
status_code = 400
|
||||
|
||||
return {
|
||||
"id": f"batch_req_{uuid.uuid4().hex[:12]}",
|
||||
"custom_id": custom_id,
|
||||
"response": {
|
||||
"status_code": status_code,
|
||||
"request_id": f"req_{uuid.uuid4().hex[:12]}",
|
||||
"body": response_body,
|
||||
},
|
||||
"error": None,
|
||||
}
|
||||
|
||||
|
||||
async def create_output_file(batch_id: str):
|
||||
results = batch_results.get(batch_id, [])
|
||||
output_lines = [json.dumps(result) for result in results]
|
||||
output_content = "\n".join(output_lines)
|
||||
|
||||
output_file_id = f"file-batch-output-{uuid.uuid4().hex[:12]}"
|
||||
file_storage[output_file_id] = {
|
||||
"content": output_content,
|
||||
"filename": f"batch_output_{batch_id}.jsonl",
|
||||
"purpose": "batch_output",
|
||||
"bytes": len(output_content.encode()),
|
||||
"created_at": int(time.time()),
|
||||
}
|
||||
|
||||
batch = batch_storage[batch_id]
|
||||
batch.output_file_id = output_file_id
|
||||
logger.info(f"Created output file {output_file_id} for batch {batch_id}")
|
||||
|
||||
|
||||
def validate_batch_input(content: str) -> tuple[bool, str, List[Dict]]:
|
||||
requests = []
|
||||
custom_ids = set()
|
||||
|
||||
lines = content.strip().split("\n")
|
||||
if not lines or all(not line.strip() for line in lines):
|
||||
return False, "empty_batch", []
|
||||
|
||||
for line_num, line in enumerate(lines, 1):
|
||||
if not line.strip():
|
||||
continue
|
||||
try:
|
||||
req = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
return False, "invalid_json_line", []
|
||||
|
||||
for field in ["custom_id", "method", "url", "body"]:
|
||||
if field not in req:
|
||||
return False, "invalid_request", []
|
||||
|
||||
if req["custom_id"] in custom_ids:
|
||||
return False, "duplicate_custom_id", []
|
||||
custom_ids.add(req["custom_id"])
|
||||
|
||||
requests.append(req)
|
||||
|
||||
if len(requests) > 100000:
|
||||
return False, "too_many_tasks", []
|
||||
|
||||
return True, "", requests
|
||||
|
||||
|
||||
def setup_batch_routes(app: FastAPI):
|
||||
# Files endpoints (OpenAI and Azure paths)
|
||||
@app.post("/openai/v1/files")
|
||||
@app.post("/openai/files")
|
||||
@app.post("/v1/files")
|
||||
@app.post("/files")
|
||||
async def create_file(request: Request):
|
||||
form = await request.form()
|
||||
logger.info(f"File upload form fields: {list(form.keys())}")
|
||||
|
||||
file: UploadFile = form.get("file")
|
||||
purpose: str = form.get("purpose", "batch")
|
||||
|
||||
if not file:
|
||||
raise HTTPException(status_code=400, detail="No file provided")
|
||||
|
||||
logger.info(f"Uploading file: {file.filename}, purpose: {purpose}")
|
||||
|
||||
content = await file.read()
|
||||
content_str = content.decode("utf-8")
|
||||
|
||||
file_id = f"file-{uuid.uuid4().hex[:24]}"
|
||||
created_at = int(time.time())
|
||||
|
||||
expires_at = None
|
||||
expires_after_seconds = form.get("expires_after[seconds]")
|
||||
if expires_after_seconds:
|
||||
try:
|
||||
seconds = int(expires_after_seconds)
|
||||
logger.info(f"expires_after[seconds] = {seconds}")
|
||||
if seconds < 259200 or seconds > 2592000:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": {
|
||||
"code": "invalidPayload",
|
||||
"message": "Value for Seconds must be between 259200 and 2592000.",
|
||||
},
|
||||
},
|
||||
)
|
||||
expires_at = created_at + seconds
|
||||
logger.info(f"Calculated expires_at: {expires_at}")
|
||||
except ValueError as e:
|
||||
logger.warning(f"Failed to parse expires_after[seconds]: {e}")
|
||||
|
||||
file_storage[file_id] = {
|
||||
"content": content_str,
|
||||
"filename": file.filename or "batch_input.jsonl",
|
||||
"purpose": purpose,
|
||||
"bytes": len(content),
|
||||
"created_at": created_at,
|
||||
"expires_at": expires_at,
|
||||
}
|
||||
|
||||
logger.info(f"Created file {file_id}, expires_at={expires_at}")
|
||||
return FileObject(
|
||||
id=file_id,
|
||||
bytes=len(content),
|
||||
created_at=created_at,
|
||||
filename=file.filename or "batch_input.jsonl",
|
||||
purpose=purpose,
|
||||
expires_at=expires_at,
|
||||
).model_dump()
|
||||
|
||||
@app.get("/openai/v1/files/{file_id}")
|
||||
@app.get("/openai/files/{file_id}")
|
||||
@app.get("/v1/files/{file_id}")
|
||||
@app.get("/files/{file_id}")
|
||||
async def get_file(file_id: str):
|
||||
logger.info(f"Getting file: {file_id}")
|
||||
if file_id not in file_storage:
|
||||
raise HTTPException(status_code=404, detail="File not found")
|
||||
|
||||
file_data = file_storage[file_id]
|
||||
return FileObject(
|
||||
id=file_id,
|
||||
bytes=file_data["bytes"],
|
||||
created_at=file_data["created_at"],
|
||||
filename=file_data["filename"],
|
||||
purpose=file_data["purpose"],
|
||||
expires_at=file_data.get("expires_at"),
|
||||
).model_dump()
|
||||
|
||||
@app.get("/openai/v1/files/{file_id}/content")
|
||||
@app.get("/openai/files/{file_id}/content")
|
||||
@app.get("/v1/files/{file_id}/content")
|
||||
@app.get("/files/{file_id}/content")
|
||||
async def get_file_content(file_id: str):
|
||||
logger.info(f"Getting file content: {file_id}")
|
||||
if file_id not in file_storage:
|
||||
raise HTTPException(status_code=404, detail="File not found")
|
||||
|
||||
file_data = file_storage[file_id]
|
||||
content = file_data["content"]
|
||||
|
||||
return StreamingResponse(
|
||||
io.StringIO(content),
|
||||
media_type="application/octet-stream",
|
||||
headers={
|
||||
"Content-Disposition": f"attachment; filename={file_data['filename']}",
|
||||
},
|
||||
)
|
||||
|
||||
@app.delete("/openai/v1/files/{file_id}")
|
||||
@app.delete("/openai/files/{file_id}")
|
||||
@app.delete("/v1/files/{file_id}")
|
||||
@app.delete("/files/{file_id}")
|
||||
async def delete_file(file_id: str):
|
||||
logger.info(f"Deleting file: {file_id}")
|
||||
if file_id not in file_storage:
|
||||
raise HTTPException(status_code=404, detail="File not found")
|
||||
|
||||
del file_storage[file_id]
|
||||
return {"id": file_id, "object": "file", "deleted": True}
|
||||
|
||||
@app.get("/openai/v1/files")
|
||||
@app.get("/openai/files")
|
||||
@app.get("/v1/files")
|
||||
@app.get("/files")
|
||||
async def list_files(
|
||||
purpose: Optional[str] = None,
|
||||
limit: int = Query(10000, le=10000),
|
||||
):
|
||||
logger.info(f"Listing files, purpose: {purpose}, limit: {limit}")
|
||||
files = []
|
||||
for file_id, file_data in file_storage.items():
|
||||
if purpose is None or file_data.get("purpose") == purpose:
|
||||
files.append(
|
||||
FileObject(
|
||||
id=file_id,
|
||||
bytes=file_data["bytes"],
|
||||
created_at=file_data["created_at"],
|
||||
filename=file_data["filename"],
|
||||
purpose=file_data["purpose"],
|
||||
expires_at=file_data.get("expires_at"),
|
||||
).model_dump(),
|
||||
)
|
||||
return {"object": "list", "data": files[:limit]}
|
||||
|
||||
# Batches endpoints (OpenAI and Azure paths)
|
||||
@app.post("/openai/v1/batches")
|
||||
@app.post("/openai/batches")
|
||||
@app.post("/v1/batches")
|
||||
@app.post("/batches")
|
||||
async def create_batch(request_data: dict):
|
||||
input_file_id = request_data.get("input_file_id")
|
||||
endpoint = request_data.get("endpoint", "/v1/chat/completions")
|
||||
completion_window = request_data.get("completion_window", "24h")
|
||||
metadata = request_data.get("metadata", {})
|
||||
output_expires_after = request_data.get("output_expires_after")
|
||||
|
||||
logger.info(
|
||||
f"Creating batch with input_file: {input_file_id}, endpoint: {endpoint}, output_expires_after: {output_expires_after}",
|
||||
)
|
||||
|
||||
if not input_file_id or input_file_id not in file_storage:
|
||||
raise HTTPException(status_code=400, detail="Input file not found")
|
||||
|
||||
input_file = file_storage[input_file_id]
|
||||
is_valid, error_code, _ = validate_batch_input(input_file["content"])
|
||||
if not is_valid:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": {
|
||||
"code": error_code,
|
||||
"message": f"Validation failed: {error_code}",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
batch_id = f"batch_{uuid.uuid4()}"
|
||||
created_at = int(time.time())
|
||||
|
||||
if output_expires_after:
|
||||
seconds = (
|
||||
output_expires_after.get("seconds", 0)
|
||||
if isinstance(output_expires_after, dict)
|
||||
else 0
|
||||
)
|
||||
expires_at = created_at + seconds
|
||||
logger.info(
|
||||
f"Using output_expires_after: {seconds}s, expires_at: {expires_at}",
|
||||
)
|
||||
elif completion_window == "24h":
|
||||
expires_at = created_at + (24 * 60 * 60)
|
||||
else:
|
||||
expires_at = created_at + (24 * 60 * 60)
|
||||
|
||||
batch = BatchObject(
|
||||
id=batch_id,
|
||||
endpoint=endpoint,
|
||||
input_file_id=input_file_id,
|
||||
completion_window=completion_window,
|
||||
status="validating",
|
||||
created_at=created_at,
|
||||
expires_at=expires_at,
|
||||
request_counts={"total": 0, "completed": 0, "failed": 0},
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
batch_storage[batch_id] = batch
|
||||
logger.info(f"Created batch {batch_id}")
|
||||
|
||||
asyncio.create_task(process_batch(batch_id))
|
||||
|
||||
return batch.model_dump()
|
||||
|
||||
@app.get("/openai/v1/batches/{batch_id}")
|
||||
@app.get("/openai/batches/{batch_id}")
|
||||
@app.get("/v1/batches/{batch_id}")
|
||||
@app.get("/batches/{batch_id}")
|
||||
async def get_batch(batch_id: str):
|
||||
logger.info(f"Getting batch: {batch_id}")
|
||||
if batch_id not in batch_storage:
|
||||
raise HTTPException(status_code=404, detail="Batch not found")
|
||||
|
||||
return batch_storage[batch_id].model_dump()
|
||||
|
||||
@app.get("/openai/v1/batches")
|
||||
@app.get("/openai/batches")
|
||||
@app.get("/v1/batches")
|
||||
@app.get("/batches")
|
||||
async def list_batches(
|
||||
after: Optional[str] = Query(None),
|
||||
limit: int = Query(20, le=100),
|
||||
):
|
||||
logger.info(f"Listing batches, after: {after}, limit: {limit}")
|
||||
batches = list(batch_storage.values())
|
||||
batches.sort(key=lambda x: x.created_at, reverse=True)
|
||||
|
||||
if after:
|
||||
after_index = next((i for i, b in enumerate(batches) if b.id == after), -1)
|
||||
if after_index >= 0:
|
||||
batches = batches[after_index + 1 :]
|
||||
|
||||
batches = batches[:limit]
|
||||
|
||||
return BatchListResponse(
|
||||
data=[batch.model_dump() for batch in batches],
|
||||
first_id=batches[0].id if batches else None,
|
||||
last_id=batches[-1].id if batches else None,
|
||||
has_more=len(batches) == limit,
|
||||
).model_dump()
|
||||
|
||||
@app.post("/openai/v1/batches/{batch_id}/cancel")
|
||||
@app.post("/openai/batches/{batch_id}/cancel")
|
||||
@app.post("/v1/batches/{batch_id}/cancel")
|
||||
@app.post("/batches/{batch_id}/cancel")
|
||||
async def cancel_batch(batch_id: str):
|
||||
logger.info(f"Cancelling batch: {batch_id}")
|
||||
if batch_id not in batch_storage:
|
||||
raise HTTPException(status_code=404, detail="Batch not found")
|
||||
|
||||
batch = batch_storage[batch_id]
|
||||
if batch.status in ["completed", "failed", "cancelled", "expired"]:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Cannot cancel batch in {batch.status} status",
|
||||
)
|
||||
|
||||
batch.status = "cancelled"
|
||||
batch.cancelled_at = int(time.time())
|
||||
logger.info(f"Batch {batch_id} cancelled")
|
||||
|
||||
return batch.model_dump()
|
||||
|
||||
# Debug endpoints
|
||||
@app.get("/debug/batches")
|
||||
async def debug_list_batches():
|
||||
return {
|
||||
"batches": {
|
||||
batch_id: batch.model_dump()
|
||||
for batch_id, batch in batch_storage.items()
|
||||
},
|
||||
"files": {
|
||||
file_id: {k: v for k, v in data.items() if k != "content"}
|
||||
for file_id, data in file_storage.items()
|
||||
},
|
||||
}
|
||||
|
||||
@app.post("/reset")
|
||||
@app.post("/debug/clear")
|
||||
async def reset_all():
|
||||
file_storage.clear()
|
||||
batch_storage.clear()
|
||||
batch_results.clear()
|
||||
logger.info("All data cleared")
|
||||
return {"message": "All data cleared"}
|
||||
|
||||
@app.get("/debug/status")
|
||||
async def debug_status():
|
||||
return {
|
||||
"files_count": len(file_storage),
|
||||
"batches_count": len(batch_storage),
|
||||
"batch_statuses": {bid: b.status for bid, b in batch_storage.items()},
|
||||
}
|
||||
|
|
@ -1,124 +0,0 @@
|
|||
import json
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
|
||||
def get_request_details(request: Request, body: dict = None) -> str:
|
||||
details = {
|
||||
"method": request.method,
|
||||
"url": str(request.url),
|
||||
"path": request.url.path,
|
||||
"headers": dict(request.headers),
|
||||
"query_params": dict(request.query_params),
|
||||
}
|
||||
return json.dumps(details, indent=2)
|
||||
|
||||
|
||||
def data_generator(response_details: str, model: str):
|
||||
response_id = uuid.uuid4().hex
|
||||
content = response_details
|
||||
chunk_size = 50
|
||||
for i in range(0, len(content), chunk_size):
|
||||
text_chunk = content[i : i + chunk_size]
|
||||
chunk = {
|
||||
"id": f"chatcmpl-{response_id}",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": int(time.time()),
|
||||
"model": model,
|
||||
"choices": [{"index": 0, "delta": {"content": text_chunk}}],
|
||||
}
|
||||
yield f"data: {json.dumps(chunk)}\n\n"
|
||||
final_chunk = {
|
||||
"id": f"chatcmpl-{response_id}",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": int(time.time()),
|
||||
"model": model,
|
||||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
||||
}
|
||||
yield f"data: {json.dumps(final_chunk)}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
|
||||
def setup_chat_routes(app: FastAPI):
|
||||
@app.post("/chat/completions")
|
||||
@app.post("/v1/chat/completions")
|
||||
@app.post("/openai/deployments/{model:path}/chat/completions")
|
||||
async def completion(request: Request):
|
||||
data = await request.json()
|
||||
model = data.get("model", "unknown")
|
||||
request_details = get_request_details(request, data)
|
||||
timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
response_details = f"Request:{request_details}, Canned Response:{timestamp}"
|
||||
|
||||
if data.get("stream"):
|
||||
return StreamingResponse(
|
||||
content=data_generator(response_details, model),
|
||||
media_type="text/event-stream",
|
||||
)
|
||||
else:
|
||||
response_id = uuid.uuid4().hex
|
||||
response = {
|
||||
"id": f"chatcmpl-{response_id}",
|
||||
"object": "chat.completion",
|
||||
"created": int(time.time()),
|
||||
"model": model,
|
||||
"system_fingerprint": "fp_mock_server",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": response_details,
|
||||
},
|
||||
"logprobs": None,
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 9,
|
||||
"completion_tokens": 12,
|
||||
"total_tokens": 21,
|
||||
},
|
||||
}
|
||||
return response
|
||||
|
||||
@app.post("/completions")
|
||||
@app.post("/v1/completions")
|
||||
async def text_completion(request: Request):
|
||||
data = await request.json()
|
||||
model = data.get("model", "unknown")
|
||||
request_details = get_request_details(request, data)
|
||||
timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
response_details = f"Request:{request_details}, Canned Response:{timestamp}"
|
||||
|
||||
if data.get("stream"):
|
||||
return StreamingResponse(
|
||||
content=data_generator(response_details, model),
|
||||
media_type="text/event-stream",
|
||||
)
|
||||
else:
|
||||
response = {
|
||||
"id": f"cmpl-{uuid.uuid4().hex}",
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
"logprobs": None,
|
||||
"text": response_details,
|
||||
},
|
||||
],
|
||||
"created": int(time.time()),
|
||||
"model": model,
|
||||
"object": "text_completion",
|
||||
"system_fingerprint": None,
|
||||
"usage": {
|
||||
"completion_tokens": 16,
|
||||
"prompt_tokens": 10,
|
||||
"total_tokens": 26,
|
||||
},
|
||||
}
|
||||
return response
|
||||
|
|
@ -1,23 +0,0 @@
|
|||
from fastapi import FastAPI, Request
|
||||
|
||||
|
||||
def setup_embeddings_routes(app: FastAPI):
|
||||
@app.post("/embeddings")
|
||||
@app.post("/v1/embeddings")
|
||||
@app.post("/openai/deployments/{model:path}/embeddings")
|
||||
async def embeddings(request: Request):
|
||||
data = await request.json()
|
||||
model = data.get("model", "unknown")
|
||||
_small_embedding = [
|
||||
-0.006929283495992422,
|
||||
-0.005336422007530928,
|
||||
-4.547132266452536e-05,
|
||||
-0.024047505110502243,
|
||||
]
|
||||
big_embedding = _small_embedding * 100
|
||||
return {
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": big_embedding}],
|
||||
"model": model,
|
||||
"usage": {"prompt_tokens": 5, "total_tokens": 5},
|
||||
}
|
||||
|
|
@ -1,170 +0,0 @@
|
|||
import json
|
||||
import re
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastapi import FastAPI, Request, HTTPException
|
||||
|
||||
|
||||
# Header to identify which model/deployment this request targets (simulates Azure model-specific encryption).
|
||||
# When set, the mock validates that encrypted_content in input was produced by this model.
|
||||
MOCK_AZURE_MODEL_HEADER = "X-Mock-Azure-Model"
|
||||
|
||||
# Prefix we use in mock encrypted_content: gAAA_model_<model_id>_<32hex uuid>
|
||||
# Model id can contain underscores (e.g. gpt-5.1-codex-openai-2).
|
||||
ENCRYPTED_CONTENT_MODEL_PREFIX = re.compile(r"^gAAA_model_(.+)_[0-9a-f]{32}$")
|
||||
|
||||
|
||||
def _extract_model_from_encrypted_content(encrypted: str) -> str | None:
|
||||
"""Extract model id from our mock encrypted_content format, or None if not our format."""
|
||||
if not isinstance(encrypted, str) or not encrypted.startswith("gAAA"):
|
||||
return None
|
||||
m = ENCRYPTED_CONTENT_MODEL_PREFIX.match(encrypted)
|
||||
return m.group(1) if m else None
|
||||
|
||||
|
||||
def _collect_encrypted_contents(obj, out: list[str]) -> None:
|
||||
"""Recursively collect all encrypted_content string values from input structure."""
|
||||
if isinstance(obj, dict):
|
||||
if "encrypted_content" in obj and obj["encrypted_content"]:
|
||||
out.append(obj["encrypted_content"])
|
||||
for v in obj.values():
|
||||
_collect_encrypted_contents(v, out)
|
||||
elif isinstance(obj, list):
|
||||
for item in obj:
|
||||
_collect_encrypted_contents(item, out)
|
||||
|
||||
|
||||
def _validate_encrypted_content_model(request_model: str | None, input_data: Any) -> str | None:
|
||||
"""
|
||||
If request_model is set, check that all encrypted_content in input was produced by this model.
|
||||
Returns error message if validation fails, else None.
|
||||
Content with our format (gAAA_model_<id>_) must match request_model.
|
||||
"""
|
||||
if not request_model:
|
||||
return None
|
||||
encrypted_values: list[str] = []
|
||||
_collect_encrypted_contents(input_data, encrypted_values)
|
||||
for enc in encrypted_values:
|
||||
content_model = _extract_model_from_encrypted_content(enc)
|
||||
if content_model is not None and content_model != request_model:
|
||||
err = enc[:50] + "..." if len(enc) > 50 else enc
|
||||
return f"The encrypted content {err} could not be verified."
|
||||
return None
|
||||
|
||||
|
||||
def get_request_details(request: Request, body: dict = None) -> str:
|
||||
details = {
|
||||
"method": request.method,
|
||||
"url": str(request.url),
|
||||
"path": request.url.path,
|
||||
"headers": dict(request.headers),
|
||||
"query_params": dict(request.query_params),
|
||||
}
|
||||
return json.dumps(details, indent=2)
|
||||
|
||||
|
||||
def setup_responses_routes(app: FastAPI):
|
||||
@app.post("/responses")
|
||||
@app.post("/v1/responses")
|
||||
@app.post("/openai/responses")
|
||||
async def responses_api(request: Request):
|
||||
data = await request.json()
|
||||
model = data.get("model", "unknown")
|
||||
|
||||
# Simulate Azure: encrypted content from one model cannot be verified by another.
|
||||
input_data = data.get("input")
|
||||
err_msg = _validate_encrypted_content_model(model, input_data)
|
||||
if err_msg is not None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": {
|
||||
"message": err_msg,
|
||||
"type": "invalid_request_error",
|
||||
"param": None,
|
||||
"code": "invalid_encrypted_content",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
request_details = get_request_details(request, data)
|
||||
timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
response_details = f"Request:{request_details}, Canned Response:{timestamp}"
|
||||
response_id = uuid.uuid4().hex
|
||||
message_id = f"msg_{uuid.uuid4().hex[:34]}"
|
||||
reasoning_id = f"rs_{uuid.uuid4().hex[:34]}"
|
||||
|
||||
output_items: list[dict[str, Any]] = [
|
||||
{
|
||||
"id": message_id,
|
||||
"content": [
|
||||
{
|
||||
"annotations": [],
|
||||
"text": response_details,
|
||||
"type": "output_text",
|
||||
"logprobs": [],
|
||||
},
|
||||
],
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"type": "message",
|
||||
},
|
||||
]
|
||||
|
||||
if model:
|
||||
output_items.append(
|
||||
{
|
||||
"id": reasoning_id,
|
||||
"type": "reasoning",
|
||||
"status": "completed",
|
||||
"encrypted_content": f"gAAA_model_{model}_{uuid.uuid4().hex}",
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"id": f"resp_{response_id}",
|
||||
"created_at": int(time.time()),
|
||||
"error": None,
|
||||
"incomplete_details": None,
|
||||
"instructions": None,
|
||||
"metadata": {},
|
||||
"model": model,
|
||||
"object": "response",
|
||||
"output": output_items,
|
||||
"parallel_tool_calls": True,
|
||||
"temperature": data.get("temperature", 1.0),
|
||||
"tool_choice": data.get("tool_choice", "auto"),
|
||||
"tools": data.get("tools", []),
|
||||
"top_p": data.get("top_p", 1.0),
|
||||
"max_output_tokens": data.get("max_output_tokens"),
|
||||
"previous_response_id": None,
|
||||
"reasoning": {"effort": None, "summary": None},
|
||||
"status": "completed",
|
||||
"text": {"format": {"type": "text"}, "verbosity": "medium"},
|
||||
"truncation": "disabled",
|
||||
"usage": {
|
||||
"input_tokens": 11,
|
||||
"input_tokens_details": {
|
||||
"audio_tokens": None,
|
||||
"cached_tokens": 0,
|
||||
"text_tokens": None,
|
||||
},
|
||||
"output_tokens": 19,
|
||||
"output_tokens_details": {"reasoning_tokens": 0, "text_tokens": None},
|
||||
"total_tokens": 30,
|
||||
"cost": None,
|
||||
},
|
||||
"user": None,
|
||||
"store": True,
|
||||
"background": False,
|
||||
"content_filters": None,
|
||||
"max_tool_calls": None,
|
||||
"prompt_cache_key": None,
|
||||
"safety_identifier": None,
|
||||
"service_tier": "default",
|
||||
"top_logprobs": 0,
|
||||
}
|
||||
|
|
@ -1,98 +0,0 @@
|
|||
"""
|
||||
Mock S3 callback receiver for testing LiteLLM S3 callbacks.
|
||||
|
||||
This module provides S3-compatible endpoints that capture callback data
|
||||
sent by LiteLLM's s3_v2 callback handler after batch completion.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from pydantic import BaseModel
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class S3CallbackRecord(BaseModel):
|
||||
key: str
|
||||
bucket: str
|
||||
content: Dict[str, Any]
|
||||
timestamp: int
|
||||
content_type: Optional[str] = None
|
||||
|
||||
|
||||
callback_storage: List[S3CallbackRecord] = []
|
||||
|
||||
|
||||
def setup_s3_callback_routes(app: FastAPI):
|
||||
@app.put("/{bucket}/{key:path}")
|
||||
async def s3_put_object(bucket: str, key: str, request: Request):
|
||||
content_type = request.headers.get("content-type", "application/json")
|
||||
body = await request.body()
|
||||
|
||||
try:
|
||||
content = json.loads(body.decode("utf-8"))
|
||||
except (json.JSONDecodeError, UnicodeDecodeError):
|
||||
content = {"raw": body.decode("utf-8", errors="replace")}
|
||||
|
||||
record = S3CallbackRecord(
|
||||
key=key,
|
||||
bucket=bucket,
|
||||
content=content,
|
||||
timestamp=int(time.time()),
|
||||
content_type=content_type,
|
||||
)
|
||||
callback_storage.append(record)
|
||||
|
||||
logger.info(f"S3 callback received: bucket={bucket}, key={key}")
|
||||
logger.debug(f"Callback content: {json.dumps(content, indent=2)[:500]}")
|
||||
|
||||
return {
|
||||
"ETag": f'"{hash(body)}"',
|
||||
"VersionId": None,
|
||||
}
|
||||
|
||||
@app.get("/mock-s3/callbacks")
|
||||
async def list_callbacks(
|
||||
bucket: Optional[str] = None,
|
||||
key_prefix: Optional[str] = None,
|
||||
limit: int = 100,
|
||||
):
|
||||
results = callback_storage
|
||||
|
||||
if bucket:
|
||||
results = [r for r in results if r.bucket == bucket]
|
||||
|
||||
if key_prefix:
|
||||
results = [r for r in results if r.key.startswith(key_prefix)]
|
||||
|
||||
return {
|
||||
"count": len(results),
|
||||
"callbacks": [r.model_dump() for r in results[-limit:]],
|
||||
}
|
||||
|
||||
@app.get("/mock-s3/callbacks/count")
|
||||
async def count_callbacks(bucket: Optional[str] = None):
|
||||
if bucket:
|
||||
count = sum(1 for r in callback_storage if r.bucket == bucket)
|
||||
else:
|
||||
count = len(callback_storage)
|
||||
|
||||
return {"count": count}
|
||||
|
||||
@app.get("/mock-s3/callbacks/latest")
|
||||
async def get_latest_callback():
|
||||
if not callback_storage:
|
||||
return {"callback": None}
|
||||
return {"callback": callback_storage[-1].model_dump()}
|
||||
|
||||
@app.delete("/mock-s3/callbacks")
|
||||
async def clear_callbacks():
|
||||
count = len(callback_storage)
|
||||
callback_storage.clear()
|
||||
logger.info(f"Cleared {count} S3 callbacks")
|
||||
return {"cleared": count}
|
||||
|
|
@ -1,33 +0,0 @@
|
|||
from fastapi import FastAPI, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
|
||||
from .mock_azure_batch import setup_batch_routes
|
||||
from .mock_chat import setup_chat_routes
|
||||
from .mock_embeddings import setup_embeddings_routes
|
||||
from .mock_responses import setup_responses_routes
|
||||
from .mock_s3_callback import setup_s3_callback_routes
|
||||
|
||||
|
||||
def create_mock_azure_batch_server() -> FastAPI:
|
||||
"""Create a FastAPI app that mocks Azure Batch API and S3 callbacks."""
|
||||
app = FastAPI()
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "ok"}
|
||||
|
||||
setup_chat_routes(app)
|
||||
setup_responses_routes(app)
|
||||
setup_embeddings_routes(app)
|
||||
setup_batch_routes(app)
|
||||
setup_s3_callback_routes(app)
|
||||
|
||||
return app
|
||||
|
|
@ -1,12 +0,0 @@
|
|||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
from fixtures.mock_azure_batch_server import create_mock_azure_batch_server
|
||||
import uvicorn
|
||||
|
||||
if __name__ == "__main__":
|
||||
app = create_mock_azure_batch_server()
|
||||
uvicorn.run(app, host="0.0.0.0", port=8090, log_level="info", access_log=False)
|
||||
|
|
@ -1,41 +0,0 @@
|
|||
"""
|
||||
Smoke test to verify fixtures start and stop correctly.
|
||||
Run this first to ensure the infrastructure works before running full E2E tests.
|
||||
"""
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
|
||||
pytestmark = pytest.mark.usefixtures("mock_azure_server", "litellm_proxy_server")
|
||||
|
||||
|
||||
def test_mock_server_health(mock_azure_server):
|
||||
"""Verify mock Azure server is running and healthy."""
|
||||
response = httpx.get(f"{mock_azure_server}/health", timeout=5.0)
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"status": "ok"}
|
||||
print(f"✓ Mock Azure server is healthy at {mock_azure_server}")
|
||||
|
||||
|
||||
def test_litellm_proxy_health(litellm_proxy_server):
|
||||
"""Verify LiteLLM proxy is running and healthy."""
|
||||
response = httpx.get(f"{litellm_proxy_server}/health", timeout=5.0)
|
||||
assert response.status_code == 200
|
||||
print(f"✓ LiteLLM proxy is healthy at {litellm_proxy_server}")
|
||||
|
||||
|
||||
def test_litellm_proxy_model_list(litellm_proxy_server):
|
||||
"""Verify LiteLLM proxy can list models."""
|
||||
response = httpx.get(
|
||||
f"{litellm_proxy_server}/v1/models",
|
||||
headers={"Authorization": "Bearer sk-1234"},
|
||||
timeout=5.0,
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert "data" in data
|
||||
models = [m["id"] for m in data["data"]]
|
||||
print(f"✓ LiteLLM proxy has {len(models)} models configured")
|
||||
assert "azure-fake-gpt-5-batch-2025-08-07" in models
|
||||
print(f"✓ Azure batch model is configured")
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -1,324 +0,0 @@
|
|||
import base64
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import warnings
|
||||
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
from tenacity import RetryError
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
from base_integration_test import (
|
||||
get_mock_server_base_url,
|
||||
model_id,
|
||||
use_mock_models,
|
||||
UserKeyTestMixin,
|
||||
)
|
||||
from test_managed_files_base import (
|
||||
ManagedFilesBase,
|
||||
MIN_EXPIRY_SECONDS,
|
||||
get_batch_model_names,
|
||||
)
|
||||
|
||||
MANAGED_FILE_ID_PREFIX = "litellm_proxy"
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.usefixtures("mock_azure_server", "litellm_proxy_server"),
|
||||
pytest.mark.skipif(
|
||||
os.environ.get("SKIP_E2E_TESTS", "false").lower() == "true",
|
||||
reason="E2E tests disabled via SKIP_E2E_TESTS env var"
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def is_managed_id(file_id: str) -> bool:
|
||||
"""Check if a file ID is a base64-encoded LiteLLM managed/unified ID."""
|
||||
try:
|
||||
padded = file_id + "=" * (-len(file_id) % 4)
|
||||
decoded = base64.urlsafe_b64decode(padded).decode()
|
||||
return decoded.startswith(MANAGED_FILE_ID_PREFIX)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def assert_managed_id(file_id: str, label: str):
|
||||
assert is_managed_id(file_id), f"{label} should be a managed ID, got raw: {file_id}"
|
||||
|
||||
|
||||
def wip_features_enabled() -> bool:
|
||||
return os.environ.get("WIP_FEATURES", "").lower() == "true"
|
||||
|
||||
|
||||
class TestManagedFilesAPI(ManagedFilesBase, UserKeyTestMixin):
|
||||
@classmethod
|
||||
def setup_class(cls):
|
||||
super().setup_class()
|
||||
cls.setup_admin_client()
|
||||
|
||||
@classmethod
|
||||
def teardown_class(cls):
|
||||
cls.teardown_admin_client()
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def setup_test(self):
|
||||
print(
|
||||
f"\nBase URL: {self.base_url}, Using mock models: {use_mock_models()}",
|
||||
)
|
||||
self.clear_s3_callbacks()
|
||||
|
||||
user_id, api_key, user_email, client = self.create_user_key_and_client(
|
||||
"e2e-batch",
|
||||
)
|
||||
self.test_user_id = user_id
|
||||
self.openai_client = client
|
||||
print(f"Using user {user_email} (id={user_id})")
|
||||
|
||||
def _create_and_verify_batch_input_file(self, tmp_path, model_name):
|
||||
request_file = self.create_batch_request_file_on_disk(tmp_path, model_name)
|
||||
|
||||
print("Creating batch input file...")
|
||||
batch_input_file = self.create_batch_input_file(
|
||||
self.openai_client,
|
||||
request_file,
|
||||
MIN_EXPIRY_SECONDS,
|
||||
target_model_names=model_name,
|
||||
)
|
||||
print(f"Created batch input file: {self.shorten_id(batch_input_file.id)}")
|
||||
assert_managed_id(batch_input_file.id, "batch_input_file.id")
|
||||
|
||||
print("Retrieving batch input file metadata...")
|
||||
metadata = self.openai_client.files.retrieve(batch_input_file.id)
|
||||
assert_managed_id(metadata.id, "files.retrieve(input).id")
|
||||
assert metadata.id == batch_input_file.id, (
|
||||
f"Input file ID mismatch: retrieve returned '{metadata.id}' but expected '{batch_input_file.id}'"
|
||||
)
|
||||
assert metadata.object == "file"
|
||||
assert metadata.bytes > 0, "bytes not set"
|
||||
assert metadata.filename == "modified_file.jsonl"
|
||||
assert metadata.purpose == "batch"
|
||||
assert metadata.status in ["uploaded", "processed", "error"]
|
||||
assert metadata.created_at > 0
|
||||
if wip_features_enabled():
|
||||
assert metadata.expires_at > 0, "expires_at not set"
|
||||
self.print_file_metadata(metadata, "Input file")
|
||||
|
||||
return batch_input_file
|
||||
|
||||
def _create_and_verify_batch(self, input_file_id):
|
||||
print("\nCreating batch...")
|
||||
batch = self.create_batch(
|
||||
self.openai_client,
|
||||
input_file_id,
|
||||
MIN_EXPIRY_SECONDS,
|
||||
)
|
||||
print(f"Created batch: {self.shorten_id(batch.id)}")
|
||||
|
||||
assert batch.id, "No batch ID returned"
|
||||
assert_managed_id(batch.id, "batch.id")
|
||||
assert_managed_id(batch.input_file_id, "batch.input_file_id")
|
||||
assert batch.input_file_id == input_file_id, "batch.input_file_id mismatch"
|
||||
assert batch.status in ["validating", "in_progress", "finalizing", "completed"]
|
||||
if not batch.expires_at:
|
||||
warnings.warn("batch expires_at not set")
|
||||
else:
|
||||
assert batch.expires_at > 0
|
||||
if not batch.endpoint:
|
||||
warnings.warn("batch.endpoint empty - Azure API quirk, not a bug")
|
||||
else:
|
||||
assert batch.endpoint == "/v1/chat/completions"
|
||||
assert batch.completion_window == "24h"
|
||||
assert batch.created_at > 0
|
||||
self.print_batch_metadata(batch)
|
||||
|
||||
return batch
|
||||
|
||||
def _list_batches(self, batch_id, model_name):
|
||||
if not wip_features_enabled():
|
||||
return
|
||||
print("\nListing batches...")
|
||||
try:
|
||||
batches_list = self.wait_for_batch_list(
|
||||
model_name,
|
||||
max_seconds=30,
|
||||
wait_seconds=5,
|
||||
)
|
||||
batch_ids = [b.id for b in (batches_list.data if batches_list else [])]
|
||||
if batch_id not in batch_ids:
|
||||
warnings.warn(
|
||||
f"Batch {batch_id} not found in list. "
|
||||
f"batches.list returns raw IDs, not encoded IDs. raw IDs: {batch_ids}",
|
||||
)
|
||||
except openai.APIError as e:
|
||||
pytest.fail(f"batches.list() failed: {e}")
|
||||
|
||||
def _wait_for_batch_completion(self, batch_id, tracker):
|
||||
print(f"\nWaiting for batch {self.shorten_id(batch_id)} to complete...")
|
||||
try:
|
||||
batch_response = self.wait_for_batch_state(
|
||||
self.openai_client,
|
||||
batch_id,
|
||||
"completed",
|
||||
max_seconds=25 * 60,
|
||||
wait_seconds=15,
|
||||
state_tracker=tracker,
|
||||
)
|
||||
except RetryError:
|
||||
tracker.print_state("Timeout waiting for batch completion")
|
||||
raise TimeoutError("Timed out waiting for batch to be in state: completed")
|
||||
|
||||
assert_managed_id(batch_response.id, "batch_response.id")
|
||||
assert batch_response.id == batch_id, (
|
||||
f"batch_response.id mismatch: got '{batch_response.id}' but expected '{batch_id}'"
|
||||
)
|
||||
assert_managed_id(batch_response.input_file_id, "batch_response.input_file_id")
|
||||
assert_managed_id(
|
||||
batch_response.output_file_id,
|
||||
"batch_response.output_file_id",
|
||||
)
|
||||
|
||||
return batch_response
|
||||
|
||||
def _get_and_verify_batch_output(self, output_file_id):
|
||||
print("\nRetrieving batch output file metadata...")
|
||||
metadata = self.openai_client.files.retrieve(output_file_id)
|
||||
assert_managed_id(metadata.id, "files.retrieve(output_file_id).id")
|
||||
assert metadata.id == output_file_id, (
|
||||
f"Output file ID mismatch: retrieve returned '{metadata.id}' but expected '{output_file_id}'"
|
||||
)
|
||||
assert metadata.object == "file"
|
||||
assert metadata.bytes > 0, "bytes not set"
|
||||
assert metadata.filename, "filename not set"
|
||||
assert metadata.purpose in ["batch_output", "batch"]
|
||||
assert metadata.created_at > 0
|
||||
self.print_file_metadata(metadata, "Output file")
|
||||
|
||||
print("\nFetching batch output file content...")
|
||||
content = self.openai_client.files.content(output_file_id)
|
||||
assert content.text, "No batch file content returned"
|
||||
assert len(content.text) > 0, "Batch file content is empty"
|
||||
print(f"Output file content ({len(content.text)} bytes):")
|
||||
for line in content.text.strip().split("\n")[:3]:
|
||||
print(f"\t{line}")
|
||||
|
||||
return metadata
|
||||
|
||||
def _delete_file(self, file_id, label, max_retries=10, retry_delay=5):
|
||||
print(f"\nDeleting {label}: {self.shorten_id(file_id)}")
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
self.openai_client.files.delete(file_id)
|
||||
return
|
||||
except openai.BadRequestError as e:
|
||||
if "batch_processed" in str(e) and attempt < max_retries - 1:
|
||||
print(
|
||||
f" File still referenced by unprocessed batch, "
|
||||
f"retrying in {retry_delay}s ({attempt + 1}/{max_retries})"
|
||||
)
|
||||
time.sleep(retry_delay)
|
||||
else:
|
||||
pytest.fail(f"files.delete({label}) failed: {e}")
|
||||
except openai.APIError as e:
|
||||
pytest.fail(f"files.delete({label}) failed: {e}")
|
||||
|
||||
def _verify_file_deleted(self, file_id, label):
|
||||
print(f"Verifying {label} is deleted...")
|
||||
try:
|
||||
self.openai_client.files.content(file_id)
|
||||
assert False, f"{label} {file_id} still accessible after deletion"
|
||||
except openai.NotFoundError:
|
||||
print(f"{label} correctly not accessible after deletion")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Tests
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.flaky(reruns=2)
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
get_batch_model_names(),
|
||||
ids=model_id,
|
||||
)
|
||||
def test_e2e_managed_batch(self, tmp_path, model_name):
|
||||
print(
|
||||
f"\n\nStarting test with base_url={self.base_url} and model_name={model_name}\n",
|
||||
)
|
||||
self.reset_mock_server()
|
||||
tracker = self.create_state_tracker()
|
||||
|
||||
batch_input_file = self._create_and_verify_batch_input_file(
|
||||
tmp_path,
|
||||
model_name,
|
||||
)
|
||||
tracker.set_file_id(batch_input_file.id)
|
||||
tracker.print_state("After creating batch input file")
|
||||
|
||||
batch = self._create_and_verify_batch(batch_input_file.id)
|
||||
tracker.set_batch_id(batch.id)
|
||||
tracker.print_state("After creating batch")
|
||||
|
||||
self._list_batches(batch.id, model_name)
|
||||
|
||||
batch_response = self._wait_for_batch_completion(batch.id, tracker)
|
||||
tracker.print_state("After batch completed")
|
||||
|
||||
self._get_and_verify_batch_output(batch_response.output_file_id)
|
||||
tracker.print_state("After retrieving output file")
|
||||
|
||||
tracker.print_state("Final state after cleanup")
|
||||
tracker.wait_and_print_s3_callbacks()
|
||||
tracker.assert_batch_cost_callback()
|
||||
|
||||
self._delete_file(batch_input_file.id, "input file")
|
||||
self._delete_file(batch_response.output_file_id, "output file")
|
||||
|
||||
self._verify_file_deleted(batch_input_file.id, "input file")
|
||||
self._verify_file_deleted(batch_response.output_file_id, "output file")
|
||||
|
||||
def cleanup_batches_in_database(self):
|
||||
import psycopg2
|
||||
|
||||
print("Cleaning up stale batch records from database...")
|
||||
try:
|
||||
conn = psycopg2.connect(
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="litellm",
|
||||
user="llmproxy",
|
||||
password="dbpassword9090",
|
||||
)
|
||||
with conn.cursor() as cur:
|
||||
cur.execute("""
|
||||
DELETE FROM "LiteLLM_ManagedObjectTable"
|
||||
WHERE file_purpose = 'batch' AND status = 'validating'
|
||||
""")
|
||||
deleted = cur.rowcount
|
||||
conn.commit()
|
||||
if deleted > 0:
|
||||
print(f"Deleted {deleted} stale batch records")
|
||||
conn.close()
|
||||
except Exception as e:
|
||||
print(f"Warning: Could not clean up database: {e}")
|
||||
|
||||
def clear_s3_callbacks(self):
|
||||
clear_response = httpx.delete(f"{get_mock_server_base_url()}/mock-s3/callbacks")
|
||||
assert clear_response.status_code == 200, (
|
||||
f"Failed to clear callbacks: {clear_response.text}"
|
||||
)
|
||||
return clear_response.json()
|
||||
|
||||
@pytest.mark.skipif(
|
||||
True,
|
||||
reason="Skipping managed files test till managed files feature is available",
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
get_batch_model_names(),
|
||||
ids=model_id,
|
||||
)
|
||||
def test_error_files(self, tmp_path, model_name):
|
||||
raise NotImplementedError(
|
||||
"To implement. Fail a batch and retrieve the error file.",
|
||||
)
|
||||
|
|
@ -1,119 +0,0 @@
|
|||
#!/usr/bin/env python
|
||||
"""
|
||||
Validation script for Azure Batch E2E test setup.
|
||||
Run this before running the actual tests to verify all components are accessible.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
def check_imports():
|
||||
"""Verify all required imports work."""
|
||||
print("Checking imports...")
|
||||
try:
|
||||
from base_integration_test import (
|
||||
get_mock_server_base_url,
|
||||
get_litellm_base_url,
|
||||
get_litellm_api_key,
|
||||
)
|
||||
print(" ✓ base_integration_test imports OK")
|
||||
|
||||
from test_managed_files_base import ManagedFilesBase, get_batch_model_names
|
||||
print(" ✓ test_managed_files_base imports OK")
|
||||
|
||||
from fixtures.mock_azure_batch_server import create_mock_azure_batch_server
|
||||
print(" ✓ mock_azure_batch_server imports OK")
|
||||
|
||||
import httpx
|
||||
import openai
|
||||
import psycopg2
|
||||
import uvicorn
|
||||
print(" ✓ All external dependencies OK")
|
||||
|
||||
return True
|
||||
except ImportError as e:
|
||||
print(f" ✗ Import error: {e}")
|
||||
return False
|
||||
|
||||
|
||||
def check_config_file():
|
||||
"""Verify config file exists."""
|
||||
print("\nChecking config file...")
|
||||
config_path = Path(__file__).parent / "fixtures" / "config.yml"
|
||||
if config_path.exists():
|
||||
print(f" ✓ Config file found: {config_path}")
|
||||
return True
|
||||
else:
|
||||
print(f" ✗ Config file not found: {config_path}")
|
||||
return False
|
||||
|
||||
|
||||
def check_database():
|
||||
"""Verify database connection."""
|
||||
print("\nChecking database connection...")
|
||||
try:
|
||||
import psycopg2
|
||||
conn = psycopg2.connect(
|
||||
host="localhost",
|
||||
port=5432,
|
||||
database="litellm",
|
||||
user="llmproxy",
|
||||
password="dbpassword9090",
|
||||
)
|
||||
conn.close()
|
||||
print(" ✓ Database connection OK")
|
||||
return True
|
||||
except Exception as e:
|
||||
print(f" ✗ Database connection failed: {e}")
|
||||
print(" Start PostgreSQL with:")
|
||||
print(" docker run --name litellm-postgres -e POSTGRES_USER=llmproxy \\")
|
||||
print(" -e POSTGRES_PASSWORD=dbpassword9090 -e POSTGRES_DB=litellm \\")
|
||||
print(" -p 5432:5432 -d postgres:15")
|
||||
return False
|
||||
|
||||
|
||||
def check_ports():
|
||||
"""Check if required ports are available."""
|
||||
print("\nChecking ports...")
|
||||
import socket
|
||||
|
||||
for port, name in [(4000, "LiteLLM Proxy"), (8090, "Mock Server")]:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
try:
|
||||
s.bind(("localhost", port))
|
||||
print(f" ✓ Port {port} ({name}) is available")
|
||||
except OSError:
|
||||
print(f" ⚠ Port {port} ({name}) is in use (will reuse if healthy)")
|
||||
return True
|
||||
|
||||
|
||||
def main():
|
||||
print("=" * 70)
|
||||
print("Azure Batch E2E Test Setup Validation")
|
||||
print("=" * 70)
|
||||
|
||||
checks = [
|
||||
check_imports(),
|
||||
check_config_file(),
|
||||
check_database(),
|
||||
check_ports(),
|
||||
]
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
if all(checks):
|
||||
print("✓ All checks passed! Ready to run E2E tests.")
|
||||
print("\nRun tests with:")
|
||||
print(" cd litellm")
|
||||
print(" export DATABASE_URL='postgresql://llmproxy:dbpassword9090@localhost:5432/litellm'")
|
||||
print(" poetry run pytest tests/proxy_e2e_azure_batches_tests/test_proxy_e2e_azure_batches.py -vv")
|
||||
return 0
|
||||
else:
|
||||
print("✗ Some checks failed. Please fix the issues above.")
|
||||
return 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
|
|
@ -47,11 +47,15 @@ class TestCheckResponsesCost:
|
|||
CheckResponsesCost,
|
||||
)
|
||||
|
||||
return CheckResponsesCost(
|
||||
instance = CheckResponsesCost(
|
||||
proxy_logging_obj=mock_proxy_logging_obj,
|
||||
prisma_client=mock_prisma_client,
|
||||
llm_router=mock_llm_router,
|
||||
)
|
||||
# Mock _expire_stale_rows (raw SQL) so _cleanup_stale_managed_objects
|
||||
# succeeds without a real DB. Individual tests can override this.
|
||||
instance._expire_stale_rows = AsyncMock(return_value=0)
|
||||
return instance
|
||||
|
||||
def test_initialization(self, check_responses_cost_instance):
|
||||
"""Test that CheckResponsesCost initializes correctly"""
|
||||
|
|
@ -67,9 +71,6 @@ class TestCheckResponsesCost:
|
|||
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
|
||||
return_value=[]
|
||||
)
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
|
||||
return_value=0
|
||||
)
|
||||
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
|
|
@ -86,24 +87,20 @@ class TestCheckResponsesCost:
|
|||
async def test_cleanup_stale_managed_objects(
|
||||
self, check_responses_cost_instance, mock_prisma_client
|
||||
):
|
||||
"""Stale rows (older than cutoff) are bulk-updated to stale_expired before polling."""
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
|
||||
return_value=5
|
||||
)
|
||||
"""Stale rows are expired via _expire_stale_rows before polling."""
|
||||
from litellm.constants import STALE_OBJECT_CLEANUP_BATCH_SIZE
|
||||
|
||||
check_responses_cost_instance._expire_stale_rows = AsyncMock(return_value=5)
|
||||
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
|
||||
return_value=[]
|
||||
)
|
||||
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
# The first update_many call should be the stale-row cleanup scoped to "response"
|
||||
calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
|
||||
stale_call = calls[0]
|
||||
assert stale_call[1]["data"] == {"status": "stale_expired"}
|
||||
where = stale_call[1]["where"]
|
||||
assert where["file_purpose"] == "response"
|
||||
assert "stale_expired" in where["status"]["not_in"]
|
||||
assert "created_at" in where
|
||||
# _expire_stale_rows should have been called with a cutoff datetime and batch size
|
||||
check_responses_cost_instance._expire_stale_rows.assert_called_once()
|
||||
call_args = check_responses_cost_instance._expire_stale_rows.call_args
|
||||
assert call_args[0][1] == STALE_OBJECT_CLEANUP_BATCH_SIZE
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_with_completed_response(
|
||||
|
|
@ -145,10 +142,10 @@ class TestCheckResponsesCost:
|
|||
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
# calls[0] = stale cleanup, calls[1] = job completion
|
||||
# update_many should only contain the job completion call
|
||||
calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
|
||||
assert len(calls) == 2
|
||||
completion_call = calls[1]
|
||||
assert len(calls) == 1
|
||||
completion_call = calls[0]
|
||||
assert completion_call[1]["data"]["status"] == "completed"
|
||||
assert completion_call[1]["where"]["id"]["in"] == ["job-123"]
|
||||
|
||||
|
|
@ -188,10 +185,10 @@ class TestCheckResponsesCost:
|
|||
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
# calls[0] = stale cleanup, calls[1] = job completion
|
||||
# update_many should only contain the job completion call
|
||||
calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
|
||||
assert len(calls) == 2
|
||||
assert calls[1][1]["data"]["status"] == "completed"
|
||||
assert len(calls) == 1
|
||||
assert calls[0][1]["data"]["status"] == "completed"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_with_cancelled_response(
|
||||
|
|
@ -229,10 +226,10 @@ class TestCheckResponsesCost:
|
|||
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
# calls[0] = stale cleanup, calls[1] = job completion
|
||||
# update_many should only contain the job completion call
|
||||
calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
|
||||
assert len(calls) == 2
|
||||
assert calls[1][1]["data"]["status"] == "completed"
|
||||
assert len(calls) == 1
|
||||
assert calls[0][1]["data"]["status"] == "completed"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_with_in_progress_response(
|
||||
|
|
@ -270,10 +267,11 @@ class TestCheckResponsesCost:
|
|||
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
# Only the stale-cleanup call should have fired — no completion update
|
||||
# No job completion update_many — response is still in progress
|
||||
calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
|
||||
assert len(calls) == 1
|
||||
assert calls[0][1]["data"] == {"status": "stale_expired"}
|
||||
assert len(calls) == 0
|
||||
# Stale cleanup still ran via _expire_stale_rows
|
||||
check_responses_cost_instance._expire_stale_rows.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_with_queued_response(
|
||||
|
|
@ -311,10 +309,11 @@ class TestCheckResponsesCost:
|
|||
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
# Only the stale-cleanup call should have fired — no completion update
|
||||
# No job completion update_many — response is still queued
|
||||
calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
|
||||
assert len(calls) == 1
|
||||
assert calls[0][1]["data"] == {"status": "stale_expired"}
|
||||
assert len(calls) == 0
|
||||
# Stale cleanup still ran via _expire_stale_rows
|
||||
check_responses_cost_instance._expire_stale_rows.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_with_exception(
|
||||
|
|
@ -345,10 +344,11 @@ class TestCheckResponsesCost:
|
|||
# Should not raise, just skip the job
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
# Only the stale-cleanup call should have fired — no completion update
|
||||
# No job completion update_many — exception skipped the job
|
||||
calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
|
||||
assert len(calls) == 1
|
||||
assert calls[0][1]["data"] == {"status": "stale_expired"}
|
||||
assert len(calls) == 0
|
||||
# Stale cleanup still ran via _expire_stale_rows
|
||||
check_responses_cost_instance._expire_stale_rows.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_multiple_jobs(
|
||||
|
|
@ -424,10 +424,10 @@ class TestCheckResponsesCost:
|
|||
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
# calls[0] = stale cleanup, calls[1] = completion of 2 finished jobs
|
||||
# update_many should only contain the job completion call
|
||||
calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
|
||||
assert len(calls) == 2
|
||||
completion_call = calls[1]
|
||||
assert len(calls) == 1
|
||||
completion_call = calls[0]
|
||||
assert len(completion_call[1]["where"]["id"]["in"]) == 2
|
||||
assert "job-1" in completion_call[1]["where"]["id"]["in"]
|
||||
assert "job-3" in completion_call[1]["where"]["id"]["in"]
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from litellm.proxy._experimental.mcp_server import rest_endpoints
|
|||
from litellm.proxy._experimental.mcp_server.auth import (
|
||||
user_api_key_auth_mcp as auth_mcp,
|
||||
)
|
||||
from litellm.proxy._types import NewMCPServerRequest, UserAPIKeyAuth
|
||||
from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
|
|
@ -156,7 +156,6 @@ class TestExecuteWithMcpClient:
|
|||
"Authorization": "STATIC token",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_m2m_credentials_forwarded_to_server_model(self, monkeypatch):
|
||||
"""M2M OAuth credentials (client_id, client_secret) from the nested
|
||||
|
|
@ -199,9 +198,7 @@ class TestExecuteWithMcpClient:
|
|||
},
|
||||
)
|
||||
|
||||
result = await rest_endpoints._execute_with_mcp_client(
|
||||
payload, ok_operation
|
||||
)
|
||||
result = await rest_endpoints._execute_with_mcp_client(payload, ok_operation)
|
||||
|
||||
assert result["status"] == "ok"
|
||||
server = captured["server"]
|
||||
|
|
@ -262,7 +259,10 @@ class TestExecuteWithMcpClient:
|
|||
assert result["status"] == "ok"
|
||||
# The incoming Authorization must be dropped — extra_headers should
|
||||
# contain no oauth2 headers (only static_headers, which are None here).
|
||||
assert captured["extra_headers"] is None or "Authorization" not in captured["extra_headers"]
|
||||
assert (
|
||||
captured["extra_headers"] is None
|
||||
or "Authorization" not in captured["extra_headers"]
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catches_exception_group(self, monkeypatch):
|
||||
|
|
@ -300,9 +300,7 @@ class TestExecuteWithMcpClient:
|
|||
auth_type=MCPAuth.none,
|
||||
)
|
||||
|
||||
result = await rest_endpoints._execute_with_mcp_client(
|
||||
payload, ok_operation
|
||||
)
|
||||
result = await rest_endpoints._execute_with_mcp_client(payload, ok_operation)
|
||||
|
||||
assert result["status"] == "error"
|
||||
assert result["error"] is True
|
||||
|
|
@ -365,8 +363,12 @@ class TestTestToolsList:
|
|||
credentials={"auth_value": "secret-key"},
|
||||
)
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
result = await rest_endpoints.test_tools_list(
|
||||
request, payload, user_api_key_dict=UserAPIKeyAuth()
|
||||
request,
|
||||
payload,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert result["message"] == "Successfully retrieved tools"
|
||||
|
|
@ -419,8 +421,12 @@ class TestTestToolsList:
|
|||
auth_type=MCPAuth.oauth2,
|
||||
)
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
result = await rest_endpoints.test_tools_list(
|
||||
request, payload, user_api_key_dict=UserAPIKeyAuth()
|
||||
request,
|
||||
payload,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert result["message"] == "Successfully retrieved tools"
|
||||
|
|
@ -484,7 +490,11 @@ class TestListToolsRestAPI:
|
|||
captured = {"called": False}
|
||||
|
||||
async def fake_get_tools(
|
||||
server, server_auth_header, raw_headers=None, user_api_key_auth=None, extra_headers=None
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
extra_headers=None,
|
||||
):
|
||||
captured["called"] = True
|
||||
captured["server"] = server
|
||||
|
|
@ -555,27 +565,47 @@ class TestListToolsRestAPI:
|
|||
|
||||
captured = {"called": False, "server_arg": None}
|
||||
|
||||
async def fake_get_tools(server, server_auth_header, raw_headers=None, user_api_key_auth=None, extra_headers=None):
|
||||
async def fake_get_tools(
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
extra_headers=None,
|
||||
):
|
||||
captured["called"] = True
|
||||
captured["server_arg"] = server
|
||||
return ["tool-x"]
|
||||
|
||||
monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers",
|
||||
fake_get_allowed_mcp_servers, raising=False,
|
||||
rest_endpoints,
|
||||
"build_effective_auth_contexts",
|
||||
fake_contexts,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager, "get_mcp_server_by_name",
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_allowed_mcp_servers",
|
||||
fake_get_allowed_mcp_servers,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_mcp_server_by_name",
|
||||
lambda name: stub_server if name == "my-server" else None,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager, "get_mcp_server_by_id",
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_mcp_server_by_id",
|
||||
lambda sid: stub_server if sid == "uuid-abc-123" else None,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(rest_endpoints, "_get_tools_for_single_server", fake_get_tools, raising=False)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints,
|
||||
"_get_tools_for_single_server",
|
||||
fake_get_tools,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
request = _build_request(path="/mcp-rest/tools/list", method="GET")
|
||||
result = await rest_endpoints.list_tool_rest_api(
|
||||
|
|
@ -609,18 +639,27 @@ class TestListToolsRestAPI:
|
|||
async def fake_get_allowed_mcp_servers(*args, **kwargs):
|
||||
return []
|
||||
|
||||
monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers",
|
||||
fake_get_allowed_mcp_servers, raising=False,
|
||||
rest_endpoints,
|
||||
"build_effective_auth_contexts",
|
||||
fake_contexts,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager, "get_mcp_server_by_name",
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_allowed_mcp_servers",
|
||||
fake_get_allowed_mcp_servers,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_mcp_server_by_name",
|
||||
lambda name: stub_server if name == "restricted-server" else None,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager, "get_mcp_server_by_id",
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_mcp_server_by_id",
|
||||
lambda sid: stub_server if sid == "uuid-xyz-999" else None,
|
||||
raising=False,
|
||||
)
|
||||
|
|
@ -662,31 +701,54 @@ class TestListToolsRestAPI:
|
|||
|
||||
oauth_headers = {"Authorization": "Bearer user-oauth-token"}
|
||||
|
||||
async def fake_get_user_oauth_extra_headers(server, user_api_key_dict, prefetched_creds=None):
|
||||
async def fake_get_user_oauth_extra_headers(
|
||||
server, user_api_key_dict, prefetched_creds=None
|
||||
):
|
||||
return oauth_headers
|
||||
|
||||
captured = {}
|
||||
|
||||
async def fake_get_tools(server, server_auth_header, raw_headers=None, user_api_key_auth=None, extra_headers=None):
|
||||
async def fake_get_tools(
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
extra_headers=None,
|
||||
):
|
||||
captured["server"] = server
|
||||
captured["auth_header"] = server_auth_header
|
||||
return ["oauth-tool"]
|
||||
|
||||
monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers",
|
||||
fake_get_allowed_mcp_servers, raising=False,
|
||||
rest_endpoints,
|
||||
"build_effective_auth_contexts",
|
||||
fake_contexts,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager, "get_mcp_server_by_id",
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_allowed_mcp_servers",
|
||||
fake_get_allowed_mcp_servers,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_mcp_server_by_id",
|
||||
lambda sid: stub_server if sid == "oauth-server-id" else None,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints, "_get_user_oauth_extra_headers",
|
||||
fake_get_user_oauth_extra_headers, raising=False,
|
||||
rest_endpoints,
|
||||
"_get_user_oauth_extra_headers",
|
||||
fake_get_user_oauth_extra_headers,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints,
|
||||
"_get_tools_for_single_server",
|
||||
fake_get_tools,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(rest_endpoints, "_get_tools_for_single_server", fake_get_tools, raising=False)
|
||||
|
||||
request = _build_request(path="/mcp-rest/tools/list", method="GET")
|
||||
result = await rest_endpoints.list_tool_rest_api(
|
||||
|
|
@ -1124,3 +1186,189 @@ class TestGetToolsForSingleServer:
|
|||
assert "tool3" in tool_names
|
||||
assert "tool1" not in tool_names
|
||||
assert "tool4" not in tool_names
|
||||
|
||||
|
||||
class TestStdioCommandAllowlist:
|
||||
"""Tests for MCP stdio command allowlist validation."""
|
||||
|
||||
def test_allowed_command_passes_validation(self):
|
||||
"""npx, uvx, python, etc. should be accepted."""
|
||||
req = NewMCPServerRequest(
|
||||
server_name="test",
|
||||
transport="stdio",
|
||||
command="npx",
|
||||
args=["-y", "@modelcontextprotocol/server-filesystem"],
|
||||
)
|
||||
assert req.command == "npx"
|
||||
|
||||
def test_disallowed_command_raises(self):
|
||||
"""Arbitrary commands like bash should be rejected."""
|
||||
with pytest.raises(ValueError, match="not in the allowed commands list"):
|
||||
NewMCPServerRequest(
|
||||
server_name="test",
|
||||
transport="stdio",
|
||||
command="bash",
|
||||
args=["-c", "echo pwned"],
|
||||
)
|
||||
|
||||
def test_sh_command_raises(self):
|
||||
"""sh should be rejected."""
|
||||
with pytest.raises(ValueError, match="not in the allowed commands list"):
|
||||
NewMCPServerRequest(
|
||||
server_name="test",
|
||||
transport="stdio",
|
||||
command="sh",
|
||||
args=["-c", "id > /tmp/output.txt"],
|
||||
)
|
||||
|
||||
def test_absolute_path_bypass_blocked(self):
|
||||
"""/bin/bash should be blocked (basename is 'bash')."""
|
||||
with pytest.raises(ValueError, match="not in the allowed commands list"):
|
||||
NewMCPServerRequest(
|
||||
server_name="test",
|
||||
transport="stdio",
|
||||
command="/bin/bash",
|
||||
args=["-c", "echo pwned"],
|
||||
)
|
||||
|
||||
def test_absolute_path_to_allowed_command_works(self):
|
||||
"""/usr/bin/python3 should pass (basename is 'python3')."""
|
||||
req = NewMCPServerRequest(
|
||||
server_name="test",
|
||||
transport="stdio",
|
||||
command="/usr/bin/python3",
|
||||
args=["-m", "some_module"],
|
||||
)
|
||||
assert req.command == "/usr/bin/python3"
|
||||
|
||||
def test_http_transport_ignores_allowlist(self):
|
||||
"""HTTP/SSE transport should not trigger command validation."""
|
||||
req = NewMCPServerRequest(
|
||||
server_name="test",
|
||||
transport="sse",
|
||||
url="https://example.com/mcp",
|
||||
)
|
||||
assert req.transport == "sse"
|
||||
|
||||
def test_uvx_command_passes(self):
|
||||
req = NewMCPServerRequest(
|
||||
server_name="test",
|
||||
transport="stdio",
|
||||
command="uvx",
|
||||
args=["mcp-server-sqlite"],
|
||||
)
|
||||
assert req.command == "uvx"
|
||||
|
||||
def test_node_command_passes(self):
|
||||
req = NewMCPServerRequest(
|
||||
server_name="test",
|
||||
transport="stdio",
|
||||
command="node",
|
||||
args=["server.js"],
|
||||
)
|
||||
assert req.command == "node"
|
||||
|
||||
def test_update_request_disallowed_command_raises(self):
|
||||
"""UpdateMCPServerRequest should also block non-allowlisted commands."""
|
||||
with pytest.raises(ValueError, match="not in the allowed commands list"):
|
||||
UpdateMCPServerRequest(
|
||||
server_id="some-id",
|
||||
transport="stdio",
|
||||
command="bash",
|
||||
args=["-c", "echo pwned"],
|
||||
)
|
||||
|
||||
|
||||
class TestEndpointRoleChecks:
|
||||
"""Tests for PROXY_ADMIN role checks on MCP test endpoints."""
|
||||
|
||||
def test_test_connection_has_auth_dependency(self):
|
||||
route = _get_route("/mcp-rest/test/connection", "POST")
|
||||
assert _route_has_dependency(route, user_api_key_auth)
|
||||
|
||||
def test_test_tools_list_has_auth_dependency(self):
|
||||
route = _get_route("/mcp-rest/test/tools/list", "POST")
|
||||
assert _route_has_dependency(route, user_api_key_auth)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_test_connection_rejects_non_admin(self):
|
||||
"""Non-admin users should get 403 from test_connection."""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="test",
|
||||
url="https://example.com/mcp",
|
||||
auth_type=MCPAuth.none,
|
||||
)
|
||||
user_key = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="non_admin",
|
||||
api_key="sk-test",
|
||||
)
|
||||
request = _build_request()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await rest_endpoints.test_connection(
|
||||
request=request,
|
||||
new_mcp_server_request=payload,
|
||||
user_api_key_dict=user_key,
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_test_tools_list_rejects_non_admin(self):
|
||||
"""Non-admin users should get 403 from test_tools_list."""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="test",
|
||||
url="https://example.com/mcp",
|
||||
auth_type=MCPAuth.none,
|
||||
)
|
||||
user_key = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="non_admin",
|
||||
api_key="sk-test",
|
||||
)
|
||||
request = _build_request()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await rest_endpoints.test_tools_list(
|
||||
request=request,
|
||||
new_mcp_server_request=payload,
|
||||
user_api_key_dict=user_key,
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_test_connection_allows_admin(self, monkeypatch):
|
||||
"""PROXY_ADMIN should pass the role check."""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
async def fake_execute(*args, **kwargs):
|
||||
return {"status": "ok"}
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints,
|
||||
"_execute_with_mcp_client",
|
||||
fake_execute,
|
||||
)
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="test",
|
||||
url="https://example.com/mcp",
|
||||
auth_type=MCPAuth.none,
|
||||
)
|
||||
user_key = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id="admin",
|
||||
api_key="sk-admin",
|
||||
)
|
||||
request = _build_request()
|
||||
|
||||
result = await rest_endpoints.test_connection(
|
||||
request=request,
|
||||
new_mcp_server_request=payload,
|
||||
user_api_key_dict=user_key,
|
||||
)
|
||||
assert result["status"] == "ok"
|
||||
|
|
|
|||
|
|
@ -713,6 +713,51 @@ class TestJWTOAuth2Coexistence:
|
|||
mock_jwt_auth.assert_not_called()
|
||||
assert result.user_id == "machine-client-1"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth2_path_requires_premium_user(self):
|
||||
"""
|
||||
OAuth2 token validation should fail when enterprise premium is disabled.
|
||||
"""
|
||||
opaque_token = "some-opaque-m2m-oauth2-token"
|
||||
general_settings = {
|
||||
"enable_oauth2_auth": True,
|
||||
"enable_jwt_auth": True,
|
||||
}
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/v1/chat/completions"
|
||||
mock_request.headers = {"authorization": f"Bearer {opaque_token}"}
|
||||
mock_request.query_params = {}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings", general_settings
|
||||
), patch("litellm.proxy.proxy_server.premium_user", False), patch(
|
||||
"litellm.proxy.proxy_server.master_key", "sk-master"
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", None
|
||||
), patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_oauth2:
|
||||
litellm.proxy.proxy_server.jwt_handler.update_environment(
|
||||
prisma_client=None,
|
||||
user_api_key_cache=DualCache(),
|
||||
litellm_jwtauth=LiteLLM_JWTAuth(),
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {opaque_token}",
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.auth_error
|
||||
assert (
|
||||
"Oauth2 token validation is only available for premium users"
|
||||
in exc_info.value.message
|
||||
)
|
||||
mock_oauth2.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_both_enabled_jwt_token_skips_oauth2(self):
|
||||
"""
|
||||
|
|
@ -974,6 +1019,248 @@ class TestJWTOAuth2Coexistence:
|
|||
mock_jwt_auth.assert_not_called()
|
||||
assert result.user_id == "machine-client-aud-list"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_routing_override_routes_jwt_to_oauth2_when_oauth2_globally_disabled(
|
||||
self,
|
||||
):
|
||||
"""
|
||||
If enable_oauth2_auth is false, JWT tokens matching routing_overrides
|
||||
should still route to OAuth2 introspection.
|
||||
"""
|
||||
jwt_token = (
|
||||
"eyJhbGciOiJSUzI1NiJ9."
|
||||
"eyJpc3MiOiJtYWNoaW5lLWlzc3Vlci5leGFtcGxlLmNvbSIsImNsaWVudF9pZCI6Ik1JRF9MSVRFTExNIn0."
|
||||
"c2ln"
|
||||
)
|
||||
general_settings = {
|
||||
"enable_oauth2_auth": False,
|
||||
"enable_jwt_auth": True,
|
||||
}
|
||||
mock_oauth2_response = UserAPIKeyAuth(
|
||||
api_key=jwt_token,
|
||||
user_id="machine-client-override-oauth2-off",
|
||||
)
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/v1/chat/completions"
|
||||
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
|
||||
mock_request.query_params = {}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings", general_settings
|
||||
), patch("litellm.proxy.proxy_server.premium_user", True), patch(
|
||||
"litellm.proxy.proxy_server.master_key", "sk-master"
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", None
|
||||
), patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_oauth2_response,
|
||||
) as mock_oauth2, patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_jwt_auth:
|
||||
litellm.proxy.proxy_server.jwt_handler.update_environment(
|
||||
prisma_client=None,
|
||||
user_api_key_cache=DualCache(),
|
||||
litellm_jwtauth=LiteLLM_JWTAuth(
|
||||
routing_overrides=[
|
||||
JWTRoutingOverride(
|
||||
iss="machine-issuer.example.com",
|
||||
client_id="MID_LITELLM",
|
||||
path="oauth2",
|
||||
)
|
||||
]
|
||||
),
|
||||
)
|
||||
|
||||
result = await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {jwt_token}",
|
||||
)
|
||||
|
||||
mock_oauth2.assert_called_once_with(token=jwt_token)
|
||||
mock_jwt_auth.assert_not_called()
|
||||
assert result.user_id == "machine-client-override-oauth2-off"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_opaque_token_does_not_use_oauth2_when_oauth2_globally_disabled(
|
||||
self,
|
||||
):
|
||||
"""
|
||||
With enable_oauth2_auth=false, opaque tokens must not be sent to OAuth2.
|
||||
"""
|
||||
opaque_token = "sk-ui-session-token"
|
||||
general_settings = {
|
||||
"enable_oauth2_auth": False,
|
||||
"enable_jwt_auth": True,
|
||||
}
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/v1/chat/completions"
|
||||
mock_request.headers = {"authorization": f"Bearer {opaque_token}"}
|
||||
mock_request.query_params = {}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings", general_settings
|
||||
), patch("litellm.proxy.proxy_server.premium_user", True), patch(
|
||||
"litellm.proxy.proxy_server.master_key", "sk-master"
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", None
|
||||
), patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_oauth2:
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {opaque_token}",
|
||||
)
|
||||
|
||||
assert exc_info.value.type in (
|
||||
ProxyErrorTypes.auth_error,
|
||||
ProxyErrorTypes.no_db_connection,
|
||||
)
|
||||
mock_oauth2.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_routing_override_on_info_route_uses_oauth2_when_oauth2_globally_disabled(
|
||||
self,
|
||||
):
|
||||
"""
|
||||
With enable_oauth2_auth=false, a JWT matching routing_overrides should
|
||||
still route to OAuth2 on info routes.
|
||||
"""
|
||||
jwt_token = (
|
||||
"eyJhbGciOiJSUzI1NiJ9."
|
||||
"eyJpc3MiOiJtYWNoaW5lLWlzc3Vlci5leGFtcGxlLmNvbSIsImNsaWVudF9pZCI6Ik1JRF9MSVRFTExNIn0."
|
||||
"c2ln"
|
||||
)
|
||||
general_settings = {
|
||||
"enable_oauth2_auth": False,
|
||||
"enable_jwt_auth": True,
|
||||
}
|
||||
mock_oauth2_response = UserAPIKeyAuth(
|
||||
api_key=jwt_token,
|
||||
user_id="machine-client-info-override-oauth2-off",
|
||||
)
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/team/list"
|
||||
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
|
||||
mock_request.query_params = {}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings", general_settings
|
||||
), patch("litellm.proxy.proxy_server.premium_user", True), patch(
|
||||
"litellm.proxy.proxy_server.master_key", "sk-master"
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", None
|
||||
), patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_oauth2_response,
|
||||
) as mock_oauth2, patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_jwt_auth:
|
||||
litellm.proxy.proxy_server.jwt_handler.update_environment(
|
||||
prisma_client=None,
|
||||
user_api_key_cache=DualCache(),
|
||||
litellm_jwtauth=LiteLLM_JWTAuth(
|
||||
routing_overrides=[
|
||||
JWTRoutingOverride(
|
||||
iss="machine-issuer.example.com",
|
||||
client_id="MID_LITELLM",
|
||||
path="oauth2",
|
||||
)
|
||||
]
|
||||
),
|
||||
)
|
||||
|
||||
result = await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {jwt_token}",
|
||||
)
|
||||
|
||||
mock_oauth2.assert_called_once_with(token=jwt_token)
|
||||
mock_jwt_auth.assert_not_called()
|
||||
assert result.user_id == "machine-client-info-override-oauth2-off"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_routing_override_on_management_route_does_not_use_oauth2(self):
|
||||
"""
|
||||
JWT routing_overrides should not force OAuth2 on management routes.
|
||||
"""
|
||||
jwt_token = (
|
||||
"eyJhbGciOiJSUzI1NiJ9."
|
||||
"eyJpc3MiOiJtYWNoaW5lLWlzc3Vlci5leGFtcGxlLmNvbSIsImNsaWVudF9pZCI6Ik1JRF9MSVRFTExNIn0."
|
||||
"c2ln"
|
||||
)
|
||||
general_settings = {
|
||||
"enable_oauth2_auth": False,
|
||||
"enable_jwt_auth": True,
|
||||
}
|
||||
mock_jwt_result = {
|
||||
"is_proxy_admin": True,
|
||||
"team_object": None,
|
||||
"user_object": None,
|
||||
"end_user_object": None,
|
||||
"org_object": None,
|
||||
"token": jwt_token,
|
||||
"team_id": None,
|
||||
"user_id": "jwt-admin-user",
|
||||
"end_user_id": None,
|
||||
"org_id": None,
|
||||
"team_membership": None,
|
||||
"jwt_claims": {
|
||||
"iss": "machine-issuer.example.com",
|
||||
"client_id": "MID_LITELLM",
|
||||
},
|
||||
}
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/key/generate"
|
||||
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
|
||||
mock_request.query_params = {}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings", general_settings
|
||||
), patch("litellm.proxy.proxy_server.premium_user", True), patch(
|
||||
"litellm.proxy.proxy_server.master_key", "sk-master"
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", None
|
||||
), patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_oauth2, patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_jwt_result,
|
||||
) as mock_jwt_auth:
|
||||
litellm.proxy.proxy_server.jwt_handler.update_environment(
|
||||
prisma_client=None,
|
||||
user_api_key_cache=DualCache(),
|
||||
litellm_jwtauth=LiteLLM_JWTAuth(
|
||||
routing_overrides=[
|
||||
JWTRoutingOverride(
|
||||
iss="machine-issuer.example.com",
|
||||
client_id="MID_LITELLM",
|
||||
path="oauth2",
|
||||
)
|
||||
]
|
||||
),
|
||||
)
|
||||
|
||||
result = await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {jwt_token}",
|
||||
)
|
||||
|
||||
mock_oauth2.assert_not_called()
|
||||
mock_jwt_auth.assert_called_once()
|
||||
assert result.user_id == "jwt-admin-user"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_only_oauth2_enabled_handles_all_tokens(self):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import sys
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../")
|
||||
|
|
@ -485,3 +486,327 @@ class TestSafeDbOverrides:
|
|||
from litellm.constants import LITELLM_SETTINGS_SAFE_DB_OVERRIDES
|
||||
|
||||
assert "default_internal_user_params" in LITELLM_SETTINGS_SAFE_DB_OVERRIDES
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# POST /team/permissions/bulk_update
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestBulkUpdateTeamMemberPermissions:
|
||||
"""Tests for the bulk_update_team_member_permissions endpoint."""
|
||||
|
||||
def _make_team(self, team_id: str, permissions: list):
|
||||
"""Create a mock team object."""
|
||||
team = MagicMock()
|
||||
team.team_id = team_id
|
||||
team.team_member_permissions = permissions
|
||||
return team
|
||||
|
||||
def _admin_key_dict(self):
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
return UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN.value,
|
||||
api_key="sk-1234",
|
||||
)
|
||||
|
||||
def _non_admin_key_dict(self):
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
return UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER.value,
|
||||
api_key="sk-user",
|
||||
)
|
||||
|
||||
# --- apply_to_all_teams tests ---
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_all_teams_appends_preserving_existing(self, monkeypatch):
|
||||
"""apply_to_all_teams: permissions are merged, not overwritten."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
bulk_update_team_member_permissions,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
team_a = self._make_team("team-a", ["/key/generate"])
|
||||
team_b = self._make_team("team-b", ["/key/delete", "/key/update"])
|
||||
|
||||
mock_batcher = MagicMock()
|
||||
mock_batcher.commit = AsyncMock(return_value=None)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_a, team_b])
|
||||
mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
data = BulkUpdateTeamMemberPermissionsRequest(
|
||||
permissions=["/team/daily/activity"], apply_to_all_teams=True
|
||||
)
|
||||
result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict())
|
||||
|
||||
assert result["teams_updated"] == 2
|
||||
calls = mock_batcher.litellm_teamtable.update.call_args_list
|
||||
assert len(calls) == 2
|
||||
|
||||
team_a_call = [c for c in calls if c.kwargs["where"]["team_id"] == "team-a"][0]
|
||||
assert "/key/generate" in team_a_call.kwargs["data"]["team_member_permissions"]
|
||||
assert "/team/daily/activity" in team_a_call.kwargs["data"]["team_member_permissions"]
|
||||
|
||||
team_b_call = [c for c in calls if c.kwargs["where"]["team_id"] == "team-b"][0]
|
||||
assert "/key/delete" in team_b_call.kwargs["data"]["team_member_permissions"]
|
||||
assert "/key/update" in team_b_call.kwargs["data"]["team_member_permissions"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_all_teams_skips_teams_that_already_have_permission(self, monkeypatch):
|
||||
"""apply_to_all_teams: teams that already have the permission are skipped."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
bulk_update_team_member_permissions,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
team_has = self._make_team("team-has", ["/team/daily/activity", "/key/update"])
|
||||
team_missing = self._make_team("team-missing", ["/key/generate"])
|
||||
|
||||
mock_batcher = MagicMock()
|
||||
mock_batcher.commit = AsyncMock(return_value=None)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_has, team_missing])
|
||||
mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
data = BulkUpdateTeamMemberPermissionsRequest(
|
||||
permissions=["/team/daily/activity"], apply_to_all_teams=True
|
||||
)
|
||||
result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict())
|
||||
|
||||
assert result["teams_updated"] == 1
|
||||
calls = mock_batcher.litellm_teamtable.update.call_args_list
|
||||
assert len(calls) == 1
|
||||
assert calls[0].kwargs["where"]["team_id"] == "team-missing"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_all_teams_pagination(self, monkeypatch):
|
||||
"""apply_to_all_teams: cursor-based pagination processes multiple pages."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
bulk_update_team_member_permissions,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
page1 = [self._make_team(f"team-{i}", []) for i in range(500)]
|
||||
page2 = [self._make_team(f"team-{i}", []) for i in range(500, 502)]
|
||||
|
||||
mock_batcher = MagicMock()
|
||||
mock_batcher.commit = AsyncMock(return_value=None)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(side_effect=[page1, page2])
|
||||
mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
data = BulkUpdateTeamMemberPermissionsRequest(
|
||||
permissions=["/team/daily/activity"], apply_to_all_teams=True
|
||||
)
|
||||
result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict())
|
||||
|
||||
assert result["teams_updated"] == 502
|
||||
find_calls = mock_prisma.db.litellm_teamtable.find_many.call_args_list
|
||||
assert len(find_calls) == 2
|
||||
assert find_calls[1].kwargs["cursor"] == {"team_id": "team-499"}
|
||||
assert mock_batcher.commit.call_count == 2
|
||||
|
||||
# --- team_ids tests ---
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_ids_updates_only_specified_teams(self, monkeypatch):
|
||||
"""team_ids: only the specified teams are fetched and updated."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
bulk_update_team_member_permissions,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
team_a = self._make_team("team-a", ["/key/generate"])
|
||||
team_b = self._make_team("team-b", ["/key/delete"])
|
||||
|
||||
mock_batcher = MagicMock()
|
||||
mock_batcher.commit = AsyncMock(return_value=None)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_a, team_b])
|
||||
mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
data = BulkUpdateTeamMemberPermissionsRequest(
|
||||
permissions=["/team/daily/activity"], team_ids=["team-a", "team-b"]
|
||||
)
|
||||
result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict())
|
||||
|
||||
assert result["teams_updated"] == 2
|
||||
|
||||
# Verify find_many was called with the team_ids filter
|
||||
find_call = mock_prisma.db.litellm_teamtable.find_many.call_args
|
||||
assert find_call.kwargs["where"] == {"team_id": {"in": ["team-a", "team-b"]}}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_ids_skips_teams_that_already_have_permission(self, monkeypatch):
|
||||
"""team_ids: teams that already have the permission are skipped."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
bulk_update_team_member_permissions,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
team_has = self._make_team("team-has", ["/team/daily/activity"])
|
||||
team_missing = self._make_team("team-missing", [])
|
||||
|
||||
mock_batcher = MagicMock()
|
||||
mock_batcher.commit = AsyncMock(return_value=None)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_has, team_missing])
|
||||
mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
data = BulkUpdateTeamMemberPermissionsRequest(
|
||||
permissions=["/team/daily/activity"], team_ids=["team-has", "team-missing"]
|
||||
)
|
||||
result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict())
|
||||
|
||||
assert result["teams_updated"] == 1
|
||||
calls = mock_batcher.litellm_teamtable.update.call_args_list
|
||||
assert calls[0].kwargs["where"]["team_id"] == "team-missing"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_ids_returns_404_for_missing_teams(self, monkeypatch):
|
||||
"""If any provided team_ids don't exist, return 404."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
bulk_update_team_member_permissions,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
team_a = self._make_team("team-a", ["/key/generate"])
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
# Only team-a exists, team-b does not
|
||||
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_a])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
data = BulkUpdateTeamMemberPermissionsRequest(
|
||||
permissions=["/team/daily/activity"], team_ids=["team-a", "team-b"]
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict())
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
assert "team-b" in str(exc_info.value.detail)
|
||||
|
||||
# --- validation tests ---
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rejects_when_no_team_ids_and_no_apply_all(self, monkeypatch):
|
||||
"""Must provide team_ids or set apply_to_all_teams=True."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
bulk_update_team_member_permissions,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
data = BulkUpdateTeamMemberPermissionsRequest(permissions=["/team/daily/activity"])
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict())
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rejects_when_both_team_ids_and_apply_all(self, monkeypatch):
|
||||
"""Cannot set both team_ids and apply_to_all_teams."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
bulk_update_team_member_permissions,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
data = BulkUpdateTeamMemberPermissionsRequest(
|
||||
permissions=["/team/daily/activity"],
|
||||
team_ids=["team-a"],
|
||||
apply_to_all_teams=True,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict())
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_permissions_list_is_noop(self, monkeypatch):
|
||||
"""Passing an empty permissions list returns immediately with 0 updated."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
bulk_update_team_member_permissions,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
data = BulkUpdateTeamMemberPermissionsRequest(permissions=[])
|
||||
result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict())
|
||||
|
||||
assert result["teams_updated"] == 0
|
||||
mock_prisma.db.litellm_teamtable.find_many.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_admin_gets_403(self, monkeypatch):
|
||||
"""Non-admin users are rejected with 403."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
bulk_update_team_member_permissions,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
data = BulkUpdateTeamMemberPermissionsRequest(
|
||||
permissions=["/team/daily/activity"], apply_to_all_teams=True
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._non_admin_key_dict())
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
def test_invalid_permission_rejected_by_pydantic(self):
|
||||
"""Invalid permission strings are rejected at the type level by Pydantic."""
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
BulkUpdateTeamMemberPermissionsRequest(permissions=["/not/a/real/permission"])
|
||||
|
|
|
|||
|
|
@ -148,7 +148,10 @@ async def test_get_prompt_info_by_base_id():
|
|||
)
|
||||
|
||||
# Mock In-Memory Registry
|
||||
# Patch prisma_client to None to avoid leaking state from other tests
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", None
|
||||
), patch(
|
||||
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
|
||||
) as mock_registry:
|
||||
# Setup mocks behavior
|
||||
|
|
|
|||
40
tests/ui_e2e_tests/constants.ts
Normal file
40
tests/ui_e2e_tests/constants.ts
Normal file
|
|
@ -0,0 +1,40 @@
|
|||
export const ADMIN_STORAGE_PATH = "admin.storageState.json";
|
||||
|
||||
// Page enum — maps to ?page= query parameter values in the UI
|
||||
export enum Page {
|
||||
ApiKeys = "api-keys",
|
||||
Teams = "teams",
|
||||
AdminSettings = "settings",
|
||||
}
|
||||
|
||||
// Test user credentials — all users have password "test" (hashed in seed.sql)
|
||||
export enum Role {
|
||||
ProxyAdmin = "proxy_admin",
|
||||
ProxyAdminViewer = "proxy_admin_viewer",
|
||||
InternalUser = "internal_user",
|
||||
InternalUserViewer = "internal_user_viewer",
|
||||
TeamAdmin = "team_admin",
|
||||
}
|
||||
|
||||
export const users: Record<Role, { email: string; password: string }> = {
|
||||
[Role.ProxyAdmin]: {
|
||||
email: "admin",
|
||||
password: process.env.LITELLM_MASTER_KEY || "sk-1234",
|
||||
},
|
||||
[Role.ProxyAdminViewer]: {
|
||||
email: "adminviewer@test.local",
|
||||
password: "test",
|
||||
},
|
||||
[Role.InternalUser]: {
|
||||
email: "internal@test.local",
|
||||
password: "test",
|
||||
},
|
||||
[Role.InternalUserViewer]: {
|
||||
email: "viewer@test.local",
|
||||
password: "test",
|
||||
},
|
||||
[Role.TeamAdmin]: {
|
||||
email: "teamadmin@test.local",
|
||||
password: "test",
|
||||
},
|
||||
};
|
||||
16
tests/ui_e2e_tests/fixtures/config.yml
Normal file
16
tests/ui_e2e_tests/fixtures/config.yml
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
model_list:
|
||||
- model_name: fake-openai-gpt-4
|
||||
litellm_params:
|
||||
model: openai/fake-gpt-4
|
||||
api_base: os.environ/MOCK_LLM_URL
|
||||
api_key: fake-key
|
||||
- model_name: fake-anthropic-claude
|
||||
litellm_params:
|
||||
model: openai/fake-claude
|
||||
api_base: os.environ/MOCK_LLM_URL
|
||||
api_key: fake-key
|
||||
|
||||
general_settings:
|
||||
master_key: os.environ/LITELLM_MASTER_KEY
|
||||
database_url: os.environ/DATABASE_URL
|
||||
store_prompts_in_spend_logs: true
|
||||
118
tests/ui_e2e_tests/fixtures/mock_llm_server/server.py
Normal file
118
tests/ui_e2e_tests/fixtures/mock_llm_server/server.py
Normal file
|
|
@ -0,0 +1,118 @@
|
|||
"""
|
||||
Mock LLM server for UI e2e tests.
|
||||
Responds to OpenAI-format endpoints with canned responses.
|
||||
"""
|
||||
|
||||
import time
|
||||
import json
|
||||
import uuid
|
||||
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
|
||||
app = FastAPI(title="Mock LLM Server")
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "ok"}
|
||||
|
||||
|
||||
@app.get("/v1/models")
|
||||
@app.get("/models")
|
||||
async def list_models():
|
||||
return {
|
||||
"object": "list",
|
||||
"data": [
|
||||
{"id": "fake-gpt-4", "object": "model", "owned_by": "mock"},
|
||||
{"id": "fake-claude", "object": "model", "owned_by": "mock"},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
@app.post("/chat/completions")
|
||||
async def chat_completions(request: Request):
|
||||
body = await request.json()
|
||||
model = body.get("model", "mock-model")
|
||||
stream = body.get("stream", False)
|
||||
|
||||
response_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||
created = int(time.time())
|
||||
|
||||
if stream:
|
||||
async def stream_generator():
|
||||
chunk = {
|
||||
"id": response_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"role": "assistant", "content": "This is a mock response."},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
yield f"data: {json.dumps(chunk)}\n\n"
|
||||
|
||||
done_chunk = {
|
||||
"id": response_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
||||
}
|
||||
yield f"data: {json.dumps(done_chunk)}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(
|
||||
stream_generator(), media_type="text/event-stream"
|
||||
)
|
||||
|
||||
return {
|
||||
"id": response_id,
|
||||
"object": "chat.completion",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "This is a mock response."},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18},
|
||||
}
|
||||
|
||||
|
||||
@app.post("/v1/embeddings")
|
||||
@app.post("/embeddings")
|
||||
async def embeddings(request: Request):
|
||||
body = await request.json()
|
||||
inputs = body.get("input", [""])
|
||||
if isinstance(inputs, str):
|
||||
inputs = [inputs]
|
||||
return {
|
||||
"object": "list",
|
||||
"data": [
|
||||
{"object": "embedding", "index": i, "embedding": [0.0] * 1536}
|
||||
for i in range(len(inputs))
|
||||
],
|
||||
"model": body.get("model", "mock-embedding"),
|
||||
"usage": {"prompt_tokens": 5, "total_tokens": 5},
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
uvicorn.run(app, host="127.0.0.1", port=8090)
|
||||
103
tests/ui_e2e_tests/fixtures/seed.sql
Normal file
103
tests/ui_e2e_tests/fixtures/seed.sql
Normal file
|
|
@ -0,0 +1,103 @@
|
|||
-- UI E2E Test Database Seed
|
||||
-- Run with: psql $DATABASE_URL -f seed.sql
|
||||
|
||||
-- ============================================================
|
||||
-- 1. Budget Table (must be first — referenced by org FK)
|
||||
-- ============================================================
|
||||
INSERT INTO "LiteLLM_BudgetTable" (
|
||||
budget_id, max_budget, created_by, updated_by
|
||||
) VALUES (
|
||||
'e2e-budget-org', 1000.0, 'e2e-proxy-admin', 'e2e-proxy-admin'
|
||||
) ON CONFLICT (budget_id) DO NOTHING;
|
||||
|
||||
-- ============================================================
|
||||
-- 2. Organization
|
||||
-- ============================================================
|
||||
INSERT INTO "LiteLLM_OrganizationTable" (
|
||||
organization_id, organization_alias, budget_id, metadata, models, spend,
|
||||
model_spend, created_by, updated_by
|
||||
) VALUES (
|
||||
'e2e-org-main', 'E2E Organization', 'e2e-budget-org', '{}'::jsonb,
|
||||
ARRAY[]::text[], 0.0, '{}'::jsonb, 'e2e-proxy-admin', 'e2e-proxy-admin'
|
||||
) ON CONFLICT (organization_id) DO NOTHING;
|
||||
|
||||
-- ============================================================
|
||||
-- 3. Users (password is scrypt hash of "test")
|
||||
-- ============================================================
|
||||
INSERT INTO "LiteLLM_UserTable" (
|
||||
user_id, user_email, user_role, password, teams, models, metadata,
|
||||
spend, model_spend, model_max_budget
|
||||
) VALUES
|
||||
(
|
||||
'e2e-proxy-admin', 'admin@test.local', 'proxy_admin', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr',
|
||||
ARRAY['e2e-team-crud']::text[], ARRAY[]::text[], '{}'::jsonb,
|
||||
0.0, '{}'::jsonb, '{}'::jsonb
|
||||
),
|
||||
(
|
||||
'e2e-admin-viewer', 'adminviewer@test.local', 'proxy_admin_viewer', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr',
|
||||
ARRAY[]::text[], ARRAY[]::text[], '{}'::jsonb,
|
||||
0.0, '{}'::jsonb, '{}'::jsonb
|
||||
),
|
||||
(
|
||||
'e2e-internal-user', 'internal@test.local', 'internal_user', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr',
|
||||
ARRAY['e2e-team-crud', 'e2e-team-org']::text[], ARRAY[]::text[], '{}'::jsonb,
|
||||
0.0, '{}'::jsonb, '{}'::jsonb
|
||||
),
|
||||
(
|
||||
'e2e-internal-viewer', 'viewer@test.local', 'internal_user_viewer', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr',
|
||||
ARRAY[]::text[], ARRAY[]::text[], '{}'::jsonb,
|
||||
0.0, '{}'::jsonb, '{}'::jsonb
|
||||
),
|
||||
(
|
||||
'e2e-team-admin', 'teamadmin@test.local', 'internal_user', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr',
|
||||
ARRAY['e2e-team-crud', 'e2e-team-delete']::text[], ARRAY[]::text[], '{}'::jsonb,
|
||||
0.0, '{}'::jsonb, '{}'::jsonb
|
||||
)
|
||||
ON CONFLICT (user_id) DO NOTHING;
|
||||
|
||||
-- ============================================================
|
||||
-- 4. Teams
|
||||
-- ============================================================
|
||||
INSERT INTO "LiteLLM_TeamTable" (
|
||||
team_id, team_alias, organization_id, admins, members,
|
||||
members_with_roles, metadata, models, spend, model_spend,
|
||||
model_max_budget, blocked
|
||||
) VALUES
|
||||
(
|
||||
'e2e-team-crud', 'E2E Team CRUD', NULL,
|
||||
ARRAY['e2e-team-admin']::text[],
|
||||
ARRAY['e2e-team-admin', 'e2e-internal-user']::text[],
|
||||
'[{"role": "admin", "user_id": "e2e-team-admin"}, {"role": "user", "user_id": "e2e-internal-user"}]'::jsonb,
|
||||
'{}'::jsonb,
|
||||
ARRAY['fake-openai-gpt-4', 'fake-anthropic-claude']::text[],
|
||||
0.0, '{}'::jsonb, '{}'::jsonb, false
|
||||
),
|
||||
(
|
||||
'e2e-team-delete', 'E2E Team Delete', NULL,
|
||||
ARRAY['e2e-team-admin']::text[],
|
||||
ARRAY['e2e-team-admin']::text[],
|
||||
'[{"role": "admin", "user_id": "e2e-team-admin"}]'::jsonb,
|
||||
'{}'::jsonb,
|
||||
ARRAY['fake-openai-gpt-4']::text[],
|
||||
0.0, '{}'::jsonb, '{}'::jsonb, false
|
||||
),
|
||||
(
|
||||
'e2e-team-org', 'E2E Team In Org', 'e2e-org-main',
|
||||
ARRAY[]::text[],
|
||||
ARRAY['e2e-internal-user']::text[],
|
||||
'[{"role": "user", "user_id": "e2e-internal-user"}]'::jsonb,
|
||||
'{}'::jsonb,
|
||||
ARRAY['fake-openai-gpt-4']::text[],
|
||||
0.0, '{}'::jsonb, '{}'::jsonb, false
|
||||
)
|
||||
ON CONFLICT (team_id) DO NOTHING;
|
||||
|
||||
-- ============================================================
|
||||
-- 5. Team Memberships
|
||||
-- ============================================================
|
||||
INSERT INTO "LiteLLM_TeamMembership" (user_id, team_id, spend) VALUES
|
||||
('e2e-team-admin', 'e2e-team-crud', 0.0),
|
||||
('e2e-internal-user', 'e2e-team-crud', 0.0),
|
||||
('e2e-team-admin', 'e2e-team-delete', 0.0),
|
||||
('e2e-internal-user', 'e2e-team-org', 0.0)
|
||||
ON CONFLICT (user_id, team_id) DO NOTHING;
|
||||
32
tests/ui_e2e_tests/globalSetup.ts
Normal file
32
tests/ui_e2e_tests/globalSetup.ts
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
import { chromium, expect } from "@playwright/test";
|
||||
import { users, Role, ADMIN_STORAGE_PATH } from "./constants";
|
||||
import * as fs from "fs";
|
||||
|
||||
async function globalSetup() {
|
||||
const browser = await chromium.launch();
|
||||
const page = await browser.newPage();
|
||||
await page.goto("http://localhost:4000/ui/login");
|
||||
await page.getByPlaceholder("Enter your username").fill(users[Role.ProxyAdmin].email);
|
||||
await page.getByPlaceholder("Enter your password").fill(users[Role.ProxyAdmin].password);
|
||||
await page.getByRole("button", { name: "Login", exact: true }).click();
|
||||
try {
|
||||
// Wait for navigation away from login page into the dashboard
|
||||
await page.waitForURL(
|
||||
(url) => url.pathname.startsWith("/ui") && !url.pathname.includes("/login"),
|
||||
{ timeout: 30_000 },
|
||||
);
|
||||
// Wait for sidebar to render as a signal that the dashboard is ready
|
||||
await expect(page.getByRole("menuitem", { name: "Virtual Keys" })).toBeVisible({ timeout: 30_000 });
|
||||
} catch (e) {
|
||||
// Save a screenshot for debugging before re-throwing
|
||||
fs.mkdirSync("test-results", { recursive: true });
|
||||
await page.screenshot({ path: "test-results/global-setup-failure.png", fullPage: true });
|
||||
console.error("Global setup failed. Screenshot saved to test-results/global-setup-failure.png");
|
||||
console.error("Current URL:", page.url());
|
||||
throw e;
|
||||
}
|
||||
await page.context().storageState({ path: ADMIN_STORAGE_PATH });
|
||||
await browser.close();
|
||||
}
|
||||
|
||||
export default globalSetup;
|
||||
16
tests/ui_e2e_tests/helpers/login.ts
Normal file
16
tests/ui_e2e_tests/helpers/login.ts
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
import { Page as PlaywrightPage, expect } from "@playwright/test";
|
||||
import { users, Role } from "../constants";
|
||||
|
||||
export async function loginAs(page: PlaywrightPage, role: Role) {
|
||||
const user = users[role];
|
||||
await page.goto("/ui/login");
|
||||
await page.getByPlaceholder("Enter your username").fill(user.email);
|
||||
await page.getByPlaceholder("Enter your password").fill(user.password);
|
||||
await page.getByRole("button", { name: "Login", exact: true }).click();
|
||||
// Wait for navigation away from login page into the dashboard
|
||||
await page.waitForURL((url) => url.pathname.startsWith("/ui") && !url.pathname.includes("/login"), {
|
||||
timeout: 30_000,
|
||||
});
|
||||
// Wait for sidebar to render as a signal that the dashboard is ready
|
||||
await expect(page.getByRole("menuitem", { name: "Virtual Keys" })).toBeVisible({ timeout: 30_000 });
|
||||
}
|
||||
6
tests/ui_e2e_tests/helpers/navigation.ts
Normal file
6
tests/ui_e2e_tests/helpers/navigation.ts
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
import { Page as PlaywrightPage } from "@playwright/test";
|
||||
import { Page } from "../constants";
|
||||
|
||||
export async function navigateToPage(page: PlaywrightPage, targetPage: Page) {
|
||||
await page.goto(`/ui?page=${targetPage}`);
|
||||
}
|
||||
76
tests/ui_e2e_tests/package-lock.json
generated
Normal file
76
tests/ui_e2e_tests/package-lock.json
generated
Normal file
|
|
@ -0,0 +1,76 @@
|
|||
{
|
||||
"name": "litellm-ui-e2e-tests",
|
||||
"lockfileVersion": 3,
|
||||
"requires": true,
|
||||
"packages": {
|
||||
"": {
|
||||
"name": "litellm-ui-e2e-tests",
|
||||
"devDependencies": {
|
||||
"@playwright/test": "^1.50.0"
|
||||
}
|
||||
},
|
||||
"node_modules/@playwright/test": {
|
||||
"version": "1.59.1",
|
||||
"resolved": "https://registry.npmjs.org/@playwright/test/-/test-1.59.1.tgz",
|
||||
"integrity": "sha512-PG6q63nQg5c9rIi4/Z5lR5IVF7yU5MqmKaPOe0HSc0O2cX1fPi96sUQu5j7eo4gKCkB2AnNGoWt7y4/Xx3Kcqg==",
|
||||
"dev": true,
|
||||
"license": "Apache-2.0",
|
||||
"dependencies": {
|
||||
"playwright": "1.59.1"
|
||||
},
|
||||
"bin": {
|
||||
"playwright": "cli.js"
|
||||
},
|
||||
"engines": {
|
||||
"node": ">=18"
|
||||
}
|
||||
},
|
||||
"node_modules/fsevents": {
|
||||
"version": "2.3.2",
|
||||
"resolved": "https://registry.npmjs.org/fsevents/-/fsevents-2.3.2.tgz",
|
||||
"integrity": "sha512-xiqMQR4xAeHTuB9uWm+fFRcIOgKBMiOBP+eXiyT7jsgVCq1bkVygt00oASowB7EdtpOHaaPgKt812P9ab+DDKA==",
|
||||
"dev": true,
|
||||
"hasInstallScript": true,
|
||||
"license": "MIT",
|
||||
"optional": true,
|
||||
"os": [
|
||||
"darwin"
|
||||
],
|
||||
"engines": {
|
||||
"node": "^8.16.0 || ^10.6.0 || >=11.0.0"
|
||||
}
|
||||
},
|
||||
"node_modules/playwright": {
|
||||
"version": "1.59.1",
|
||||
"resolved": "https://registry.npmjs.org/playwright/-/playwright-1.59.1.tgz",
|
||||
"integrity": "sha512-C8oWjPR3F81yljW9o5OxcWzfh6avkVwDD2VYdwIGqTkl+OGFISgypqzfu7dOe4QNLL2aqcWBmI3PMtLIK233lw==",
|
||||
"dev": true,
|
||||
"license": "Apache-2.0",
|
||||
"dependencies": {
|
||||
"playwright-core": "1.59.1"
|
||||
},
|
||||
"bin": {
|
||||
"playwright": "cli.js"
|
||||
},
|
||||
"engines": {
|
||||
"node": ">=18"
|
||||
},
|
||||
"optionalDependencies": {
|
||||
"fsevents": "2.3.2"
|
||||
}
|
||||
},
|
||||
"node_modules/playwright-core": {
|
||||
"version": "1.59.1",
|
||||
"resolved": "https://registry.npmjs.org/playwright-core/-/playwright-core-1.59.1.tgz",
|
||||
"integrity": "sha512-HBV/RJg81z5BiiZ9yPzIiClYV/QMsDCKUyogwH9p3MCP6IYjUFu/MActgYAvK0oWyV9NlwM3GLBjADyWgydVyg==",
|
||||
"dev": true,
|
||||
"license": "Apache-2.0",
|
||||
"bin": {
|
||||
"playwright-core": "cli.js"
|
||||
},
|
||||
"engines": {
|
||||
"node": ">=18"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
12
tests/ui_e2e_tests/package.json
Normal file
12
tests/ui_e2e_tests/package.json
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
{
|
||||
"name": "litellm-ui-e2e-tests",
|
||||
"private": true,
|
||||
"devDependencies": {
|
||||
"@playwright/test": "^1.50.0"
|
||||
},
|
||||
"scripts": {
|
||||
"e2e": "playwright test",
|
||||
"e2e:headed": "playwright test --headed",
|
||||
"e2e:ui": "playwright test --ui"
|
||||
}
|
||||
}
|
||||
31
tests/ui_e2e_tests/playwright.config.ts
Normal file
31
tests/ui_e2e_tests/playwright.config.ts
Normal file
|
|
@ -0,0 +1,31 @@
|
|||
import { defineConfig, devices } from "@playwright/test";
|
||||
|
||||
const isCI = !!process.env.CI;
|
||||
|
||||
export default defineConfig({
|
||||
testDir: "./tests",
|
||||
testMatch: "**/*.spec.ts",
|
||||
globalSetup: "./globalSetup.ts",
|
||||
fullyParallel: false,
|
||||
forbidOnly: isCI,
|
||||
retries: isCI ? 2 : 0,
|
||||
workers: 1,
|
||||
reporter: isCI ? [["html", { open: "never" }]] : [["html"]],
|
||||
timeout: 4 * 60 * 1000,
|
||||
expect: {
|
||||
timeout: 10_000,
|
||||
},
|
||||
use: {
|
||||
baseURL: "http://localhost:4000",
|
||||
trace: "on-first-retry",
|
||||
screenshot: "only-on-failure",
|
||||
actionTimeout: 15_000,
|
||||
navigationTimeout: 30_000,
|
||||
},
|
||||
projects: [
|
||||
{
|
||||
name: "chromium",
|
||||
use: { ...devices["Desktop Chrome"] },
|
||||
},
|
||||
],
|
||||
});
|
||||
162
tests/ui_e2e_tests/run_e2e.sh
Executable file
162
tests/ui_e2e_tests/run_e2e.sh
Executable file
|
|
@ -0,0 +1,162 @@
|
|||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
# ================================================================
|
||||
# UI E2E Test Runner
|
||||
# Starts postgres, seeds DB, starts mock + proxy, runs Playwright.
|
||||
# All credentials are generated per run — nothing is stored on disk.
|
||||
#
|
||||
# In CI (CI=true), expects:
|
||||
# - PostgreSQL already running on 127.0.0.1:5432
|
||||
# - DATABASE_URL already set
|
||||
# - Python/Poetry already installed
|
||||
# - Node.js/npx already available
|
||||
# ================================================================
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||
REPO_ROOT="$(cd "$SCRIPT_DIR/../.." && pwd)"
|
||||
IS_CI="${CI:-false}"
|
||||
CONTAINER_NAME="litellm-e2e-postgres-$$"
|
||||
MOCK_PID=""
|
||||
PROXY_PID=""
|
||||
|
||||
# --- Ensure common tool paths are available (local dev only) ---
|
||||
if [ "$IS_CI" = "false" ]; then
|
||||
for p in /usr/local/bin /opt/homebrew/bin "$HOME/.local/bin" /opt/homebrew/opt/postgresql@14/bin /opt/homebrew/opt/libpq/bin; do
|
||||
[ -d "$p" ] && export PATH="$p:$PATH"
|
||||
done
|
||||
[ -s "$HOME/.nvm/nvm.sh" ] && source "$HOME/.nvm/nvm.sh"
|
||||
fi
|
||||
|
||||
# --- Cleanup on exit ---
|
||||
cleanup() {
|
||||
echo "Cleaning up..."
|
||||
[ -n "$MOCK_PID" ] && kill "$MOCK_PID" 2>/dev/null || true
|
||||
[ -n "$PROXY_PID" ] && kill "$PROXY_PID" 2>/dev/null || true
|
||||
if [ "$IS_CI" = "false" ]; then
|
||||
docker stop "$CONTAINER_NAME" 2>/dev/null || true
|
||||
fi
|
||||
echo "Done."
|
||||
}
|
||||
trap cleanup EXIT INT TERM
|
||||
|
||||
# --- Pre-flight checks ---
|
||||
for cmd in python3 npx poetry; do
|
||||
command -v "$cmd" >/dev/null 2>&1 || { echo "Error: $cmd not found."; exit 1; }
|
||||
done
|
||||
|
||||
# --- Database setup ---
|
||||
if [ "$IS_CI" = "false" ]; then
|
||||
# Local: spin up a postgres container
|
||||
for cmd in docker psql; do
|
||||
command -v "$cmd" >/dev/null 2>&1 || { echo "Error: $cmd not found."; exit 1; }
|
||||
done
|
||||
for port in 4000 5432 8090; do
|
||||
if lsof -ti ":$port" >/dev/null 2>&1; then
|
||||
echo "Error: port $port is in use"
|
||||
exit 1
|
||||
fi
|
||||
done
|
||||
|
||||
export POSTGRES_USER="e2euser"
|
||||
export POSTGRES_PASSWORD="$(openssl rand -hex 32)"
|
||||
export POSTGRES_DB="litellm_e2e"
|
||||
export DATABASE_URL="postgresql://${POSTGRES_USER}:${POSTGRES_PASSWORD}@127.0.0.1:5432/${POSTGRES_DB}"
|
||||
|
||||
echo "=== Starting PostgreSQL ==="
|
||||
docker run -d --rm --name "$CONTAINER_NAME" \
|
||||
-e POSTGRES_USER -e POSTGRES_PASSWORD -e POSTGRES_DB \
|
||||
-p 127.0.0.1:5432:5432 \
|
||||
postgres:16
|
||||
|
||||
echo "Waiting for PostgreSQL..."
|
||||
for i in $(seq 1 30); do
|
||||
if PGPASSWORD="$POSTGRES_PASSWORD" pg_isready -h 127.0.0.1 -U "$POSTGRES_USER" -d "$POSTGRES_DB" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
else
|
||||
# CI: postgres is already running as a service container
|
||||
echo "=== Using CI PostgreSQL service ==="
|
||||
: "${DATABASE_URL:?DATABASE_URL must be set in CI}"
|
||||
fi
|
||||
|
||||
# --- Credentials ---
|
||||
export LITELLM_MASTER_KEY="sk-e2e-$(openssl rand -hex 32)"
|
||||
export MOCK_LLM_URL="http://127.0.0.1:8090/v1"
|
||||
export DISABLE_SCHEMA_UPDATE="true"
|
||||
|
||||
# --- Python environment ---
|
||||
echo "=== Setting up Python environment ==="
|
||||
cd "$REPO_ROOT"
|
||||
if ! poetry run python3 -c "import prisma" 2>/dev/null; then
|
||||
echo "Installing Python dependencies (first run)..."
|
||||
poetry install --with dev,proxy-dev --extras "proxy" --quiet
|
||||
poetry run pip install nodejs-wheel-binaries 2>/dev/null || true
|
||||
poetry run prisma generate --schema litellm/proxy/schema.prisma
|
||||
fi
|
||||
|
||||
echo "=== Pushing Prisma schema to database ==="
|
||||
poetry run prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
|
||||
|
||||
# --- Mock LLM server ---
|
||||
echo "=== Starting mock LLM server ==="
|
||||
poetry run python3 "$SCRIPT_DIR/fixtures/mock_llm_server/server.py" &
|
||||
MOCK_PID=$!
|
||||
|
||||
for i in $(seq 1 15); do
|
||||
if curl -sf http://127.0.0.1:8090/health >/dev/null 2>&1; then break; fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
# --- LiteLLM proxy ---
|
||||
echo "=== Starting LiteLLM proxy ==="
|
||||
cd "$REPO_ROOT"
|
||||
poetry run python3 -m litellm.proxy.proxy_cli \
|
||||
--config "$SCRIPT_DIR/fixtures/config.yml" \
|
||||
--port 4000 &
|
||||
PROXY_PID=$!
|
||||
|
||||
echo "Waiting for proxy..."
|
||||
PROXY_READY=0
|
||||
for i in $(seq 1 180); do
|
||||
if ! kill -0 "$PROXY_PID" 2>/dev/null; then
|
||||
echo "Error: proxy process exited unexpectedly"
|
||||
exit 1
|
||||
fi
|
||||
HTTP_CODE=$(curl -s -o /dev/null -w "%{http_code}" http://127.0.0.1:4000/health -H "Authorization: Bearer $LITELLM_MASTER_KEY" 2>/dev/null || true)
|
||||
if [ "$HTTP_CODE" = "200" ]; then
|
||||
PROXY_READY=1
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
if [ "$PROXY_READY" -ne 1 ]; then
|
||||
echo "Error: proxy did not become healthy within 180 seconds"
|
||||
exit 1
|
||||
fi
|
||||
echo "Proxy is ready."
|
||||
|
||||
# --- Seed database ---
|
||||
echo "=== Seeding database ==="
|
||||
# Extract credentials from DATABASE_URL for psql
|
||||
DB_USER=$(echo "$DATABASE_URL" | sed -n 's|.*://\([^:]*\):.*|\1|p')
|
||||
DB_PASS=$(echo "$DATABASE_URL" | sed -n 's|.*://[^:]*:\([^@]*\)@.*|\1|p')
|
||||
DB_HOST=$(echo "$DATABASE_URL" | sed -n 's|.*@\([^:]*\):.*|\1|p')
|
||||
DB_PORT=$(echo "$DATABASE_URL" | sed -n 's|.*:\([0-9]*\)/.*|\1|p')
|
||||
DB_NAME=$(echo "$DATABASE_URL" | sed -n 's|.*/\([^?]*\).*|\1|p')
|
||||
|
||||
PGPASSWORD="$DB_PASS" psql -h "$DB_HOST" -p "$DB_PORT" -U "$DB_USER" -d "$DB_NAME" \
|
||||
-f "$SCRIPT_DIR/fixtures/seed.sql"
|
||||
|
||||
# --- Playwright ---
|
||||
echo "=== Installing Playwright dependencies ==="
|
||||
cd "$SCRIPT_DIR"
|
||||
npm install --silent
|
||||
|
||||
echo "=== Running Playwright tests ==="
|
||||
npx playwright test "$@"
|
||||
EXIT_CODE=$?
|
||||
|
||||
exit $EXIT_CODE
|
||||
10
tests/ui_e2e_tests/tests/roles/admin-viewer.spec.ts
Normal file
10
tests/ui_e2e_tests/tests/roles/admin-viewer.spec.ts
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
import { test, expect } from "@playwright/test";
|
||||
import { Role } from "../../constants";
|
||||
import { loginAs } from "../../helpers/login";
|
||||
|
||||
test.describe("Admin Viewer Role", () => {
|
||||
test("Should not see Test Key page", async ({ page }) => {
|
||||
await loginAs(page, Role.ProxyAdminViewer);
|
||||
await expect(page.getByRole("menuitem", { name: "Test Key" })).not.toBeVisible();
|
||||
});
|
||||
});
|
||||
12
tests/ui_e2e_tests/tests/roles/internal-user.spec.ts
Normal file
12
tests/ui_e2e_tests/tests/roles/internal-user.spec.ts
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
import { test, expect } from "@playwright/test";
|
||||
import { Page, Role } from "../../constants";
|
||||
import { loginAs } from "../../helpers/login";
|
||||
import { navigateToPage } from "../../helpers/navigation";
|
||||
|
||||
test.describe("Internal User Role", () => {
|
||||
test("Should not see litellm-dashboard keys", async ({ page }) => {
|
||||
await loginAs(page, Role.InternalUser);
|
||||
await navigateToPage(page, Page.ApiKeys);
|
||||
await expect(page.getByText("litellm-dashboard")).not.toBeVisible();
|
||||
});
|
||||
});
|
||||
28
tests/ui_e2e_tests/tests/roles/internal-viewer.spec.ts
Normal file
28
tests/ui_e2e_tests/tests/roles/internal-viewer.spec.ts
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
import { test, expect } from "@playwright/test";
|
||||
import { Page, Role } from "../../constants";
|
||||
import { loginAs } from "../../helpers/login";
|
||||
import { navigateToPage } from "../../helpers/navigation";
|
||||
|
||||
test.describe("Internal User Viewer Role", () => {
|
||||
test("Can only see allowed pages", async ({ page }) => {
|
||||
await loginAs(page, Role.InternalUserViewer);
|
||||
await expect(page.getByRole("menuitem", { name: "Virtual Keys" })).toBeVisible();
|
||||
await expect(page.getByRole("menuitem", { name: "Admin Settings" })).not.toBeVisible();
|
||||
});
|
||||
|
||||
test("Cannot create keys", async ({ page }) => {
|
||||
await loginAs(page, Role.InternalUserViewer);
|
||||
await navigateToPage(page, Page.ApiKeys);
|
||||
await expect(page.getByRole("button", { name: /Create New Key/i })).not.toBeVisible();
|
||||
});
|
||||
|
||||
test("Cannot edit or delete keys", async ({ page }) => {
|
||||
await loginAs(page, Role.InternalUserViewer);
|
||||
await navigateToPage(page, Page.ApiKeys);
|
||||
// Ensure the keys table has loaded before asserting absence of actions
|
||||
await expect(page.getByRole("menuitem", { name: "Virtual Keys" })).toBeVisible();
|
||||
await expect(page.getByRole("button", { name: /Edit Key/i })).not.toBeVisible();
|
||||
await expect(page.getByRole("button", { name: /Delete Key/i })).not.toBeVisible();
|
||||
await expect(page.getByRole("button", { name: /Regenerate Key/i })).not.toBeVisible();
|
||||
});
|
||||
});
|
||||
21
tests/ui_e2e_tests/tests/roles/proxy-admin.spec.ts
Normal file
21
tests/ui_e2e_tests/tests/roles/proxy-admin.spec.ts
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
import { test, expect } from "@playwright/test";
|
||||
import { ADMIN_STORAGE_PATH, Page, Role, users } from "../../constants";
|
||||
import { navigateToPage } from "../../helpers/navigation";
|
||||
|
||||
test.describe("Proxy Admin Role", () => {
|
||||
test.use({ storageState: ADMIN_STORAGE_PATH });
|
||||
|
||||
test("Can create keys", async ({ page }) => {
|
||||
await navigateToPage(page, Page.ApiKeys);
|
||||
await expect(page.getByRole("button", { name: /Create New Key/i })).toBeVisible();
|
||||
});
|
||||
|
||||
test("Can list teams via API", async ({ page }) => {
|
||||
const response = await page.request.get("/team/list", {
|
||||
headers: {
|
||||
Authorization: `Bearer ${users[Role.ProxyAdmin].password}`,
|
||||
},
|
||||
});
|
||||
expect(response.status()).toBe(200);
|
||||
});
|
||||
});
|
||||
13
tests/ui_e2e_tests/tests/roles/team-admin.spec.ts
Normal file
13
tests/ui_e2e_tests/tests/roles/team-admin.spec.ts
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
import { test, expect } from "@playwright/test";
|
||||
import { Page, Role } from "../../constants";
|
||||
import { loginAs } from "../../helpers/login";
|
||||
import { navigateToPage } from "../../helpers/navigation";
|
||||
|
||||
test.describe("Team Admin Role", () => {
|
||||
test("Can view team keys but not admin settings", async ({ page }) => {
|
||||
await loginAs(page, Role.TeamAdmin);
|
||||
await navigateToPage(page, Page.ApiKeys);
|
||||
await expect(page.getByRole("menuitem", { name: "Virtual Keys" })).toBeVisible();
|
||||
await expect(page.getByRole("menuitem", { name: "Admin Settings" })).not.toBeVisible();
|
||||
});
|
||||
});
|
||||
18
tests/ui_e2e_tests/tests/security/login-logout.spec.ts
Normal file
18
tests/ui_e2e_tests/tests/security/login-logout.spec.ts
Normal file
|
|
@ -0,0 +1,18 @@
|
|||
import { test, expect } from "@playwright/test";
|
||||
import { users, Role } from "../../constants";
|
||||
|
||||
test.describe("Authentication", () => {
|
||||
test("Login with valid admin credentials", async ({ page }) => {
|
||||
await page.goto("/ui/login");
|
||||
await page.getByPlaceholder("Enter your username").fill(users[Role.ProxyAdmin].email);
|
||||
await page.getByPlaceholder("Enter your password").fill(users[Role.ProxyAdmin].password);
|
||||
await page.getByRole("button", { name: "Login", exact: true }).click();
|
||||
await expect(page.getByRole("menuitem", { name: "Virtual Keys" })).toBeVisible();
|
||||
});
|
||||
|
||||
test("Unauthenticated user is redirected to login", async ({ page }) => {
|
||||
await page.goto("/ui");
|
||||
await page.waitForURL(/\/ui\/login/);
|
||||
await expect(page.getByRole("heading", { name: /Login/i })).toBeVisible();
|
||||
});
|
||||
});
|
||||
11
tests/ui_e2e_tests/tsconfig.json
Normal file
11
tests/ui_e2e_tests/tsconfig.json
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
{
|
||||
"compilerOptions": {
|
||||
"target": "ES2020",
|
||||
"module": "commonjs",
|
||||
"strict": true,
|
||||
"esModuleInterop": true,
|
||||
"outDir": "./dist",
|
||||
"rootDir": "."
|
||||
},
|
||||
"include": ["**/*.ts"]
|
||||
}
|
||||
|
|
@ -1,7 +0,0 @@
|
|||
# js-yaml CVE-2025-64718
|
||||
# This vulnerability is not applicable because we've forced js-yaml to version 4.1.1
|
||||
# via npm overrides in package.json. Trivy incorrectly reports this based on
|
||||
# dependency requirements in the lockfile, but the actual installed version is 4.1.1.
|
||||
# Verified with: npm list js-yaml
|
||||
CVE-2025-64718
|
||||
|
||||
|
|
@ -1,6 +1,22 @@
|
|||
// Storage state paths for each role
|
||||
export const ADMIN_STORAGE_PATH = "admin.storageState.json";
|
||||
export const ADMIN_VIEWER_STORAGE_PATH = "adminViewer.storageState.json";
|
||||
export const INTERNAL_USER_STORAGE_PATH = "internalUser.storageState.json";
|
||||
export const INTERNAL_VIEWER_STORAGE_PATH = "internalViewer.storageState.json";
|
||||
export const TEAM_ADMIN_STORAGE_PATH = "teamAdmin.storageState.json";
|
||||
|
||||
export const E2E_UPDATE_LIMITS_KEY_ID_PREFIX = "102c";
|
||||
export const E2E_DELETE_KEY_ID_PREFIX = "94a5";
|
||||
export const E2E_DELETE_KEY_NAME = "e2eDeleteKey";
|
||||
export const E2E_REGENERATE_KEY_ID_PREFIX = "593a";
|
||||
// Key aliases for seeded test keys (match seed.sql)
|
||||
export const E2E_UPDATE_LIMITS_KEY_ALIAS = "e2eUpdateLimitsKey";
|
||||
export const E2E_DELETE_KEY_ALIAS = "e2eDeleteKey";
|
||||
export const E2E_REGENERATE_KEY_ALIAS = "e2eRegenerateKey";
|
||||
export const E2E_INTERNAL_USER_KEY_ALIAS = "e2eInternalUserKey";
|
||||
export const E2E_VIEWER_KEY_ALIAS = "e2eViewerKey";
|
||||
|
||||
// Team identifiers (match seed.sql)
|
||||
export const E2E_TEAM_CRUD_ID = "e2e-team-crud";
|
||||
export const E2E_TEAM_CRUD_ALIAS = "E2E Team CRUD";
|
||||
export const E2E_TEAM_DELETE_ID = "e2e-team-delete";
|
||||
export const E2E_TEAM_DELETE_ALIAS = "E2E Team Delete";
|
||||
export const E2E_TEAM_ORG_ID = "e2e-team-org";
|
||||
export const E2E_TEAM_NO_ADMIN_ID = "e2e-team-no-admin";
|
||||
export const E2E_TEAM_NO_ADMIN_ALIAS = "E2E Team No Admin";
|
||||
|
|
|
|||
16
ui/litellm-dashboard/e2e_tests/fixtures/config.yml
Normal file
16
ui/litellm-dashboard/e2e_tests/fixtures/config.yml
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
model_list:
|
||||
- model_name: fake-openai-gpt-4
|
||||
litellm_params:
|
||||
model: openai/fake-gpt-4
|
||||
api_base: os.environ/MOCK_LLM_URL
|
||||
api_key: fake-key
|
||||
- model_name: fake-anthropic-claude
|
||||
litellm_params:
|
||||
model: openai/fake-claude
|
||||
api_base: os.environ/MOCK_LLM_URL
|
||||
api_key: fake-key
|
||||
|
||||
general_settings:
|
||||
master_key: os.environ/LITELLM_MASTER_KEY
|
||||
database_url: os.environ/DATABASE_URL
|
||||
store_prompts_in_spend_logs: true
|
||||
|
|
@ -0,0 +1,120 @@
|
|||
"""
|
||||
Mock LLM server for UI e2e tests.
|
||||
Responds to OpenAI-format endpoints with canned responses.
|
||||
"""
|
||||
|
||||
import time
|
||||
import json
|
||||
import uuid
|
||||
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
|
||||
app = FastAPI(title="Mock LLM Server")
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "ok"}
|
||||
|
||||
|
||||
@app.get("/v1/models")
|
||||
@app.get("/models")
|
||||
async def list_models():
|
||||
return {
|
||||
"object": "list",
|
||||
"data": [
|
||||
{"id": "fake-gpt-4", "object": "model", "owned_by": "mock"},
|
||||
{"id": "fake-claude", "object": "model", "owned_by": "mock"},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
@app.post("/chat/completions")
|
||||
async def chat_completions(request: Request):
|
||||
body = await request.json()
|
||||
model = body.get("model", "mock-model")
|
||||
stream = body.get("stream", False)
|
||||
|
||||
response_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||
created = int(time.time())
|
||||
|
||||
if stream:
|
||||
|
||||
async def stream_generator():
|
||||
chunk = {
|
||||
"id": response_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"role": "assistant",
|
||||
"content": "This is a mock response.",
|
||||
},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
yield f"data: {json.dumps(chunk)}\n\n"
|
||||
|
||||
done_chunk = {
|
||||
"id": response_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
||||
}
|
||||
yield f"data: {json.dumps(done_chunk)}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(stream_generator(), media_type="text/event-stream")
|
||||
|
||||
return {
|
||||
"id": response_id,
|
||||
"object": "chat.completion",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "This is a mock response."},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18},
|
||||
}
|
||||
|
||||
|
||||
@app.post("/v1/embeddings")
|
||||
@app.post("/embeddings")
|
||||
async def embeddings(request: Request):
|
||||
body = await request.json()
|
||||
inputs = body.get("input", [""])
|
||||
if isinstance(inputs, str):
|
||||
inputs = [inputs]
|
||||
return {
|
||||
"object": "list",
|
||||
"data": [
|
||||
{"object": "embedding", "index": i, "embedding": [0.0] * 1536}
|
||||
for i in range(len(inputs))
|
||||
],
|
||||
"model": body.get("model", "mock-embedding"),
|
||||
"usage": {"prompt_tokens": 5, "total_tokens": 5},
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
uvicorn.run(app, host="127.0.0.1", port=8090)
|
||||
84
ui/litellm-dashboard/e2e_tests/fixtures/seed.sql
Normal file
84
ui/litellm-dashboard/e2e_tests/fixtures/seed.sql
Normal file
|
|
@ -0,0 +1,84 @@
|
|||
-- E2E Test Seed Data
|
||||
-- Idempotent: deletes all e2e-* rows then re-inserts deterministic data.
|
||||
|
||||
-- 1. Clean up in dependency order
|
||||
DELETE FROM "LiteLLM_TeamMembership" WHERE "user_id" LIKE 'e2e-%';
|
||||
DELETE FROM "LiteLLM_VerificationToken" WHERE token LIKE 'e2e-%';
|
||||
DELETE FROM "LiteLLM_TeamTable" WHERE "team_id" LIKE 'e2e-%';
|
||||
DELETE FROM "LiteLLM_OrganizationTable" WHERE "organization_id" LIKE 'e2e-%';
|
||||
DELETE FROM "LiteLLM_UserTable" WHERE "user_id" LIKE 'e2e-%';
|
||||
DELETE FROM "LiteLLM_BudgetTable" WHERE "budget_id" LIKE 'e2e-%';
|
||||
|
||||
-- 2. Budget (created_by and updated_by are NOT NULL)
|
||||
INSERT INTO "LiteLLM_BudgetTable" ("budget_id", "max_budget", "created_by", "updated_by")
|
||||
VALUES ('e2e-budget-org', 1000, 'e2e-proxy-admin', 'e2e-proxy-admin');
|
||||
|
||||
-- 3. Organization (created_by and updated_by are NOT NULL)
|
||||
INSERT INTO "LiteLLM_OrganizationTable" (
|
||||
"organization_id", "organization_alias", "budget_id",
|
||||
"metadata", "models", "spend", "model_spend",
|
||||
"created_by", "updated_by"
|
||||
) VALUES (
|
||||
'e2e-org-main', 'E2E Organization', 'e2e-budget-org',
|
||||
'{}'::jsonb, ARRAY[]::text[], 0.0, '{}'::jsonb,
|
||||
'e2e-proxy-admin', 'e2e-proxy-admin'
|
||||
);
|
||||
|
||||
-- 4. Users (password hash is scrypt of "test")
|
||||
INSERT INTO "LiteLLM_UserTable" ("user_id", "user_email", "user_role", "teams", "password")
|
||||
VALUES
|
||||
('e2e-proxy-admin', 'admin@test.local', 'proxy_admin', '{"e2e-team-crud"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-admin-viewer', 'adminviewer@test.local', 'proxy_admin_viewer', '{}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-internal-user', 'internal@test.local', 'internal_user', '{"e2e-team-crud","e2e-team-org"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-internal-viewer', 'viewer@test.local', 'internal_user_viewer', '{"e2e-team-crud"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-team-admin', 'teamadmin@test.local', 'internal_user', '{"e2e-team-crud","e2e-team-delete"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-invitable-user', 'invitable@test.local', 'internal_user', '{}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-removable-member', 'removable@test.local', 'internal_user', '{"e2e-team-crud"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr');
|
||||
|
||||
-- 5. Teams (members_with_roles is required JSON)
|
||||
INSERT INTO "LiteLLM_TeamTable" (
|
||||
"team_id", "team_alias", "organization_id", "admins", "members",
|
||||
"members_with_roles", "metadata", "models", "spend", "model_spend", "model_max_budget", "blocked"
|
||||
) VALUES
|
||||
('e2e-team-crud', 'E2E Team CRUD', NULL,
|
||||
'{"e2e-team-admin"}',
|
||||
'{"e2e-team-admin","e2e-internal-user","e2e-internal-viewer","e2e-removable-member"}',
|
||||
'[{"role":"admin","user_id":"e2e-team-admin"},{"role":"user","user_id":"e2e-internal-user"},{"role":"user","user_id":"e2e-internal-viewer"},{"role":"user","user_id":"e2e-removable-member"}]'::jsonb,
|
||||
'{}'::jsonb, '{"fake-openai-gpt-4","fake-anthropic-claude"}', 0.0, '{}'::jsonb, '{}'::jsonb, false),
|
||||
|
||||
('e2e-team-delete', 'E2E Team Delete', NULL,
|
||||
'{"e2e-team-admin"}', '{"e2e-team-admin"}',
|
||||
'[{"role":"admin","user_id":"e2e-team-admin"}]'::jsonb,
|
||||
'{}'::jsonb, '{"fake-openai-gpt-4"}', 0.0, '{}'::jsonb, '{}'::jsonb, false),
|
||||
|
||||
('e2e-team-org', 'E2E Team In Org', 'e2e-org-main',
|
||||
'{}', '{"e2e-internal-user"}',
|
||||
'[{"role":"user","user_id":"e2e-internal-user"}]'::jsonb,
|
||||
'{}'::jsonb, '{"fake-openai-gpt-4"}', 0.0, '{}'::jsonb, '{}'::jsonb, false),
|
||||
|
||||
('e2e-team-no-admin', 'E2E Team No Admin', NULL,
|
||||
'{}', '{"e2e-invitable-user"}',
|
||||
'[{"role":"user","user_id":"e2e-invitable-user"}]'::jsonb,
|
||||
'{}'::jsonb, '{"fake-openai-gpt-4"}', 0.0, '{}'::jsonb, '{}'::jsonb, false);
|
||||
|
||||
-- 6. Team Memberships (only user_id, team_id, spend — no created_at/updated_at)
|
||||
INSERT INTO "LiteLLM_TeamMembership" ("user_id", "team_id", "spend")
|
||||
VALUES
|
||||
('e2e-team-admin', 'e2e-team-crud', 0.0),
|
||||
('e2e-internal-user', 'e2e-team-crud', 0.0),
|
||||
('e2e-internal-viewer', 'e2e-team-crud', 0.0),
|
||||
('e2e-removable-member', 'e2e-team-crud', 0.0),
|
||||
('e2e-team-admin', 'e2e-team-delete', 0.0),
|
||||
('e2e-internal-user', 'e2e-team-org', 0.0),
|
||||
('e2e-invitable-user', 'e2e-team-no-admin', 0.0);
|
||||
|
||||
-- 7. Verification Tokens (API Keys)
|
||||
INSERT INTO "LiteLLM_VerificationToken" (
|
||||
"token", "key_name", "key_alias", "user_id", "team_id",
|
||||
"models", "spend", "max_budget", "expires", "metadata"
|
||||
) VALUES
|
||||
('e2e-key-update-limits', 'sk-e2e-update', 'e2eUpdateLimitsKey', 'e2e-proxy-admin', 'e2e-team-crud', '{"fake-openai-gpt-4"}', 0.0, NULL, NULL, '{}'::jsonb),
|
||||
('e2e-key-delete', 'sk-e2e-delete', 'e2eDeleteKey', 'e2e-proxy-admin', 'e2e-team-crud', '{"fake-openai-gpt-4"}', 0.0, NULL, NULL, '{}'::jsonb),
|
||||
('e2e-key-regenerate', 'sk-e2e-regen', 'e2eRegenerateKey', 'e2e-proxy-admin', 'e2e-team-crud', '{"fake-openai-gpt-4"}', 0.0, NULL, NULL, '{}'::jsonb),
|
||||
('e2e-key-internal-user', 'sk-e2e-internal', 'e2eInternalUserKey', 'e2e-internal-user', 'e2e-team-crud', '{"fake-openai-gpt-4"}', 0.0, NULL, NULL, '{}'::jsonb),
|
||||
('e2e-key-viewer', 'sk-e2e-viewer', 'e2eViewerKey', 'e2e-internal-viewer', NULL, '{"fake-openai-gpt-4"}', 0.0, NULL, NULL, '{}'::jsonb);
|
||||
|
|
@ -1,10 +1,38 @@
|
|||
import { Role } from "./roles";
|
||||
export enum Role {
|
||||
ProxyAdmin = "proxy_admin",
|
||||
ProxyAdminViewer = "proxy_admin_viewer",
|
||||
InternalUser = "internal_user",
|
||||
InternalUserViewer = "internal_user_viewer",
|
||||
TeamAdmin = "team_admin",
|
||||
}
|
||||
|
||||
const isCI = !!process.env.CI;
|
||||
|
||||
export const users = {
|
||||
export const users: Record<Role, { email: string; password: string }> = {
|
||||
[Role.ProxyAdmin]: {
|
||||
email: "admin",
|
||||
password: isCI ? "gm" : "sk-1234",
|
||||
password: process.env.LITELLM_MASTER_KEY || "sk-1234",
|
||||
},
|
||||
[Role.ProxyAdminViewer]: {
|
||||
email: "adminviewer@test.local",
|
||||
password: "test",
|
||||
},
|
||||
[Role.InternalUser]: {
|
||||
email: "internal@test.local",
|
||||
password: "test",
|
||||
},
|
||||
[Role.InternalUserViewer]: {
|
||||
email: "viewer@test.local",
|
||||
password: "test",
|
||||
},
|
||||
[Role.TeamAdmin]: {
|
||||
email: "teamadmin@test.local",
|
||||
password: "test",
|
||||
},
|
||||
};
|
||||
|
||||
export const STORAGE_PATHS: Record<Role, string> = {
|
||||
[Role.ProxyAdmin]: "admin.storageState.json",
|
||||
[Role.ProxyAdminViewer]: "adminViewer.storageState.json",
|
||||
[Role.InternalUser]: "internalUser.storageState.json",
|
||||
[Role.InternalUserViewer]: "internalViewer.storageState.json",
|
||||
[Role.TeamAdmin]: "teamAdmin.storageState.json",
|
||||
};
|
||||
|
|
|
|||
|
|
@ -1,17 +1,40 @@
|
|||
import { chromium } from "@playwright/test";
|
||||
import { users } from "./fixtures/users";
|
||||
import { Role } from "./fixtures/roles";
|
||||
import { chromium, expect } from "@playwright/test";
|
||||
import { users, Role, STORAGE_PATHS } from "./fixtures/users";
|
||||
import * as fs from "fs";
|
||||
|
||||
async function globalSetup() {
|
||||
const browser = await chromium.launch();
|
||||
const page = await browser.newPage();
|
||||
await page.goto("http://localhost:4000/ui/login");
|
||||
await page.getByPlaceholder("Enter your username").fill(users[Role.ProxyAdmin].email);
|
||||
await page.getByPlaceholder("Enter your password").fill(users[Role.ProxyAdmin].password);
|
||||
const loginButton = page.getByRole("button", { name: "Login", exact: true });
|
||||
await loginButton.click();
|
||||
await page.waitForSelector("text=Virtual Keys");
|
||||
await page.context().storageState({ path: "admin.storageState.json" });
|
||||
|
||||
for (const role of Object.values(Role)) {
|
||||
const { email, password } = users[role];
|
||||
const storagePath = STORAGE_PATHS[role];
|
||||
const page = await browser.newPage();
|
||||
try {
|
||||
await page.goto("http://localhost:4000/ui/login");
|
||||
await page.getByPlaceholder("Enter your username").fill(email);
|
||||
await page.getByPlaceholder("Enter your password").fill(password);
|
||||
await page.getByRole("button", { name: "Login", exact: true }).click();
|
||||
await page.waitForURL(
|
||||
(url) => url.pathname.startsWith("/ui") && !url.pathname.includes("/login"),
|
||||
{ timeout: 30_000 },
|
||||
);
|
||||
await expect(page.locator("a", { hasText: "Virtual Keys" })).toBeVisible({ timeout: 30_000 });
|
||||
// Dismiss feedback popup if present
|
||||
const dismiss = page.getByText("Don't ask me again");
|
||||
if (await dismiss.isVisible({ timeout: 1_500 }).catch(() => false)) {
|
||||
await dismiss.click();
|
||||
}
|
||||
await page.context().storageState({ path: storagePath });
|
||||
} catch (e) {
|
||||
fs.mkdirSync("test-results", { recursive: true });
|
||||
await page.screenshot({ path: `test-results/global-setup-${role}-failure.png`, fullPage: true });
|
||||
console.error(`Global setup failed for role ${role}. Screenshot saved. URL: ${page.url()}`);
|
||||
throw e;
|
||||
} finally {
|
||||
await page.close();
|
||||
}
|
||||
}
|
||||
|
||||
await browser.close();
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,12 +1,25 @@
|
|||
import { Page } from "../fixtures/pages";
|
||||
import { Page as PlaywrightPage } from "@playwright/test";
|
||||
import { Page as PlaywrightPage, expect } from "@playwright/test";
|
||||
|
||||
/**
|
||||
* Navigates to a specific page using the page query parameter.
|
||||
* Uses relative path which will be resolved against the baseURL configured in playwright.config.ts
|
||||
* @param page - The Playwright page object
|
||||
* @param pageEnum - The page enum value to navigate to
|
||||
* Waits for the sidebar to be visible before returning.
|
||||
*/
|
||||
export async function navigateToPage(page: PlaywrightPage, pageEnum: Page): Promise<void> {
|
||||
await page.goto(`/ui?page=${pageEnum}`);
|
||||
await page.waitForLoadState("networkidle");
|
||||
// Dismiss the "Quick feedback" popup if it appears
|
||||
await dismissFeedbackPopup(page);
|
||||
}
|
||||
|
||||
/**
|
||||
* Dismiss the "Quick feedback" popup that may appear on any page.
|
||||
*/
|
||||
export async function dismissFeedbackPopup(page: PlaywrightPage): Promise<void> {
|
||||
const dismissButton = page.getByText("Don't ask me again");
|
||||
if (await dismissButton.isVisible({ timeout: 1_500 }).catch(() => false)) {
|
||||
await dismissButton.click();
|
||||
// Wait for the popup to disappear
|
||||
await expect(dismissButton).not.toBeVisible({ timeout: 2_000 }).catch(() => {});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -36,11 +36,6 @@ export default defineConfig({
|
|||
name: "chromium",
|
||||
use: { ...devices["Desktop Chrome"] },
|
||||
},
|
||||
|
||||
{
|
||||
name: "firefox",
|
||||
use: { ...devices["Desktop Firefox"] },
|
||||
},
|
||||
],
|
||||
|
||||
/* Timeout settings */
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue