diff --git a/.circleci/config.yml b/.circleci/config.yml
index d998ee467b8..16019fe047b 100644
--- a/.circleci/config.yml
+++ b/.circleci/config.yml
@@ -406,119 +406,6 @@ jobs:
# Store test results
- store_test_results:
path: test-results
- caching_unit_tests:
- docker:
- - image: cimg/python:3.11
- auth:
- username: ${DOCKERHUB_USERNAME}
- password: ${DOCKERHUB_PASSWORD}
- resource_class: large
- working_directory: ~/project
- parallelism: 2
-
- steps:
- - checkout
- - setup_google_dns
- - run:
- name: DNS lookup for Redis host
- command: |
- sudo apt-get update
- sudo apt-get install -y dnsutils
- dig redis-19899.c239.us-east-1-2.ec2.redns.redis-cloud.com +short
- - run:
- name: Show git commit hash
- command: |
- echo "Git commit hash: $CIRCLE_SHA1"
-
- - restore_cache:
- keys:
- - v2-caching-deps-{{ checksum ".circleci/requirements.txt" }}
- - v2-caching-deps-
- - run:
- name: Install Dependencies
- command: |
- python -m pip install --upgrade pip
- python -m pip install -r .circleci/requirements.txt
- pip install "pytest==7.3.1"
- pip install "pytest-retry==1.6.3"
- pip install "pytest-asyncio==0.21.1"
- pip install "pytest-cov==5.0.0"
- pip install "mypy==1.18.2"
- pip install "google-generativeai==0.3.2"
- pip install "google-cloud-aiplatform==1.43.0"
- pip install pyarrow
- pip install "boto3==1.36.0"
- pip install "aioboto3==13.4.0"
- pip install langchain
- pip install lunary==0.2.5
- pip install "azure-identity==1.16.1"
- pip install "langfuse==2.59.7"
- pip install "logfire==0.29.0"
- pip install numpydoc
- pip install traceloop-sdk==0.21.1
- pip install opentelemetry-api==1.25.0
- pip install opentelemetry-sdk==1.25.0
- pip install opentelemetry-exporter-otlp==1.25.0
- pip install openai==1.100.1
- pip install prisma==0.11.0
- pip install "detect_secrets==1.5.0"
- pip install "httpx==0.24.1"
- pip install "respx==0.22.0"
- pip install fastapi
- pip install "gunicorn==21.2.0"
- pip install "anyio==4.2.0"
- pip install "aiodynamo==23.10.1"
- pip install "asyncio==3.4.3"
- pip install "apscheduler==3.10.4"
- pip install "PyGithub==1.59.1"
- pip install argon2-cffi
- pip install "pytest-mock==3.12.0"
- pip install python-multipart
- pip install google-cloud-aiplatform
- pip install prometheus-client==0.20.0
- pip install "pydantic==2.10.2"
- pip install "diskcache==5.6.1"
- pip install "Pillow==10.3.0"
- pip install "jsonschema==4.22.0"
- pip install "websockets==13.1.0"
- pip install "pytest-xdist==3.6.1"
- - setup_litellm_enterprise_pip
- - save_cache:
- paths:
- - /home/circleci/.pyenv/versions
- - /home/circleci/.local
- key: v2-caching-deps-{{ checksum ".circleci/requirements.txt" }}
- - run:
- name: Run prisma ./docker/entrypoint.sh
- command: |
- set +e
- chmod +x docker/entrypoint.sh
- ./docker/entrypoint.sh
- set -e
-
- # Run pytest and generate JUnit XML report
- - run:
- name: Run tests
- command: |
- pwd
- ls
- mkdir -p test-results
-
- TEST_FILES=$(circleci tests glob "tests/local_testing/**/test_*.py")
-
- echo "$TEST_FILES" | circleci tests run \
- --split-by=timings \
- --verbose \
- --command="xargs python -m pytest \
- -v \
- --junitxml=test-results/junit.xml \
- --durations=5 \
- -k 'caching or cache'"
- no_output_timeout: 15m
-
- # Store test results
- - store_test_results:
- path: test-results
auth_ui_unit_tests:
docker:
- image: cimg/python:3.11
@@ -664,376 +551,6 @@ jobs:
# Store test results
- store_test_results:
path: test-results
- litellm_security_tests:
- docker:
- - image: cimg/python:3.13
- auth:
- username: ${DOCKERHUB_USERNAME}
- password: ${DOCKERHUB_PASSWORD}
- - image: cimg/postgres:14.0
- environment:
- POSTGRES_USER: postgres
- POSTGRES_PASSWORD: postgres
- POSTGRES_DB: circle_test
- resource_class: xlarge
- working_directory: ~/project
- environment:
- DATABASE_URL: "postgresql://postgres:postgres@localhost:5432/circle_test"
- steps:
- - checkout
- - setup_google_dns
- - run:
- name: Show git commit hash
- command: |
- echo "Git commit hash: $CIRCLE_SHA1"
- - setup_remote_docker:
- docker_layer_caching: true
- - restore_cache:
- keys:
- - v3-litellm-uv-deps-{{ checksum "requirements.txt" }}-{{ checksum ".circleci/config.yml" }}
- - 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" "pytest-retry==1.6.3" "pytest-mock==3.12.0" \
- "pytest-asyncio==0.21.1" "pytest-cov==5.0.0"
- - save_cache:
- paths:
- - ~/.local/lib
- - ~/.local/bin
- - ~/.cache/uv
- key: v3-litellm-uv-deps-{{ checksum "requirements.txt" }}-{{ checksum ".circleci/config.yml" }}
- - 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: Wait for PostgreSQL to be ready
- command: dockerize -wait tcp://localhost:5432 -timeout 1m
- - run:
- name: Run Security Scans
- command: |
- chmod +x ci_cd/security_scans.sh
- ./ci_cd/security_scans.sh
- - run:
- name: Run prisma ./docker/entrypoint.sh
- command: |
- set +e
- chmod +x docker/entrypoint.sh
- ./docker/entrypoint.sh
- set -e
- # Run pytest and generate JUnit XML report
- - run:
- name: Run tests
- command: |
- python -m pytest tests/proxy_security_tests -v -x --junitxml=test-results/junit.xml --durations=5
- no_output_timeout: 15m
- # Store test results
- - store_test_results:
- path: test-results
- # Split proxy unit tests into 3 jobs for faster execution and better debugging
- # test_key_generate_prisma runs separately without parallel execution to avoid event loop issues with logging worker
- litellm_proxy_unit_testing_key_generation:
- docker:
- - image: cimg/python:3.11
- auth:
- username: ${DOCKERHUB_USERNAME}
- password: ${DOCKERHUB_PASSWORD}
- working_directory: ~/project
- resource_class: medium
- steps:
- - checkout
- - setup_google_dns
- - run:
- name: Show git commit hash
- command: |
- echo "Git commit hash: $CIRCLE_SHA1"
- - run:
- name: Install PostgreSQL
- command: |
- sudo apt-get update
- sudo apt-get install -y postgresql-14 postgresql-contrib-14
- - restore_cache:
- keys:
- - v1-dependencies-{{ checksum ".circleci/requirements.txt" }}
- - run:
- name: Install Dependencies
- command: |
- python -m pip install --upgrade pip
- python -m pip install -r .circleci/requirements.txt
- pip install "pytest==7.3.1"
- pip install "pytest-retry==1.6.3"
- pip install "pytest-asyncio==0.21.1"
- pip install "pytest-cov==5.0.0"
- pip install "pytest-timeout==2.2.0"
- pip install "pytest-forked==1.6.0"
- pip install "mypy==1.18.2"
- pip install "google-generativeai==0.3.2"
- pip install "google-cloud-aiplatform==1.43.0"
- pip install "google-genai==1.22.0"
- pip install pyarrow
- pip install "boto3==1.36.0"
- pip install "aioboto3==13.4.0"
- pip install langchain
- pip install lunary==0.2.5
- pip install "azure-identity==1.16.1"
- pip install "langfuse==2.59.7"
- pip install "logfire==0.29.0"
- pip install numpydoc
- pip install traceloop-sdk==0.21.1
- pip install opentelemetry-api==1.25.0
- pip install opentelemetry-sdk==1.25.0
- pip install opentelemetry-exporter-otlp==1.25.0
- pip install openai==1.100.1
- pip install prisma==0.11.0
- pip install "detect_secrets==1.5.0"
- pip install "httpx==0.24.1"
- pip install "respx==0.22.0"
- pip install fastapi
- pip install "gunicorn==21.2.0"
- pip install "anyio==4.2.0"
- pip install "aiodynamo==23.10.1"
- pip install "asyncio==3.4.3"
- pip install "apscheduler==3.10.4"
- pip install "PyGithub==1.59.1"
- pip install argon2-cffi
- pip install "pytest-mock==3.12.0"
- pip install python-multipart
- pip install google-cloud-aiplatform
- pip install prometheus-client==0.20.0
- pip install "pydantic==2.10.2"
- pip install "diskcache==5.6.1"
- pip install "Pillow==10.3.0"
- pip install "jsonschema==4.22.0"
- pip install "pytest-postgresql==7.0.1"
- pip install "fakeredis==2.28.1"
- - setup_litellm_enterprise_pip
- - save_cache:
- paths:
- - ./venv
- key: v1-dependencies-{{ checksum ".circleci/requirements.txt" }}
- - run:
- name: Run prisma ./docker/entrypoint.sh
- command: |
- set +e
- chmod +x docker/entrypoint.sh
- ./docker/entrypoint.sh
- set -e
- - run:
- name: Run key generation tests (no parallel execution to avoid event loop issues)
- command: |
- pwd
- ls
- # Run without -n flag to avoid pytest-xdist event loop conflicts with logging worker
- python -m pytest tests/proxy_unit_tests/test_key_generate_prisma.py --cov=litellm --cov-report=xml --junitxml=test-results/junit-key-generation.xml --durations=10 --timeout=300 -vv --log-cli-level=INFO
- no_output_timeout: 15m
- - run:
- name: Rename the coverage files
- command: |
- mv coverage.xml litellm_proxy_unit_tests_key_generation_coverage.xml
- mv .coverage litellm_proxy_unit_tests_key_generation_coverage
- - store_test_results:
- path: test-results
- - persist_to_workspace:
- root: .
- paths:
- - litellm_proxy_unit_tests_key_generation_coverage.xml
- - litellm_proxy_unit_tests_key_generation_coverage
- litellm_proxy_unit_testing_part1:
- docker:
- - image: cimg/python:3.11
- auth:
- username: ${DOCKERHUB_USERNAME}
- password: ${DOCKERHUB_PASSWORD}
- working_directory: ~/project
- resource_class: xlarge
- steps:
- - checkout
- - setup_google_dns
- - run:
- name: Show git commit hash
- command: |
- echo "Git commit hash: $CIRCLE_SHA1"
- - run:
- name: Install PostgreSQL
- command: |
- sudo apt-get update
- sudo apt-get install -y postgresql-14 postgresql-contrib-14
- - restore_cache:
- keys:
- - v1-dependencies-{{ checksum ".circleci/requirements.txt" }}
- - run:
- name: Install Dependencies
- command: |
- python -m pip install --upgrade pip
- python -m pip install -r .circleci/requirements.txt
- pip install "pytest==7.3.1"
- pip install "pytest-retry==1.6.3"
- pip install "pytest-asyncio==0.21.1"
- pip install "pytest-cov==5.0.0"
- pip install "pytest-timeout==2.2.0"
- pip install "pytest-forked==1.6.0"
- pip install "mypy==1.18.2"
- pip install "google-generativeai==0.3.2"
- pip install "google-cloud-aiplatform==1.43.0"
- pip install "google-genai==1.22.0"
- pip install pyarrow
- pip install "boto3==1.36.0"
- pip install "aioboto3==13.4.0"
- pip install langchain
- pip install lunary==0.2.5
- pip install "azure-identity==1.16.1"
- pip install "langfuse==2.59.7"
- pip install "logfire==0.29.0"
- pip install numpydoc
- pip install traceloop-sdk==0.21.1
- pip install opentelemetry-api==1.25.0
- pip install opentelemetry-sdk==1.25.0
- pip install opentelemetry-exporter-otlp==1.25.0
- pip install openai==1.100.1
- pip install prisma==0.11.0
- pip install "detect_secrets==1.5.0"
- pip install "httpx==0.24.1"
- pip install "respx==0.22.0"
- pip install fastapi
- pip install "gunicorn==21.2.0"
- pip install "anyio==4.2.0"
- pip install "aiodynamo==23.10.1"
- pip install "asyncio==3.4.3"
- pip install "apscheduler==3.10.4"
- pip install "PyGithub==1.59.1"
- pip install argon2-cffi
- pip install "pytest-mock==3.12.0"
- pip install python-multipart
- pip install google-cloud-aiplatform
- pip install prometheus-client==0.20.0
- pip install "pydantic==2.10.2"
- pip install "diskcache==5.6.1"
- pip install "Pillow==10.3.0"
- pip install "jsonschema==4.22.0"
- pip install "pytest-postgresql==7.0.1"
- pip install "fakeredis==2.28.1"
- pip install "pytest-xdist==3.6.1"
- - setup_litellm_enterprise_pip
- - save_cache:
- paths:
- - ./venv
- key: v1-dependencies-{{ checksum ".circleci/requirements.txt" }}
- - run:
- name: Run prisma ./docker/entrypoint.sh
- command: |
- set +e
- chmod +x docker/entrypoint.sh
- ./docker/entrypoint.sh
- set -e
- - run:
- name: Run proxy unit tests (part 1 - auth checks)
- command: |
- pwd
- ls
- python -m pytest tests/proxy_unit_tests/test_auth_checks.py tests/proxy_unit_tests/test_user_api_key_auth.py --junitxml=test-results/junit-part1.xml --durations=10 -n 8 --timeout=300 -v
- no_output_timeout: 15m
- - store_test_results:
- path: test-results
- litellm_proxy_unit_testing_part2:
- docker:
- - image: cimg/python:3.11
- auth:
- username: ${DOCKERHUB_USERNAME}
- password: ${DOCKERHUB_PASSWORD}
- working_directory: ~/project
- resource_class: xlarge
- steps:
- - checkout
- - setup_google_dns
- - run:
- name: Show git commit hash
- command: |
- echo "Git commit hash: $CIRCLE_SHA1"
- - run:
- name: Install PostgreSQL
- command: |
- sudo apt-get update
- sudo apt-get install -y postgresql-14 postgresql-contrib-14
- - restore_cache:
- keys:
- - v1-dependencies-{{ checksum ".circleci/requirements.txt" }}
- - run:
- name: Install Dependencies
- command: |
- python -m pip install --upgrade pip
- python -m pip install -r .circleci/requirements.txt
- pip install "pytest==7.3.1"
- pip install "pytest-retry==1.6.3"
- pip install "pytest-asyncio==0.21.1"
- pip install "pytest-cov==5.0.0"
- pip install "pytest-timeout==2.2.0"
- pip install "pytest-forked==1.6.0"
- pip install "mypy==1.18.2"
- pip install "google-generativeai==0.3.2"
- pip install "google-cloud-aiplatform==1.43.0"
- pip install "google-genai==1.22.0"
- pip install pyarrow
- pip install "boto3==1.36.0"
- pip install "aioboto3==13.4.0"
- pip install langchain
- pip install lunary==0.2.5
- pip install "azure-identity==1.16.1"
- pip install "langfuse==2.59.7"
- pip install "logfire==0.29.0"
- pip install numpydoc
- pip install traceloop-sdk==0.21.1
- pip install opentelemetry-api==1.25.0
- pip install opentelemetry-sdk==1.25.0
- pip install opentelemetry-exporter-otlp==1.25.0
- pip install openai==1.100.1
- pip install prisma==0.11.0
- pip install "detect_secrets==1.5.0"
- pip install "httpx==0.24.1"
- pip install "respx==0.22.0"
- pip install fastapi
- pip install "gunicorn==21.2.0"
- pip install "anyio==4.2.0"
- pip install "aiodynamo==23.10.1"
- pip install "asyncio==3.4.3"
- pip install "apscheduler==3.10.4"
- pip install "PyGithub==1.59.1"
- pip install argon2-cffi
- pip install "pytest-mock==3.12.0"
- pip install python-multipart
- pip install google-cloud-aiplatform
- pip install prometheus-client==0.20.0
- pip install "pydantic==2.10.2"
- pip install "diskcache==5.6.1"
- pip install "Pillow==10.3.0"
- pip install "jsonschema==4.22.0"
- pip install "pytest-postgresql==7.0.1"
- pip install "fakeredis==2.28.1"
- pip install "pytest-xdist==3.6.1"
- - setup_litellm_enterprise_pip
- - save_cache:
- paths:
- - ./venv
- key: v1-dependencies-{{ checksum ".circleci/requirements.txt" }}
- - run:
- name: Run prisma ./docker/entrypoint.sh
- command: |
- set +e
- chmod +x docker/entrypoint.sh
- ./docker/entrypoint.sh
- set -e
- - run:
- name: Run proxy unit tests (part 2 - remaining tests)
- command: |
- pwd
- ls
- python -m pytest tests/proxy_unit_tests --ignore=tests/proxy_unit_tests/test_key_generate_prisma.py --ignore=tests/proxy_unit_tests/test_auth_checks.py --ignore=tests/proxy_unit_tests/test_user_api_key_auth.py --junitxml=test-results/junit-part2.xml --durations=10 -n 8 --timeout=300 -v
- no_output_timeout: 15m
- - store_test_results:
- path: test-results
litellm_assistants_api_testing: # Runs all tests with the "assistants" keyword
docker:
- image: cimg/python:3.13.1
@@ -1507,101 +1024,6 @@ jobs:
no_output_timeout: 15m
- store_test_results:
path: test-results
- litellm_mapped_tests_llms:
- docker:
- - image: cimg/python:3.11
- auth:
- username: ${DOCKERHUB_USERNAME}
- password: ${DOCKERHUB_PASSWORD}
- working_directory: ~/project
- resource_class: large
- steps:
- - setup_litellm_test_deps
- - run:
- name: Run LLM provider tests
- command: |
- python -m pytest tests/test_litellm/llms --junitxml=test-results/junit-llms.xml --durations=10 -n 4 --maxfail=5 --timeout=300 -vv --log-cli-level=WARNING
- no_output_timeout: 15m
- - store_test_results:
- path: test-results
- litellm_mapped_tests_core:
- docker:
- - image: cimg/python:3.11
- auth:
- username: ${DOCKERHUB_USERNAME}
- password: ${DOCKERHUB_PASSWORD}
- working_directory: ~/project
- resource_class: large
- steps:
- - setup_litellm_test_deps
- - run:
- name: Run core tests
- command: |
- python -m pytest tests/test_litellm --ignore=tests/test_litellm/proxy --ignore=tests/test_litellm/llms --ignore=tests/test_litellm/integrations --ignore=tests/test_litellm/litellm_core_utils --ignore=tests/test_litellm/experimental_mcp_client --junitxml=test-results/junit-core.xml --durations=10 -n 4 --maxfail=5 --timeout=300 -vv --log-cli-level=WARNING
- no_output_timeout: 15m
- - store_test_results:
- path: test-results
- litellm_mapped_tests_litellm_core_utils:
- docker:
- - image: cimg/python:3.11
- auth:
- username: ${DOCKERHUB_USERNAME}
- password: ${DOCKERHUB_PASSWORD}
- working_directory: ~/project
- resource_class: large
- steps:
- - setup_litellm_test_deps
- - run:
- name: Run litellm_core_utils tests
- command: |
- python -m pytest tests/test_litellm/litellm_core_utils --junitxml=test-results/junit-litellm-core-utils.xml --durations=10 -n 4 --maxfail=5 --timeout=300 -vv --log-cli-level=WARNING
- no_output_timeout: 15m
- - store_test_results:
- path: test-results
- litellm_mapped_tests_mcps:
- docker:
- - image: cimg/python:3.11
- auth:
- username: ${DOCKERHUB_USERNAME}
- password: ${DOCKERHUB_PASSWORD}
- working_directory: ~/project
- resource_class: medium
- steps:
- - setup_litellm_test_deps
- - run:
- name: Run MCP client tests
- command: |
- python -m pytest tests/test_litellm/experimental_mcp_client --cov=litellm --cov-report=xml --junitxml=test-results/junit-mcps.xml --durations=10 -n 2 --maxfail=5 --timeout=300 -vv --log-cli-level=WARNING
- no_output_timeout: 15m
- - run:
- name: Rename the coverage files
- command: |
- mv coverage.xml litellm_mcps_tests_coverage.xml
- mv .coverage litellm_mcps_tests_coverage
- - store_test_results:
- path: test-results
- - persist_to_workspace:
- root: .
- paths:
- - litellm_mcps_tests_coverage.xml
- - litellm_mcps_tests_coverage
- litellm_mapped_tests_integrations:
- docker:
- - image: cimg/python:3.11
- auth:
- username: ${DOCKERHUB_USERNAME}
- password: ${DOCKERHUB_PASSWORD}
- working_directory: ~/project
- resource_class: large
- steps:
- - setup_litellm_test_deps
- - run:
- name: Run integrations tests
- command: |
- python -m pytest tests/test_litellm/integrations --junitxml=test-results/junit-integrations.xml --durations=10 -n 4 --maxfail=5 --timeout=300 -vv --log-cli-level=WARNING
- no_output_timeout: 15m
- - store_test_results:
- path: test-results
litellm_mapped_enterprise_tests:
docker:
- image: cimg/python:3.11
@@ -4002,34 +3424,14 @@ workflows:
only:
- main
- /litellm_.*/
- - caching_unit_tests:
- filters:
- branches:
- only:
- main
- /litellm_.*/
- - litellm_proxy_unit_testing_key_generation:
- filters:
- branches:
- only:
- main
- /litellm_.*/
- - litellm_proxy_unit_testing_part1:
- filters:
- branches:
- only:
- main
- /litellm_.*/
- - litellm_proxy_unit_testing_part2:
- filters:
- branches:
- only:
- main
- /litellm_.*/
- - litellm_security_tests:
- filters:
- branches:
- only:
- main
- /litellm_.*/
- litellm_assistants_api_testing:
@@ -4265,34 +3667,14 @@ workflows:
only:
- main
- /litellm_.*/
- - litellm_mapped_tests_llms:
- filters:
- branches:
- only:
- main
- /litellm_.*/
- - litellm_mapped_tests_core:
- filters:
- branches:
- only:
- main
- /litellm_.*/
- - litellm_mapped_tests_mcps:
- filters:
- branches:
- only:
- main
- /litellm_.*/
- - litellm_mapped_tests_integrations:
- filters:
- branches:
- only:
- main
- /litellm_.*/
- - litellm_mapped_tests_litellm_core_utils:
- filters:
- branches:
- only:
- main
- /litellm_.*/
- batches_testing:
@@ -4342,11 +3724,6 @@ workflows:
- search_testing
- litellm_mapped_tests_proxy_part1
- litellm_mapped_tests_proxy_part2
- - litellm_mapped_tests_llms
- - litellm_mapped_tests_core
- - litellm_mapped_tests_mcps
- - litellm_mapped_tests_integrations
- - litellm_mapped_tests_litellm_core_utils
- litellm_mapped_enterprise_tests
- batches_testing
- litellm_utils_testing
@@ -4354,8 +3731,6 @@ workflows:
- image_gen_testing
- logging_testing
- audio_testing
- - caching_unit_tests
- - litellm_proxy_unit_testing_key_generation
- langfuse_logging_unit_tests
- local_testing_part1
- local_testing_part2
diff --git a/.github/actions/helm-oci-chart-releaser/action.yml b/.github/actions/helm-oci-chart-releaser/action.yml
index 1823e262832..454c591d436 100644
--- a/.github/actions/helm-oci-chart-releaser/action.yml
+++ b/.github/actions/helm-oci-chart-releaser/action.yml
@@ -41,32 +41,54 @@ runs:
using: composite
steps:
- name: Helm | Setup
- uses: azure/setup-helm@v4
+ uses: azure/setup-helm@1a275c3b69536ee54be43f2070a358922e12c8d4 # v4.3.1
with:
version: v3.20.0
- name: Helm | Login
shell: bash
- run: echo ${{ inputs.registry_password }} | helm registry login -u ${{ inputs.registry_username }} --password-stdin ${{ inputs.registry }}
+ env:
+ REGISTRY_PASSWORD: ${{ inputs.registry_password }}
+ REGISTRY_USERNAME: ${{ inputs.registry_username }}
+ REGISTRY: ${{ inputs.registry }}
+ run: echo "$REGISTRY_PASSWORD" | helm registry login -u "$REGISTRY_USERNAME" --password-stdin "$REGISTRY"
- name: Helm | Dependency
if: inputs.update_dependencies == 'true'
shell: bash
- run: helm dependency update ${{ inputs.path == null && format('{0}/{1}', 'charts', inputs.name) || inputs.path }}
+ env:
+ CHART_PATH: ${{ inputs.path == null && format('{0}/{1}', 'charts', inputs.name) || inputs.path }}
+ run: helm dependency update "$CHART_PATH"
- name: Helm | Package
shell: bash
- run: helm package ${{ inputs.path == null && format('{0}/{1}', 'charts', inputs.name) || inputs.path }} --version ${{ inputs.tag }} --app-version ${{ inputs.app_version }}
+ env:
+ CHART_PATH: ${{ inputs.path == null && format('{0}/{1}', 'charts', inputs.name) || inputs.path }}
+ TAG: ${{ inputs.tag }}
+ APP_VERSION: ${{ inputs.app_version }}
+ run: helm package "$CHART_PATH" --version "$TAG" --app-version "$APP_VERSION"
- name: Helm | Push
shell: bash
- run: helm push ${{ inputs.name }}-${{ inputs.tag }}.tgz oci://${{ inputs.registry }}/${{ inputs.repository }}
+ env:
+ NAME: ${{ inputs.name }}
+ TAG: ${{ inputs.tag }}
+ REGISTRY: ${{ inputs.registry }}
+ REPOSITORY: ${{ inputs.repository }}
+ run: helm push "${NAME}-${TAG}.tgz" "oci://${REGISTRY}/${REPOSITORY}"
- name: Helm | Logout
shell: bash
- run: helm registry logout ${{ inputs.registry }}
+ env:
+ REGISTRY: ${{ inputs.registry }}
+ run: helm registry logout "$REGISTRY"
- name: Helm | Output
id: output
shell: bash
- run: echo "image=${{ inputs.registry }}/${{ inputs.repository }}/${{ inputs.name }}:${{ inputs.tag }}" >> $GITHUB_OUTPUT
+ env:
+ REGISTRY: ${{ inputs.registry }}
+ REPOSITORY: ${{ inputs.repository }}
+ NAME: ${{ inputs.name }}
+ TAG: ${{ inputs.tag }}
+ run: echo "image=${REGISTRY}/${REPOSITORY}/${NAME}:${TAG}" >> $GITHUB_OUTPUT
diff --git a/.github/codeql/codeql-config.yml b/.github/codeql/codeql-config.yml
index 20807685e12..36d70c1d746 100644
--- a/.github/codeql/codeql-config.yml
+++ b/.github/codeql/codeql-config.yml
@@ -1,22 +1,21 @@
name: "LiteLLM CodeQL config"
-# Use security-extended suite instead of security-and-quality to avoid
-# result sets > 2 GiB on this codebase that cause fatal OOM failures.
queries:
- - uses: security-extended
+ - uses: security-and-quality
-# These two queries are security queries included in security-extended that
-# individually produce result sets > 2 GiB on this codebase, causing fatal
-# OOM failures. Exclude them as a safety net until CI confirms they no longer
-# OOM; drop these exclusions in a follow-up once verified.
+# Known OOM queries on large Python codebases:
+# CodeQL builds a full data flow graph in memory. These two queries trace
+# sensitive data through every log call / regex pattern, causing combinatorial
+# path explosion on codebases with extensive logging like LiteLLM (>2 GiB
+# result sets). This is a known CodeQL scaling limitation, not a code issue.
+# Re-test periodically as CodeQL improves or the codebase refactors logging.
query-filters:
- exclude:
- id: py/clear-text-logging-sensitive-data # CWE-312 — > 2 GiB result set
+ id: py/clear-text-logging-sensitive-data # CWE-312
- exclude:
- id: py/polynomial-redos # CWE-730 — > 2 GiB result set
+ id: py/polynomial-redos # CWE-730
paths-ignore:
- tests
- docs
- "**/*.md"
- - litellm/proxy/_experimental/out
diff --git a/.github/dependabot.yaml b/.github/dependabot.yaml
index 58e7cfe10da..c49882a8d62 100644
--- a/.github/dependabot.yaml
+++ b/.github/dependabot.yaml
@@ -4,6 +4,9 @@ updates:
directory: "/"
schedule:
interval: "daily"
+ cooldown:
+ default-days: 7
+ semver-major-days: 14
groups:
github-actions:
patterns:
diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml
new file mode 100644
index 00000000000..f1ae30e67d7
--- /dev/null
+++ b/.github/workflows/_test-unit-base.yml
@@ -0,0 +1,96 @@
+name: _Unit Test Base (Reusable)
+
+on:
+ workflow_call:
+ inputs:
+ test-path:
+ description: "Pytest path(s) to run"
+ required: true
+ type: string
+ workers:
+ description: "Number of pytest-xdist workers"
+ required: false
+ type: number
+ default: 2
+ reruns:
+ description: "Number of reruns for flaky tests"
+ required: false
+ type: number
+ default: 2
+ timeout-minutes:
+ description: "Job timeout in minutes"
+ required: false
+ type: number
+ default: 20
+ max-failures:
+ description: "Stop after this many failures"
+ required: false
+ type: number
+ default: 10
+
+permissions:
+ contents: read
+
+jobs:
+ run:
+ name: Run tests
+ runs-on: ubuntu-latest
+ timeout-minutes: ${{ inputs.timeout-minutes }}
+
+ 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.3.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"
+ 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
+ env:
+ TEST_PATH: ${{ inputs.test-path }}
+ MAX_FAILURES: ${{ inputs.max-failures }}
+ WORKERS: ${{ inputs.workers }}
+ RERUNS: ${{ inputs.reruns }}
+ run: |
+ poetry run pytest ${TEST_PATH:?} \
+ --tb=short -vv \
+ --maxfail="${MAX_FAILURES}" \
+ -n "${WORKERS}" \
+ --reruns "${RERUNS}" \
+ --reruns-delay 1 \
+ --dist=loadscope \
+ --durations=20
diff --git a/.github/workflows/_test-unit-services-base.yml b/.github/workflows/_test-unit-services-base.yml
new file mode 100644
index 00000000000..d53a9e8822a
--- /dev/null
+++ b/.github/workflows/_test-unit-services-base.yml
@@ -0,0 +1,164 @@
+name: _Unit Test Services Base (Reusable)
+
+on:
+ workflow_call:
+ inputs:
+ test-path:
+ description: "Pytest path(s) to run"
+ required: true
+ type: string
+ workers:
+ description: "Number of pytest-xdist workers (0 = no parallelism)"
+ required: false
+ type: number
+ default: 2
+ reruns:
+ description: "Number of reruns for flaky tests"
+ required: false
+ type: number
+ default: 2
+ timeout-minutes:
+ description: "Job timeout in minutes"
+ required: false
+ type: number
+ default: 20
+ max-failures:
+ description: "Stop after this many failures"
+ 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
+ secrets:
+ REDIS_HOST:
+ required: false
+ REDIS_PORT:
+ required: false
+ REDIS_PASSWORD:
+ required: false
+ DATABASE_URL:
+ required: false
+ POSTGRES_USER:
+ required: false
+ POSTGRES_PASSWORD:
+ required: false
+
+permissions:
+ contents: read
+
+jobs:
+ run:
+ name: Run tests
+ runs-on: ubuntu-latest
+ timeout-minutes: ${{ inputs.timeout-minutes }}
+ # Environment is derived from the enable-* flags, not caller-controllable.
+ # This prevents callers from passing arbitrary environment names to bypass secret scoping.
+ # 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' ||
+ ''
+ }}
+
+ services:
+ postgres:
+ image: postgres@sha256:705a5d5b5836f3fcba0d02c4d281e6a7dd9ed2dd4078640f08a1e1e9896e097d # postgres:14
+ env:
+ POSTGRES_USER: ${{ secrets.POSTGRES_USER }}
+ POSTGRES_PASSWORD: ${{ secrets.POSTGRES_PASSWORD }}
+ POSTGRES_DB: litellm_test
+ 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.3.0
+ with:
+ path: |
+ ~/.cache/pypoetry
+ ~/.cache/pip
+ .venv
+ key: ${{ runner.os }}-poetry-services-${{ hashFiles('poetry.lock') }}
+ restore-keys: |
+ ${{ runner.os }}-poetry-services-
+
+ - name: Install dependencies
+ run: |
+ poetry config virtualenvs.in-project true
+ poetry install --with dev,proxy-dev --extras "proxy semantic-router"
+ 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 Prisma migrations
+ if: ${{ inputs.enable-postgres }}
+ env:
+ DATABASE_URL: ${{ secrets.DATABASE_URL }}
+ run: |
+ poetry run prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
+
+ - name: Run tests
+ env:
+ TEST_PATH: ${{ inputs.test-path }}
+ MAX_FAILURES: ${{ inputs.max-failures }}
+ 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:?} \
+ --tb=short -vv \
+ --maxfail="${MAX_FAILURES}" \
+ --reruns "${RERUNS}" \
+ --reruns-delay 1 \
+ --durations=20
+ else
+ poetry run pytest ${TEST_PATH:?} \
+ --tb=short -vv \
+ --maxfail="${MAX_FAILURES}" \
+ -n "${WORKERS}" \
+ --reruns "${RERUNS}" \
+ --reruns-delay 1 \
+ --dist=loadscope \
+ --durations=20
+ fi
diff --git a/.github/workflows/auto_update_price_and_context_window.yml b/.github/workflows/auto_update_price_and_context_window.yml
index 8265e0d09c5..60e89936219 100644
--- a/.github/workflows/auto_update_price_and_context_window.yml
+++ b/.github/workflows/auto_update_price_and_context_window.yml
@@ -5,12 +5,18 @@ on:
- cron: "0 0 * * 0" # Run every Sundays at midnight
#- cron: "0 0 * * *" # Run daily at midnight
+permissions:
+ contents: write
+ pull-requests: write
+
jobs:
auto_update_price_and_context_window:
if: github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Install Dependencies
run: |
pip install 'aiohttp==3.13.3'
diff --git a/.github/workflows/check-schema-sync.yml b/.github/workflows/check-schema-sync.yml
new file mode 100644
index 00000000000..0e5e2804e60
--- /dev/null
+++ b/.github/workflows/check-schema-sync.yml
@@ -0,0 +1,58 @@
+name: Check Schema Sync
+
+on:
+ pull_request:
+ paths:
+ - 'schema.prisma'
+ - 'litellm/proxy/schema.prisma'
+ - 'litellm-proxy-extras/litellm_proxy_extras/schema.prisma'
+
+permissions:
+ contents: read
+
+jobs:
+ check-sync:
+ name: Verify schema.prisma copies match root
+ runs-on: ubuntu-latest
+ timeout-minutes: 5
+ steps:
+ - name: Checkout PR
+ uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
+
+ - name: Reject symlinked schema files
+ run: |
+ for f in schema.prisma litellm/proxy/schema.prisma litellm-proxy-extras/litellm_proxy_extras/schema.prisma; do
+ if [ -L "$f" ]; then
+ echo "::error file=$f::$f is a symlink, which is not allowed"
+ exit 1
+ fi
+ done
+
+ - name: Check all schemas match root
+ run: |
+ EXIT=0
+
+ diff schema.prisma litellm/proxy/schema.prisma || {
+ echo "::error file=litellm/proxy/schema.prisma::litellm/proxy/schema.prisma differs from root schema.prisma"
+ EXIT=1
+ }
+
+ diff schema.prisma litellm-proxy-extras/litellm_proxy_extras/schema.prisma || {
+ echo "::error file=litellm-proxy-extras/litellm_proxy_extras/schema.prisma::litellm-proxy-extras/litellm_proxy_extras/schema.prisma differs from root schema.prisma"
+ EXIT=1
+ }
+
+ if [ "$EXIT" -ne 0 ]; then
+ echo ""
+ echo "Schema files are out of sync."
+ echo "The root schema.prisma is the source of truth."
+ echo ""
+ echo "To fix, run from the repo root:"
+ echo " cp schema.prisma litellm/proxy/schema.prisma"
+ echo " cp schema.prisma litellm-proxy-extras/litellm_proxy_extras/schema.prisma"
+ exit 1
+ fi
+
+ echo "All schema copies are in sync with root."
diff --git a/.github/workflows/check_duplicate_issues.yml b/.github/workflows/check_duplicate_issues.yml
index 539290bfeaf..289d78880ad 100644
--- a/.github/workflows/check_duplicate_issues.yml
+++ b/.github/workflows/check_duplicate_issues.yml
@@ -33,6 +33,7 @@ jobs:
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
sparse-checkout: .github/scripts
+ persist-credentials: false
- name: Set up Python
if: github.event.action == 'opened'
diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml
index 0e49b3af138..e86fca17c7a 100644
--- a/.github/workflows/codeql.yml
+++ b/.github/workflows/codeql.yml
@@ -6,8 +6,8 @@ on:
pull_request:
branches: [main]
schedule:
- # Run weekly on Sundays at 04:00 UTC
- - cron: "0 4 * * 0"
+ # Run daily at 04:00 UTC
+ - cron: "0 4 * * *"
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
@@ -39,6 +39,8 @@ jobs:
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Initialize CodeQL
uses: github/codeql-action/init@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3
diff --git a/.github/workflows/codspeed.yml b/.github/workflows/codspeed.yml
index 749242b1cd6..52d64addea9 100644
--- a/.github/workflows/codspeed.yml
+++ b/.github/workflows/codspeed.yml
@@ -26,6 +26,8 @@ jobs:
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
diff --git a/.github/workflows/create_daily_staging_branch.yml b/.github/workflows/create_daily_staging_branch.yml
index 0df5f4f92ea..424d8de0a41 100644
--- a/.github/workflows/create_daily_staging_branch.yml
+++ b/.github/workflows/create_daily_staging_branch.yml
@@ -9,12 +9,15 @@ jobs:
create-staging-branch:
if: github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
+ permissions:
+ contents: write
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 0
+ persist-credentials: false
- name: Create daily staging branch
env:
@@ -46,12 +49,15 @@ jobs:
create-internal-dev-branch:
if: github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
+ permissions:
+ contents: write
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 0
+ persist-credentials: false
- name: Create internal dev branch
env:
diff --git a/.github/workflows/helm_unit_test.yml b/.github/workflows/helm_unit_test.yml
index 416523f241b..06836b1d1cd 100644
--- a/.github/workflows/helm_unit_test.yml
+++ b/.github/workflows/helm_unit_test.yml
@@ -6,12 +6,17 @@ on:
branches:
- main
+permissions:
+ contents: read
+
jobs:
unit-test:
runs-on: ubuntu-latest
steps:
- name: Checkout
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Set up Helm 3.11.1
uses: azure/setup-helm@1a275c3b69536ee54be43f2070a358922e12c8d4 # v4.3.1
diff --git a/.github/workflows/issue-keyword-labeler.yml b/.github/workflows/issue-keyword-labeler.yml
index 59b8fd9cf9f..7e2693209b6 100644
--- a/.github/workflows/issue-keyword-labeler.yml
+++ b/.github/workflows/issue-keyword-labeler.yml
@@ -14,6 +14,8 @@ jobs:
steps:
- name: Checkout code
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Scan for provider keywords
id: scan
diff --git a/.github/workflows/llm-translation-testing.yml b/.github/workflows/llm-translation-testing.yml
index a6b643dd92e..922013c4b54 100644
--- a/.github/workflows/llm-translation-testing.yml
+++ b/.github/workflows/llm-translation-testing.yml
@@ -11,6 +11,9 @@ on:
tags:
- "v*-rc*" # Triggers on release candidate tags like v1.0.0-rc1
+permissions:
+ contents: read
+
jobs:
run-llm-translation-tests:
runs-on: ubuntu-latest
@@ -20,6 +23,7 @@ jobs:
- name: Checkout code
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
+ persist-credentials: false
ref: ${{ github.event.inputs.release_candidate_tag || github.ref }}
- name: Set up Python
@@ -33,8 +37,8 @@ jobs:
poetry config virtualenvs.create true
poetry config virtualenvs.in-project true
- - name: Cache Poetry dependencies
- uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.0.0
+ - name: Restore Poetry dependencies cache
+ uses: actions/cache/restore@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.0.0
with:
path: |
~/.cache/pypoetry
@@ -60,11 +64,12 @@ jobs:
AZURE_API_KEY: ${{ secrets.AZURE_API_KEY }}
AZURE_API_BASE: ${{ secrets.AZURE_API_BASE }}
AZURE_API_VERSION: ${{ secrets.AZURE_API_VERSION }}
- # Add other API keys as needed
+ RC_TAG: ${{ github.event.inputs.release_candidate_tag || github.ref_name }}
+ COMMIT_SHA: ${{ github.sha }}
run: |
python .github/workflows/run_llm_translation_tests.py \
- --tag "${{ github.event.inputs.release_candidate_tag || github.ref_name }}" \
- --commit "${{ github.sha }}" \
+ --tag "$RC_TAG" \
+ --commit "$COMMIT_SHA" \
|| true # Continue even if tests fail
- name: Display test summary
diff --git a/.github/workflows/read_pyproject_version.yml b/.github/workflows/read_pyproject_version.yml
index a9d16ca4413..04b4a38ce19 100644
--- a/.github/workflows/read_pyproject_version.yml
+++ b/.github/workflows/read_pyproject_version.yml
@@ -5,6 +5,9 @@ on:
branches:
- main # Change this to the default branch of your repository
+permissions:
+ contents: read
+
jobs:
read-version:
runs-on: ubuntu-latest
@@ -12,6 +15,8 @@ jobs:
steps:
- name: Checkout code
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Read version from pyproject.toml
id: read-version
diff --git a/.github/workflows/run_observatory_tests.yml b/.github/workflows/run_observatory_tests.yml
index b0706b7b716..a25b96766d7 100644
--- a/.github/workflows/run_observatory_tests.yml
+++ b/.github/workflows/run_observatory_tests.yml
@@ -34,6 +34,8 @@ jobs:
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Validate tag input
env:
@@ -49,11 +51,12 @@ jobs:
TAG: ${{ inputs.tag }}
AZURE_API_KEY: ${{ secrets.AZURE_API_KEY }}
AZURE_API_BASE: ${{ secrets.AZURE_API_BASE }}
+ WORKSPACE: ${{ github.workspace }}
run: |
docker run -d \
--name litellm-rc \
-p 4000:4000 \
- -v "${{ github.workspace }}/.github/observatory/litellm_config.yaml:/app/config.yaml" \
+ -v "${WORKSPACE}/.github/observatory/litellm_config.yaml:/app/config.yaml" \
-e LITELLM_MASTER_KEY="${LITELLM_MASTER_KEY}" \
-e AZURE_API_KEY="${AZURE_API_KEY}" \
-e AZURE_API_BASE="${AZURE_API_BASE}" \
@@ -104,11 +107,11 @@ jobs:
- name: Verify tunnel connectivity
run: |
- echo "Testing tunnel at ${{ env.TUNNEL_URL }}..."
+ echo "Testing tunnel at ${TUNNEL_URL}..."
# Quick tunnels need time for DNS propagation; retry to avoid
# transient NXDOMAIN (curl exit code 6) on first attempt.
for i in $(seq 1 10); do
- if curl -sf "${{ env.TUNNEL_URL }}/health/liveliness" > /dev/null 2>&1; then
+ if curl -sf "${TUNNEL_URL}/health/liveliness" > /dev/null 2>&1; then
echo "Tunnel is working (attempt $i)"
exit 0
fi
@@ -222,5 +225,5 @@ jobs:
- name: Cleanup
if: always()
run: |
- kill "${{ env.CLOUDFLARED_PID }}" 2>/dev/null || true
+ kill "$CLOUDFLARED_PID" 2>/dev/null || true
docker rm -f litellm-rc 2>/dev/null || true
diff --git a/.github/workflows/scan_duplicate_issues.yml b/.github/workflows/scan_duplicate_issues.yml
index 6c88a54554e..222ff11f304 100644
--- a/.github/workflows/scan_duplicate_issues.yml
+++ b/.github/workflows/scan_duplicate_issues.yml
@@ -24,6 +24,7 @@ jobs:
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
sparse-checkout: .github/scripts
+ persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
diff --git a/.github/workflows/stale.yml b/.github/workflows/stale.yml
index e7ce48b9fbd..c905bb12312 100644
--- a/.github/workflows/stale.yml
+++ b/.github/workflows/stale.yml
@@ -5,6 +5,10 @@ on:
- cron: "0 0 * * *" # Runs daily at midnight UTC
workflow_dispatch:
+permissions:
+ issues: write
+ pull-requests: write
+
jobs:
stale:
if: github.repository == 'BerriAI/litellm'
diff --git a/.github/workflows/sync-schema.yml b/.github/workflows/sync-schema.yml
new file mode 100644
index 00000000000..72a5c56293e
--- /dev/null
+++ b/.github/workflows/sync-schema.yml
@@ -0,0 +1,73 @@
+name: Sync schema.prisma copies
+
+on:
+ pull_request:
+ paths:
+ - 'schema.prisma'
+
+# Scoped to ONLY the permissions needed:
+# - contents:write to push the sync commit to the PR branch
+# - pull-requests:read is implicit (needed to check out the PR)
+permissions:
+ contents: write
+
+jobs:
+ sync:
+ name: Copy root schema to proxy and proxy-extras
+ runs-on: ubuntu-latest
+ timeout-minutes: 5
+ # Only run on PRs from branches in THIS repo (not forks).
+ # Fork PRs cannot push back to the head branch with GITHUB_TOKEN,
+ # and pull_request events from forks have read-only tokens anyway.
+ # Also reject PRs from branches named after protected branches to
+ # prevent pushing directly to main/master.
+ if: >-
+ github.event.pull_request.head.repo.full_name == github.repository
+ && github.head_ref != 'main'
+ && github.head_ref != 'master'
+ steps:
+ - name: Checkout PR branch by SHA
+ uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ # Use the merge commit SHA for safety — github.head_ref is an
+ # attacker-controlled string (the branch name) and could contain
+ # unusual characters that cause unexpected git behavior.
+ ref: ${{ github.event.pull_request.head.sha }}
+ persist-credentials: true # needed for git push
+
+ - name: Reject symlinked schema files
+ run: |
+ for f in schema.prisma litellm/proxy/schema.prisma litellm-proxy-extras/litellm_proxy_extras/schema.prisma; do
+ if [ -L "$f" ]; then
+ echo "::error file=$f::$f is a symlink, which is not allowed"
+ exit 1
+ fi
+ done
+
+ - name: Copy root schema to other locations
+ run: |
+ cp schema.prisma litellm/proxy/schema.prisma
+ cp schema.prisma litellm-proxy-extras/litellm_proxy_extras/schema.prisma
+
+ - name: Check for changes
+ id: diff
+ run: |
+ if git diff --quiet -- litellm/proxy/schema.prisma litellm-proxy-extras/litellm_proxy_extras/schema.prisma; then
+ echo "changed=false" >> "$GITHUB_OUTPUT"
+ echo "Schemas already in sync. Nothing to do."
+ else
+ echo "changed=true" >> "$GITHUB_OUTPUT"
+ echo "Schema copies need updating."
+ fi
+
+ - name: Commit synced schemas
+ if: steps.diff.outputs.changed == 'true'
+ run: |
+ # Push to the PR's head branch (need the branch name for git push).
+ # We checked out by SHA above for safety, so configure the push target explicitly.
+ git config user.name "github-actions[bot]"
+ git config user.email "41898282+github-actions[bot]@users.noreply.github.com"
+ git checkout -B "$GITHUB_HEAD_REF"
+ git add -- litellm/proxy/schema.prisma litellm-proxy-extras/litellm_proxy_extras/schema.prisma
+ git commit -m "chore: sync schema.prisma copies from root"
+ git push origin "HEAD:$GITHUB_HEAD_REF"
diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml
index 1424088eaa2..5bb85716a17 100644
--- a/.github/workflows/test-linting.yml
+++ b/.github/workflows/test-linting.yml
@@ -4,6 +4,9 @@ on:
pull_request:
branches: [main]
+permissions:
+ contents: read
+
jobs:
lint:
runs-on: ubuntu-latest
@@ -14,6 +17,7 @@ jobs:
with:
fetch-depth: 0
clean: true
+ persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
@@ -87,6 +91,7 @@ jobs:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 0
+ persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
diff --git a/.github/workflows/test-litellm-matrix.yml b/.github/workflows/test-litellm-matrix.yml
index 37c55fc4bf0..860d25636c5 100644
--- a/.github/workflows/test-litellm-matrix.yml
+++ b/.github/workflows/test-litellm-matrix.yml
@@ -4,6 +4,9 @@ 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 }}
@@ -118,6 +121,8 @@ jobs:
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
@@ -151,7 +156,10 @@ jobs:
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 }}
diff --git a/.github/workflows/test-litellm-ui-build.yml b/.github/workflows/test-litellm-ui-build.yml
index 4ac268e4e7a..6b0b3a413a6 100644
--- a/.github/workflows/test-litellm-ui-build.yml
+++ b/.github/workflows/test-litellm-ui-build.yml
@@ -17,6 +17,8 @@ jobs:
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Setup Node.js
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0
diff --git a/.github/workflows/test-litellm.yml b/.github/workflows/test-litellm.yml
index 57051e9fd6c..0c040b3ebe7 100644
--- a/.github/workflows/test-litellm.yml
+++ b/.github/workflows/test-litellm.yml
@@ -8,6 +8,9 @@ on:
# pull_request:
# branches: [ main ]
+permissions:
+ contents: read
+
jobs:
test:
runs-on: ubuntu-latest
@@ -15,6 +18,8 @@ jobs:
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Thank You Message
run: |
diff --git a/.github/workflows/test-mcp.yml b/.github/workflows/test-mcp.yml
index 3b8f2920aba..1b228ab76bb 100644
--- a/.github/workflows/test-mcp.yml
+++ b/.github/workflows/test-mcp.yml
@@ -4,6 +4,9 @@ on:
pull_request:
branches: [main]
+permissions:
+ contents: read
+
jobs:
test:
runs-on: ubuntu-latest
@@ -11,6 +14,8 @@ jobs:
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Thank You Message
run: |
diff --git a/.github/workflows/test-model-map.yaml b/.github/workflows/test-model-map.yaml
index 2874a001e5f..429f9e1ce0a 100644
--- a/.github/workflows/test-model-map.yaml
+++ b/.github/workflows/test-model-map.yaml
@@ -4,11 +4,16 @@ on:
pull_request:
branches: [main]
+permissions:
+ contents: read
+
jobs:
validate-model-prices-json:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Validate model_prices_and_context_window.json
run: |
diff --git a/.github/workflows/test-proxy-e2e-azure-batches.yml b/.github/workflows/test-proxy-e2e-azure-batches.yml
index b579b37de9c..7cbbe0b338f 100644
--- a/.github/workflows/test-proxy-e2e-azure-batches.yml
+++ b/.github/workflows/test-proxy-e2e-azure-batches.yml
@@ -9,6 +9,9 @@ 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
@@ -31,6 +34,8 @@ jobs:
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
@@ -63,7 +68,10 @@ jobs:
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
diff --git a/.github/workflows/test-unit-caching-redis.yml b/.github/workflows/test-unit-caching-redis.yml
new file mode 100644
index 00000000000..ca274324f2f
--- /dev/null
+++ b/.github/workflows/test-unit-caching-redis.yml
@@ -0,0 +1,38 @@
+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 }}
diff --git a/.github/workflows/test-unit-core-utils.yml b/.github/workflows/test-unit-core-utils.yml
new file mode 100644
index 00000000000..2f3698fdf60
--- /dev/null
+++ b/.github/workflows/test-unit-core-utils.yml
@@ -0,0 +1,20 @@
+name: "Unit Tests: Core Utilities"
+
+on:
+ pull_request:
+ branches: [main]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ core-utils:
+ uses: ./.github/workflows/_test-unit-base.yml
+ with:
+ test-path: "tests/test_litellm/litellm_core_utils"
+ workers: 2
+ reruns: 1
diff --git a/.github/workflows/test-unit-documentation.yml b/.github/workflows/test-unit-documentation.yml
new file mode 100644
index 00000000000..d8b30de6844
--- /dev/null
+++ b/.github/workflows/test-unit-documentation.yml
@@ -0,0 +1,67 @@
+name: "Unit Tests: Documentation Validation"
+
+on:
+ pull_request:
+ branches: [main]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ documentation:
+ runs-on: ubuntu-latest
+ timeout-minutes: 10
+
+ 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.3.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"
+ 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
+
+ # Run the same documentation tests that CircleCI ran (as direct Python scripts)
+ - name: Run documentation validation tests
+ run: |
+ poetry run python ./tests/documentation_tests/test_env_keys.py
+ poetry run python ./tests/documentation_tests/test_router_settings.py
+ poetry run python ./tests/documentation_tests/test_api_docs.py
+ poetry run python ./tests/documentation_tests/test_circular_imports.py
diff --git a/.github/workflows/test-unit-enterprise-routing.yml b/.github/workflows/test-unit-enterprise-routing.yml
new file mode 100644
index 00000000000..13ae3efedba
--- /dev/null
+++ b/.github/workflows/test-unit-enterprise-routing.yml
@@ -0,0 +1,24 @@
+name: "Unit Tests: Enterprise, Google GenAI & Routing"
+
+on:
+ pull_request:
+ branches: [main]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ enterprise-routing:
+ uses: ./.github/workflows/_test-unit-base.yml
+ with:
+ test-path: >-
+ tests/test_litellm/enterprise
+ tests/test_litellm/google_genai
+ tests/test_litellm/router_utils
+ tests/test_litellm/router_strategy
+ workers: 2
+ reruns: 2
diff --git a/.github/workflows/test-unit-integrations.yml b/.github/workflows/test-unit-integrations.yml
new file mode 100644
index 00000000000..2789f99d81c
--- /dev/null
+++ b/.github/workflows/test-unit-integrations.yml
@@ -0,0 +1,20 @@
+name: "Unit Tests: Integrations (Callbacks & Logging)"
+
+on:
+ pull_request:
+ branches: [main]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ integrations:
+ uses: ./.github/workflows/_test-unit-base.yml
+ with:
+ test-path: "tests/test_litellm/integrations"
+ workers: 2
+ reruns: 3
diff --git a/.github/workflows/test-unit-llm-providers.yml b/.github/workflows/test-unit-llm-providers.yml
new file mode 100644
index 00000000000..6c00272b0c8
--- /dev/null
+++ b/.github/workflows/test-unit-llm-providers.yml
@@ -0,0 +1,29 @@
+name: "Unit Tests: LLM Provider Transformations"
+
+on:
+ pull_request:
+ branches: [main]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ vertex-ai:
+ name: Vertex AI
+ uses: ./.github/workflows/_test-unit-base.yml
+ with:
+ test-path: "tests/test_litellm/llms/vertex_ai"
+ workers: 1
+ reruns: 2
+
+ other-providers:
+ name: All Other Providers
+ 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
diff --git a/.github/workflows/test-unit-misc.yml b/.github/workflows/test-unit-misc.yml
new file mode 100644
index 00000000000..9228decd7cc
--- /dev/null
+++ b/.github/workflows/test-unit-misc.yml
@@ -0,0 +1,31 @@
+name: "Unit Tests: MCP, Secrets, Containers & Misc"
+
+on:
+ pull_request:
+ branches: [main]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ misc:
+ uses: ./.github/workflows/_test-unit-base.yml
+ with:
+ test-path: >-
+ 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
+ tests/test_litellm/test_*.py
+ workers: 2
+ reruns: 2
diff --git a/.github/workflows/test-unit-proxy-auth.yml b/.github/workflows/test-unit-proxy-auth.yml
new file mode 100644
index 00000000000..e71821db701
--- /dev/null
+++ b/.github/workflows/test-unit-proxy-auth.yml
@@ -0,0 +1,20 @@
+name: "Unit Tests: Proxy Auth & Key Management"
+
+on:
+ pull_request:
+ branches: [main]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ proxy-auth:
+ uses: ./.github/workflows/_test-unit-base.yml
+ with:
+ 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
diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml
new file mode 100644
index 00000000000..bdfb6efeef1
--- /dev/null
+++ b/.github/workflows/test-unit-proxy-db.yml
@@ -0,0 +1,45 @@
+name: "Unit Tests: Proxy DB Operations"
+
+# Uses DATABASE_URL secret — only runs on trusted branches, not PRs.
+on:
+ push:
+ branches: [main, "litellm_*"]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ proxy-db:
+ strategy:
+ fail-fast: false
+ matrix:
+ include:
+ # Key generation tests must NOT run in parallel (event loop conflicts with logging worker)
+ - test-group: key-generation
+ test-path: "tests/proxy_unit_tests/test_key_generate_prisma.py"
+ workers: 0
+ timeout: 30
+ - test-group: auth-checks
+ test-path: "tests/proxy_unit_tests/test_auth_checks.py tests/proxy_unit_tests/test_user_api_key_auth.py"
+ workers: 8
+ timeout: 20
+ - test-group: remaining
+ test-path: "tests/proxy_unit_tests --ignore=tests/proxy_unit_tests/test_key_generate_prisma.py --ignore=tests/proxy_unit_tests/test_auth_checks.py --ignore=tests/proxy_unit_tests/test_user_api_key_auth.py"
+ workers: 8
+ timeout: 20
+ uses: ./.github/workflows/_test-unit-services-base.yml
+ with:
+ test-path: ${{ matrix.test-path }}
+ workers: ${{ matrix.workers }}
+ reruns: 2
+ timeout-minutes: ${{ matrix.timeout }}
+ enable-redis: false
+ enable-postgres: true
+ secrets:
+ DATABASE_URL: ${{ secrets.DATABASE_URL }}
+ POSTGRES_USER: ${{ secrets.POSTGRES_USER }}
+ POSTGRES_PASSWORD: ${{ secrets.POSTGRES_PASSWORD }}
diff --git a/.github/workflows/test-unit-proxy-endpoints.yml b/.github/workflows/test-unit-proxy-endpoints.yml
new file mode 100644
index 00000000000..caff3b3ae06
--- /dev/null
+++ b/.github/workflows/test-unit-proxy-endpoints.yml
@@ -0,0 +1,35 @@
+name: "Unit Tests: Proxy API Endpoints"
+
+on:
+ pull_request:
+ branches: [main]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ proxy-endpoints:
+ uses: ./.github/workflows/_test-unit-base.yml
+ with:
+ test-path: >-
+ tests/test_litellm/proxy/management_endpoints
+ tests/test_litellm/proxy/guardrails
+ tests/test_litellm/proxy/management_helpers
+ tests/test_litellm/proxy/anthropic_endpoints
+ tests/test_litellm/proxy/google_endpoints
+ tests/test_litellm/proxy/openai_files_endpoint
+ tests/test_litellm/proxy/response_api_endpoints
+ tests/test_litellm/proxy/image_endpoints
+ tests/test_litellm/proxy/vector_store_endpoints
+ tests/test_litellm/proxy/agent_endpoints
+ tests/test_litellm/proxy/discovery_endpoints
+ tests/test_litellm/proxy/health_endpoints
+ tests/test_litellm/proxy/public_endpoints
+ tests/test_litellm/proxy/prompts
+ tests/test_litellm/proxy/ui_crud_endpoints
+ workers: 2
+ reruns: 2
diff --git a/.github/workflows/test-unit-proxy-infra.yml b/.github/workflows/test-unit-proxy-infra.yml
new file mode 100644
index 00000000000..4dfbbe317ed
--- /dev/null
+++ b/.github/workflows/test-unit-proxy-infra.yml
@@ -0,0 +1,28 @@
+name: "Unit Tests: Proxy Infrastructure"
+
+on:
+ pull_request:
+ branches: [main]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ proxy-infra:
+ uses: ./.github/workflows/_test-unit-base.yml
+ with:
+ test-path: >-
+ tests/test_litellm/proxy/db
+ tests/test_litellm/proxy/middleware
+ tests/test_litellm/proxy/spend_tracking
+ tests/test_litellm/proxy/pass_through_endpoints
+ tests/test_litellm/proxy/_experimental
+ tests/test_litellm/proxy/experimental
+ tests/test_litellm/proxy/common_utils
+ tests/test_litellm/proxy/test_*.py
+ workers: 2
+ reruns: 2
diff --git a/.github/workflows/test-unit-proxy-legacy.yml b/.github/workflows/test-unit-proxy-legacy.yml
new file mode 100644
index 00000000000..a9391137263
--- /dev/null
+++ b/.github/workflows/test-unit-proxy-legacy.yml
@@ -0,0 +1,96 @@
+name: "Unit Tests: Proxy Legacy Tests"
+
+on:
+ pull_request:
+ branches: [main]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ test:
+ runs-on: ubuntu-latest
+ timeout-minutes: 20
+ strategy:
+ fail-fast: false
+ matrix:
+ test-group:
+ - name: "auth-and-jwt"
+ path: "tests/proxy_unit_tests/test_[a-j]*.py"
+ - name: "key-generation"
+ path: "tests/proxy_unit_tests/test_[k-o]*.py"
+ - name: "proxy-config"
+ 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"
+ - name: "proxy-server"
+ path: "tests/proxy_unit_tests/test_proxy_server.py"
+ - name: "proxy-server-extras"
+ path: "tests/proxy_unit_tests/test_proxy_server_*.py tests/proxy_unit_tests/test_proxy_setting_guardrails.py"
+ - name: "proxy-utils"
+ path: "tests/proxy_unit_tests/test_proxy_utils.py"
+ - name: "proxy-token-counter"
+ path: "tests/proxy_unit_tests/test_proxy_token_counter.py"
+ - name: "proxy-response-and-misc"
+ path: "tests/proxy_unit_tests/test_[r-t]*.py"
+ - name: "proxy-user-auth-and-spend"
+ path: "tests/proxy_unit_tests/test_[u-z]*.py"
+
+ name: ${{ 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.3.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"
+ 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 }}
+ env:
+ TEST_PATH: ${{ matrix.test-group.path }}
+ run: |
+ poetry run pytest ${TEST_PATH} \
+ --tb=short -vv \
+ --maxfail=10 \
+ -n 2 \
+ --reruns 1 \
+ --reruns-delay 1 \
+ --dist=loadscope \
+ --durations=20
diff --git a/.github/workflows/test-unit-responses-caching-types.yml b/.github/workflows/test-unit-responses-caching-types.yml
new file mode 100644
index 00000000000..7f3acac2803
--- /dev/null
+++ b/.github/workflows/test-unit-responses-caching-types.yml
@@ -0,0 +1,20 @@
+name: "Unit Tests: Responses, Caching & Types"
+
+on:
+ pull_request:
+ branches: [main]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ responses-caching-types:
+ uses: ./.github/workflows/_test-unit-base.yml
+ with:
+ test-path: "tests/test_litellm/responses tests/test_litellm/caching tests/test_litellm/types"
+ workers: 2
+ reruns: 2
diff --git a/.github/workflows/test-unit-security.yml b/.github/workflows/test-unit-security.yml
new file mode 100644
index 00000000000..b38c82b1c24
--- /dev/null
+++ b/.github/workflows/test-unit-security.yml
@@ -0,0 +1,28 @@
+name: "Unit Tests: Security"
+
+# Uses DATABASE_URL secret — only runs on trusted branches, not PRs.
+on:
+ push:
+ branches: [main, "litellm_*"]
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ security:
+ uses: ./.github/workflows/_test-unit-services-base.yml
+ with:
+ test-path: "tests/proxy_security_tests/"
+ workers: 1
+ reruns: 2
+ timeout-minutes: 20
+ enable-redis: false
+ enable-postgres: true
+ secrets:
+ DATABASE_URL: ${{ secrets.DATABASE_URL }}
+ POSTGRES_USER: ${{ secrets.POSTGRES_USER }}
+ POSTGRES_PASSWORD: ${{ secrets.POSTGRES_PASSWORD }}
diff --git a/.github/workflows/test_server_root_path.yml b/.github/workflows/test_server_root_path.yml
index 647d81849c6..47636ce8e92 100644
--- a/.github/workflows/test_server_root_path.yml
+++ b/.github/workflows/test_server_root_path.yml
@@ -18,6 +18,8 @@ jobs:
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
+ with:
+ persist-credentials: false
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3.12
diff --git a/docs/my-website/blog/security_townhall_updates/index.md b/docs/my-website/blog/security_townhall_updates/index.md
new file mode 100644
index 00000000000..b997de9c185
--- /dev/null
+++ b/docs/my-website/blog/security_townhall_updates/index.md
@@ -0,0 +1,190 @@
+---
+slug: security-townhall-updates
+title: "Security Townhall Updates"
+date: 2026-03-27T12:00:00
+authors:
+ - krrish
+ - ishaan-alt
+description: "What happened, what we've done, and what comes next for LiteLLM's release and security processes."
+tags: [security, incident-report]
+hide_table_of_contents: false
+---
+
+import Image from '@theme/IdealImage';
+
+Thank you to everyone who joined our town hall.
+
+We wanted to use that time to walk through what we know, what we've done so far, and how we're improving LiteLLM's release and security processes going forward. This post is a written version of that update. [Slides available here](https://drive.google.com/file/d/17hsSG7nk-OYL7VRCTbTa7McrWREtS9OO/view?usp=sharing)
+
+{/* truncate */}
+
+## What happened
+
+On March 24, 2026 at 10:39 UTC, LiteLLM v1.82.7 was pushed to PyPI. Version v1.82.8 was published soon after. Those packages were live for about 40 minutes before being quarantined by PyPI. By 16:00 UTC, the LiteLLM team had worked with PyPI to delete the affected packages.
+
+At this point, our understanding is that this was a supply-chain incident affecting those two published versions.
+
+## How did this happen?
+
+Our understanding is that the issue came from the [compromised Trivy security scanner](https://www.aquasec.com/blog/trivy-supply-chain-attack-what-you-need-to-know/) dependency in our CI/CD pipeline.
+
+
+
+There were three major contributing factors:
+
+### 1. Shared CI/CD environment
+
+At the time, everything was running on CircleCI, and all steps shared a common environment. That increased blast radius: if one component was compromised, it could potentially access credentials or context intended for other parts of the pipeline.
+
+### 2. Static credentials in environment variables
+
+Release credentials, including credentials for PyPI, GHCR, and Docker publishing, were available as static secrets in the environment. That meant a compromised step could access long-lived release credentials.
+
+### 3. Unpinned Trivy dependency
+
+In our security scanning component, we had an unpinned Trivy dependency. Our present understanding is that a compromised Trivy package ran during the scan, had access to environment variables, and enabled attackers to obtain those credentials.
+
+**In summary:** a compromised package in CI had access to secrets it should not have had, and those secrets were then used in the release path.
+
+## What we've already done
+
+
+In the last 3 days, we've taken the following steps:
+
+### 1. Minimize Scope of Impact
+
+#### Prevented further key abuse
+
+We deleted or rotated all impacted or adjacent secret keys, including PyPI, GitHub, Docker, and related credentials. Out of an abundance of caution, we've also rotated LiteLLM maintainer accounts.
+
+#### Prevent branch attacks
+
+We removed roughly 6,000 open branches and added an auto-deletion policy for branches merged into `main`. This reduces the surface area for branch-based abuse.
+
+#### Pinned CI/CD dependencies
+
+We've pinned all Github Actions, and are working on pinning all CircleCI dependencies as well.
+
+#### Paused releases
+
+We've paused new releases until we've confirmed codebase security and put stronger release controls in place.
+
+### 2. Secured LiteLLM
+
+#### Forensic analysis
+
+We are working with Google's Mandiant cybersecurity team to confirm the source of the attack and verify the security of the codebase. We also confirmed that no malicious code was pushed to `main`.
+
+#### Confirm Application Security
+
+In parallel, we are working with whitehat hackers at [Veria Labs](https://verialabs.com/) to verify application security and review improvements to our CI/CD process.
+
+We have also confirmed that the last 20 LiteLLM releases contain no indicators of compromise, and that no unauthenticated attacks can be made against LiteLLM Proxy based on our current investigation. [Check Security Blog for release verification.](https://docs.litellm.ai/blog/security-update-march-2026#verified-safe-versions)
+
+#### Created a security working group
+
+We created a new security working group inside LiteLLM focused on:
+
+- Building threat models
+- Auditing the build process and dependencies
+
+If you're interested in joining the security working group, please file an issue [here](https://github.com/BerriAI/litellm-security-wg).
+
+### 3. Improved CI/CD
+
+We've already begun making structural changes to how releases are built and published. These align with our goals (covered in the next section) around isolated environments, ephemeral credentials, and release auditing.
+
+## Roadmap
+
+We plan on following 4 guiding principles for our new CI/CD pipeline:
+
+1. **Limit** what each package can access
+2. **Reduce** the number of sensitive environment variables
+3. **Avoid** compromised packages
+4. **Prevent** release tampering
+
+
+### Isolated environments
+
+
+
+We are breaking our CI/CD into 4 semantic concepts:
+
+1. Unit tests
+2. Integration tests
+3. Security scans
+4. Release publishing
+
+And will be running each of these in isolated environments.
+
+This will limit the damage that any single compromised component can cause.
+
+### Ephemeral credentials
+
+We plan to move to ephemeral credentials for PyPI (Trusted Publisher) and GHCR (Token-based authentication) releases. This will reduce the risk of credentials being leaked or compromised.
+
+We have already begun doing this:
+
+- PyPI Trusted Publisher on GitHub Actions [PR](https://github.com/BerriAI/litellm/pull/24654)
+- GHCR Token-based authentication on GitHub Actions [PR](https://github.com/BerriAI/litellm/pull/24683)
+
+### Release auditing
+
+Our goal is to allow users to independently verify that a release came from us and prevent silent modifications of releases after they are published.
+
+This will ensure, your releases are safe, even when:
+- Stolen PyPI/GHCR credentials are used to publish malicious releases
+- 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).
+
+
+### Avoid Compromised Packages
+
+- Move to pinned, verified SHAs for packages and actions used in CI/CD, avoiding `latest` wherever possible.
+- Add a cooldown period before upgrading to a new version of a package - allows more time to investigate and verify the new version.
+
+We've added zizmor to help us catch issues such as unpinned dependencies and credential leakage. [commit](https://github.com/BerriAI/litellm/commit/a671275f5c5b0e1fb1adacdf3b6ef779aaa5d56c).
+
+
+## Frequently Asked Questions
+
+**Q: Did you observe any lateral movement into your corporate environment during this incident?**
+
+A: No. Our investigation to date, conducted in coordination with external security experts, has found no evidence of lateral movement into our internal corporate systems. The incident was isolated to the CI/CD pipeline and the release path for specific versions (v1.82.7 and v1.82.8). As a proactive measure, we have rotated all potentially impacted or adjacent secrets—including PyPI, GitHub, and Docker credentials—and updated maintainer account security to ensure continued isolation.
+
+**Q: Do you expect delays in future product releases due to these new security measures?**
+
+A: We are committed to balancing security with speed. While we have temporarily paused releases to implement stronger controls, we are moving quickly to automate our new security protocols. We are currently implementing isolated CI/CD environments, ephemeral credentials (via Trusted Publishers), and release auditing with Cosign. These improvements are designed to be integrated into our automated pipeline, allowing us to maintain a fast release cadence while ensuring every package is verified and secure.
+
+**Q: Were older packages impacted?**
+
+Our current findings show no indicators of compromise in the last 20 versions of LiteLLM. This was manually verified by our team and independently reviewed by Veria Labs.
+
+We have also published the verified versions for users to use. [Check Security Blog for release verification.](https://docs.litellm.ai/blog/security-update-march-2026#verified-safe-versions)
+
+
+
+## Questions & Support
+
+If you believe your systems may be affected, contact us immediately:
+
+- **Security:** security@berri.ai
+- **Support:** support@berri.ai
+- **Slack:** Reach out to the LiteLLM team directly [here](https://join.slack.com/t/litellmossslack/shared_invite/zt-3o7nkuyfr-p_kbNJj8taRfXGgQI1~YyA)
+
+## Hiring
+
+We are currently hiring for:
+
+- DevOps Engineer - to keep ci/cd secure and running smoothly
+- Security Engineer - to keep the application secure
+
+If you're interest in joining, please apply [here](https://jobs.ashbyhq.com/litellm)
\ No newline at end of file
diff --git a/docs/my-website/blog/security_townhall_updates/shared_ci_cd_environment.png b/docs/my-website/blog/security_townhall_updates/shared_ci_cd_environment.png
new file mode 100644
index 00000000000..29ec195b7fb
Binary files /dev/null and b/docs/my-website/blog/security_townhall_updates/shared_ci_cd_environment.png differ
diff --git a/docs/my-website/blog/security_update_march_2026/index.md b/docs/my-website/blog/security_update_march_2026/index.md
index aaadd27b1bd..93ae7daac3a 100644
--- a/docs/my-website/blog/security_update_march_2026/index.md
+++ b/docs/my-website/blog/security_update_march_2026/index.md
@@ -19,7 +19,7 @@ import VersionVerificationTable from '@site/src/components/VersionVerificationTa
> **Update (March 27):** Added [Verified safe versions](#verified-safe-versions) section with SHA-256 checksums for all audited PyPI and Docker releases.
-> **Update (March 26):** Added `checkmarx[.]zone` to [Indicators of compromise](#indicators-of-compromise-iocs).
+> **Update (March 26):** Added `checkmarx[.]zone` to [Indicators of compromise](#indicators-of-compromise-iocs)
> **Update (March 25):** Added community-contributed scripts for scanning GitHub Actions and GitLab CI pipelines for the compromised versions. See [How to check if you are affected](#how-to-check-if-you-are-affected). s/o [@Zach Fury](https://www.linkedin.com/in/fryware/) for these scripts.
diff --git a/docs/my-website/docs/learn/sdk_quickstart.md b/docs/my-website/docs/learn/sdk_quickstart.md
index bdf7b63eb5d..0fb8c3f02a5 100644
--- a/docs/my-website/docs/learn/sdk_quickstart.md
+++ b/docs/my-website/docs/learn/sdk_quickstart.md
@@ -11,7 +11,7 @@ Use this path if you are integrating LiteLLM directly into application code.
## 1. Install LiteLLM
```bash
-pip install litellm
+pip install litellm==1.82.6
```
## 2. Set Provider Credentials
diff --git a/docs/my-website/docs/providers/gemini.md b/docs/my-website/docs/providers/gemini.md
index c8c9114ea87..87ab5ad40f4 100644
--- a/docs/my-website/docs/providers/gemini.md
+++ b/docs/my-website/docs/providers/gemini.md
@@ -11,6 +11,7 @@ import TabItem from '@theme/TabItem';
| Provider Doc | [Google AI Studio ↗](https://aistudio.google.com/) |
| API Endpoint for Provider | https://generativelanguage.googleapis.com |
| Supported OpenAI Endpoints | `/chat/completions`, [`/embeddings`](../embedding/supported_embedding#gemini-ai-embedding-models), `/completions`, [`/videos`](./gemini/videos.md), [`/images/edits`](../image_edits.md) |
+| Lyria (music) | [Cost map & notes](./gemini/music.md) |
| Pass-through Endpoint | [Supported](../pass_through/google_ai_studio.md) |
diff --git a/docs/my-website/docs/providers/gemini/music.md b/docs/my-website/docs/providers/gemini/music.md
new file mode 100644
index 00000000000..f3968f2db39
--- /dev/null
+++ b/docs/my-website/docs/providers/gemini/music.md
@@ -0,0 +1,28 @@
+# Gemini — Lyria (music generation)
+
+Google Lyria 3 preview models are listed in LiteLLM’s [model cost map](https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json) under the `gemini/` provider for metadata and spend tracking.
+
+| Property | Details |
+|----------|---------|
+| Provider route | `gemini/` |
+| Models | `gemini/lyria-3-clip-preview`, `gemini/lyria-3-pro-preview` |
+| Provider docs | [Gemini API pricing / models ↗](https://ai.google.dev/gemini-api/docs/pricing) |
+
+## Models
+
+| Model | Notes |
+|-------|--------|
+| `gemini/lyria-3-clip-preview` | ~30s clip; paid tier listed as per generated song in Google’s pricing |
+| `gemini/lyria-3-pro-preview` | Full song; paid tier listed as per generated song in Google’s pricing |
+
+Input context limit in the cost map: **131,072** tokens. For modalities, limits, and features, see [Google’s Gemini API docs ↗](https://ai.google.dev/gemini-api/docs/models).
+
+## LiteLLM behavior
+
+- **Cost map**: Per-song paid pricing is stored as `output_cost_per_image` on those entries (flat per generation unit). Token-based completion cost may not reflect music billing until a dedicated path exists.
+- **API calls**: Use the Gemini API as documented by Google. LiteLLM does not ship a separate `music_generation` helper like Veo’s `video_generation`.
+
+## Auth
+
+Same as other Gemini API models: `GEMINI_API_KEY` or `GOOGLE_API_KEY`.
+
diff --git a/docs/my-website/docs/providers/openai.md b/docs/my-website/docs/providers/openai.md
index 2907cdf9f47..1f4a1687e8b 100644
--- a/docs/my-website/docs/providers/openai.md
+++ b/docs/my-website/docs/providers/openai.md
@@ -581,6 +581,90 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
See [OpenAI Reasoning documentation](https://platform.openai.com/docs/guides/reasoning) for more details on organization verification requirements.
+### Multi-turn Conversations with `reasoning_items`
+
+For multi-turn conversations you need `reasoning_items`: structured blocks that include the `encrypted_content` token OpenAI uses to restore reasoning state on the next request. Pass `include=["reasoning.encrypted_content"]` on every call where you want that token returned.
+
+
+
+
+```python showLineNumbers title="Non-streaming: round-trip reasoning_items"
+import litellm
+
+messages = [{"role": "user", "content": "Solve this step by step: 2 + 2"}]
+
+# Turn 1 — get reasoning_items (encrypted_content);
+response = litellm.completion(
+ model="openai/responses/gpt-5-mini",
+ messages=messages,
+ reasoning_effort="low",
+ include=["reasoning.encrypted_content"],
+)
+
+assistant_msg = response.choices[0].message
+
+# Turn 2 — pass reasoning_items back; LiteLLM converts to the correct Responses API format
+messages.append({
+ "role": "assistant",
+ "content": assistant_msg.content,
+ "reasoning_items": assistant_msg.reasoning_items,
+})
+messages.append({"role": "user", "content": "Now summarize your reasoning."})
+
+response2 = litellm.completion(
+ model="openai/responses/gpt-5-mini",
+ messages=messages,
+ reasoning_effort="low",
+ include=["reasoning.encrypted_content"],
+)
+```
+
+
+
+
+`reasoning_items` (with `encrypted_content`) arrive on the final chunk when the full response completes:
+
+```python showLineNumbers title="Streaming: collect and round-trip reasoning_items"
+import litellm
+
+messages = [{"role": "user", "content": "Solve this step by step: 2 + 2"}]
+
+collected_content = []
+collected_reasoning_items = []
+
+stream = litellm.completion(
+ model="openai/responses/gpt-5-mini",
+ messages=messages,
+ stream=True,
+ reasoning_effort="low",
+ include=["reasoning.encrypted_content"],
+)
+
+for chunk in stream:
+ delta = chunk.choices[0].delta
+ if delta.content:
+ collected_content.append(delta.content)
+ if getattr(delta, "reasoning_items", None):
+ collected_reasoning_items.extend(delta.reasoning_items)
+
+messages.append({
+ "role": "assistant",
+ "content": "".join(collected_content),
+ "reasoning_items": collected_reasoning_items or None,
+})
+messages.append({"role": "user", "content": "Continue the conversation."})
+
+response2 = litellm.completion(
+ model="openai/responses/gpt-5-mini",
+ messages=messages,
+ reasoning_effort="low",
+ include=["reasoning.encrypted_content"],
+)
+```
+
+
+
+
### Verbosity Control for GPT-5 Models
The `verbosity` parameter controls the length and detail of responses from GPT-5 family models. It accepts three values: `"low"`, `"medium"`, or `"high"`.
diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md
index 02f5c2be9c7..cc9090c2de6 100644
--- a/docs/my-website/docs/proxy/config_settings.md
+++ b/docs/my-website/docs/proxy/config_settings.md
@@ -279,6 +279,33 @@ router_settings:
| forward_client_headers_to_llm_api | boolean | If true, forwards the client headers (any `x-` headers and `anthropic-beta` headers) to the backend LLM call |
| maximum_spend_logs_retention_period | str | Used to set the max retention time for spend logs in the db, after which they will be auto-purged |
| maximum_spend_logs_retention_interval | str | Used to set the interval in which the spend log cleanup task should run in. |
+| alert_type_config | dict | Configuration mapping alert types to their handler settings |
+| always_include_stream_usage | boolean | If true, includes usage metrics in every streaming response chunk |
+| auto_redirect_ui_login_to_sso | boolean | If true, automatically redirects UI login page to SSO provider |
+| control_plane_url | string | URL of the control plane for cross-instance state sharing |
+| custom_auth_run_common_checks | boolean | If true, runs standard auth validation checks alongside custom auth handlers |
+| custom_ui_sso_sign_in_handler | string | Custom handler for SSO sign-in logic in the UI |
+| database_connection_pool_timeout | integer | Database connection pool timeout in seconds |
+| disable_error_logs | boolean | If true, suppresses error tracking and storage in the database |
+| enable_health_check_routing | boolean | If true, enables health check-driven request routing to avoid unhealthy deployments |
+| enable_mcp_registry | boolean | If true, enables access to the centralized MCP server registry |
+| enforce_rbac | boolean | If true, enables role-based access control (RBAC) for all proxy operations |
+| forward_llm_provider_auth_headers | boolean | If true, forwards provider-specific auth headers to LLM API calls |
+| health_check_concurrency | integer | Maximum number of concurrent health check operations |
+| health_check_staleness_threshold | integer | Maximum age in seconds for health check results before marking deployments as stale |
+| maximum_spend_logs_cleanup_cron | string | Cron expression for scheduling automatic spend log cleanup tasks |
+| mcp_client_side_auth_header_name | string | HTTP header name for client-side MCP server credentials |
+| mcp_internal_ip_ranges | list | CIDR ranges considered internal for non-public MCP server access control |
+| mcp_required_fields | list | List of required field names for MCP server submissions |
+| mcp_trusted_proxy_ranges | list | CIDR ranges of proxies trusted to forward X-Forwarded-For headers for MCP |
+| require_end_user_mcp_access_defined | boolean | If true, requires end users to have explicit MCP access permissions defined |
+| role_permissions | list | List of role-based permission configurations |
+| search_tools | list | List of search tool configurations for enabling web search capabilities |
+| token_rate_limit_type | string | Rate limit counting method: "total", "output", or "input" tokens |
+| use_redis_transaction_buffer | boolean | If true, buffers database transactions in Redis before writing |
+| use_shared_health_check | boolean | If true, uses Redis-backed shared health check state across multiple proxy instances |
+| user_header_mappings | dict | Map custom request headers to user IDs using lookup rules |
+| user_header_name | string | HTTP header name to extract user identity from requests |
### router_settings - Reference
@@ -367,6 +394,8 @@ router_settings:
| ignore_invalid_deployments | boolean | If true, ignores invalid deployments. Default for proxy is True - to prevent invalid models from blocking other models from being loaded. |
| search_tools | List[SearchToolTypedDict] | List of search tool configurations for Search API integration. Each tool specifies a search_tool_name and litellm_params with search_provider, api_key, api_base, etc. [Further Docs](../search/index.md) |
| guardrail_list | List[GuardrailTypedDict] | List of guardrail configurations for guardrail load balancing. Enables load balancing across multiple guardrail deployments with the same guardrail_name. [Further Docs](./guardrails/guardrail_load_balancing.md) |
+| enable_health_check_routing | boolean | If true, enables health check-driven deployment filtering to avoid routing requests to unhealthy deployments |
+| health_check_staleness_threshold | integer | Maximum age in seconds for cached health check results before marking deployments as stale |
### environment variables - Reference
@@ -804,6 +833,7 @@ router_settings:
| LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS | Optionally enable semantic logs for OTEL
| LITELLM_OTEL_INTEGRATION_ENABLE_METRICS | Optionally enable emantic metrics for OTEL
| LITELLM_ENABLE_PYROSCOPE | If true, enables Pyroscope CPU profiling. Profiles are sent to PYROSCOPE_SERVER_ADDRESS. Off by default. See [Pyroscope profiling](/proxy/pyroscope_profiling).
+| LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS | When `true`, if a team's legacy `model_aliases` entry maps a public model name to an internal `model_name__` deployment, pre-call handling can skip that rewrite when team-scoped sibling deployments exist for the public name—so load balancing / `order` apply across siblings. Default is `false` for backwards compatibility. See [Team-scoped models and legacy aliases](./load_balancing#team-scoped-models-and-legacy-model_aliases). When stale aliases are detected and this flag is off, the proxy may log a one-time warning.
| PYROSCOPE_APP_NAME | Application name reported to Pyroscope. Required when LITELLM_ENABLE_PYROSCOPE is true. No default.
| PYROSCOPE_SERVER_ADDRESS | Pyroscope server URL to send profiles to. Required when LITELLM_ENABLE_PYROSCOPE is true. No default.
| PYROSCOPE_SAMPLE_RATE | Optional. Sample rate for Pyroscope profiling (integer). No default; when unset, the pyroscope-io library default is used.
diff --git a/docs/my-website/docs/proxy/health.md b/docs/my-website/docs/proxy/health.md
index 2764a6f0d4f..530bea3d06b 100644
--- a/docs/my-website/docs/proxy/health.md
+++ b/docs/my-website/docs/proxy/health.md
@@ -314,6 +314,89 @@ general_settings:
health_check_details: False
```
+## Health Check Driven Routing
+
+By default, background health checks are observability-only — they populate the `/health` endpoint but don't affect routing. Unhealthy deployments still receive traffic until request failures trigger cooldown.
+
+With `enable_health_check_routing: true`, the router **excludes deployments that failed their last background health check** before selecting a candidate. This gives you proactive failover instead of reactive cooldown.
+
+### How it works
+
+1. Background health checks run on their configured interval
+2. After each cycle, every deployment is marked healthy or unhealthy
+3. On each incoming request, the router filters out unhealthy deployments **before** cooldown filtering and load balancing
+4. If all deployments are unhealthy, the filter is bypassed (safety net — never causes a total outage)
+5. If health state is stale (older than `health_check_staleness_threshold`), it is ignored
+
+### Quick start
+
+```yaml
+model_list:
+ - model_name: gpt-4
+ litellm_params:
+ model: openai/gpt-4
+ api_key: os.environ/OPENAI_API_KEY
+ - model_name: gpt-4
+ litellm_params:
+ model: openai/gpt-4
+ api_key: os.environ/OPENAI_API_KEY_SECONDARY
+
+general_settings:
+ background_health_checks: true
+ health_check_interval: 60
+ enable_health_check_routing: true
+```
+
+### Configuration
+
+| Setting | Where | Default | Description |
+|---------|-------|---------|-------------|
+| `enable_health_check_routing` | `general_settings` | `false` | Enable/disable health-check-driven routing |
+| `health_check_staleness_threshold` | `general_settings` | `health_check_interval * 2` | Seconds before health state is considered stale and ignored |
+| `background_health_checks` | `general_settings` | `false` | Must be `true` for health check routing to work |
+| `health_check_interval` | `general_settings` | `300` | Seconds between health check cycles |
+
+### Interaction with cooldown
+
+Health check filtering and cooldown are **additive**. A deployment can be excluded by either mechanism:
+
+- **Health check filter** — proactive, runs on the configured interval, excludes deployments that failed the last check
+- **Cooldown** — reactive, triggered by request failures, excludes deployments for a short TTL
+
+This means request failures still provide fast detection between health check intervals.
+
+### Staleness
+
+If a health check result is older than `health_check_staleness_threshold`, it is ignored and the deployment is treated as eligible. This prevents stale data from permanently excluding a deployment if the health check loop stops or slows down.
+
+The default staleness threshold is `health_check_interval * 2`. For a 60s interval, health state expires after 120s.
+
+### Example: custom staleness
+
+```yaml
+general_settings:
+ background_health_checks: true
+ health_check_interval: 30
+ enable_health_check_routing: true
+ health_check_staleness_threshold: 90 # ignore health state older than 90s
+```
+
+### Debugging
+
+Run the proxy with `--detailed_debug` and look for:
+
+```
+health_check_routing_state_updated healthy=3 unhealthy=1
+```
+
+This is logged after each health check cycle when routing state is written.
+
+If the safety net triggers (all deployments unhealthy), you'll see:
+
+```
+All deployments marked unhealthy by health checks, bypassing health filter
+```
+
## Health Check Timeout
The health check timeout is set in `litellm/constants.py` and defaults to 60 seconds.
diff --git a/docs/my-website/docs/proxy/load_balancing.md b/docs/my-website/docs/proxy/load_balancing.md
index 74b3e8a5117..93f3d944340 100644
--- a/docs/my-website/docs/proxy/load_balancing.md
+++ b/docs/my-website/docs/proxy/load_balancing.md
@@ -324,17 +324,58 @@ model_list:
litellm_params:
model: azure/gpt-4-fallback
api_key: os.environ/AZURE_API_KEY_2
- order: 2 # 👈 Used when order=1 is unavailable
-
-router_settings:
- enable_pre_call_checks: true # 👈 Required for 'order' to work
+ order: 2 # 👈 Used when order=1 fails
```
-:::important
-The `order` parameter requires `enable_pre_call_checks: true` in `router_settings`.
-:::
+### How order-based fallback works
-If `order=1` deployment is unavailable (e.g., rate-limited), the router falls back to `order=2` deployments.
+When a request to an `order=1` deployment fails (connection error, 404, 429, etc.), the router automatically tries `order=2` deployments, then `order=3`, and so on. Each order level gets its own set of retries before escalating to the next.
+
+If all order levels are exhausted, the router falls through to any configured [model-level fallbacks](#fallbacks).
+
+```yaml
+model_list:
+ - model_name: gpt-4
+ litellm_params:
+ model: azure/gpt-4-primary
+ api_key: os.environ/AZURE_API_KEY
+ order: 1
+
+ - model_name: gpt-4
+ litellm_params:
+ model: azure/gpt-4-secondary
+ api_key: os.environ/AZURE_API_KEY_2
+ order: 2
+
+ - model_name: gpt-4-fallback
+ litellm_params:
+ model: openai/gpt-4
+ api_key: os.environ/OPENAI_API_KEY
+
+router_settings:
+ fallbacks:
+ - gpt-4:
+ - gpt-4-fallback # tried after all order levels fail
+```
+
+The fallback chain for the above config: `order=1` → `order=2` → `gpt-4-fallback`.
+
+For 429 (rate limit) errors specifically, the failed deployment is immediately placed on cooldown. If all `order=1` deployments are on cooldown, the router picks `order=2` deployments directly during retries without waiting for the fallback path.
+
+### Team-scoped models and legacy `model_aliases` {#team-scoped-models-and-legacy-model_aliases}
+
+Team-scoped deployments are identified by `model_info.team_id` and `model_info.team_public_model_name`. Requests should use the **public** model name; the router resolves all sibling deployments (same public name, different `api_base` / `order`, etc.) for routing, failover, and deployment `order`.
+
+For router internals: when a `team_id` is in scope, optimized lookups key off `(team_id, team_public_model_name)`. If code passes an internal deployment id (e.g. `model_name__`) instead of the public name, routing still works via the usual deployment-name paths, but the team-specific fast path applies only to the public name.
+
+**Legacy teams:** Older proxy versions could persist `model_aliases` on the team row mapping a public name to a single internal deployment id (`model_name__`). On each request, pre-call logic may still rewrite `model` to that internal name **before** routing, which collapses to one deployment and can make newer sibling deployments unreachable.
+
+**Migration options:**
+
+1. **Recommended for upgrades:** Set environment variable `LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS=true` so that when sibling team deployments exist for the public name, the stale alias rewrite is skipped and team-scoped routing (including `order` and failover) applies. See the [Environment variables](./config_settings) table in the proxy settings doc.
+2. **Data cleanup:** Remove obsolete `model_aliases` entries for team public names from the team record in the database so only `team_public_model_name` + team model list drive access.
+
+If a stale alias is detected and the bypass is **not** enabled, the proxy may emit a **one-time** warning in logs explaining that sibling deployments may be unreachable until the flag is set or aliases are cleaned up.
### When You'll See Load Balancing in Action
diff --git a/docs/my-website/docs/routing.md b/docs/my-website/docs/routing.md
index 67e7f681147..5aa655ae212 100644
--- a/docs/my-website/docs/routing.md
+++ b/docs/my-website/docs/routing.md
@@ -842,6 +842,8 @@ Traffic mirroring allows you to "mimic" production traffic to a secondary (silen
Set `order` in `litellm_params` to prioritize deployments. Lower values = higher priority. When multiple deployments share the same `order`, the routing strategy picks among them.
+When a request to an `order=1` deployment fails (connection error, 404, 429, etc.), the router automatically tries `order=2` deployments, then `order=3`, and so on. Each order level gets its own set of retries before escalating to the next. If all order levels are exhausted, the router falls through to any configured [fallbacks](#fallbacks).
+
@@ -862,18 +864,14 @@ model_list = [
"litellm_params": {
"model": "azure/gpt-4-fallback",
"api_key": os.getenv("AZURE_API_KEY_2"),
- "order": 2, # 👈 Used when order=1 is unavailable
+ "order": 2, # 👈 Tried when order=1 fails
},
},
]
-router = Router(model_list=model_list, enable_pre_call_checks=True) # 👈 Required for 'order' to work
+router = Router(model_list=model_list)
```
-:::important
-The `order` parameter requires `enable_pre_call_checks=True` to be set on the Router.
-:::
-
@@ -889,10 +887,7 @@ model_list:
litellm_params:
model: azure/gpt-4-fallback
api_key: os.environ/AZURE_API_KEY_2
- order: 2 # 👈 Used when order=1 is unavailable
-
-router_settings:
- enable_pre_call_checks: true # 👈 Required for 'order' to work
+ order: 2 # 👈 Tried when order=1 fails
```
diff --git a/docs/my-website/img/isolated_ci_cd_environments.png b/docs/my-website/img/isolated_ci_cd_environments.png
new file mode 100644
index 00000000000..347523f0fab
Binary files /dev/null and b/docs/my-website/img/isolated_ci_cd_environments.png differ
diff --git a/docs/my-website/img/shared_ci_cd_environment.png b/docs/my-website/img/shared_ci_cd_environment.png
new file mode 100644
index 00000000000..e54e11faa85
Binary files /dev/null and b/docs/my-website/img/shared_ci_cd_environment.png differ
diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js
index 6446e227d99..cc2d7800a21 100644
--- a/docs/my-website/sidebars.js
+++ b/docs/my-website/sidebars.js
@@ -850,6 +850,7 @@ const sidebars = {
items: [
"providers/gemini",
"providers/gemini/videos",
+ "providers/gemini/music",
"providers/google_ai_studio/files",
"providers/google_ai_studio/image_gen",
"providers/google_ai_studio/realtime",
diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py
index cbe8d449b42..356f6ecd4b5 100644
--- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py
+++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py
@@ -3,7 +3,7 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if t
"""
from datetime import datetime, timedelta, timezone
-from typing import TYPE_CHECKING, Optional
+from typing import TYPE_CHECKING, List, Optional, Tuple
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@@ -118,6 +118,15 @@ class CheckBatchCost:
get_model_id_from_unified_batch_id,
)
+ try:
+ from litellm.integrations.prometheus import PrometheusLogger
+ prom_logger = PrometheusLogger.get_instance()
+ except Exception as e:
+ verbose_proxy_logger.error(f"CheckBatchCost: could not get Prometheus logger: {e}")
+ prom_logger = None
+
+ processed_models: List[Tuple[Optional[str], Optional[str]]] = []
+
try:
await self._cleanup_stale_managed_objects()
except Exception as cleanup_err:
@@ -172,6 +181,8 @@ class CheckBatchCost:
verbose_proxy_logger.info(
f"Skipping job {unified_object_id} because it is not a valid unified object id"
)
+ if prom_logger:
+ prom_logger.record_check_batch_cost_error("invalid_unified_id")
continue
else:
unified_object_id = decoded_unified_object_id
@@ -183,6 +194,8 @@ class CheckBatchCost:
verbose_proxy_logger.info(
f"Skipping job {unified_object_id} because it is not a valid model id"
)
+ if prom_logger:
+ prom_logger.record_check_batch_cost_error("invalid_model_id")
continue
verbose_proxy_logger.info(
@@ -202,6 +215,8 @@ class CheckBatchCost:
verbose_proxy_logger.info(
f"Skipping job {unified_object_id} because of error querying model ID: {model_id} for cost and usage of batch ID: {batch_id}: {e}"
)
+ if prom_logger:
+ prom_logger.record_check_batch_cost_error("provider_retrieval_error")
continue
## RETRIEVE THE BATCH JOB OUTPUT FILE
@@ -257,11 +272,25 @@ class CheckBatchCost:
content_bytes # type: ignore[arg-type]
)
+ # Record output file size
+ if prom_logger and content_bytes:
+ try:
+ prom_logger.record_managed_file_size(
+ size_bytes=len(content_bytes), # type: ignore
+ purpose="batch",
+ file_type="output",
+ model=model_id,
+ )
+ except Exception:
+ pass
+
deployment_info = self.llm_router.get_deployment(model_id=model_id)
if deployment_info is None:
verbose_proxy_logger.info(
f"Skipping job {unified_object_id} because it is not a valid deployment info"
)
+ if prom_logger:
+ prom_logger.record_check_batch_cost_error("deployment_not_found")
continue
custom_llm_provider = deployment_info.litellm_params.custom_llm_provider
litellm_model_name = deployment_info.litellm_params.model
@@ -318,6 +347,19 @@ class CheckBatchCost:
batch_models=batch_models,
)
+ # Record batch duration (completed_at - created_at)
+ if prom_logger and response.completed_at and response.created_at:
+ duration_seconds = float(response.completed_at - response.created_at)
+ if duration_seconds >= 0:
+ prom_logger.record_managed_batch_duration(
+ duration_seconds=duration_seconds,
+ model=model_name,
+ api_provider=str(llm_provider) if llm_provider else None,
+ )
+
+ # Track this job for the final metrics summary
+ processed_models.append((model_name, str(llm_provider) if llm_provider else None))
+
# mark the job as complete
try:
update_data: dict = {
@@ -334,3 +376,10 @@ class CheckBatchCost:
verbose_proxy_logger.error(
f"CheckBatchCost: failed to mark job {job.id} complete in DB: {db_err}"
)
+
+ # Record polling run metrics (always, even if nothing was processed)
+ if prom_logger:
+ prom_logger.record_check_batch_cost_run(
+ jobs_polled=len(jobs),
+ processed_models=processed_models if processed_models else None,
+ )
diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py
index dc14937d46b..60c564072a0 100644
--- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py
+++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py
@@ -74,6 +74,13 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
self.internal_usage_cache = internal_usage_cache
self.prisma_client = prisma_client
+ @staticmethod
+ def _get_prometheus_logger():
+ """Find PrometheusLogger from litellm.callbacks, if registered."""
+ from litellm.integrations.prometheus import PrometheusLogger
+
+ return PrometheusLogger.get_instance()
+
async def store_unified_file_id(
self,
file_id: str,
@@ -905,6 +912,31 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
model_mappings=model_mappings,
user_api_key_dict=user_api_key_dict,
)
+
+ # Emit Prometheus metrics for managed file creation
+ prom_logger = self._get_prometheus_logger()
+ if prom_logger:
+ first_model = target_model_names_list[0] if target_model_names_list else None
+ first_provider = ""
+ if responses:
+ first_provider = getattr(responses[0], "_hidden_params", {}).get("custom_llm_provider") or ""
+ prom_logger.record_managed_file_created(
+ model=first_model or "",
+ api_provider=first_provider,
+ user=user_api_key_dict.user_id or "",
+ user_email=getattr(user_api_key_dict, "user_email", None) or "",
+ api_key_alias=user_api_key_dict.key_alias or "",
+ )
+ if response.bytes and response.bytes > 0:
+ prom_logger.record_managed_file_size(
+ size_bytes=response.bytes,
+ purpose=response.purpose or "batch",
+ file_type="input",
+ model=first_model,
+ api_provider=first_provider,
+ user=user_api_key_dict.user_id,
+ )
+
return response
@staticmethod
@@ -1083,6 +1115,31 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
file_purpose="batch",
user_api_key_dict=user_api_key_dict,
)
+
+ # Only record batch creation metric on actual create (not retrieve/cancel).
+ # unified_file_id in _hidden_params is only set by the create_batch endpoint.
+ original_unified_file_id = response._hidden_params.get("unified_file_id")
+ if original_unified_file_id:
+ prom_logger = self._get_prometheus_logger()
+ if prom_logger:
+ batch_provider = ""
+ if model_name:
+ try:
+ from litellm.litellm_core_utils.get_llm_provider_logic import (
+ get_llm_provider,
+ )
+ _, batch_provider, _, _ = get_llm_provider(model=model_name)
+ except Exception:
+ if "/" in model_name:
+ batch_provider = model_name.split("/")[0]
+ prom_logger.record_managed_batch_created(
+ model=model_name or "",
+ api_provider=batch_provider,
+ user=user_api_key_dict.user_id or "",
+ user_email=getattr(user_api_key_dict, "user_email", None) or "",
+ api_key_alias=user_api_key_dict.key_alias or "",
+ )
+
elif isinstance(response, LiteLLMFineTuningJob):
## Check if unified_file_id is in the response
unified_file_id = response._hidden_params.get(
@@ -1332,6 +1389,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
f"Alternatively, wait for all batches to complete and for cost to be computed (batch_processed=true)."
)
+ # Record blocked deletion metric
+ prom_logger = self._get_prometheus_logger()
+ if prom_logger:
+ prom_logger.record_managed_file_deleted(result="blocked")
+
raise HTTPException(
status_code=400,
detail=error_message,
@@ -1365,6 +1427,12 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
file_id, litellm_parent_otel_span
)
+ # Record successful deletion metric only on actual success
+ if stored_file_object or delete_response:
+ prom_logger = self._get_prometheus_logger()
+ if prom_logger:
+ prom_logger.record_managed_file_deleted(result="success")
+
if stored_file_object:
return stored_file_object
elif delete_response:
diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260311180521_schema_sync/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260311180521_schema_sync/migration.sql
deleted file mode 100644
index 84eb70ce097..00000000000
--- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260311180521_schema_sync/migration.sql
+++ /dev/null
@@ -1,11 +0,0 @@
--- DropIndex
-DROP INDEX IF EXISTS "LiteLLM_MCPServerTable_approval_status_idx";
-
--- AlterTable
-ALTER TABLE "LiteLLM_MCPServerTable" DROP COLUMN IF EXISTS "approval_status",
-DROP COLUMN IF EXISTS "review_notes",
-DROP COLUMN IF EXISTS "reviewed_at",
-DROP COLUMN IF EXISTS "source_url",
-DROP COLUMN IF EXISTS "submitted_at",
-DROP COLUMN IF EXISTS "submitted_by";
-
diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
index a2c83295403..46be6b31e1f 100644
--- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
+++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
@@ -320,11 +320,15 @@ model LiteLLM_MCPServerTable {
is_byok Boolean @default(false)
byok_description String[] @default([])
byok_api_key_help_url String?
- approval_status String @default("approved")
- submitted_by String?
- submitted_at DateTime?
- reviewed_at DateTime?
- review_notes String?
+ source_url String?
+ // BYOM submission lifecycle
+ approval_status String? @default("active")
+ submitted_by String?
+ submitted_at DateTime?
+ reviewed_at DateTime?
+ review_notes String?
+
+ @@index([approval_status])
}
// Per-user BYOK credentials for MCP servers
diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml
index 14253ad15db..f03cdf61c2a 100644
--- a/litellm-proxy-extras/pyproject.toml
+++ b/litellm-proxy-extras/pyproject.toml
@@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm-proxy-extras"
-version = "0.4.60"
+version = "0.4.61"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
authors = ["BerriAI"]
readme = "README.md"
@@ -22,7 +22,7 @@ requires = ["poetry-core"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
-version = "0.4.60"
+version = "0.4.61"
version_files = [
"pyproject.toml:version",
"../requirements.txt:litellm-proxy-extras==",
diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py
index ee4cdbcdf36..6e6070f7f2c 100644
--- a/litellm/completion_extras/litellm_responses_transformation/transformation.py
+++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py
@@ -32,6 +32,7 @@ from litellm.llms.base_llm.bridges.completion_transformation import (
)
from litellm.types.llms.openai import (
ChatCompletionAnnotation,
+ ChatCompletionReasoningItem,
ChatCompletionToolParamFunctionChunk,
Reasoning,
ResponsesAPIOptionalRequestParams,
@@ -55,6 +56,49 @@ if TYPE_CHECKING:
)
+def _build_reasoning_item(
+ item_id: str,
+ encrypted_content: Optional[str],
+ summary_raw: Any,
+) -> Dict[str, Any]:
+ """Build a ChatCompletionReasoningItem-shaped dict from raw response data.
+
+ Handles both pydantic objects (attribute access) and plain dicts.
+ """
+ summary: List[Dict[str, Any]] = []
+ for s in summary_raw or []:
+ if isinstance(s, dict):
+ summary.append(
+ {"type": s.get("type", "summary_text"), "text": s.get("text", "")}
+ )
+ else:
+ summary.append(
+ {
+ "type": getattr(s, "type", "summary_text"),
+ "text": getattr(s, "text", ""),
+ }
+ )
+ return {
+ "id": item_id,
+ "type": "reasoning",
+ "encrypted_content": encrypted_content,
+ "summary": summary,
+ }
+
+
+def _reasoning_item_to_response_input(r_item: Dict[str, Any]) -> Dict[str, Any]:
+ """Convert a stored ChatCompletionReasoningItem back to a Responses API input item."""
+ r_input: Dict[str, Any] = {
+ "type": "reasoning",
+ "id": r_item.get("id") or f"rs_{id(r_item)}",
+ # summary is always required by the Responses API, even when empty
+ "summary": r_item.get("summary") or [],
+ }
+ if r_item.get("encrypted_content"):
+ r_input["encrypted_content"] = r_item["encrypted_content"]
+ return r_input
+
+
class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
"""
Handler for transforming /chat/completions api requests to litellm.responses requests
@@ -202,10 +246,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
}
)
elif role == "assistant" and tool_calls and isinstance(tool_calls, list):
+ for r_item in msg.get("reasoning_items") or []:
+ input_items.append(_reasoning_item_to_response_input(r_item))
for tool_call in tool_calls:
function = tool_call.get("function")
if function:
- input_tool_call = {
+ input_tool_call: Dict[str, Any] = {
"type": "function_call",
"call_id": tool_call["id"],
}
@@ -217,7 +263,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
else:
raise ValueError(f"tool call not supported: {tool_call}")
elif content is not None:
- # Regular user/assistant message
+ if role == "assistant":
+ for r_item in msg.get("reasoning_items") or []:
+ input_items.append(_reasoning_item_to_response_input(r_item))
input_items.append(
{
"type": "message",
@@ -411,6 +459,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
choices: List[Choices] = []
index = 0
reasoning_content: Optional[str] = None
+ pending_reasoning_item: Optional[Dict[str, Any]] = None
# Collect all tool calls to put them in a single choice
# (Chat Completions API expects all tool calls in one message)
@@ -419,9 +468,16 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
for item in output_items:
if isinstance(item, ResponseReasoningItem):
- for summary_item in item.summary:
- response_text = getattr(summary_item, "text", "")
- reasoning_content = response_text if response_text else ""
+ pending_reasoning_item = _build_reasoning_item(
+ item_id=item.id,
+ encrypted_content=getattr(item, "encrypted_content", None),
+ summary_raw=item.summary,
+ )
+ reasoning_content = " ".join(
+ s["text"]
+ for s in pending_reasoning_item["summary"]
+ if s.get("text")
+ )
elif isinstance(item, ResponseOutputMessage):
for content in item.content:
@@ -436,6 +492,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
content=response_text if response_text else "",
reasoning_content=reasoning_content,
annotations=annotations,
+ reasoning_items=cast(
+ Optional[List[ChatCompletionReasoningItem]],
+ [pending_reasoning_item]
+ if pending_reasoning_item is not None
+ else None,
+ ),
)
choices.append(
@@ -446,7 +508,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
)
)
- reasoning_content = None # flush reasoning content
+ reasoning_content = None # flush
+ pending_reasoning_item = None # flush
index += 1
elif isinstance(item, ResponseFunctionToolCall):
@@ -489,11 +552,18 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
content=None,
tool_calls=accumulated_tool_calls,
reasoning_content=reasoning_content,
+ reasoning_items=cast(
+ Optional[List[ChatCompletionReasoningItem]],
+ [pending_reasoning_item]
+ if pending_reasoning_item is not None
+ else None,
+ ),
)
choices.append(
Choices(message=msg, finish_reason="tool_calls", index=index)
)
reasoning_content = None
+ pending_reasoning_item = None
return choices
@@ -1232,6 +1302,25 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
finish_reason = "tool_calls" if has_function_calls else "stop"
+ # Extract reasoning items with encrypted_content for round-tripping
+ completed_reasoning_items: Optional[List[Dict[str, Any]]] = None
+ for item in output_items:
+ if not isinstance(item, dict) or item.get("type") != "reasoning":
+ continue
+ if completed_reasoning_items is None:
+ completed_reasoning_items = []
+ completed_reasoning_items.append(
+ _build_reasoning_item(
+ item_id=item.get("id", ""),
+ encrypted_content=item.get("encrypted_content"),
+ summary_raw=item.get("summary"),
+ )
+ )
+ completed_reasoning_items_typed = cast(
+ Optional[List[ChatCompletionReasoningItem]],
+ completed_reasoning_items,
+ )
+
usage = None
if response_data.get("usage"):
from litellm.responses.utils import ResponseAPILoggingUtils
@@ -1245,7 +1334,10 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
choices=[
StreamingChoices(
index=0,
- delta=Delta(content=""),
+ delta=Delta(
+ content="",
+ reasoning_items=completed_reasoning_items_typed,
+ ),
finish_reason=finish_reason,
)
],
diff --git a/litellm/constants.py b/litellm/constants.py
index 423f01afac1..252068bd7b0 100644
--- a/litellm/constants.py
+++ b/litellm/constants.py
@@ -1402,6 +1402,9 @@ DEFAULT_SHARED_HEALTH_CHECK_TTL = int(
DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL = int(
os.getenv("DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL", 60)
) # 1 minute - TTL for health check lock
+DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER = (
+ 2 # health state is stale after interval * this
+)
PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS = int(
os.getenv("PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS", 9)
)
diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py
index 90306c11a42..294b9d464f5 100644
--- a/litellm/integrations/prometheus.py
+++ b/litellm/integrations/prometheus.py
@@ -65,6 +65,17 @@ def _get_cached_end_user_id_for_cost_tracking():
class PrometheusLogger(CustomLogger):
# Class variables or attributes
+
+ @staticmethod
+ def get_instance() -> Optional["PrometheusLogger"]:
+ """Find the PrometheusLogger instance from litellm.callbacks, if registered."""
+ import litellm
+
+ for cb in litellm.callbacks:
+ if isinstance(cb, PrometheusLogger):
+ return cb
+ return None
+
def __init__( # noqa: PLR0915
self,
**kwargs,
@@ -180,6 +191,31 @@ class PrometheusLogger(CustomLogger):
),
)
+ # Remaining Budget for Org
+ self.litellm_remaining_org_budget_metric = self._gauge_factory(
+ "litellm_remaining_org_budget_metric",
+ "Remaining budget for org",
+ labelnames=self.get_labels_for_metric(
+ "litellm_remaining_org_budget_metric"
+ ),
+ )
+
+ # Max Budget for Org
+ self.litellm_org_max_budget_metric = self._gauge_factory(
+ "litellm_org_max_budget_metric",
+ "Maximum budget set for org",
+ labelnames=self.get_labels_for_metric("litellm_org_max_budget_metric"),
+ )
+
+ # Org Budget Reset At
+ self.litellm_org_budget_remaining_hours_metric = self._gauge_factory(
+ "litellm_org_budget_remaining_hours_metric",
+ "Remaining hours for org budget to be reset",
+ labelnames=self.get_labels_for_metric(
+ "litellm_org_budget_remaining_hours_metric"
+ ),
+ )
+
# Remaining Budget for API Key
self.litellm_remaining_api_key_budget_metric = self._gauge_factory(
"litellm_remaining_api_key_budget_metric",
@@ -440,6 +476,76 @@ class PrometheusLogger(CustomLogger):
labelnames=[],
)
+ ########################################
+ # Managed Batch Metrics
+ ########################################
+ self.litellm_managed_batch_created_total = self._counter_factory(
+ name="litellm_managed_batch_created_total",
+ documentation="Total number of managed batches created",
+ labelnames=[
+ "model",
+ "api_provider",
+ "user",
+ "user_email",
+ "api_key_alias",
+ ],
+ )
+
+ self.litellm_managed_file_size_bytes = self._gauge_factory(
+ "litellm_managed_file_size_bytes",
+ "Size of the most recent managed batch file in bytes (last-seen value per label combination)",
+ labelnames=["purpose", "file_type", "model", "api_provider", "user"],
+ )
+
+ self.litellm_managed_batch_duration_seconds = self._histogram_factory(
+ "litellm_managed_batch_duration_seconds",
+ "Duration of completed managed batches in seconds (completed_at - created_at)",
+ labelnames=["model", "api_provider"],
+ buckets=BATCH_DURATION_BUCKETS,
+ )
+
+ self.litellm_managed_file_created_total = self._counter_factory(
+ name="litellm_managed_file_created_total",
+ documentation="Total number of managed files created",
+ labelnames=[
+ "model",
+ "api_provider",
+ "user",
+ "user_email",
+ "api_key_alias",
+ ],
+ )
+
+ self.litellm_managed_file_deleted_total = self._counter_factory(
+ name="litellm_managed_file_deleted_total",
+ documentation="Total number of managed file deletions (success or blocked)",
+ labelnames=["result"],
+ )
+
+ self.litellm_check_batch_cost_jobs_polled = self._gauge_factory(
+ "litellm_check_batch_cost_jobs_polled",
+ "Number of unprocessed batches found by the last CheckBatchCost poll",
+ labelnames=[],
+ )
+
+ self.litellm_check_batch_cost_jobs_processed_total = self._counter_factory(
+ name="litellm_check_batch_cost_jobs_processed_total",
+ documentation="Total number of batches successfully cost-tracked by CheckBatchCost",
+ labelnames=["model", "api_provider"],
+ )
+
+ self.litellm_check_batch_cost_errors_total = self._counter_factory(
+ name="litellm_check_batch_cost_errors_total",
+ documentation="Total number of errors in CheckBatchCost by error type",
+ labelnames=["error_type"],
+ )
+
+ self.litellm_check_batch_cost_last_run_timestamp = self._gauge_factory(
+ "litellm_check_batch_cost_last_run_timestamp",
+ "Unix timestamp of the last CheckBatchCost job run",
+ labelnames=[],
+ )
+
except Exception as e:
print_verbose(f"Got exception on init prometheus client {str(e)}")
raise e
@@ -922,6 +1028,9 @@ class PrometheusLogger(CustomLogger):
user_api_team_alias = standard_logging_payload["metadata"][
"user_api_key_team_alias"
]
+ user_api_key_org_id = standard_logging_payload["metadata"].get(
+ "user_api_key_org_id"
+ )
output_tokens = standard_logging_payload["completion_tokens"]
tokens_used = standard_logging_payload["total_tokens"]
response_cost = standard_logging_payload["response_cost"]
@@ -1030,6 +1139,7 @@ class PrometheusLogger(CustomLogger):
litellm_params=litellm_params,
response_cost=response_cost,
user_id=user_id,
+ user_api_key_org_id=user_api_key_org_id,
)
# set proxy virtual key rpm/tpm metrics
@@ -1185,6 +1295,7 @@ class PrometheusLogger(CustomLogger):
litellm_params: dict,
response_cost: float,
user_id: Optional[str] = None,
+ user_api_key_org_id: Optional[str] = None,
):
_metadata = litellm_params.get("metadata") or {}
_team_spend = _metadata.get("user_api_key_team_spend", None)
@@ -1217,12 +1328,16 @@ class PrometheusLogger(CustomLogger):
user_max_budget=_user_max_budget,
response_cost=response_cost,
),
+ self._set_org_budget_metrics_after_api_request(
+ org_id=user_api_key_org_id,
+ response_cost=response_cost,
+ ),
return_exceptions=True,
)
for i, r in enumerate(results):
if isinstance(r, Exception):
verbose_logger.debug(
- f"[Non-Blocking] Prometheus: Budget metric lookup {['key', 'team', 'user'][i]} failed: {r}"
+ f"[Non-Blocking] Prometheus: Budget metric lookup {['key', 'team', 'user', 'org'][i]} failed: {r}"
)
def _increment_top_level_request_and_spend_metrics(
@@ -1411,6 +1526,9 @@ class PrometheusLogger(CustomLogger):
user_api_team_alias = standard_logging_payload["metadata"][
"user_api_key_team_alias"
]
+ user_api_key_org_id = standard_logging_payload["metadata"].get(
+ "user_api_key_org_id"
+ )
try:
self.litellm_llm_api_failed_requests_metric.labels(
@@ -1426,6 +1544,10 @@ class PrometheusLogger(CustomLogger):
),
).inc()
self.set_llm_deployment_failure_metrics(kwargs)
+ await self._set_org_budget_metrics_after_api_request(
+ org_id=user_api_key_org_id,
+ response_cost=0,
+ )
except Exception as e:
verbose_logger.exception(
"prometheus Layer Error(): Exception occured - {}".format(str(e))
@@ -2158,6 +2280,127 @@ class PrometheusLogger(CustomLogger):
except Exception as e:
verbose_logger.debug(f"Error recording guardrail metrics: {str(e)}")
+ ########################################
+ # Managed Batch Metric Recording Methods
+ ########################################
+
+ def record_managed_batch_created(
+ self,
+ model: Optional[str],
+ api_provider: Optional[str],
+ user: Optional[str],
+ user_email: Optional[str],
+ api_key_alias: Optional[str],
+ ):
+ try:
+ self.litellm_managed_batch_created_total.labels(
+ model=model,
+ api_provider=api_provider,
+ user=user,
+ user_email=user_email,
+ api_key_alias=api_key_alias,
+ ).inc()
+ except Exception as e:
+ verbose_logger.warning(f"Error recording batch created metric: {e}")
+
+ def record_managed_file_size(
+ self,
+ size_bytes: int,
+ purpose: str,
+ file_type: str,
+ model: Optional[str] = None,
+ api_provider: Optional[str] = None,
+ user: Optional[str] = None,
+ ):
+ """Record the size of a managed file. Uses a gauge (last-seen value per label combination)."""
+ try:
+ self.litellm_managed_file_size_bytes.labels(
+ purpose=purpose,
+ file_type=file_type,
+ model=model or "",
+ api_provider=api_provider or "",
+ user=user or "",
+ ).set(size_bytes)
+ except Exception as e:
+ verbose_logger.warning(f"Error recording file size metric: {e}")
+
+ def record_managed_batch_duration(
+ self,
+ duration_seconds: float,
+ model: Optional[str] = None,
+ api_provider: Optional[str] = None,
+ ):
+ try:
+ self.litellm_managed_batch_duration_seconds.labels(
+ model=model or "",
+ api_provider=api_provider or "",
+ ).observe(duration_seconds)
+ except Exception as e:
+ verbose_logger.warning(f"Error recording batch duration metric: {e}")
+
+ def record_managed_file_created(
+ self,
+ model: Optional[str],
+ api_provider: Optional[str],
+ user: Optional[str],
+ user_email: Optional[str],
+ api_key_alias: Optional[str],
+ ):
+ try:
+ self.litellm_managed_file_created_total.labels(
+ model=model,
+ api_provider=api_provider,
+ user=user,
+ user_email=user_email,
+ api_key_alias=api_key_alias,
+ ).inc()
+ except Exception as e:
+ verbose_logger.warning(f"Error recording file created metric: {e}")
+
+ def record_managed_file_deleted(self, result: str):
+ """Record a managed file deletion attempt. result is 'success' or 'blocked'."""
+ try:
+ self.litellm_managed_file_deleted_total.labels(result=result).inc()
+ except Exception as e:
+ verbose_logger.warning(f"Error recording file deleted metric: {e}")
+
+ def record_check_batch_cost_run(
+ self,
+ jobs_polled: int,
+ processed_models: Optional[List[Tuple[Optional[str], Optional[str]]]] = None,
+ ):
+ """
+ Record CheckBatchCost polling metrics.
+
+ Args:
+ jobs_polled: Number of unprocessed batches found
+ processed_models: List of (model, api_provider) tuples for processed jobs
+ """
+ import time
+
+ try:
+ self.litellm_check_batch_cost_last_run_timestamp.set(time.time())
+ self.litellm_check_batch_cost_jobs_polled.set(jobs_polled)
+
+ if processed_models:
+ for model, api_provider in processed_models:
+ self.litellm_check_batch_cost_jobs_processed_total.labels(
+ model=model or "",
+ api_provider=api_provider or "",
+ ).inc()
+ except Exception as e:
+ verbose_logger.warning(f"Error recording check batch cost metrics: {e}")
+
+ def record_check_batch_cost_error(self, error_type: str):
+ try:
+ self.litellm_check_batch_cost_errors_total.labels(
+ error_type=error_type,
+ ).inc()
+ except Exception as e:
+ verbose_logger.warning(
+ f"Error recording check batch cost error metric: {e}"
+ )
+
@staticmethod
def _get_exception_class_name(exception: Exception) -> str:
exception_class_name = ""
@@ -2534,6 +2777,37 @@ class PrometheusLogger(CustomLogger):
data_type="users",
)
+ async def _initialize_org_budget_metrics(self):
+ """
+ Initialize org budget metrics by reusing the generic pagination logic.
+ """
+ from litellm.proxy.proxy_server import prisma_client
+
+ if prisma_client is None:
+ verbose_logger.debug(
+ "Prometheus: skipping org metrics initialization, DB not initialized"
+ )
+ return
+
+ async def fetch_orgs(
+ page_size: int, page: int
+ ) -> Tuple[list, Optional[int]]:
+ skip = (page - 1) * page_size
+ orgs = await prisma_client.db.litellm_organizationtable.find_many(
+ skip=skip,
+ take=page_size,
+ order={"created_at": "desc"},
+ include={"litellm_budget_table": True},
+ )
+ total_count = await prisma_client.db.litellm_organizationtable.count()
+ return orgs, total_count
+
+ await self._initialize_budget_metrics(
+ data_fetch_function=fetch_orgs,
+ set_metrics_function=self._set_org_list_budget_metrics,
+ data_type="orgs",
+ )
+
async def initialize_remaining_budget_metrics(self):
"""
Handler for initializing remaining budget metrics for all teams to avoid metric discrepancies.
@@ -2568,10 +2842,11 @@ class PrometheusLogger(CustomLogger):
"""
Helper to initialize remaining budget metrics for all teams, API keys, and users.
"""
- verbose_logger.debug("Emitting key, team, user budget metrics....")
+ verbose_logger.debug("Emitting key, team, user, org budget metrics....")
await self._initialize_team_budget_metrics()
await self._initialize_api_key_budget_metrics()
await self._initialize_user_budget_metrics()
+ await self._initialize_org_budget_metrics()
await self._initialize_user_and_team_count_metrics()
async def _initialize_user_and_team_count_metrics(self):
@@ -2627,6 +2902,20 @@ class PrometheusLogger(CustomLogger):
for user in users:
self._set_user_budget_metrics(user)
+ async def _set_org_list_budget_metrics(self, orgs: list):
+ """Helper function to set budget metrics for a list of orgs"""
+ for org in orgs:
+ budget_table = getattr(org, "litellm_budget_table", None)
+ self._set_org_budget_metrics(
+ org_id=org.organization_id or "",
+ org_alias=org.organization_alias or "",
+ spend=org.spend or 0.0,
+ max_budget=budget_table.max_budget if budget_table else None,
+ budget_reset_at=getattr(budget_table, "budget_reset_at", None)
+ if budget_table
+ else None,
+ )
+
async def _set_team_budget_metrics_after_api_request(
self,
user_api_team: Optional[str],
@@ -2748,6 +3037,113 @@ class PrometheusLogger(CustomLogger):
)
)
+ async def _set_org_budget_metrics_after_api_request(
+ self,
+ org_id: Optional[str],
+ response_cost: float,
+ ):
+ """
+ Set org budget metrics after an LLM API request
+
+ - Fetches org info via cache (get_org_object)
+ - Sets org budget metrics
+ """
+ if not org_id:
+ return
+
+ from litellm.proxy.auth.auth_checks import get_org_object
+ from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
+
+ if prisma_client is None:
+ return
+
+ try:
+ org_info = await get_org_object(
+ org_id=org_id,
+ prisma_client=prisma_client,
+ user_api_key_cache=user_api_key_cache,
+ include_budget_table=True,
+ )
+ except Exception as e:
+ verbose_logger.debug(
+ f"[Non-Blocking] Prometheus: Error getting org info: {str(e)}"
+ )
+ return
+
+ if org_info is None:
+ return
+
+ org_alias = org_info.organization_alias or ""
+ _total_org_spend = (org_info.spend or 0.0) + response_cost
+ budget_table = org_info.litellm_budget_table
+ max_budget = budget_table.max_budget if budget_table else None
+ budget_reset_at = (
+ getattr(budget_table, "budget_reset_at", None) if budget_table else None
+ )
+
+ self._set_org_budget_metrics(
+ org_id=org_id,
+ org_alias=org_alias,
+ spend=_total_org_spend,
+ max_budget=max_budget,
+ budget_reset_at=budget_reset_at,
+ )
+
+ def _set_org_budget_metrics(
+ self,
+ org_id: str,
+ org_alias: str,
+ spend: float,
+ max_budget: Optional[float],
+ budget_reset_at: Optional[datetime],
+ ):
+ """
+ Set org budget metrics for a single org
+
+ - Remaining Budget
+ - Max Budget
+ - Budget Reset At
+ """
+ enum_values = UserAPIKeyLabelValues(
+ org_id=org_id,
+ org_alias=org_alias,
+ )
+
+ _labels = prometheus_label_factory(
+ supported_enum_labels=self.get_labels_for_metric(
+ metric_name="litellm_remaining_org_budget_metric"
+ ),
+ enum_values=enum_values,
+ )
+ self.litellm_remaining_org_budget_metric.labels(**_labels).set(
+ self._safe_get_remaining_budget(
+ max_budget=max_budget,
+ spend=spend,
+ )
+ )
+
+ if max_budget is not None:
+ _labels = prometheus_label_factory(
+ supported_enum_labels=self.get_labels_for_metric(
+ metric_name="litellm_org_max_budget_metric"
+ ),
+ enum_values=enum_values,
+ )
+ self.litellm_org_max_budget_metric.labels(**_labels).set(max_budget)
+
+ if budget_reset_at is not None:
+ _labels = prometheus_label_factory(
+ supported_enum_labels=self.get_labels_for_metric(
+ metric_name="litellm_org_budget_remaining_hours_metric"
+ ),
+ enum_values=enum_values,
+ )
+ self.litellm_org_budget_remaining_hours_metric.labels(**_labels).set(
+ self._get_remaining_hours_for_budget_reset(
+ budget_reset_at=budget_reset_at
+ )
+ )
+
def _set_key_budget_metrics(self, user_api_key_dict: UserAPIKeyAuth):
"""
Set virtual key budget metrics
diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py
index 36218417377..c26759eaf0a 100644
--- a/litellm/litellm_core_utils/get_llm_provider_logic.py
+++ b/litellm/litellm_core_utils/get_llm_provider_logic.py
@@ -158,12 +158,15 @@ def get_llm_provider( # noqa: PLR0915
): # handle scenario where model="azure/*" and custom_llm_provider="azure"
model = custom_llm_provider + "/" + model
- # Native OpenRouter models have IDs like "openrouter/free" where the
- # "openrouter/" prefix is part of the actual model name on the API.
- # When called from a bridge (e.g. anthropic_messages adapter),
- # custom_llm_provider is already resolved, so return early to prevent
- # the provider-list stripping below from removing the prefix.
+ # OpenRouter: when the router/proxy already set custom_llm_provider,
+ # the model may still carry LiteLLM's "openrouter/" routing prefix.
+ # Native IDs like "openrouter/auto" must stay intact for the API; IDs
+ # like "openrouter/anthropic/claude-3.5-sonnet" must become
+ # "anthropic/claude-3.5-sonnet" (OpenRouter expects provider/model).
if custom_llm_provider == "openrouter" and model.startswith("openrouter/"):
+ remainder = model[len("openrouter/") :]
+ if "/" in remainder:
+ return remainder, custom_llm_provider, dynamic_api_key, api_base
return model, custom_llm_provider, dynamic_api_key, api_base
if api_key and api_key.startswith("os.environ/"):
diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py
index 67e4fadf638..1bb2b99c015 100644
--- a/litellm/litellm_core_utils/streaming_handler.py
+++ b/litellm/litellm_core_utils/streaming_handler.py
@@ -831,6 +831,10 @@ class CustomStreamWrapper:
"annotations" in model_response.choices[0].delta
and model_response.choices[0].delta.annotations is not None
)
+ or (
+ getattr(model_response.choices[0].delta, "reasoning_items", None)
+ is not None
+ )
):
return True
else:
diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py
index 73d1b02c76d..9a99f9efc82 100644
--- a/litellm/llms/anthropic/chat/transformation.py
+++ b/litellm/llms/anthropic/chat/transformation.py
@@ -1421,6 +1421,16 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
):
optional_params["metadata"] = {"user_id": _litellm_metadata["user_id"]}
+ ## Ensure metadata only contains user_id (only documented field in Anthropic Messages API)
+ if "metadata" in optional_params and isinstance(
+ optional_params["metadata"], dict
+ ):
+ _user_id = optional_params["metadata"].get("user_id")
+ if _user_id is not None:
+ optional_params["metadata"] = {"user_id": _user_id}
+ else:
+ optional_params.pop("metadata")
+
# Remove internal LiteLLM parameters that should not be sent to Anthropic API
optional_params.pop("is_vertex_request", None)
diff --git a/litellm/llms/azure/fine_tuning/handler.py b/litellm/llms/azure/fine_tuning/handler.py
index 429b8349896..7e225a84454 100644
--- a/litellm/llms/azure/fine_tuning/handler.py
+++ b/litellm/llms/azure/fine_tuning/handler.py
@@ -1,10 +1,15 @@
-from typing import Optional, Union
+from typing import Any, Coroutine, Dict, Optional, Union, cast
import httpx
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
+from litellm._logging import verbose_logger
from litellm.llms.azure.common_utils import BaseAzureLLM
-from litellm.llms.openai.fine_tuning.handler import OpenAIFineTuningAPI
+from litellm.llms.openai.fine_tuning.handler import (
+ OpenAIFineTuningAPI,
+ _litellm_fine_tuning_job_from_response,
+)
+from litellm.types.utils import LiteLLMFineTuningJob
class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM):
@@ -12,6 +17,194 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM):
AzureOpenAI methods to support fine tuning, inherits from OpenAIFineTuningAPI.
"""
+ @staticmethod
+ def _ensure_training_type(create_fine_tuning_job_data: Dict[str, Any]) -> None:
+ """
+ Azure requires trainingType in extra_body. Default to 1 (supervised) if omitted.
+ """
+ extra_body = create_fine_tuning_job_data.get("extra_body") or {}
+ if not isinstance(extra_body, dict):
+ extra_body = {}
+ if extra_body.get("trainingType") is None:
+ extra_body["trainingType"] = 1
+ create_fine_tuning_job_data["extra_body"] = extra_body
+ verbose_logger.debug(
+ "Azure fine-tuning: defaulting trainingType=1 (supervised)"
+ )
+
+ async def acreate_fine_tuning_job(
+ self,
+ create_fine_tuning_job_data: dict,
+ openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI],
+ ) -> LiteLLMFineTuningJob:
+ response = await openai_client.fine_tuning.jobs.create(
+ **create_fine_tuning_job_data
+ )
+ return _litellm_fine_tuning_job_from_response(response, is_azure=True)
+
+ async def acancel_fine_tuning_job(
+ self,
+ fine_tuning_job_id: str,
+ openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI],
+ ) -> LiteLLMFineTuningJob:
+ response = await openai_client.fine_tuning.jobs.cancel(
+ fine_tuning_job_id=fine_tuning_job_id
+ )
+ return _litellm_fine_tuning_job_from_response(response, is_azure=True)
+
+ async def aretrieve_fine_tuning_job(
+ self,
+ fine_tuning_job_id: str,
+ openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI],
+ ) -> LiteLLMFineTuningJob:
+ response = await openai_client.fine_tuning.jobs.retrieve(
+ fine_tuning_job_id=fine_tuning_job_id
+ )
+ return _litellm_fine_tuning_job_from_response(response, is_azure=True)
+
+ def create_fine_tuning_job(
+ self,
+ _is_async: bool,
+ create_fine_tuning_job_data: dict,
+ api_key: Optional[str],
+ api_base: Optional[str],
+ api_version: Optional[str],
+ timeout: Union[float, httpx.Timeout],
+ max_retries: Optional[int],
+ organization: Optional[str],
+ client: Optional[
+ Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
+ ] = None,
+ ) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]:
+ self._ensure_training_type(create_fine_tuning_job_data)
+
+ openai_client: Optional[
+ Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
+ ] = self.get_openai_client(
+ api_key=api_key,
+ api_base=api_base,
+ timeout=timeout,
+ max_retries=max_retries,
+ organization=organization,
+ client=client,
+ _is_async=_is_async,
+ api_version=api_version,
+ )
+ if openai_client is None:
+ raise ValueError(
+ "Azure OpenAI client is not initialized. Make sure api_key is passed or AZURE_API_KEY is set in the environment."
+ )
+
+ if _is_async is True:
+ if not isinstance(openai_client, (AsyncOpenAI, AsyncAzureOpenAI)):
+ raise ValueError(
+ "OpenAI client is not an instance of AsyncOpenAI. Make sure you passed an AsyncOpenAI client."
+ )
+ return self.acreate_fine_tuning_job(
+ create_fine_tuning_job_data=create_fine_tuning_job_data,
+ openai_client=openai_client,
+ )
+
+ verbose_logger.debug(
+ "creating fine tuning job, args= %s", create_fine_tuning_job_data
+ )
+ response = cast(OpenAI, openai_client).fine_tuning.jobs.create(
+ **create_fine_tuning_job_data
+ )
+ return _litellm_fine_tuning_job_from_response(response, is_azure=True)
+
+ def cancel_fine_tuning_job(
+ self,
+ _is_async: bool,
+ fine_tuning_job_id: str,
+ api_key: Optional[str],
+ api_base: Optional[str],
+ api_version: Optional[str],
+ timeout: Union[float, httpx.Timeout],
+ max_retries: Optional[int],
+ organization: Optional[str],
+ client: Optional[
+ Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
+ ] = None,
+ ) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]:
+ openai_client: Optional[
+ Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
+ ] = self.get_openai_client(
+ api_key=api_key,
+ api_base=api_base,
+ timeout=timeout,
+ max_retries=max_retries,
+ organization=organization,
+ client=client,
+ _is_async=_is_async,
+ api_version=api_version,
+ )
+ if openai_client is None:
+ raise ValueError(
+ "Azure OpenAI client is not initialized. Make sure api_key is passed or AZURE_API_KEY is set in the environment."
+ )
+
+ if _is_async is True:
+ if not isinstance(openai_client, (AsyncOpenAI, AsyncAzureOpenAI)):
+ raise ValueError(
+ "OpenAI client is not an instance of AsyncOpenAI. Make sure you passed an AsyncOpenAI client."
+ )
+ return self.acancel_fine_tuning_job(
+ fine_tuning_job_id=fine_tuning_job_id,
+ openai_client=openai_client,
+ )
+
+ response = cast(OpenAI, openai_client).fine_tuning.jobs.cancel(
+ fine_tuning_job_id=fine_tuning_job_id
+ )
+ return _litellm_fine_tuning_job_from_response(response, is_azure=True)
+
+ def retrieve_fine_tuning_job(
+ self,
+ _is_async: bool,
+ fine_tuning_job_id: str,
+ api_key: Optional[str],
+ api_base: Optional[str],
+ api_version: Optional[str],
+ timeout: Union[float, httpx.Timeout],
+ max_retries: Optional[int],
+ organization: Optional[str],
+ client: Optional[
+ Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
+ ] = None,
+ ) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]:
+ openai_client: Optional[
+ Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
+ ] = self.get_openai_client(
+ api_key=api_key,
+ api_base=api_base,
+ timeout=timeout,
+ max_retries=max_retries,
+ organization=organization,
+ client=client,
+ _is_async=_is_async,
+ api_version=api_version,
+ )
+ if openai_client is None:
+ raise ValueError(
+ "Azure OpenAI client is not initialized. Make sure api_key is passed or AZURE_API_KEY is set in the environment."
+ )
+
+ if _is_async is True:
+ if not isinstance(openai_client, (AsyncOpenAI, AsyncAzureOpenAI)):
+ raise ValueError(
+ "OpenAI client is not an instance of AsyncOpenAI. Make sure you passed an AsyncOpenAI client."
+ )
+ return self.aretrieve_fine_tuning_job(
+ fine_tuning_job_id=fine_tuning_job_id,
+ openai_client=openai_client,
+ )
+
+ response = cast(OpenAI, openai_client).fine_tuning.jobs.retrieve(
+ fine_tuning_job_id=fine_tuning_job_id
+ )
+ return _litellm_fine_tuning_job_from_response(response, is_azure=True)
+
def get_openai_client(
self,
api_key: Optional[str],
diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py
index dd8b1b0a69f..90b501a9cbb 100644
--- a/litellm/llms/bedrock/chat/converse_transformation.py
+++ b/litellm/llms/bedrock/chat/converse_transformation.py
@@ -91,34 +91,6 @@ UNSUPPORTED_BEDROCK_CONVERSE_BETA_PATTERNS = [
"compact-2026-01-12", # The compact beta feature is not currently supported on the Converse and ConverseStream APIs
]
-# Models that support Bedrock's native structured outputs API (outputConfig.textFormat)
-# Uses substring matching against the Bedrock model ID
-# Ref: https://docs.aws.amazon.com/bedrock/latest/userguide/structured-output.html
-BEDROCK_NATIVE_STRUCTURED_OUTPUT_MODELS = {
- # Anthropic Claude 4.5+
- "claude-haiku-4-5",
- "claude-sonnet-4-5",
- "claude-opus-4-5",
- "claude-opus-4-6",
- # Qwen3
- "qwen3",
- # DeepSeek
- "deepseek-v3.1",
- # Gemma 3
- "gemma-3",
- # MiniMax
- "minimax-m2",
- # Mistral (magistral-small excluded: broken constrained decoding on Bedrock)
- "ministral",
- "mistral-large-3",
- "voxtral",
- # Moonshot
- "kimi-k2",
- # NVIDIA
- "nemotron-nano",
- # OpenAI (gpt-oss excluded: broken constrained decoding, works via tool-call fallback)
-}
-
class AmazonConverseConfig(BaseConfig):
"""
@@ -188,8 +160,7 @@ class AmazonConverseConfig(BaseConfig):
if isinstance(content, list):
has_guarded_text = any(
- isinstance(item, dict) and item.get("type") == "guarded_text"
- for item in content
+ isinstance(item, dict) and item.get("type") == "guarded_text" for item in content
)
if has_guarded_text:
continue # Skip this message if it already has guarded_text
@@ -350,13 +321,9 @@ class AmazonConverseConfig(BaseConfig):
# Check if the model is a Nova 2 model (matches nova-2-lite, nova-2-pro, etc.)
# Also check for nova-2/ spec prefix for imported models
- return model_without_region.startswith(
- "amazon.nova-2-"
- ) or model_without_region.startswith("nova-2/")
+ return model_without_region.startswith("amazon.nova-2-") or model_without_region.startswith("nova-2/")
- def _map_web_search_options(
- self, web_search_options: dict, model: str
- ) -> Optional[BedrockToolBlock]:
+ def _map_web_search_options(self, web_search_options: dict, model: str) -> Optional[BedrockToolBlock]:
"""
Map web_search_options to Nova grounding systemTool.
@@ -385,9 +352,7 @@ class AmazonConverseConfig(BaseConfig):
# (unlike Anthropic), so we just enable grounding with no options
return BedrockToolBlock(systemTool={"name": "nova_grounding"})
- def _transform_reasoning_effort_to_reasoning_config(
- self, reasoning_effort: str
- ) -> dict:
+ def _transform_reasoning_effort_to_reasoning_config(self, reasoning_effort: str) -> dict:
"""
Transform reasoning_effort parameter to Nova 2 reasoningConfig structure.
@@ -432,9 +397,7 @@ class AmazonConverseConfig(BaseConfig):
}
}
- def _handle_reasoning_effort_parameter(
- self, model: str, reasoning_effort: str, optional_params: dict
- ) -> None:
+ def _handle_reasoning_effort_parameter(self, model: str, reasoning_effort: str, optional_params: dict) -> None:
"""
Handle the reasoning_effort parameter based on the model type.
@@ -471,9 +434,7 @@ class AmazonConverseConfig(BaseConfig):
optional_params["reasoning_effort"] = reasoning_effort
elif self._is_nova_2_model(model):
# Nova 2 models: transform to reasoningConfig
- reasoning_config = self._transform_reasoning_effort_to_reasoning_config(
- reasoning_effort
- )
+ reasoning_config = self._transform_reasoning_effort_to_reasoning_config(reasoning_effort)
optional_params.update(reasoning_config)
else:
# Anthropic and other models: convert to thinking parameter
@@ -493,8 +454,7 @@ class AmazonConverseConfig(BaseConfig):
budget = thinking.get("budget_tokens")
if isinstance(budget, int) and budget < BEDROCK_MIN_THINKING_BUDGET_TOKENS:
verbose_logger.debug(
- "Bedrock requires thinking.budget_tokens >= %d, got %d. "
- "Clamping to minimum.",
+ "Bedrock requires thinking.budget_tokens >= %d, got %d. Clamping to minimum.",
BEDROCK_MIN_THINKING_BUDGET_TOKENS,
budget,
)
@@ -518,9 +478,7 @@ class AmazonConverseConfig(BaseConfig):
"parallel_tool_calls",
]
- if (
- "arn" in model
- ): # we can't infer the model from the arn, so just add all params
+ if "arn" in model: # we can't infer the model from the arn, so just add all params
supported_params.append("tools")
supported_params.append("tool_choice")
supported_params.append("thinking")
@@ -542,9 +500,7 @@ class AmazonConverseConfig(BaseConfig):
or base_model.startswith("meta.llama3-3")
or base_model.startswith("meta.llama4")
or base_model.startswith("amazon.nova")
- or supports_function_calling(
- model=model, custom_llm_provider=self.custom_llm_provider
- )
+ or supports_function_calling(model=model, custom_llm_provider=self.custom_llm_provider)
):
supported_params.append("tools")
@@ -554,9 +510,7 @@ class AmazonConverseConfig(BaseConfig):
if litellm.utils.supports_tool_choice(
model=model, custom_llm_provider=self.custom_llm_provider
- ) or litellm.utils.supports_tool_choice(
- model=base_model, custom_llm_provider=self.custom_llm_provider
- ):
+ ) or litellm.utils.supports_tool_choice(model=base_model, custom_llm_provider=self.custom_llm_provider):
# only anthropic and mistral support tool choice config. otherwise (E.g. cohere) will fail the call - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_ToolChoice.html
supported_params.append("tool_choice")
@@ -575,9 +529,7 @@ class AmazonConverseConfig(BaseConfig):
model=model,
custom_llm_provider=self.custom_llm_provider,
)
- or supports_reasoning(
- model=base_model, custom_llm_provider=self.custom_llm_provider
- )
+ or supports_reasoning(model=base_model, custom_llm_provider=self.custom_llm_provider)
):
supported_params.append("thinking")
supported_params.append("reasoning_effort")
@@ -602,9 +554,7 @@ class AmazonConverseConfig(BaseConfig):
return ToolChoiceValuesBlock(auto={})
elif isinstance(tool_choice, dict):
# only supported for anthropic + mistral models - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_ToolChoice.html
- specific_tool = SpecificToolChoiceBlock(
- name=tool_choice.get("function", {}).get("name", "")
- )
+ specific_tool = SpecificToolChoiceBlock(name=tool_choice.get("function", {}).get("name", ""))
return ToolChoiceValuesBlock(tool=specific_tool)
else:
raise litellm.utils.UnsupportedParamsError(
@@ -624,15 +574,9 @@ class AmazonConverseConfig(BaseConfig):
return ["mp4", "mov", "mkv", "webm", "flv", "mpeg", "mpg", "wmv", "3gp"]
def get_all_supported_content_types(self) -> List[str]:
- return (
- self.get_supported_image_types()
- + self.get_supported_document_types()
- + self.get_supported_video_types()
- )
+ return self.get_supported_image_types() + self.get_supported_document_types() + self.get_supported_video_types()
- def is_computer_use_tool_used(
- self, tools: Optional[List[OpenAIChatCompletionToolParam]], model: str
- ) -> bool:
+ def is_computer_use_tool_used(self, tools: Optional[List[OpenAIChatCompletionToolParam]], model: str) -> bool:
"""Check if computer use tools are being used in the request."""
if tools is None:
return False
@@ -645,9 +589,7 @@ class AmazonConverseConfig(BaseConfig):
return True
return False
- def _transform_computer_use_tools(
- self, computer_use_tools: List[OpenAIChatCompletionToolParam]
- ) -> List[dict]:
+ def _transform_computer_use_tools(self, computer_use_tools: List[OpenAIChatCompletionToolParam]) -> List[dict]:
"""Transform computer use tools to Bedrock format."""
transformed_tools: List[dict] = []
@@ -689,9 +631,7 @@ class AmazonConverseConfig(BaseConfig):
def _separate_computer_use_tools(
self, tools: List[OpenAIChatCompletionToolParam], model: str
- ) -> Tuple[
- List[OpenAIChatCompletionToolParam], List[OpenAIChatCompletionToolParam]
- ]:
+ ) -> Tuple[List[OpenAIChatCompletionToolParam], List[OpenAIChatCompletionToolParam]]:
"""
Separate computer use tools from regular function tools.
@@ -763,10 +703,20 @@ class AmazonConverseConfig(BaseConfig):
return _tool
@staticmethod
- def _supports_native_structured_outputs(model: str) -> bool:
- """Check if the Bedrock model supports native structured outputs (outputConfig.textFormat)."""
- return any(
- substring in model for substring in BEDROCK_NATIVE_STRUCTURED_OUTPUT_MODELS
+ def _supports_native_structured_outputs(
+ model: str, custom_llm_provider: Optional[str] = None
+ ) -> bool:
+ """Check if the Bedrock model supports native structured outputs (outputConfig.textFormat).
+
+ Delegates to the standard ``supports_native_structured_output`` utility
+ which looks up the flag in ``litellm.model_cost`` via
+ ``_get_model_info_helper``.
+ Ref: https://docs.aws.amazon.com/bedrock/latest/userguide/structured-output.html
+ """
+ from litellm.utils import supports_native_structured_output
+
+ return supports_native_structured_output(
+ model=model, custom_llm_provider=custom_llm_provider
)
@staticmethod
@@ -790,25 +740,18 @@ class AmazonConverseConfig(BaseConfig):
# Recurse into nested schemas
if "properties" in result and isinstance(result["properties"], dict):
result["properties"] = {
- k: AmazonConverseConfig._add_additional_properties_to_schema(v)
- for k, v in result["properties"].items()
+ k: AmazonConverseConfig._add_additional_properties_to_schema(v) for k, v in result["properties"].items()
}
if "items" in result and isinstance(result["items"], dict):
- result["items"] = AmazonConverseConfig._add_additional_properties_to_schema(
- result["items"]
- )
+ result["items"] = AmazonConverseConfig._add_additional_properties_to_schema(result["items"])
for defs_key in ("$defs", "definitions"):
if defs_key in result and isinstance(result[defs_key], dict):
result[defs_key] = {
- k: AmazonConverseConfig._add_additional_properties_to_schema(v)
- for k, v in result[defs_key].items()
+ k: AmazonConverseConfig._add_additional_properties_to_schema(v) for k, v in result[defs_key].items()
}
for key in ("anyOf", "allOf", "oneOf"):
if key in result and isinstance(result[key], list):
- result[key] = [
- AmazonConverseConfig._add_additional_properties_to_schema(item)
- for item in result[key]
- ]
+ result[key] = [AmazonConverseConfig._add_additional_properties_to_schema(item) for item in result[key]]
return result
@@ -838,9 +781,7 @@ class AmazonConverseConfig(BaseConfig):
}
"""
if json_schema is not None:
- json_schema = AmazonConverseConfig._add_additional_properties_to_schema(
- json_schema
- )
+ json_schema = AmazonConverseConfig._add_additional_properties_to_schema(json_schema)
schema_str = json.dumps(json_schema) if json_schema is not None else "{}"
json_schema_def: JsonSchemaDefinition = {"schema": schema_str}
if name is not None:
@@ -862,14 +803,9 @@ class AmazonConverseConfig(BaseConfig):
non_default_params: dict,
optional_params: dict,
):
- optional_params = self._add_tools_to_optional_params(
- optional_params=optional_params, tools=tools
- )
+ optional_params = self._add_tools_to_optional_params(optional_params=optional_params, tools=tools)
- if (
- "meta.llama3-3-70b-instruct-v1:0" in model
- and non_default_params.get("stream", False) is True
- ):
+ if "meta.llama3-3-70b-instruct-v1:0" in model and non_default_params.get("stream", False) is True:
optional_params["fake_stream"] = True
def map_openai_params(
@@ -913,7 +849,9 @@ class AmazonConverseConfig(BaseConfig):
)
if param == "tool_choice":
_tool_choice_value = self.map_tool_choice_values(
- model=model, tool_choice=value, drop_params=drop_params # type: ignore
+ model=model,
+ tool_choice=value,
+ drop_params=drop_params, # type: ignore
)
if _tool_choice_value is not None:
optional_params["tool_choice"] = _tool_choice_value
@@ -1006,7 +944,7 @@ class AmazonConverseConfig(BaseConfig):
if "type" in value and value["type"] == "text":
return optional_params
- if self._supports_native_structured_outputs(model) and json_schema is not None:
+ if self._supports_native_structured_outputs(model, self.custom_llm_provider) and json_schema is not None:
# Use Bedrock's native structured outputs API (outputConfig.textFormat)
# No synthetic tool injection, no fake_stream needed.
# Requires an explicit schema — json_object with no schema falls through
@@ -1024,14 +962,10 @@ class AmazonConverseConfig(BaseConfig):
json_schema=json_schema,
description=description,
)
- optional_params = self._add_tools_to_optional_params(
- optional_params=optional_params, tools=[_tool]
- )
+ optional_params = self._add_tools_to_optional_params(optional_params=optional_params, tools=[_tool])
if (
- litellm.utils.supports_tool_choice(
- model=model, custom_llm_provider=self.custom_llm_provider
- )
+ litellm.utils.supports_tool_choice(model=model, custom_llm_provider=self.custom_llm_provider)
and not is_thinking_enabled
):
optional_params["tool_choice"] = ToolChoiceValuesBlock(
@@ -1043,9 +977,7 @@ class AmazonConverseConfig(BaseConfig):
optional_params["json_mode"] = True
return optional_params
- def update_optional_params_with_thinking_tokens(
- self, non_default_params: dict, optional_params: dict
- ):
+ def update_optional_params_with_thinking_tokens(self, non_default_params: dict, optional_params: dict):
"""
Handles scenario where max tokens is not specified. For anthropic models (anthropic api/bedrock/vertex ai), this requires having the max tokens being set and being greater than the thinking token budget.
@@ -1063,13 +995,9 @@ class AmazonConverseConfig(BaseConfig):
is_thinking_enabled = self.is_thinking_enabled(optional_params)
is_max_tokens_in_request = self.is_max_tokens_in_request(non_default_params)
if is_thinking_enabled and not is_max_tokens_in_request:
- thinking_token_budget = cast(dict, optional_params["thinking"]).get(
- "budget_tokens", None
- )
+ thinking_token_budget = cast(dict, optional_params["thinking"]).get("budget_tokens", None)
if thinking_token_budget is not None:
- optional_params["maxTokens"] = (
- thinking_token_budget + DEFAULT_MAX_TOKENS
- )
+ optional_params["maxTokens"] = thinking_token_budget + DEFAULT_MAX_TOKENS
@overload
def _get_cache_point_block(
@@ -1135,23 +1063,15 @@ class AmazonConverseConfig(BaseConfig):
if message["role"] == "system":
system_prompt_indices.append(idx)
if isinstance(message["content"], str) and message["content"]:
- system_content_blocks.append(
- SystemContentBlock(text=message["content"])
- )
- cache_block = self._get_cache_point_block(
- message, block_type="system", model=model
- )
+ system_content_blocks.append(SystemContentBlock(text=message["content"]))
+ cache_block = self._get_cache_point_block(message, block_type="system", model=model)
if cache_block:
system_content_blocks.append(cache_block)
elif isinstance(message["content"], list):
for m in message["content"]:
if m.get("type") == "text" and m.get("text"):
- system_content_blocks.append(
- SystemContentBlock(text=m["text"])
- )
- cache_block = self._get_cache_point_block(
- m, block_type="system", model=model
- )
+ system_content_blocks.append(SystemContentBlock(text=m["text"]))
+ cache_block = self._get_cache_point_block(m, block_type="system", model=model)
if cache_block:
system_content_blocks.append(cache_block)
if len(system_prompt_indices) > 0:
@@ -1189,16 +1109,10 @@ class AmazonConverseConfig(BaseConfig):
# Exceptions should not be stored in optional_params (this is a defensive fix)
cleaned_params = filter_exceptions_from_params(optional_params)
inference_params = safe_deep_copy(cleaned_params)
- supported_converse_params = list(
- AmazonConverseConfig.__annotations__.keys()
- ) + ["top_k"]
+ supported_converse_params = list(AmazonConverseConfig.__annotations__.keys()) + ["top_k"]
supported_tool_call_params = ["tools", "tool_choice"]
supported_config_params = list(self.get_config_blocks().keys())
- total_supported_params = (
- supported_converse_params
- + supported_tool_call_params
- + supported_config_params
- )
+ total_supported_params = supported_converse_params + supported_tool_call_params + supported_config_params
inference_params.pop("json_mode", None) # used for handling json_schema
# Anthropic-only key. Bedrock expects `outputConfig` (camelCase) and
# will reject `output_config` if it leaks through pass-through routes.
@@ -1209,25 +1123,15 @@ class AmazonConverseConfig(BaseConfig):
if request_metadata is not None:
self._validate_request_metadata(request_metadata)
- output_config: Optional[OutputConfigBlock] = inference_params.pop(
- "outputConfig", None
- )
- inference_params.pop(
- "output_config", None
- ) # Bedrock Converse doesn't support it
+ output_config: Optional[OutputConfigBlock] = inference_params.pop("outputConfig", None)
+ inference_params.pop("output_config", None) # Bedrock Converse doesn't support it
# keep supported params in 'inference_params', and set all model-specific params in 'additional_request_params'
- additional_request_params = {
- k: v for k, v in inference_params.items() if k not in total_supported_params
- }
- inference_params = {
- k: v for k, v in inference_params.items() if k in total_supported_params
- }
+ additional_request_params = {k: v for k, v in inference_params.items() if k not in total_supported_params}
+ inference_params = {k: v for k, v in inference_params.items() if k in total_supported_params}
# Handle parallel_tool_calls configuration
- parallel_tool_use_config = additional_request_params.pop(
- "_parallel_tool_use_config", None
- )
+ parallel_tool_use_config = additional_request_params.pop("_parallel_tool_use_config", None)
if parallel_tool_use_config is not None and is_claude_4_5_on_bedrock(model):
for key, value in parallel_tool_use_config.items():
if (
@@ -1242,9 +1146,7 @@ class AmazonConverseConfig(BaseConfig):
additional_request_params.pop("parallel_tool_calls", None)
# Only set the topK value in for models that support it
- additional_request_params.update(
- self._handle_top_k_value(model, inference_params)
- )
+ additional_request_params.update(self._handle_top_k_value(model, inference_params))
# Filter out internal/MCP-related parameters that shouldn't be sent to the API
# These are LiteLLM internal parameters, not API parameters
@@ -1253,9 +1155,7 @@ class AmazonConverseConfig(BaseConfig):
# Filter out non-serializable objects (exceptions, callables, logging objects, etc.)
# from additional_request_params to prevent JSON serialization errors
# This filters: Exception objects, callable objects (functions), Logging objects, etc.
- additional_request_params = filter_exceptions_from_params(
- additional_request_params
- )
+ additional_request_params = filter_exceptions_from_params(additional_request_params)
return (
inference_params,
@@ -1302,9 +1202,7 @@ class AmazonConverseConfig(BaseConfig):
# Only separate tools if computer use tools are actually present
if filtered_tools and self.is_computer_use_tool_used(filtered_tools, model):
# Separate computer use tools from regular function tools
- computer_use_tools, regular_tools = self._separate_computer_use_tools(
- filtered_tools, model
- )
+ computer_use_tools, regular_tools = self._separate_computer_use_tools(filtered_tools, model)
# Process regular function tools using existing logic
bedrock_tools = _bedrock_tools_pt(regular_tools)
@@ -1365,9 +1263,7 @@ class AmazonConverseConfig(BaseConfig):
anthropic_beta_list.append(computer_use_header)
# Transform computer use tools to proper Bedrock format
- transformed_computer_tools = self._transform_computer_use_tools(
- computer_use_tools
- )
+ transformed_computer_tools = self._transform_computer_use_tools(computer_use_tools)
additional_request_params["tools"] = transformed_computer_tools
else:
# No computer use tools, process all tools as regular tools
@@ -1396,15 +1292,9 @@ class AmazonConverseConfig(BaseConfig):
"""
Bedrock doesn't support tool calling without `tools=` param specified.
"""
- if (
- "tools" not in optional_params
- and messages is not None
- and has_tool_call_blocks(messages)
- ):
+ if "tools" not in optional_params and messages is not None and has_tool_call_blocks(messages):
if litellm.modify_params:
- optional_params["tools"] = add_dummy_tool(
- custom_llm_provider="bedrock_converse"
- )
+ optional_params["tools"] = add_dummy_tool(custom_llm_provider="bedrock_converse")
else:
raise litellm.UnsupportedParamsError(
message="Bedrock doesn't support tool calling without `tools=` param specified. Pass `tools=` param OR set `litellm.modify_params = True` // `litellm_settings::modify_params: True` to add dummy tool to the request.",
@@ -1458,9 +1348,7 @@ class AmazonConverseConfig(BaseConfig):
bedrock_tool_config: Optional[ToolConfigBlock] = None
if len(bedrock_tools) > 0:
- tool_choice_values: ToolChoiceValuesBlock = inference_params.pop(
- "tool_choice", None
- )
+ tool_choice_values: ToolChoiceValuesBlock = inference_params.pop("tool_choice", None)
bedrock_tool_config = ToolConfigBlock(
tools=bedrock_tools,
)
@@ -1470,9 +1358,7 @@ class AmazonConverseConfig(BaseConfig):
data: CommonRequestObject = {
"additionalModelRequestFields": additional_request_params,
"system": system_content_blocks,
- "inferenceConfig": self._transform_inference_params(
- inference_params=inference_params
- ),
+ "inferenceConfig": self._transform_inference_params(inference_params=inference_params),
}
# Handle all config blocks
@@ -1502,14 +1388,10 @@ class AmazonConverseConfig(BaseConfig):
litellm_params: dict,
headers: Optional[dict] = None,
) -> RequestObject:
- messages, system_content_blocks = self._transform_system_message(
- messages, model=model
- )
+ messages, system_content_blocks = self._transform_system_message(messages, model=model)
# Convert last user message to guarded_text if guardrailConfig is present
- messages = self._convert_consecutive_user_messages_to_guarded_text(
- messages, optional_params
- )
+ messages = self._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
## TRANSFORMATION ##
_data: CommonRequestObject = self._transform_request_helper(
@@ -1520,13 +1402,11 @@ class AmazonConverseConfig(BaseConfig):
headers=headers,
)
- bedrock_messages = (
- await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
- messages=messages,
- model=model,
- llm_provider="bedrock_converse",
- user_continue_message=litellm_params.pop("user_continue_message", None),
- )
+ bedrock_messages = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
+ messages=messages,
+ model=model,
+ llm_provider="bedrock_converse",
+ user_continue_message=litellm_params.pop("user_continue_message", None),
)
data: RequestObject = {"messages": bedrock_messages, **_data}
@@ -1560,14 +1440,10 @@ class AmazonConverseConfig(BaseConfig):
litellm_params: dict,
headers: Optional[dict] = None,
) -> RequestObject:
- messages, system_content_blocks = self._transform_system_message(
- messages, model=model
- )
+ messages, system_content_blocks = self._transform_system_message(messages, model=model)
# Convert last user message to guarded_text if guardrailConfig is present
- messages = self._convert_consecutive_user_messages_to_guarded_text(
- messages, optional_params
- )
+ messages = self._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
_data: CommonRequestObject = self._transform_request_helper(
model=model,
@@ -1616,9 +1492,7 @@ class AmazonConverseConfig(BaseConfig):
encoding=encoding,
)
- def _transform_reasoning_content(
- self, reasoning_content_blocks: List[BedrockConverseReasoningContentBlock]
- ) -> str:
+ def _transform_reasoning_content(self, reasoning_content_blocks: List[BedrockConverseReasoningContentBlock]) -> str:
"""
Extract the reasoning text from the reasoning content blocks
@@ -1634,9 +1508,7 @@ class AmazonConverseConfig(BaseConfig):
self, thinking_blocks: List[BedrockConverseReasoningContentBlock]
) -> List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]:
"""Return a consistent format for thinking blocks between Anthropic and Bedrock."""
- thinking_blocks_list: List[
- Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]
- ] = []
+ thinking_blocks_list: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] = []
for block in thinking_blocks:
if "reasoningText" in block:
_thinking_block = ChatCompletionThinkingBlock(type="thinking")
@@ -1672,21 +1544,11 @@ class AmazonConverseConfig(BaseConfig):
cache_creation_input_tokens = usage["cacheWriteInputTokens"]
input_tokens += cache_creation_input_tokens
- prompt_tokens_details = PromptTokensDetailsWrapper(
- cached_tokens=cache_read_input_tokens
- )
- reasoning_tokens = (
- token_counter(text=reasoning_content, count_response_tokens=True)
- if reasoning_content
- else 0
- )
+ prompt_tokens_details = PromptTokensDetailsWrapper(cached_tokens=cache_read_input_tokens)
+ reasoning_tokens = token_counter(text=reasoning_content, count_response_tokens=True) if reasoning_content else 0
completion_tokens_details = CompletionTokensDetailsWrapper(
reasoning_tokens=reasoning_tokens,
- text_tokens=(
- output_tokens - reasoning_tokens
- if reasoning_tokens > 0
- else output_tokens
- ),
+ text_tokens=(output_tokens - reasoning_tokens if reasoning_tokens > 0 else output_tokens),
)
openai_usage = Usage(
prompt_tokens=input_tokens,
@@ -1701,9 +1563,7 @@ class AmazonConverseConfig(BaseConfig):
def get_tool_call_names(
self,
- tools: Optional[
- Union[List[ToolBlock], List[OpenAIChatCompletionToolParam]]
- ] = None,
+ tools: Optional[Union[List[ToolBlock], List[OpenAIChatCompletionToolParam]]] = None,
) -> List[str]:
if tools is None:
return []
@@ -1742,13 +1602,8 @@ class AmazonConverseConfig(BaseConfig):
try:
tool_call_names = self.get_tool_call_names(tools)
json_content = json.loads(message.content)
- if (
- json_content.get("type") == "function"
- and json_content.get("name") in tool_call_names
- ):
- tool_calls = [
- ChatCompletionMessageToolCall(function=Function(**json_content))
- ]
+ if json_content.get("type") == "function" and json_content.get("name") in tool_call_names:
+ tool_calls = [ChatCompletionMessageToolCall(function=Function(**json_content))]
message.tool_calls = tool_calls
message.content = None
@@ -1777,9 +1632,7 @@ class AmazonConverseConfig(BaseConfig):
"""
content_str = ""
tools: List[ChatCompletionToolCallChunk] = []
- reasoningContentBlocks: Optional[
- List[BedrockConverseReasoningContentBlock]
- ] = None
+ reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = None
citationsContentBlocks: Optional[List[CitationsContentBlock]] = None
for idx, content in enumerate(content_blocks):
"""
@@ -1796,9 +1649,7 @@ class AmazonConverseConfig(BaseConfig):
if "toolUse" in content:
## check tool name was formatted by litellm
_response_tool_name = content["toolUse"]["name"]
- response_tool_name = get_bedrock_tool_name(
- response_tool_name=_response_tool_name
- )
+ response_tool_name = get_bedrock_tool_name(response_tool_name=_response_tool_name)
_function_chunk = ChatCompletionToolCallFunctionChunk(
name=response_tool_name,
arguments=json.dumps(content["toolUse"]["input"]),
@@ -1849,11 +1700,7 @@ class AmazonConverseConfig(BaseConfig):
"""
try:
response_data = json.loads(json_str)
- if (
- isinstance(response_data, dict)
- and "properties" in response_data
- and len(response_data) == 1
- ):
+ if isinstance(response_data, dict) and "properties" in response_data and len(response_data) == 1:
response_data = response_data["properties"]
return json.dumps(response_data)
except json.JSONDecodeError:
@@ -1877,11 +1724,7 @@ class AmazonConverseConfig(BaseConfig):
if not json_mode or not tools:
return tools if tools else None
- json_tool_indices = [
- i
- for i, t in enumerate(tools)
- if t["function"].get("name") == RESPONSE_FORMAT_TOOL_NAME
- ]
+ json_tool_indices = [i for i, t in enumerate(tools) if t["function"].get("name") == RESPONSE_FORMAT_TOOL_NAME]
if not json_tool_indices:
# No json_tool_call found, return tools unchanged
@@ -1889,14 +1732,10 @@ class AmazonConverseConfig(BaseConfig):
if len(json_tool_indices) == len(tools):
# All tools are json_tool_call — convert first one to content
- verbose_logger.debug(
- "Processing JSON tool call response for response_format"
- )
+ verbose_logger.debug("Processing JSON tool call response for response_format")
json_mode_content_str: Optional[str] = tools[0]["function"].get("arguments")
if json_mode_content_str is not None:
- json_mode_content_str = AmazonConverseConfig._unwrap_bedrock_properties(
- json_mode_content_str
- )
+ json_mode_content_str = AmazonConverseConfig._unwrap_bedrock_properties(json_mode_content_str)
chat_completion_message["content"] = json_mode_content_str
return None
@@ -1906,13 +1745,9 @@ class AmazonConverseConfig(BaseConfig):
first_idx = json_tool_indices[0]
json_mode_args = tools[first_idx]["function"].get("arguments")
if json_mode_args is not None:
- json_mode_args = AmazonConverseConfig._unwrap_bedrock_properties(
- json_mode_args
- )
+ json_mode_args = AmazonConverseConfig._unwrap_bedrock_properties(json_mode_args)
existing = chat_completion_message.get("content") or ""
- chat_completion_message["content"] = (
- existing + json_mode_args if existing else json_mode_args
- )
+ chat_completion_message["content"] = existing + json_mode_args if existing else json_mode_args
real_tools = [t for i, t in enumerate(tools) if i not in json_tool_indices]
return real_tools if real_tools else None
@@ -1990,9 +1825,7 @@ class AmazonConverseConfig(BaseConfig):
chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"}
content_str = ""
tools: List[ChatCompletionToolCallChunk] = []
- reasoningContentBlocks: Optional[
- List[BedrockConverseReasoningContentBlock]
- ] = None
+ reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = None
citationsContentBlocks: Optional[List[CitationsContentBlock]] = None
if message is not None:
@@ -2011,17 +1844,11 @@ class AmazonConverseConfig(BaseConfig):
provider_specific_fields["citationsContent"] = citationsContentBlocks
if provider_specific_fields:
- chat_completion_message[
- "provider_specific_fields"
- ] = provider_specific_fields
+ chat_completion_message["provider_specific_fields"] = provider_specific_fields
if reasoningContentBlocks is not None:
- chat_completion_message[
- "reasoning_content"
- ] = self._transform_reasoning_content(reasoningContentBlocks)
- chat_completion_message[
- "thinking_blocks"
- ] = self._transform_thinking_blocks(reasoningContentBlocks)
+ chat_completion_message["reasoning_content"] = self._transform_reasoning_content(reasoningContentBlocks)
+ chat_completion_message["thinking_blocks"] = self._transform_thinking_blocks(reasoningContentBlocks)
chat_completion_message["content"] = content_str
filtered_tools = self._filter_json_mode_tools(
json_mode=json_mode,
diff --git a/litellm/llms/gemini/files/transformation.py b/litellm/llms/gemini/files/transformation.py
index bdfb0ee1e52..a29ed66e63d 100644
--- a/litellm/llms/gemini/files/transformation.py
+++ b/litellm/llms/gemini/files/transformation.py
@@ -5,6 +5,7 @@ For vertex ai, check out the vertex_ai/files/handler.py file.
"""
import time
from typing import Any, List, Literal, Optional
+from urllib.parse import urlparse
import httpx
from openai.types.file_deleted import FileDeleted
@@ -209,27 +210,58 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
"""
Get the URL to retrieve a file from Google AI Studio.
- We expect file_id to be the URI (e.g. https://generativelanguage.googleapis.com/v1beta/files/...)
- as returned by the upload response.
+ Endpoint:
+ GET https://generativelanguage.googleapis.com/v1beta/{name=files/*}
+
+ The URL should look like:
+ https://generativelanguage.googleapis.com/v1beta/files/{file_id}?key=API_KEY
+
+ We expect file_id to be just the file identifier (e.g., files/abc123 or abc123)
+ as returned by the upload response. (If it's a full URL, extract the file name.)
"""
api_key = litellm_params.get("api_key") or self.get_api_key()
if not api_key:
raise ValueError("api_key is required")
- if file_id.startswith("http"):
- url = "{}?key={}".format(file_id, api_key)
- else:
- # Fallback for just file name (files/...)
- api_base = (
- self.get_api_base(litellm_params.get("api_base"))
- or "https://generativelanguage.googleapis.com"
- )
- api_base = api_base.rstrip("/")
- url = "{}/v1beta/{}?key={}".format(api_base, file_id, api_key)
+ file_part = self._normalize_gemini_file_id(file_id)
+
+ api_base = (
+ self.get_api_base(litellm_params.get("api_base"))
+ or "https://generativelanguage.googleapis.com"
+ )
+ api_base = api_base.rstrip("/")
+
+ url = f"{api_base}/v1beta/{file_part}?key={api_key}"
# Return empty params dict - API key is already in URL, no query params needed
return url, {}
+ def _normalize_gemini_file_id(self, file_id: str) -> str:
+ """
+ Normalize file identifier into `files/{id}` form.
+
+ Supports:
+ - `abc123`
+ - `files/abc123`
+ - `https://generativelanguage.googleapis.com/v1beta/files/abc123`
+ """
+ if file_id.startswith(("http://", "https://")):
+ parsed = urlparse(file_id)
+ path = parsed.path.lstrip("/")
+ files_index = path.find("files/")
+ if files_index != -1:
+ normalized_file_id = path[files_index:]
+ else:
+ normalized_file_id = path
+ else:
+ normalized_file_id = file_id
+
+ normalized_file_id = normalized_file_id.strip("/")
+ if not normalized_file_id.startswith("files/"):
+ normalized_file_id = f"files/{normalized_file_id}"
+
+ return normalized_file_id
+
def transform_retrieve_file_response(
self,
raw_response: httpx.Response,
@@ -240,8 +272,9 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
Transform Gemini's file retrieval response into OpenAI-style FileObject
"""
try:
+ verbose_logger.debug(f"Retrieve file response: {raw_response.text}")
response_json = raw_response.json()
-
+ verbose_logger.debug(f"Response JSON: {response_json}")
# Map Gemini state to OpenAI status
gemini_state = response_json.get("state", "STATE_UNSPECIFIED")
# Explicitly type status as the Literal union
diff --git a/litellm/llms/openai/fine_tuning/handler.py b/litellm/llms/openai/fine_tuning/handler.py
index 9804ff3539e..c065325254e 100644
--- a/litellm/llms/openai/fine_tuning/handler.py
+++ b/litellm/llms/openai/fine_tuning/handler.py
@@ -1,4 +1,4 @@
-from typing import Any, Coroutine, Optional, Union, cast
+from typing import Any, Coroutine, Dict, Optional, Union, cast
import httpx
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
@@ -6,6 +6,55 @@ from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
from litellm._logging import verbose_logger
from litellm.types.utils import LiteLLMFineTuningJob
+_AZURE_STATUS_MAP = {
+ "pending": "queued",
+ "notRunning": "queued",
+ "running": "running",
+ "succeeded": "succeeded",
+ "failed": "failed",
+ "canceled": "cancelled",
+ "canceling": "cancelled",
+}
+# Note: Azure's "canceling" (in-progress) is mapped to "cancelled" (terminal)
+# because LiteLLMFineTuningJob schema has no intermediate cancellation state.
+
+
+def _normalize_fine_tuning_job_dict(
+ data: Dict[str, Any], is_azure: bool = False
+) -> Dict[str, Any]:
+ """
+ Normalize Azure OpenAI FineTuningJob response to match OpenAI schema.
+
+ Azure differences:
+ - organization_id: null → ""
+ - result_files: null → []
+ - status: mapped via _AZURE_STATUS_MAP
+ """
+ if not is_azure:
+ return data
+
+ normalized = data.copy()
+
+ if normalized.get("organization_id") is None:
+ normalized["organization_id"] = ""
+
+ if normalized.get("result_files") is None:
+ normalized["result_files"] = []
+
+ status = normalized.get("status")
+ if status in _AZURE_STATUS_MAP:
+ normalized["status"] = _AZURE_STATUS_MAP[status]
+
+ return normalized
+
+
+def _litellm_fine_tuning_job_from_response(
+ response: Any, is_azure: bool = False
+) -> LiteLLMFineTuningJob:
+ return LiteLLMFineTuningJob(
+ **_normalize_fine_tuning_job_dict(response.model_dump(), is_azure=is_azure)
+ )
+
class OpenAIFineTuningAPI:
"""
@@ -60,7 +109,7 @@ class OpenAIFineTuningAPI:
**create_fine_tuning_job_data
)
- return LiteLLMFineTuningJob(**response.model_dump())
+ return _litellm_fine_tuning_job_from_response(response)
def create_fine_tuning_job(
self,
@@ -108,7 +157,7 @@ class OpenAIFineTuningAPI:
response = cast(OpenAI, openai_client).fine_tuning.jobs.create(
**create_fine_tuning_job_data
)
- return LiteLLMFineTuningJob(**response.model_dump())
+ return _litellm_fine_tuning_job_from_response(response)
async def acancel_fine_tuning_job(
self,
@@ -118,7 +167,7 @@ class OpenAIFineTuningAPI:
response = await openai_client.fine_tuning.jobs.cancel(
fine_tuning_job_id=fine_tuning_job_id
)
- return LiteLLMFineTuningJob(**response.model_dump())
+ return _litellm_fine_tuning_job_from_response(response)
def cancel_fine_tuning_job(
self,
@@ -164,7 +213,7 @@ class OpenAIFineTuningAPI:
response = cast(OpenAI, openai_client).fine_tuning.jobs.cancel(
fine_tuning_job_id=fine_tuning_job_id
)
- return LiteLLMFineTuningJob(**response.model_dump())
+ return _litellm_fine_tuning_job_from_response(response)
async def alist_fine_tuning_jobs(
self,
@@ -229,7 +278,7 @@ class OpenAIFineTuningAPI:
response = await openai_client.fine_tuning.jobs.retrieve(
fine_tuning_job_id=fine_tuning_job_id
)
- return LiteLLMFineTuningJob(**response.model_dump())
+ return _litellm_fine_tuning_job_from_response(response)
def retrieve_fine_tuning_job(
self,
@@ -275,4 +324,4 @@ class OpenAIFineTuningAPI:
response = cast(OpenAI, openai_client).fine_tuning.jobs.retrieve(
fine_tuning_job_id=fine_tuning_job_id
)
- return LiteLLMFineTuningJob(**response.model_dump())
+ return _litellm_fine_tuning_job_from_response(response)
diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json
index c53ee943c58..d4f986edd9b 100644
--- a/litellm/model_prices_and_context_window_backup.json
+++ b/litellm/model_prices_and_context_window_backup.json
@@ -722,7 +722,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"anthropic.claude-haiku-4-5@20251001": {
"cache_creation_input_token_cost": 1.25e-06,
@@ -745,7 +746,8 @@
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
- "supports_native_streaming": true
+ "supports_native_streaming": true,
+ "supports_native_structured_output": true
},
"anthropic.claude-3-5-sonnet-20240620-v1:0": {
"input_cost_per_token": 3e-06,
@@ -967,22 +969,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_structured_output": true
},
"anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -997,22 +996,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"global.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1027,22 +1023,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"us.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.875e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
"cache_read_input_token_cost": 5.5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
"input_cost_per_token": 5.5e-06,
- "input_cost_per_token_above_200k_tokens": 1.1e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.75e-05,
- "output_cost_per_token_above_200k_tokens": 4.125e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1057,22 +1050,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"eu.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.875e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
"cache_read_input_token_cost": 5.5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
"input_cost_per_token": 5.5e-06,
- "input_cost_per_token_above_200k_tokens": 1.1e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.75e-05,
- "output_cost_per_token_above_200k_tokens": 4.125e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1087,22 +1077,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"au.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.875e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
"cache_read_input_token_cost": 5.5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
"input_cost_per_token": 5.5e-06,
- "input_cost_per_token_above_200k_tokens": 1.1e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.75e-05,
- "output_cost_per_token_above_200k_tokens": 4.125e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1117,22 +1104,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token": 3e-06,
- "input_cost_per_token_above_200k_tokens": 6e-06,
"litellm_provider": "bedrock_converse",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
- "output_cost_per_token_above_200k_tokens": 2.25e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1147,22 +1131,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"global.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token": 3e-06,
- "input_cost_per_token_above_200k_tokens": 6e-06,
"litellm_provider": "bedrock_converse",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
- "output_cost_per_token_above_200k_tokens": 2.25e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1177,22 +1158,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"us.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 4.125e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_read_input_token_cost": 3.3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"input_cost_per_token": 3.3e-06,
- "input_cost_per_token_above_200k_tokens": 6.6e-06,
"litellm_provider": "bedrock_converse",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
- "output_cost_per_token_above_200k_tokens": 2.475e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1207,22 +1185,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"eu.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 4.125e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_read_input_token_cost": 3.3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"input_cost_per_token": 3.3e-06,
- "input_cost_per_token_above_200k_tokens": 6.6e-06,
"litellm_provider": "bedrock_converse",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
- "output_cost_per_token_above_200k_tokens": 2.475e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1237,22 +1212,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"au.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 4.125e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_read_input_token_cost": 3.3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"input_cost_per_token": 3.3e-06,
- "input_cost_per_token_above_200k_tokens": 6.6e-06,
"litellm_provider": "bedrock_converse",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
- "output_cost_per_token_above_200k_tokens": 2.475e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1267,7 +1239,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"anthropic.claude-sonnet-4-20250514-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -1327,7 +1300,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_structured_output": true
},
"anthropic.claude-v1": {
"input_cost_per_token": 8e-06,
@@ -1577,7 +1551,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"apac.anthropic.claude-3-sonnet-20240229-v1:0": {
"input_cost_per_token": 3e-06,
@@ -1665,7 +1640,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"azure/ada": {
"input_cost_per_token": 1e-07,
@@ -1831,7 +1807,7 @@
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "azure_ai",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
@@ -8503,18 +8479,14 @@
},
"claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token": 3e-06,
- "input_cost_per_token_above_200k_tokens": 6e-06,
"litellm_provider": "anthropic",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
- "output_cost_per_token_above_200k_tokens": 2.25e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -8695,19 +8667,15 @@
},
"claude-opus-4-6": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "anthropic",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -8730,19 +8698,15 @@
},
"claude-opus-4-6-20260205": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "anthropic",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -11742,7 +11706,8 @@
"output_cost_per_token": 1.68e-06,
"supports_function_calling": true,
"supports_reasoning": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_structured_output": true
},
"deepseek.v3.2": {
"input_cost_per_token": 6.2e-07,
@@ -12187,7 +12152,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"eu.anthropic.claude-3-5-sonnet-20240620-v1:0": {
"input_cost_per_token": 3e-06,
@@ -12401,7 +12367,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"eu.meta.llama3-2-1b-instruct-v1:0": {
"input_cost_per_token": 1.3e-07,
@@ -14631,18 +14598,6 @@
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"uses_embed_content": true
},
- "vertex_ai/gemini-embedding-2-preview": {
- "input_cost_per_token": 1.5e-07,
- "litellm_provider": "vertex_ai",
- "max_input_tokens": 8192,
- "max_tokens": 8192,
- "mode": "embedding",
- "output_cost_per_token": 0,
- "output_vector_size": 3072,
- "source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal",
- "supports_multimodal": true,
- "uses_embed_content": true
- },
"gemini/gemini-embedding-001": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "gemini",
@@ -15931,6 +15886,55 @@
"supports_tool_choice": true,
"supports_vision": true
},
+ "gemini/lyria-3-clip-preview": {
+ "input_cost_per_token": 0,
+ "litellm_provider": "gemini",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_image": 0.04,
+ "output_cost_per_token": 0,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "audio"
+ ],
+ "supports_audio_input": false,
+ "supports_audio_output": true,
+ "supports_function_calling": false,
+ "supports_prompt_caching": false,
+ "supports_response_schema": false,
+ "supports_system_messages": false,
+ "supports_vision": false,
+ "supports_web_search": false
+ },
+ "gemini/lyria-3-pro-preview": {
+ "input_cost_per_token": 0,
+ "litellm_provider": "gemini",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "audio"
+ ],
+ "supports_audio_input": false,
+ "supports_audio_output": true,
+ "supports_function_calling": false,
+ "supports_prompt_caching": false,
+ "supports_response_schema": false,
+ "supports_system_messages": false,
+ "supports_vision": false,
+ "supports_web_search": false
+ },
"gemini/veo-2.0-generate-001": {
"litellm_provider": "gemini",
"max_input_tokens": 1024,
@@ -16770,7 +16774,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"global.anthropic.claude-sonnet-4-20250514-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -16822,7 +16827,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"global.amazon.nova-2-lite-v1:0": {
"cache_read_input_token_cost": 7.5e-08,
@@ -20551,7 +20557,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"jp.anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.375e-06,
@@ -20573,7 +20580,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"lambda_ai/deepseek-llama3.3-70b": {
"input_cost_per_token": 2e-07,
@@ -21268,7 +21276,8 @@
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.2e-06,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"minimax.minimax-m2.1": {
"input_cost_per_token": 3e-07,
@@ -21424,7 +21433,8 @@
"mode": "chat",
"output_cost_per_token": 2e-07,
"supports_function_calling": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral.ministral-3-3b-instruct": {
"input_cost_per_token": 1e-07,
@@ -21435,7 +21445,8 @@
"mode": "chat",
"output_cost_per_token": 1e-07,
"supports_function_calling": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral.ministral-3-8b-instruct": {
"input_cost_per_token": 1.5e-07,
@@ -21446,7 +21457,8 @@
"mode": "chat",
"output_cost_per_token": 1.5e-07,
"supports_function_calling": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral.mistral-7b-instruct-v0:2": {
"input_cost_per_token": 1.5e-07,
@@ -21488,7 +21500,8 @@
"mode": "chat",
"output_cost_per_token": 1.5e-06,
"supports_function_calling": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral.mistral-small-2402-v1:0": {
"input_cost_per_token": 1e-06,
@@ -21519,7 +21532,8 @@
"mode": "chat",
"output_cost_per_token": 4e-08,
"supports_audio_input": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral.voxtral-small-24b-2507": {
"input_cost_per_token": 1e-07,
@@ -21530,7 +21544,8 @@
"mode": "chat",
"output_cost_per_token": 3e-07,
"supports_audio_input": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral/codestral-2405": {
"input_cost_per_token": 1e-06,
@@ -22217,7 +22232,8 @@
"mode": "chat",
"output_cost_per_token": 2.5e-06,
"supports_reasoning": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"moonshotai.kimi-k2.5": {
"input_cost_per_token": 6e-07,
@@ -23092,7 +23108,8 @@
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "source": "https://aws.amazon.com/bedrock/pricing/"
+ "source": "https://aws.amazon.com/bedrock/pricing/",
+ "supports_native_structured_output": true
},
"o1": {
"cache_read_input_token_cost": 7.5e-06,
@@ -26176,7 +26193,8 @@
"output_cost_per_token": 1.8e-06,
"supports_function_calling": true,
"supports_reasoning": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-235b-a22b-2507-v1:0": {
"input_cost_per_token": 2.2e-07,
@@ -26188,7 +26206,8 @@
"output_cost_per_token": 8.8e-07,
"supports_function_calling": true,
"supports_reasoning": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-coder-30b-a3b-v1:0": {
"input_cost_per_token": 1.5e-07,
@@ -26200,7 +26219,8 @@
"output_cost_per_token": 6e-07,
"supports_function_calling": true,
"supports_reasoning": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-32b-v1:0": {
"input_cost_per_token": 1.5e-07,
@@ -26212,7 +26232,8 @@
"output_cost_per_token": 6e-07,
"supports_function_calling": true,
"supports_reasoning": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-next-80b-a3b": {
"input_cost_per_token": 1.5e-07,
@@ -26223,7 +26244,8 @@
"mode": "chat",
"output_cost_per_token": 1.2e-06,
"supports_function_calling": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-vl-235b-a22b": {
"input_cost_per_token": 5.3e-07,
@@ -26235,7 +26257,8 @@
"output_cost_per_token": 2.66e-06,
"supports_function_calling": true,
"supports_system_messages": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-coder-next": {
"input_cost_per_token": 5e-07,
@@ -28260,7 +28283,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"us.anthropic.claude-3-5-sonnet-20240620-v1:0": {
"input_cost_per_token": 3e-06,
@@ -28418,7 +28442,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "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,
@@ -28439,7 +28464,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"us.anthropic.claude-opus-4-20250514-v1:0": {
"cache_creation_input_token_cost": 1.875e-05,
@@ -28491,7 +28517,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_structured_output": true
},
"global.anthropic.claude-opus-4-5-20251101-v1:0": {
"cache_creation_input_token_cost": 6.25e-06,
@@ -28517,7 +28544,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_structured_output": true
},
"eu.anthropic.claude-opus-4-5-20251101-v1:0": {
"cache_creation_input_token_cost": 6.25e-06,
@@ -28543,7 +28571,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_structured_output": true
},
"us.anthropic.claude-sonnet-4-20250514-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -30323,18 +30352,14 @@
},
"vertex_ai/claude-opus-4-6": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -30353,18 +30378,14 @@
},
"vertex_ai/claude-opus-4-6@default": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -30409,18 +30430,14 @@
},
"vertex_ai/claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token": 3e-06,
- "input_cost_per_token_above_200k_tokens": 6e-06,
"litellm_provider": "vertex_ai-anthropic_models",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
- "output_cost_per_token_above_200k_tokens": 2.25e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@@ -36830,6 +36847,38 @@
"supports_audio_input": true,
"supports_audio_output": true
},
+ "gemini-3.1-flash-live-preview": {
+ "input_cost_per_audio_token": 3e-06,
+ "input_cost_per_image_token": 1e-06,
+ "input_cost_per_token": 7.5e-07,
+ "input_cost_per_video_per_second": 3.3333333333333335e-05,
+ "litellm_provider": "gemini",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 65536,
+ "max_tokens": 65536,
+ "mode": "chat",
+ "output_cost_per_audio_token": 1.2e-05,
+ "output_cost_per_token": 4.5e-06,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_endpoints": [
+ "/v1/realtime"
+ ],
+ "supported_modalities": [
+ "text",
+ "image",
+ "audio",
+ "video"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "audio"
+ ],
+ "supports_audio_input": true,
+ "supports_audio_output": true,
+ "supports_function_calling": true,
+ "supports_vision": true,
+ "supports_web_search": true
+ },
"gemini/gemini-2.5-flash-native-audio-latest": {
"input_cost_per_audio_token": 1e-06,
"input_cost_per_token": 3e-07,
@@ -36908,6 +36957,40 @@
"tpm": 250000,
"rpm": 10
},
+ "gemini/gemini-3.1-flash-live-preview": {
+ "input_cost_per_audio_token": 3e-06,
+ "input_cost_per_image_token": 1e-06,
+ "input_cost_per_token": 7.5e-07,
+ "input_cost_per_video_per_second": 3.3333333333333335e-05,
+ "litellm_provider": "gemini",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 65536,
+ "max_tokens": 65536,
+ "mode": "chat",
+ "output_cost_per_audio_token": 1.2e-05,
+ "output_cost_per_token": 4.5e-06,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_endpoints": [
+ "/v1/realtime"
+ ],
+ "supported_modalities": [
+ "text",
+ "image",
+ "audio",
+ "video"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "audio"
+ ],
+ "supports_audio_input": true,
+ "supports_audio_output": true,
+ "supports_function_calling": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "tpm": 250000,
+ "rpm": 10
+ },
"gemini-2.5-flash-preview-tts": {
"input_cost_per_token": 3e-07,
"litellm_provider": "gemini",
@@ -37153,18 +37236,14 @@
},
"vertex_ai/claude-sonnet-4-6@default": {
"cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token": 3e-06,
- "input_cost_per_token_above_200k_tokens": 6e-06,
"litellm_provider": "vertex_ai-anthropic_models",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
- "output_cost_per_token_above_200k_tokens": 2.25e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
diff --git a/litellm/proxy/_experimental/out/404.html b/litellm/proxy/_experimental/out/404/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/404.html
rename to litellm/proxy/_experimental/out/404/index.html
diff --git a/litellm/proxy/_experimental/out/_not-found.html b/litellm/proxy/_experimental/out/_not-found/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/_not-found.html
rename to litellm/proxy/_experimental/out/_not-found/index.html
diff --git a/litellm/proxy/_experimental/out/api-reference.html b/litellm/proxy/_experimental/out/api-reference/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/api-reference.html
rename to litellm/proxy/_experimental/out/api-reference/index.html
diff --git a/litellm/proxy/_experimental/out/chat.html b/litellm/proxy/_experimental/out/chat/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/chat.html
rename to litellm/proxy/_experimental/out/chat/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/api-playground.html b/litellm/proxy/_experimental/out/experimental/api-playground/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/api-playground.html
rename to litellm/proxy/_experimental/out/experimental/api-playground/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/budgets.html b/litellm/proxy/_experimental/out/experimental/budgets/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/budgets.html
rename to litellm/proxy/_experimental/out/experimental/budgets/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/caching.html b/litellm/proxy/_experimental/out/experimental/caching/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/caching.html
rename to litellm/proxy/_experimental/out/experimental/caching/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/claude-code-plugins.html b/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/claude-code-plugins.html
rename to litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/old-usage.html b/litellm/proxy/_experimental/out/experimental/old-usage/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/old-usage.html
rename to litellm/proxy/_experimental/out/experimental/old-usage/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/prompts.html b/litellm/proxy/_experimental/out/experimental/prompts/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/prompts.html
rename to litellm/proxy/_experimental/out/experimental/prompts/index.html
diff --git a/litellm/proxy/_experimental/out/experimental/tag-management.html b/litellm/proxy/_experimental/out/experimental/tag-management/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/experimental/tag-management.html
rename to litellm/proxy/_experimental/out/experimental/tag-management/index.html
diff --git a/litellm/proxy/_experimental/out/guardrails.html b/litellm/proxy/_experimental/out/guardrails/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/guardrails.html
rename to litellm/proxy/_experimental/out/guardrails/index.html
diff --git a/litellm/proxy/_experimental/out/login.html b/litellm/proxy/_experimental/out/login/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/login.html
rename to litellm/proxy/_experimental/out/login/index.html
diff --git a/litellm/proxy/_experimental/out/logs.html b/litellm/proxy/_experimental/out/logs/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/logs.html
rename to litellm/proxy/_experimental/out/logs/index.html
diff --git a/litellm/proxy/_experimental/out/mcp/oauth/callback.html b/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/mcp/oauth/callback.html
rename to litellm/proxy/_experimental/out/mcp/oauth/callback/index.html
diff --git a/litellm/proxy/_experimental/out/model-hub.html b/litellm/proxy/_experimental/out/model-hub/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/model-hub.html
rename to litellm/proxy/_experimental/out/model-hub/index.html
diff --git a/litellm/proxy/_experimental/out/model_hub.html b/litellm/proxy/_experimental/out/model_hub/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/model_hub.html
rename to litellm/proxy/_experimental/out/model_hub/index.html
diff --git a/litellm/proxy/_experimental/out/model_hub_table.html b/litellm/proxy/_experimental/out/model_hub_table/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/model_hub_table.html
rename to litellm/proxy/_experimental/out/model_hub_table/index.html
diff --git a/litellm/proxy/_experimental/out/models-and-endpoints.html b/litellm/proxy/_experimental/out/models-and-endpoints/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/models-and-endpoints.html
rename to litellm/proxy/_experimental/out/models-and-endpoints/index.html
diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/onboarding.html
rename to litellm/proxy/_experimental/out/onboarding/index.html
diff --git a/litellm/proxy/_experimental/out/organizations.html b/litellm/proxy/_experimental/out/organizations/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/organizations.html
rename to litellm/proxy/_experimental/out/organizations/index.html
diff --git a/litellm/proxy/_experimental/out/playground.html b/litellm/proxy/_experimental/out/playground/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/playground.html
rename to litellm/proxy/_experimental/out/playground/index.html
diff --git a/litellm/proxy/_experimental/out/policies.html b/litellm/proxy/_experimental/out/policies/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/policies.html
rename to litellm/proxy/_experimental/out/policies/index.html
diff --git a/litellm/proxy/_experimental/out/settings/admin-settings.html b/litellm/proxy/_experimental/out/settings/admin-settings/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/settings/admin-settings.html
rename to litellm/proxy/_experimental/out/settings/admin-settings/index.html
diff --git a/litellm/proxy/_experimental/out/settings/logging-and-alerts.html b/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/settings/logging-and-alerts.html
rename to litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html
diff --git a/litellm/proxy/_experimental/out/settings/router-settings.html b/litellm/proxy/_experimental/out/settings/router-settings/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/settings/router-settings.html
rename to litellm/proxy/_experimental/out/settings/router-settings/index.html
diff --git a/litellm/proxy/_experimental/out/settings/ui-theme.html b/litellm/proxy/_experimental/out/settings/ui-theme/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/settings/ui-theme.html
rename to litellm/proxy/_experimental/out/settings/ui-theme/index.html
diff --git a/litellm/proxy/_experimental/out/teams.html b/litellm/proxy/_experimental/out/teams/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/teams.html
rename to litellm/proxy/_experimental/out/teams/index.html
diff --git a/litellm/proxy/_experimental/out/test-key.html b/litellm/proxy/_experimental/out/test-key/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/test-key.html
rename to litellm/proxy/_experimental/out/test-key/index.html
diff --git a/litellm/proxy/_experimental/out/tools/mcp-servers.html b/litellm/proxy/_experimental/out/tools/mcp-servers/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/tools/mcp-servers.html
rename to litellm/proxy/_experimental/out/tools/mcp-servers/index.html
diff --git a/litellm/proxy/_experimental/out/tools/vector-stores.html b/litellm/proxy/_experimental/out/tools/vector-stores/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/tools/vector-stores.html
rename to litellm/proxy/_experimental/out/tools/vector-stores/index.html
diff --git a/litellm/proxy/_experimental/out/usage.html b/litellm/proxy/_experimental/out/usage/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/usage.html
rename to litellm/proxy/_experimental/out/usage/index.html
diff --git a/litellm/proxy/_experimental/out/users.html b/litellm/proxy/_experimental/out/users/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/users.html
rename to litellm/proxy/_experimental/out/users/index.html
diff --git a/litellm/proxy/_experimental/out/virtual-keys.html b/litellm/proxy/_experimental/out/virtual-keys/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/virtual-keys.html
rename to litellm/proxy/_experimental/out/virtual-keys/index.html
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index b59fc85d4b8..8faf36df4c6 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -525,6 +525,7 @@ class LiteLLMRoutes(enum.Enum):
# user
"/user/new",
"/user/update",
+ "/user/bulk_update",
"/user/delete",
"/user/info",
"/user/list",
diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py
index 815393467de..0efc3dab638 100644
--- a/litellm/proxy/auth/auth_checks.py
+++ b/litellm/proxy/auth/auth_checks.py
@@ -2876,7 +2876,15 @@ async def _virtual_key_max_budget_check(
Triggers a budget alert if the token is over it's max budget.
"""
- if valid_token.spend is not None and valid_token.max_budget is not None:
+ if valid_token.max_budget is not None:
+ from litellm.proxy.proxy_server import get_current_spend
+
+ # Read spend from cross-pod counter (Redis-first) or cached object (fallback)
+ spend = await get_current_spend(
+ counter_key=f"spend:key:{valid_token.token}",
+ fallback_spend=valid_token.spend or 0.0,
+ )
+
####################################
# collect information for alerting #
####################################
@@ -2888,7 +2896,7 @@ async def _virtual_key_max_budget_check(
call_info = CallInfo(
token=valid_token.token,
- spend=valid_token.spend,
+ spend=spend,
max_budget=valid_token.max_budget,
soft_budget=valid_token.soft_budget,
user_id=valid_token.user_id,
@@ -2909,9 +2917,9 @@ async def _virtual_key_max_budget_check(
# collect information for alerting #
####################################
- if valid_token.spend >= valid_token.max_budget:
+ if spend >= valid_token.max_budget:
raise litellm.BudgetExceededError(
- current_cost=valid_token.spend,
+ current_cost=spend,
max_budget=valid_token.max_budget,
)
@@ -3042,6 +3050,14 @@ async def _check_team_member_budget(
team_member_budget = team_membership.litellm_budget_table.max_budget
team_member_spend = team_membership.spend or 0.0
+ # Read from cross-pod counter (Redis-first) if available
+ from litellm.proxy.proxy_server import get_current_spend
+
+ team_member_spend = await get_current_spend(
+ counter_key=f"spend:team_member:{valid_token.user_id}:{team_object.team_id}",
+ fallback_spend=team_member_spend,
+ )
+
if team_member_spend >= team_member_budget:
raise litellm.BudgetExceededError(
current_cost=team_member_spend,
@@ -3065,33 +3081,40 @@ async def _team_max_budget_check(
if (
team_object is not None
and team_object.max_budget is not None
- and team_object.spend is not None
- and team_object.spend > team_object.max_budget
):
- if valid_token:
- call_info = CallInfo(
- token=valid_token.token,
- spend=team_object.spend,
- max_budget=team_object.max_budget,
- user_id=valid_token.user_id,
- team_id=valid_token.team_id,
- team_alias=valid_token.team_alias,
- organization_id=valid_token.org_id,
- event_group=Litellm_EntityType.TEAM,
- )
- asyncio.create_task(
- proxy_logging_obj.budget_alerts(
- type="team_budget",
- user_info=call_info,
- )
- )
+ from litellm.proxy.proxy_server import get_current_spend
- raise litellm.BudgetExceededError(
- current_cost=team_object.spend,
- max_budget=team_object.max_budget,
- message=f"Budget has been exceeded! Team={team_object.team_id} Current cost: {team_object.spend}, Max budget: {team_object.max_budget}",
+ # Read spend from cross-pod counter (Redis-first) or cached object (fallback)
+ spend = await get_current_spend(
+ counter_key=f"spend:team:{team_object.team_id}",
+ fallback_spend=team_object.spend or 0.0,
)
+ if spend > team_object.max_budget:
+ if valid_token:
+ call_info = CallInfo(
+ token=valid_token.token,
+ spend=spend,
+ max_budget=team_object.max_budget,
+ user_id=valid_token.user_id,
+ team_id=valid_token.team_id,
+ team_alias=valid_token.team_alias,
+ organization_id=valid_token.org_id,
+ event_group=Litellm_EntityType.TEAM,
+ )
+ asyncio.create_task(
+ proxy_logging_obj.budget_alerts(
+ type="team_budget",
+ user_info=call_info,
+ )
+ )
+
+ raise litellm.BudgetExceededError(
+ current_cost=spend,
+ max_budget=team_object.max_budget,
+ message=f"Budget has been exceeded! Team={team_object.team_id} Current cost: {spend}, Max budget: {team_object.max_budget}",
+ )
+
async def _team_soft_budget_check(
team_object: Optional[LiteLLM_TeamTable],
diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py
index bfad9f0c3c7..880ce3fb322 100644
--- a/litellm/proxy/auth/handle_jwt.py
+++ b/litellm/proxy/auth/handle_jwt.py
@@ -18,6 +18,7 @@ from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
+from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
from litellm.llms.custom_httpx.httpx_handler import HTTPHandler
from litellm.proxy._types import (
@@ -89,7 +90,9 @@ class JWTHandler:
self.leeway = leeway
@staticmethod
- def is_jwt(token: str):
+ def is_jwt(token: Optional[str]) -> bool:
+ if token is None:
+ return False
parts = token.split(".")
return len(parts) == 3
@@ -1324,6 +1327,7 @@ class JWTAuthManager:
jwt_valid_token: dict,
user_object: Optional[LiteLLM_UserTable],
prisma_client: Optional[PrismaClient],
+ user_api_key_cache: Optional[DualCache] = None,
) -> None:
"""
Sync user role and team memberships with JWT claims
@@ -1348,6 +1352,12 @@ class JWTAuthManager:
data={"user_role": new_role.value},
)
user_object.user_role = new_role.value
+ if user_api_key_cache is not None:
+ await user_api_key_cache.async_set_cache(
+ key=user_object.user_id,
+ value=user_object.model_dump(),
+ ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
+ )
# Sync team memberships
jwt_team_ids = set(jwt_handler.get_team_ids_from_jwt(jwt_valid_token))
@@ -1365,6 +1375,12 @@ class JWTAuthManager:
teams_ids_to_remove_user_from=list(teams_to_remove),
)
user_object.teams = list(jwt_team_ids)
+ if user_api_key_cache is not None:
+ await user_api_key_cache.async_set_cache(
+ key=user_object.user_id,
+ value=user_object.model_dump(),
+ ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
+ )
return None
@staticmethod
@@ -1536,6 +1552,7 @@ class JWTAuthManager:
jwt_valid_token=jwt_valid_token,
user_object=user_object,
prisma_client=prisma_client,
+ user_api_key_cache=user_api_key_cache,
)
## MAP USER TO TEAMS
diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py
index 53cc88e3b11..26bbdef3090 100644
--- a/litellm/proxy/auth/route_checks.py
+++ b/litellm/proxy/auth/route_checks.py
@@ -629,6 +629,7 @@ class RouteChecks:
in [
"/user/new",
"/user/delete",
+ "/user/bulk_update",
"/team/new",
"/team/update",
"/team/delete",
diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py
index eba787c63b3..9dd2ab18c8a 100644
--- a/litellm/proxy/auth/user_api_key_auth.py
+++ b/litellm/proxy/auth/user_api_key_auth.py
@@ -1303,9 +1303,21 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
team_member_info.litellm_budget_table.max_budget
)
if team_member_budget is not None and team_member_budget > 0:
- if valid_token.team_member_spend > team_member_budget:
+ # Read from cross-pod counter (Redis-first) if available
+ from litellm.proxy.proxy_server import get_current_spend
+
+ team_member_spend = valid_token.team_member_spend
+ if (
+ valid_token.user_id is not None
+ and valid_token.team_id is not None
+ ):
+ team_member_spend = await get_current_spend(
+ counter_key=f"spend:team_member:{valid_token.user_id}:{valid_token.team_id}",
+ fallback_spend=team_member_spend,
+ )
+ if team_member_spend > team_member_budget:
raise litellm.BudgetExceededError(
- current_cost=valid_token.team_member_spend,
+ current_cost=team_member_spend,
max_budget=team_member_budget,
)
diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py
index 674214b19e5..5db919659ed 100644
--- a/litellm/proxy/common_utils/reset_budget_job.py
+++ b/litellm/proxy/common_utils/reset_budget_job.py
@@ -54,16 +54,49 @@ class ResetBudgetJob:
"""
Resets the budget for all LiteLLM Team Members if their budget has expired
"""
+ budget_ids = [
+ budget.budget_id
+ for budget in budgets_to_reset
+ if budget.budget_id is not None
+ ]
+
+ # Reset spend counters for affected team members.
+ # Reset Redis directly so a transient failure doesn't leave stale
+ # counters that get_current_spend would read as authoritative.
+ try:
+ from litellm.proxy.proxy_server import spend_counter_cache
+
+ memberships = (
+ await self.prisma_client.db.litellm_teammembership.find_many(
+ where={"budget_id": {"in": budget_ids}}
+ )
+ )
+ for m in memberships:
+ counter_key = f"spend:team_member:{m.user_id}:{m.team_id}"
+ # Always reset in-memory
+ spend_counter_cache.in_memory_cache.set_cache(
+ key=counter_key, value=0.0
+ )
+ # Explicitly reset Redis with warning on failure
+ if spend_counter_cache.redis_cache is not None:
+ try:
+ await spend_counter_cache.redis_cache.async_set_cache(
+ key=counter_key, value=0.0
+ )
+ except Exception as redis_err:
+ verbose_proxy_logger.warning(
+ "Failed to reset team member spend counter in Redis %s: %s. "
+ "Budget may be over-enforced until counter expires.",
+ counter_key,
+ redis_err,
+ )
+ except Exception as e:
+ verbose_proxy_logger.warning(
+ "Failed to reset team member spend counters: %s", e
+ )
+
return await self.prisma_client.db.litellm_teammembership.update_many(
- where={
- "budget_id": {
- "in": [
- budget.budget_id
- for budget in budgets_to_reset
- if budget.budget_id is not None
- ]
- }
- },
+ where={"budget_id": {"in": budget_ids}},
data={
"spend": 0,
},
@@ -531,6 +564,39 @@ class ResetBudgetJob:
"""
try:
item.spend = 0.0
+
+ # Reset the cross-pod spend counter.
+ # Reset Redis directly (not via DualCache) so a Redis failure
+ # doesn't silently leave a stale counter that get_current_spend
+ # would read as authoritative, permanently blocking the user.
+ from litellm.proxy.proxy_server import spend_counter_cache
+
+ counter_key = None
+ if item_type == "key" and hasattr(item, "token") and item.token is not None:
+ counter_key = f"spend:key:{item.token}"
+ elif item_type == "team" and hasattr(item, "team_id") and item.team_id is not None:
+ counter_key = f"spend:team:{item.team_id}"
+
+ if counter_key is not None:
+ # Always reset in-memory (local fallback)
+ spend_counter_cache.in_memory_cache.set_cache(
+ key=counter_key, value=0.0
+ )
+ # Explicitly reset Redis with warning on failure
+ if spend_counter_cache.redis_cache is not None:
+ try:
+ await spend_counter_cache.redis_cache.async_set_cache(
+ key=counter_key, value=0.0
+ )
+ except Exception as redis_err:
+ verbose_proxy_logger.warning(
+ "Failed to reset spend counter in Redis for %s key=%s: %s. "
+ "Budget may be over-enforced until counter expires.",
+ item_type,
+ counter_key,
+ redis_err,
+ )
+
if hasattr(item, "budget_duration") and item.budget_duration is not None:
# Get standardized reset time based on budget duration
from litellm.proxy.common_utils.timezone_utils import (
diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py
index a8d0e3e9af2..3e05ee3c484 100644
--- a/litellm/proxy/health_check.py
+++ b/litellm/proxy/health_check.py
@@ -207,21 +207,65 @@ async def _perform_health_check(
for is_healthy, model in zip(results, model_list):
litellm_params = model["litellm_params"]
+ _model_id = (model.get("model_info") or {}).get("id")
if isinstance(is_healthy, dict) and "error" not in is_healthy:
- healthy_endpoints.append(
- _clean_endpoint_data({**litellm_params, **is_healthy}, details)
- )
+ cleaned = _clean_endpoint_data({**litellm_params, **is_healthy}, details)
+ if _model_id:
+ cleaned["model_id"] = _model_id
+ healthy_endpoints.append(cleaned)
elif isinstance(is_healthy, dict):
- unhealthy_endpoints.append(
- _clean_endpoint_data({**litellm_params, **is_healthy}, details)
- )
+ cleaned = _clean_endpoint_data({**litellm_params, **is_healthy}, details)
+ if _model_id:
+ cleaned["model_id"] = _model_id
+ unhealthy_endpoints.append(cleaned)
else:
- unhealthy_endpoints.append(_clean_endpoint_data(litellm_params, details))
+ cleaned = _clean_endpoint_data(litellm_params, details)
+ if _model_id:
+ cleaned["model_id"] = _model_id
+ unhealthy_endpoints.append(cleaned)
return healthy_endpoints, unhealthy_endpoints
+def build_deployment_health_states(
+ healthy_endpoints: list,
+ unhealthy_endpoints: list,
+) -> dict:
+ """
+ Build a dict mapping deployment_id -> DeploymentHealthStateValue from
+ health check endpoint results.
+
+ Each endpoint dict includes a 'model_id' field (added by _perform_health_check)
+ that maps back to the deployment's model_info.id.
+
+ Used by the background health check loop to feed health state into
+ the router's DeploymentHealthCache for health-check-driven routing.
+ """
+ now = time.time()
+ states: dict = {}
+
+ for ep in healthy_endpoints:
+ model_id = ep.get("model_id")
+ if model_id:
+ states[model_id] = {
+ "is_healthy": True,
+ "timestamp": now,
+ "reason": "",
+ }
+
+ for ep in unhealthy_endpoints:
+ model_id = ep.get("model_id")
+ if model_id:
+ states[model_id] = {
+ "is_healthy": False,
+ "timestamp": now,
+ "reason": "background_health_check_failed",
+ }
+
+ return states
+
+
def _update_litellm_params_for_health_check(
model_info: dict, litellm_params: dict
) -> dict:
diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py
index 220c5066a6d..948d6dd33af 100644
--- a/litellm/proxy/hooks/proxy_track_cost_callback.py
+++ b/litellm/proxy/hooks/proxy_track_cost_callback.py
@@ -135,7 +135,11 @@ class _ProxyDBLogger(CustomLogger):
start_time=None,
end_time=None, # start/end time for completion
):
- from litellm.proxy.proxy_server import proxy_logging_obj, update_cache
+ from litellm.proxy.proxy_server import (
+ increment_spend_counters,
+ proxy_logging_obj,
+ update_cache,
+ )
verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback")
try:
@@ -194,7 +198,17 @@ class _ProxyDBLogger(CustomLogger):
org_id=org_id,
)
- # update cache
+ # Atomically update spend counters (in-memory + Redis)
+ # for cross-pod budget enforcement.
+ await increment_spend_counters(
+ token=user_api_key,
+ team_id=team_id,
+ user_id=user_id,
+ response_cost=response_cost,
+ )
+
+ # update cache (fire-and-forget for backward compat:
+ # cached object fields, soft budget alerts, etc.)
asyncio.create_task(
update_cache(
token=user_api_key,
diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py
index 4ca0d876a1c..ba9577f35d7 100644
--- a/litellm/proxy/litellm_pre_call_utils.py
+++ b/litellm/proxy/litellm_pre_call_utils.py
@@ -1,6 +1,7 @@
import asyncio
import copy
import time
+from collections import OrderedDict
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from fastapi import Request
@@ -26,6 +27,7 @@ _SPECIAL_HEADERS_CACHE = frozenset(
v.value.lower() for v in SpecialHeaders._member_map_.values()
)
from litellm.router import Router
+from litellm.secret_managers.main import get_secret_bool
from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS
from litellm.types.services import ServiceTypes
from litellm.types.utils import (
@@ -36,6 +38,11 @@ from litellm.types.utils import (
)
service_logger_obj = ServiceLogging() # used for tracking latency on OTEL
+# Bounded dedup for stale-alias warnings (FIFO eviction when over cap).
+_MAX_STALE_ALIAS_WARNING_KEYS = 10_000
+_STALE_TEAM_ALIAS_WARNING_KEYS: OrderedDict[str, None] = OrderedDict()
+# Cache the stale alias bypass flag at module load to avoid hot-path secret lookups
+_ENABLE_TEAM_STALE_ALIAS_BYPASS: Optional[bool] = None
if TYPE_CHECKING:
@@ -1296,6 +1303,10 @@ def _update_model_if_team_alias_exists(
"gpt-4o": "gpt-4o-team-1"
}
- requested_model = "gpt-4o-team-1"
+
+ Note: model_aliases for team models are deprecated. This function only applies
+ to legacy non-team-scoped aliases. Team-scoped deployments use team_public_model_name
+ and are resolved via map_team_model in route_llm_request.
"""
_model = data.get("model")
if (
@@ -1303,7 +1314,54 @@ def _update_model_if_team_alias_exists(
and user_api_key_dict.team_model_aliases
and _model in user_api_key_dict.team_model_aliases
):
- data["model"] = user_api_key_dict.team_model_aliases[_model]
+ from litellm.proxy.proxy_server import llm_router
+
+ # Skip alias rewrite if this model resolves to team-specific deployments
+ # (team models use team_public_model_name, not model_aliases)
+ aliased_target = user_api_key_dict.team_model_aliases[_model]
+
+ # Optional bypass for stale aliases from pre-PR deployments:
+ # only enabled via feature flag to preserve backwards compatibility.
+ # Cached at module level to avoid hot-path secret lookups on every request.
+ global _ENABLE_TEAM_STALE_ALIAS_BYPASS
+ if _ENABLE_TEAM_STALE_ALIAS_BYPASS is None:
+ _ENABLE_TEAM_STALE_ALIAS_BYPASS = get_secret_bool(
+ "LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", False
+ )
+ enable_stale_alias_bypass = _ENABLE_TEAM_STALE_ALIAS_BYPASS
+ # Check if the alias points to a team-scoped UUID name
+ # (format: "model_name_{team_id}_{uuid}")
+ is_stale_team_alias = aliased_target.startswith(
+ f"model_name_{user_api_key_dict.team_id}_"
+ )
+ if is_stale_team_alias and llm_router:
+ # This is a stale alias from pre-PR deployments.
+ # Check if current team deployments exist for the public name.
+ key = (user_api_key_dict.team_id, _model)
+ if key in llm_router.team_model_to_deployment_indices:
+ if enable_stale_alias_bypass:
+ # Team deployments exist; skip stale alias
+ return
+ warning_key = f"{user_api_key_dict.team_id}:{_model}:{aliased_target}"
+ if warning_key not in _STALE_TEAM_ALIAS_WARNING_KEYS:
+ _STALE_TEAM_ALIAS_WARNING_KEYS[warning_key] = None
+ while (
+ len(_STALE_TEAM_ALIAS_WARNING_KEYS)
+ > _MAX_STALE_ALIAS_WARNING_KEYS
+ ):
+ _STALE_TEAM_ALIAS_WARNING_KEYS.popitem(last=False)
+ verbose_proxy_logger.warning(
+ "Stale team model alias detected for model='%s', team_id='%s'. "
+ "New sibling deployments may be unreachable. "
+ "Set LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS=true to enable "
+ "team-scoped sibling routing.",
+ str(_model).replace("\n", "").replace("\r", ""),
+ str(user_api_key_dict.team_id)
+ .replace("\n", "")
+ .replace("\r", ""),
+ )
+
+ data["model"] = aliased_target
return
diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py
index 20c7f9ec412..37f13269b1c 100644
--- a/litellm/proxy/management_endpoints/budget_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py
@@ -15,6 +15,7 @@ All /budget management endpoints
from datetime import timedelta
from fastapi import APIRouter, Depends, HTTPException
+from prisma.errors import UniqueViolationError
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy._types import *
@@ -90,13 +91,21 @@ async def new_budget(
budget_obj_json = budget_obj.model_dump(exclude_none=True)
budget_obj_jsonified = jsonify_object(budget_obj_json) # json dump any dictionaries
- response = await prisma_client.db.litellm_budgettable.create(
- data={
- **budget_obj_jsonified, # type: ignore
- "created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
- "updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
- } # type: ignore
- )
+ try:
+ response = await prisma_client.db.litellm_budgettable.create(
+ data={
+ **budget_obj_jsonified, # type: ignore
+ "created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
+ "updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
+ } # type: ignore
+ )
+ except UniqueViolationError:
+ raise HTTPException(
+ status_code=400,
+ detail={
+ "error": f"Budget with id '{budget_obj.budget_id}' already exists."
+ },
+ )
return response
diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py
index ca8c345f46c..7b459cd8502 100644
--- a/litellm/proxy/management_endpoints/internal_user_endpoints.py
+++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py
@@ -6,6 +6,7 @@ These are members of a Team on LiteLLM
/user/new
/user/update
+/user/bulk_update
/user/delete
/user/info
/user/list
@@ -24,13 +25,13 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.proxy._types import *
+from litellm.proxy.auth.auth_checks import get_team_object, get_user_object
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventHooks
from litellm.proxy.management_endpoints.common_daily_activity import (
get_daily_activity,
get_daily_activity_aggregated,
)
-from litellm.proxy.auth.auth_checks import get_team_object, get_user_object
from litellm.proxy.management_endpoints.common_utils import (
_is_user_team_admin,
_user_has_admin_view,
@@ -557,6 +558,18 @@ def get_team_from_list(
return None
+def _is_valid_user_id(user_id: str) -> bool:
+ """Validate that a decoded user_id is safe to use downstream."""
+ MAX_USER_ID_LENGTH = 512
+ if len(user_id) > MAX_USER_ID_LENGTH:
+ return False
+ # Reject ASCII control characters (U+0000–U+001F)
+ for ch in user_id:
+ if ord(ch) < 0x20:
+ return False
+ return True
+
+
def get_user_id_from_request(request: Request) -> Optional[str]:
"""
Get the user id from the request
@@ -573,7 +586,8 @@ def get_user_id_from_request(request: Request) -> Optional[str]:
if match:
# Use unquote instead of unquote_plus to preserve + characters
raw_user_id = unquote(match.group(1))
- user_id = raw_user_id
+ if _is_valid_user_id(raw_user_id):
+ user_id = raw_user_id
return user_id
diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py
index 44d41097833..2ca8e3daba3 100644
--- a/litellm/proxy/management_endpoints/model_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/model_management_endpoints.py
@@ -13,13 +13,13 @@ model/{model_id}/update - PATCH endpoint for model update.
import asyncio
import datetime
import json
-from litellm._uuid import uuid
from typing import Dict, List, Literal, Optional, Tuple, Union, cast
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel, ConfigDict, Field
from litellm._logging import verbose_proxy_logger
+from litellm._uuid import uuid
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
from litellm.proxy._types import (
CommonProxyErrors,
@@ -32,7 +32,7 @@ from litellm.proxy._types import (
ProxyErrorTypes,
ProxyException,
TeamModelAddRequest,
- UpdateTeamRequest,
+ TeamModelDeleteRequest,
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@@ -40,7 +40,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helpe
from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin
from litellm.proxy.management_endpoints.team_endpoints import (
team_model_add,
- update_team,
+ team_model_delete,
+)
+from litellm.proxy.management_endpoints.team_endpoints import (
+ update_team as _legacy_update_team,
)
from litellm.proxy.management_helpers.audit_logs import create_object_audit_log
from litellm.proxy.utils import PrismaClient
@@ -58,6 +61,14 @@ from litellm.utils import get_utc_datetime
router = APIRouter()
+async def update_team(*args, **kwargs):
+ """
+ Backward-compatible shim for tests/legacy call sites that patch this symbol.
+ Team model management now uses team_model_add/team_model_delete directly.
+ """
+ return await _legacy_update_team(*args, **kwargs)
+
+
class UpdatePublicModelGroupsRequest(BaseModel):
"""Request model for updating public model groups"""
@@ -324,17 +335,24 @@ async def _add_team_model_to_db(
- generate a unique 'model_name' for the model (e.g. 'model_name_{team_id}_{uuid})
- store the model in the db with the unique 'model_name'
- - store a team model alias mapping {"model_name": "model_name_{team_id}_{uuid}"}
+ - add the public model name to the team's allowed models list
"""
_team_id = model_params.model_info.team_id
if _team_id is None:
return None
+
+ # Capture the original public name FIRST, before any mutations
original_model_name = model_params.model_name
+
+ # Set team_public_model_name in model_info using the captured original_model_name
+ # This must happen BEFORE mutating model_params.model_name so _add_model_to_db
+ # serializes the correct team_public_model_name (not the internal UUID name)
if original_model_name:
model_params.model_info.team_public_model_name = original_model_name
+ # Generate and assign unique internal model_name LAST
+ # (after team_public_model_name is safely stored)
unique_model_name = f"model_name_{_team_id}_{uuid.uuid4()}"
-
model_params.model_name = unique_model_name
## CREATE MODEL IN DB ##
@@ -344,25 +362,15 @@ async def _add_team_model_to_db(
prisma_client=prisma_client,
)
- ## CREATE MODEL ALIAS IN DB ##
- await update_team(
- data=UpdateTeamRequest(
- team_id=_team_id,
- model_aliases={original_model_name: unique_model_name},
- ),
- user_api_key_dict=user_api_key_dict,
- http_request=Request(scope={"type": "http"}),
- )
-
- # add model to team object
- await team_model_add(
- data=TeamModelAddRequest(
- team_id=_team_id,
- models=[original_model_name],
- ),
- http_request=Request(scope={"type": "http"}),
- user_api_key_dict=user_api_key_dict,
- )
+ if original_model_name:
+ await team_model_add(
+ data=TeamModelAddRequest(
+ team_id=_team_id,
+ models=[original_model_name],
+ ),
+ http_request=Request(scope={"type": "http"}),
+ user_api_key_dict=user_api_key_dict,
+ )
return model_response
@@ -428,6 +436,7 @@ async def _update_team_model_in_db(
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
+ prisma_client=prisma_client,
)
return update_db_model(db_model=db_model, updated_patch=patch_data)
@@ -453,19 +462,10 @@ async def _setup_new_team_model_assignment(
patch_data: updateDeployment,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
- """Set up a new team model with unique name, alias, and team membership."""
+ """Set up a new team model with unique name and team membership."""
unique_model_name = f"model_name_{team_id}_{uuid.uuid4()}"
patch_data.model_name = unique_model_name
- await update_team(
- data=UpdateTeamRequest(
- team_id=team_id,
- model_aliases={public_model_name: unique_model_name},
- ),
- user_api_key_dict=user_api_key_dict,
- http_request=Request(scope={"type": "http"}),
- )
-
await team_model_add(
data=TeamModelAddRequest(
team_id=team_id,
@@ -476,30 +476,119 @@ async def _setup_new_team_model_assignment(
)
+async def _get_team_deployments(
+ team_id: str, prisma_client: PrismaClient
+) -> List[LiteLLM_ProxyModelTable]:
+ """
+ Fetch all deployments for a given team_id from the database.
+
+ Centralizes team deployment queries to ensure consistent filtering and error handling.
+ This is the established helper pattern for team deployment DB access in this module.
+
+ Note: Direct Prisma call is intentional here as this IS the helper function that
+ encapsulates the DB access pattern for team deployments.
+ """
+ response = await prisma_client.db.litellm_proxymodeltable.find_many(
+ where={
+ "model_info": {
+ "path": ["team_id"],
+ "equals": team_id,
+ }
+ }
+ )
+ return response if response else []
+
+
async def _update_existing_team_model_assignment(
team_id: str,
public_model_name: str,
db_model: Deployment,
patch_data: updateDeployment,
user_api_key_dict: UserAPIKeyAuth,
+ prisma_client: Optional[PrismaClient],
) -> None:
- """Update an existing team model if the public name changed."""
+ """Update an existing team model if the public name changed.
+
+ Note on DB scan: Prisma's JSON filtering does not support compound AND conditions
+ across multiple JSON paths, so we fetch all deployments for the team and filter
+ team_public_model_name in Python. For teams with many deployments this scan grows
+ linearly; if team deployment counts become large this should be revisited.
+ """
+
+ def _get_team_public_model_name(
+ model_info: Optional[Union[dict, str]]
+ ) -> Optional[str]:
+ if isinstance(model_info, dict):
+ value = model_info.get("team_public_model_name")
+ return value if isinstance(value, str) else None
+ if isinstance(model_info, str):
+ try:
+ parsed = json.loads(model_info)
+ except (TypeError, ValueError):
+ return None
+ if isinstance(parsed, dict):
+ value = parsed.get("team_public_model_name")
+ return value if isinstance(value, str) else None
+ return None
+
old_public_name = (
db_model.model_info.team_public_model_name if db_model.model_info else None
)
- # Update alias only if public name changed
if old_public_name and public_model_name != old_public_name:
- await update_team(
- data=UpdateTeamRequest(
+ # Clear user-supplied public name from patch before any early return so the
+ # caller does not overwrite the internal UUID-based model_name in the DB.
+ patch_data.model_name = None
+ if prisma_client is None:
+ verbose_proxy_logger.warning(
+ "prisma_client not initialized; skipping public name update entirely to avoid orphaned entries"
+ )
+ return
+
+ # Query DB for all team deployments to check for sibling deployments
+ team_deployments = await _get_team_deployments(team_id, prisma_client)
+ other_deployments_with_old_name = [
+ d
+ for d in team_deployments
+ if d.model_name != db_model.model_name
+ and _get_team_public_model_name(d.model_info) == old_public_name
+ ]
+
+ # Add new name first, then delete old name to prevent access loss on partial failure
+ await team_model_add(
+ data=TeamModelAddRequest(
team_id=team_id,
- model_aliases={public_model_name: db_model.model_name},
+ models=[public_model_name],
),
- user_api_key_dict=user_api_key_dict,
http_request=Request(scope={"type": "http"}),
+ user_api_key_dict=user_api_key_dict,
)
- # Keep existing unique model_name
+ if not other_deployments_with_old_name:
+ await team_model_delete(
+ data=TeamModelDeleteRequest(
+ team_id=team_id,
+ models=[old_public_name],
+ ),
+ http_request=Request(scope={"type": "http"}),
+ user_api_key_dict=user_api_key_dict,
+ )
+ elif not old_public_name and public_model_name:
+ # First-time assignment of public name on an existing team deployment:
+ # ensure the team's models list is updated so team routing can resolve it.
+ await team_model_add(
+ data=TeamModelAddRequest(
+ team_id=team_id,
+ models=[public_model_name],
+ ),
+ http_request=Request(scope={"type": "http"}),
+ user_api_key_dict=user_api_key_dict,
+ )
+ # else: old_public_name == public_model_name (no rename needed)
+ # No team_model_add/delete calls required; public name is already registered
+
+ # Always clear patch_data.model_name to prevent caller from overwriting
+ # the internal UUID-based model_name in the DB with the user-supplied public name
patch_data.model_name = None
diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py
index 79a9e9bdf2b..a5a62813b9f 100644
--- a/litellm/proxy/management_endpoints/ui_sso.py
+++ b/litellm/proxy/management_endpoints/ui_sso.py
@@ -204,7 +204,7 @@ def process_sso_jwt_access_token(
sso_jwt_handler: Optional[JWTHandler],
result: Union[OpenID, dict, None],
role_mappings: Optional["RoleMappings"] = None,
-) -> None:
+) -> Optional[dict]:
"""
Process SSO JWT access token and extract team IDs and user role if available.
@@ -218,6 +218,12 @@ def process_sso_jwt_access_token(
sso_jwt_handler: SSO-specific JWT handler for team ID extraction
result: The SSO result object to update with team IDs and role
role_mappings: Optional role mappings configuration for group-based role determination
+
+ Returns:
+ The decoded access token payload dict, or None if decoding failed or
+ inputs were missing. Callers can pass this to _sync_user_role_from_jwt_role_map
+ so it has access to custom role claims (e.g. custom_roles) that are
+ encoded inside the JWT but stripped from received_response.
"""
if access_token_str and result:
import jwt
@@ -230,7 +236,7 @@ def process_sso_jwt_access_token(
verbose_proxy_logger.debug(
"Access token is not a valid JWT (possibly an opaque token), skipping JWT-based extraction"
)
- return
+ return None
# Extract team IDs from access token if sso_jwt_handler is available
if sso_jwt_handler:
@@ -306,6 +312,10 @@ def process_sso_jwt_access_token(
f"Set user_role='{user_role}' from JWT access token"
)
+ return access_token_payload
+
+ return None
+
@router.get("/sso/key/generate", tags=["experimental"], include_in_schema=False)
async def google_login(
@@ -817,7 +827,7 @@ async def get_generic_sso_response(
], # sso specific jwt handler - used for restricted sso group access control
generic_client_id: str,
redirect_url: str,
-) -> Tuple[Union[OpenID, dict], Optional[dict]]: # return received response
+) -> Tuple[Union[OpenID, dict], Optional[dict], Optional[dict]]: # (result, received_response, access_token_payload)
# make generic sso provider
from fastapi_sso.sso.base import DiscoveryDocument
from fastapi_sso.sso.generic import create_provider
@@ -872,6 +882,7 @@ async def get_generic_sso_response(
code_verifier: Optional[
str
] = None # assigned inside try; initialized for type tracking
+ access_token_payload: Optional[dict] = None # decoded JWT access token claims
try:
token_exchange_params = (
@@ -958,7 +969,7 @@ async def get_generic_sso_response(
)
access_token_str = generic_sso.access_token
- process_sso_jwt_access_token(
+ access_token_payload = process_sso_jwt_access_token(
access_token_str, sso_jwt_handler, result, role_mappings=role_mappings
)
# Delete the single-use PKCE verifier only after all downstream processing
@@ -976,7 +987,7 @@ async def get_generic_sso_response(
additional_generic_sso_headers_dict,
)
verbose_proxy_logger.debug("generic result: %s", result)
- return result or {}, received_response
+ return result or {}, received_response, access_token_payload
async def create_team_member_add_task(team_id, user_info):
@@ -1176,6 +1187,56 @@ def _build_sso_user_update_data(
return update_data
+async def _sync_user_role_from_jwt_role_map(
+ jwt_handler: Optional[JWTHandler],
+ received_response: Optional[dict],
+ user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]],
+ prisma_client: PrismaClient,
+ user_api_key_cache: DualCache,
+ user_defined_values: Optional[SSOUserDefinedValues],
+) -> None:
+ """
+ Apply jwt_litellm_role_map during SSO login.
+
+ When jwt_litellm_role_map is configured with sync_user_role_and_teams=True,
+ this ensures SSO users get the same role mapping as API/JWT users. Without
+ this, the SSO path falls back to INTERNAL_USER_VIEW_ONLY for roles that
+ don't directly match LitellmUserRoles enum values.
+ """
+ if jwt_handler is None or received_response is None:
+ return
+ if not jwt_handler.litellm_jwtauth.sync_user_role_and_teams:
+ return
+ if not jwt_handler.litellm_jwtauth.jwt_litellm_role_map:
+ return
+
+ mapped_role = jwt_handler.map_jwt_role_to_litellm_role(received_response)
+ if mapped_role is None:
+ return
+
+ verbose_proxy_logger.info(
+ f"SSO jwt_litellm_role_map matched role: {mapped_role.value}"
+ )
+
+ # Update user_defined_values so downstream code uses the mapped role
+ if user_defined_values is not None:
+ user_defined_values["user_role"] = mapped_role.value
+
+ # Update existing DB record if role differs
+ if user_info is not None and user_info.user_role != mapped_role.value:
+ await prisma_client.db.litellm_usertable.update(
+ where={"user_id": user_info.user_id},
+ data={"user_role": mapped_role.value},
+ )
+ user_info.user_role = mapped_role.value
+ await user_api_key_cache.async_set_cache(
+ key=user_info.user_id,
+ value=user_info.model_dump()
+ if hasattr(user_info, "model_dump")
+ else dict(user_info),
+ )
+
+
def apply_user_info_values_to_sso_user_defined_values(
user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]],
user_defined_values: Optional[SSOUserDefinedValues],
@@ -1279,6 +1340,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
google_client_id = os.getenv("GOOGLE_CLIENT_ID", None)
generic_client_id = os.getenv("GENERIC_CLIENT_ID", None)
received_response: Optional[dict] = None
+ access_token_payload: Optional[dict] = None
# get url from request
if master_key is None:
raise ProxyException(
@@ -1307,7 +1369,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
)
elif generic_client_id is not None:
- result, received_response = await get_generic_sso_response(
+ result, received_response, access_token_payload = await get_generic_sso_response(
request=request,
jwt_handler=jwt_handler,
generic_client_id=generic_client_id,
@@ -1345,6 +1407,8 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
received_response=received_response,
generic_client_id=generic_client_id,
ui_access_mode=ui_access_mode,
+ access_token_payload=access_token_payload,
+ jwt_handler=jwt_handler,
return_to=cp_return_to,
)
@@ -2417,6 +2481,8 @@ class SSOAuthenticationHandler:
received_response: Optional[dict] = None,
generic_client_id: Optional[str] = None,
ui_access_mode: Optional[Dict] = None,
+ access_token_payload: Optional[dict] = None,
+ jwt_handler: Optional[JWTHandler] = None,
return_to: Optional[str] = None,
) -> RedirectResponse:
import jwt
@@ -2498,6 +2564,20 @@ class SSOAuthenticationHandler:
alternate_user_id=user_id,
)
+ # Sync user role from JWT claims via jwt_litellm_role_map (if configured).
+ # This ensures SSO users get the same role mapping as API/JWT users.
+ # Use the decoded access_token_payload (not received_response) because
+ # custom role claims (e.g. custom_roles) are encoded inside the JWT
+ # access token, which is stripped from received_response.
+ await _sync_user_role_from_jwt_role_map(
+ jwt_handler=jwt_handler,
+ received_response=access_token_payload or received_response,
+ user_info=user_info,
+ prisma_client=prisma_client,
+ user_api_key_cache=user_api_key_cache,
+ user_defined_values=user_defined_values,
+ )
+
user_defined_values = apply_user_info_values_to_sso_user_defined_values(
user_info=user_info, user_defined_values=user_defined_values
)
@@ -3703,7 +3783,7 @@ async def debug_sso_callback(request: Request):
)
elif generic_client_id is not None:
- result, _ = await get_generic_sso_response(
+ result, _, _ = await get_generic_sso_response(
request=request,
jwt_handler=jwt_handler,
generic_client_id=generic_client_id,
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index dc3d1c54fdd..ab73d44acca 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -480,11 +480,11 @@ from litellm.proxy.search_endpoints.search_tool_management import (
router as search_tool_management_router,
)
from litellm.proxy.spend_tracking.cloudzero_endpoints import router as cloudzero_router
-from litellm.proxy.spend_tracking.vantage_endpoints import router as vantage_router
from litellm.proxy.spend_tracking.spend_management_endpoints import (
router as spend_management_router,
)
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
+from litellm.proxy.spend_tracking.vantage_endpoints import router as vantage_router
from litellm.proxy.types_utils.utils import get_instance_fn
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
router as ui_crud_endpoints_router,
@@ -1530,6 +1530,9 @@ shared_aiohttp_session: Optional[
user_api_key_cache = DualCache(
default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value
)
+spend_counter_cache = DualCache(
+ default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value
+)
model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(
dual_cache=user_api_key_cache
)
@@ -1694,6 +1697,134 @@ def cost_tracking():
)
+async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
+ """
+ Read current spend from the cross-pod spend counter.
+
+ Reads Redis FIRST (authoritative cross-pod value), not DualCache's
+ async_get_cache which returns in-memory first. This is critical:
+ DualCache.async_get_cache returns stale per-pod values because each
+ pod's in-memory cache is only updated by that pod's own increments.
+
+ Fallback chain:
+ 1. Redis counter (cross-pod, authoritative)
+ 2. In-memory counter (single-instance or Redis failure)
+ 3. Cached object's .spend from DB (cold start, no counter yet)
+ """
+ # 1. Try Redis first (cross-pod authoritative)
+ if spend_counter_cache.redis_cache is not None:
+ try:
+ val = await spend_counter_cache.redis_cache.async_get_cache(
+ key=counter_key
+ )
+ if val is not None:
+ return float(val)
+ except Exception as e:
+ verbose_proxy_logger.debug(
+ "get_current_spend: Redis read failed for %s, falling back to in-memory: %s",
+ counter_key,
+ e,
+ )
+
+ # 2. Fall back to in-memory counter (single-instance or Redis failure)
+ val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
+ if val is not None:
+ return float(val)
+
+ # 3. Final fallback: cached object's spend from DB
+ return fallback_spend
+
+
+async def increment_spend_counters(
+ token: Optional[str],
+ team_id: Optional[str],
+ user_id: Optional[str],
+ response_cost: Optional[float],
+):
+ """
+ Atomically increment spend counters for budget enforcement.
+
+ Uses spend_counter_cache (DualCache with Redis backend when available)
+ so counters are shared across all pods. Budget check functions read
+ from these counters via get_current_spend() (Redis-first).
+
+ Awaited (not create_task) in the cost callback, so the counter is
+ updated before the next request's auth check runs.
+ """
+ if response_cost is None or response_cost == 0:
+ return
+
+ if token is not None:
+ # token arrives pre-hashed from metadata["user_api_key"] (auth flow
+ # hashes raw "sk-..." keys before they reach the callback). The
+ # startswith("sk-") check is a safety net matching update_cache —
+ # if a raw key somehow arrives, hash it; otherwise use as-is to
+ # avoid double-hashing (budget checks read valid_token.token which
+ # is single-hashed).
+ hashed_token = (
+ hash_token(token=token)
+ if isinstance(token, str) and token.startswith("sk-")
+ else token
+ )
+ await _init_and_increment_spend_counter(
+ counter_key=f"spend:key:{hashed_token}",
+ source_cache_key=hashed_token,
+ increment=response_cost,
+ )
+
+ if team_id is not None:
+ await _init_and_increment_spend_counter(
+ counter_key=f"spend:team:{team_id}",
+ source_cache_key=f"team_id:{team_id}",
+ increment=response_cost,
+ )
+
+ if user_id is not None and team_id is not None:
+ await _init_and_increment_spend_counter(
+ counter_key=f"spend:team_member:{user_id}:{team_id}",
+ source_cache_key=f"team_membership:{user_id}:{team_id}",
+ increment=response_cost,
+ )
+
+
+async def _init_and_increment_spend_counter(
+ counter_key: str,
+ source_cache_key: str,
+ increment: float,
+):
+ """
+ Initialize counter from cached object's DB-loaded spend if not yet set,
+ then atomically increment in both in-memory and Redis.
+
+ On first access per pod:
+ 1. Check spend_counter_cache (in-memory -> Redis via DualCache for init check)
+ 2. If not found anywhere, read base spend from user_api_key_cache (DB-loaded object)
+ 3. Seed counter via async_increment_cache (not async_set_cache) to avoid a
+ check-then-set race: if two pods cold-start simultaneously, both may see
+ the counter as absent and seed it. Using increment instead of set means
+ the worst case is over-counting (conservative — blocks slightly early)
+ rather than under-counting (would allow overspend).
+ 4. Increment atomically (both in-memory + Redis)
+ """
+ current = await spend_counter_cache.async_get_cache(key=counter_key)
+ if current is None:
+ source = await user_api_key_cache.async_get_cache(key=source_cache_key)
+ base_spend = 0.0
+ if source is not None:
+ if isinstance(source, dict):
+ base_spend = source.get("spend", 0.0) or 0.0
+ else:
+ base_spend = getattr(source, "spend", 0.0) or 0.0
+ if base_spend > 0:
+ await spend_counter_cache.async_increment_cache(
+ key=counter_key, value=base_spend
+ )
+
+ await spend_counter_cache.async_increment_cache(
+ key=counter_key, value=increment
+ )
+
+
async def update_cache( # noqa: PLR0915
token: Optional[str],
user_id: Optional[str],
@@ -2112,6 +2243,37 @@ def _schedule_background_health_check_db_save(
)
+def _write_health_state_to_router_cache(
+ healthy_endpoints: list,
+ unhealthy_endpoints: list,
+) -> None:
+ """
+ Write deployment health states to the router's health state cache
+ for health-check-driven routing. No-op if the feature is disabled.
+ """
+ from litellm.proxy.health_check import build_deployment_health_states
+
+ try:
+ if llm_router is None or not llm_router.enable_health_check_routing:
+ return
+
+ states = build_deployment_health_states(
+ healthy_endpoints=healthy_endpoints,
+ unhealthy_endpoints=unhealthy_endpoints,
+ )
+ if states:
+ llm_router.health_state_cache.set_deployment_health_states(states)
+ verbose_proxy_logger.debug(
+ "health_check_routing_state_updated healthy=%d unhealthy=%d",
+ sum(1 for s in states.values() if s.get("is_healthy")),
+ sum(1 for s in states.values() if not s.get("is_healthy")),
+ )
+ except Exception as e:
+ verbose_proxy_logger.warning(
+ "Failed to write health state to router cache: %s", str(e)
+ )
+
+
async def _run_background_health_check():
"""
Periodically run health checks in the background on the endpoints.
@@ -2281,6 +2443,9 @@ async def _run_background_health_check():
unhealthy_endpoints,
)
+ # Write health state to router cache for health-check-driven routing
+ _write_health_state_to_router_cache(healthy_endpoints, unhealthy_endpoints)
+
await asyncio.sleep(health_check_interval)
@@ -2528,6 +2693,7 @@ class ProxyConfig:
):
## INIT PROXY REDIS USAGE CLIENT ##
redis_usage_cache = litellm.cache.cache
+ spend_counter_cache.redis_cache = redis_usage_cache
# Note: PKCE verifier storage uses redis_usage_cache directly (not
# user_api_key_cache) to avoid routing all API-key lookups through Redis.
@@ -2695,12 +2861,36 @@ class ProxyConfig:
return search_tools_parsed if search_tools_parsed else None
+ # Environment variable keys that must not be overridden via config because
+ # they can alter process execution, library loading, or network routing.
+ _BLOCKED_ENV_KEYS: Set[str] = {
+ "PATH",
+ "LD_PRELOAD",
+ "LD_LIBRARY_PATH",
+ "DYLD_LIBRARY_PATH",
+ "DYLD_INSERT_LIBRARIES",
+ "PYTHONPATH",
+ "PYTHONSTARTUP",
+ "PYTHONHOME",
+ "HOME",
+ "USER",
+ "SHELL",
+ "LOGNAME",
+ "NO_PROXY",
+ "no_proxy",
+ }
+
def _load_environment_variables(self, config: dict):
## ENVIRONMENT VARIABLES
global premium_user
environment_variables = config.get("environment_variables", None)
if environment_variables:
for key, value in environment_variables.items():
+ if key in self._BLOCKED_ENV_KEYS:
+ verbose_proxy_logger.warning(
+ "Skipping blocked environment variable key: %s", key
+ )
+ continue
#########################################################
# handles this scenario:
# ```yaml
@@ -3048,6 +3238,8 @@ class ProxyConfig:
general_settings = config.get("general_settings", {})
if general_settings is None:
general_settings = {}
+ _enable_hc_routing = False
+ _hc_staleness = None
if general_settings:
### LOAD KEY MANAGEMENT SETTINGS FIRST (needed for custom secret manager) ###
key_management_settings = general_settings.get(
@@ -3227,13 +3419,21 @@ class ProxyConfig:
"health_check_concurrency", None
)
health_check_details = general_settings.get("health_check_details", True)
+ # Health-check-driven routing (opt-in, passes through to Router later)
+ _enable_hc_routing = general_settings.get(
+ "enable_health_check_routing", False
+ )
+ _hc_staleness = general_settings.get(
+ "health_check_staleness_threshold", None
+ )
verbose_proxy_logger.info(
- "background_health_check_config enabled=%s shared=%s interval_seconds=%s max_concurrency=%s details=%s",
+ "background_health_check_config enabled=%s shared=%s interval_seconds=%s max_concurrency=%s details=%s health_check_routing=%s",
use_background_health_checks,
use_shared_health_check,
health_check_interval,
health_check_concurrency,
health_check_details,
+ _enable_hc_routing,
)
### RBAC ###
@@ -3263,6 +3463,11 @@ class ProxyConfig:
"cache_responses": litellm.cache
is not None, # cache if user passed in cache values
}
+ # Health-check-driven routing params (from general_settings)
+ if _enable_hc_routing:
+ router_params["enable_health_check_routing"] = True
+ if _hc_staleness is not None:
+ router_params["health_check_staleness_threshold"] = _hc_staleness
## MODEL LIST
model_list = config.get("model_list", None)
if model_list:
diff --git a/litellm/router.py b/litellm/router.py
index 25e5c9cb5d9..6cc6bad9def 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -54,7 +54,11 @@ from litellm.caching.caching import (
RedisCache,
RedisClusterCache,
)
-from litellm.constants import DEFAULT_MAX_LRU_CACHE_SIZE
+from litellm.constants import (
+ DEFAULT_HEALTH_CHECK_INTERVAL,
+ DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER,
+ DEFAULT_MAX_LRU_CACHE_SIZE,
+)
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.core_helpers import (
@@ -113,6 +117,7 @@ from litellm.router_utils.handle_error import (
async_raise_no_deployment_exception,
send_llm_exception_alert,
)
+from litellm.router_utils.health_state_cache import DeploymentHealthCache
from litellm.router_utils.pre_call_checks.deployment_affinity_check import (
DeploymentAffinityCheck,
)
@@ -303,6 +308,8 @@ class Router:
deployment_affinity_ttl_seconds: int = 3600,
model_group_affinity_config: Optional[Dict[str, List[str]]] = None,
ignore_invalid_deployments: bool = False,
+ enable_health_check_routing: bool = False,
+ health_check_staleness_threshold: Optional[int] = None,
) -> None:
"""
Initialize the Router class with the given parameters for caching, reliability, and routing strategy.
@@ -467,6 +474,8 @@ class Router:
# Initialize model name to deployment indices mapping for O(1) lookups
# Maps model_name -> list of indices in model_list
self.model_name_to_deployment_indices: Dict[str, List[int]] = {}
+ # Maps (team_id, team_public_model_name) -> list of indices in model_list
+ self.team_model_to_deployment_indices: Dict[Tuple[str, str], List[int]] = {}
if model_list is not None:
# set_model_list will build indices automatically
@@ -491,6 +500,13 @@ class Router:
cache=self.cache, default_cooldown_time=self.cooldown_time
)
self.disable_cooldowns = disable_cooldowns
+ self.enable_health_check_routing = enable_health_check_routing
+ _staleness = health_check_staleness_threshold or (
+ DEFAULT_HEALTH_CHECK_INTERVAL * DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER
+ )
+ self.health_state_cache = DeploymentHealthCache(
+ cache=self.cache, staleness_threshold=float(_staleness)
+ )
self.failed_calls = (
InMemoryCache()
) # cache to track failed call per deployment, if num failed calls within 1 minute > allowed fails, then add it to cooldown
@@ -5288,6 +5304,64 @@ class Router:
if "fallback_depth" not in input_kwargs:
input_kwargs["fallback_depth"] = 0
+ # ORDER-BASED FALLBACKS: prepend higher order levels to the fallback list
+ # Skip for error types that have their own dedicated fallback handlers
+ _skip_order_fallback = isinstance(
+ e,
+ (litellm.ContextWindowExceededError, litellm.ContentPolicyViolationError),
+ )
+ all_deployments = self._get_all_deployments(model_name=original_model_group)
+ _order_set: set = {
+ d.get("litellm_params", {}).get("order")
+ for d in all_deployments
+ if d.get("litellm_params", {}).get("order") is not None
+ }
+ order_values: list = sorted(_order_set)
+ if len(order_values) > 1 and not _skip_order_fallback:
+ # Determine which order levels have already been tried
+ current_target = kwargs.get("_target_order")
+ skip_up_to = (
+ current_target if current_target is not None else order_values[0]
+ )
+ # Build order-based fallback entries (skip already-tried levels)
+ order_fallback_entries: List = [
+ {"model": original_model_group, "_target_order": o}
+ for o in order_values
+ if o > skip_up_to
+ ]
+ # Get external fallbacks — handle both standard and non-standard formats
+ external_fallback_group: Optional[List] = None
+ if fallbacks is not None and model_group is not None:
+ if _check_non_standard_fallback_format(fallbacks=fallbacks):
+ # Non-standard formats (e.g. ["claude-3-haiku"] or
+ # [{"model": "...", "messages": [...]}]) are passed through directly
+ external_fallback_group = fallbacks
+ else:
+ external_fallback_group, generic_idx = get_fallback_model_group(
+ fallbacks=fallbacks,
+ model_group=cast(str, model_group),
+ )
+ if external_fallback_group is None and generic_idx is not None:
+ external_fallback_group = fallbacks[generic_idx]["*"]
+
+ # Combined list: order fallbacks first, then external
+ combined_fallbacks = order_fallback_entries + (
+ external_fallback_group or []
+ )
+
+ if combined_fallbacks:
+ input_kwargs.update(
+ {
+ "fallback_model_group": combined_fallbacks,
+ "original_model_group": original_model_group,
+ }
+ )
+ response = await run_async_fallback(
+ *args,
+ **input_kwargs,
+ )
+ return response
+
try:
verbose_router_logger.info("Trying to fallback b/w models")
@@ -6835,6 +6909,7 @@ class Router:
self.model_list = []
self.model_id_to_deployment_index_map = {} # Reset the index
self.model_name_to_deployment_indices = {} # Reset the model_name index
+ self.team_model_to_deployment_indices = {} # Reset the team_model index
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
# we add api_base/api_key each model so load balancing between azure/gpt on api_base1 and api_base2 works
@@ -7132,16 +7207,17 @@ class Router:
# Update model_name_to_deployment_indices
for model_name, indices in list(self.model_name_to_deployment_indices.items()):
- # Remove the deleted index
- if removal_idx in indices:
- indices.remove(removal_idx)
-
- # Decrement all indices greater than removal_idx
+ # Build new list without mutating the original
updated_indices = []
for idx in indices:
- if idx > removal_idx:
+ if idx == removal_idx:
+ # Skip the removed index
+ continue
+ elif idx > removal_idx:
+ # Decrement indices after removal
updated_indices.append(idx - 1)
else:
+ # Keep indices before removal unchanged
updated_indices.append(idx)
# Update or remove the entry
@@ -7150,6 +7226,46 @@ class Router:
else:
del self.model_name_to_deployment_indices[model_name]
+ # Update team_model_to_deployment_indices
+ for key, indices in list(self.team_model_to_deployment_indices.items()):
+ # Build new list without mutating the original
+ updated_indices = []
+ for idx in indices:
+ if idx == removal_idx:
+ # Skip the removed index
+ continue
+ elif idx > removal_idx:
+ # Decrement indices after removal
+ updated_indices.append(idx - 1)
+ else:
+ # Keep indices before removal unchanged
+ updated_indices.append(idx)
+
+ # Update or remove the entry
+ if len(updated_indices) > 0:
+ self.team_model_to_deployment_indices[key] = updated_indices
+ else:
+ del self.team_model_to_deployment_indices[key]
+
+ def _update_team_model_index(self, model: dict, idx: int) -> None:
+ """
+ Helper to update team_model_to_deployment_indices for a single deployment.
+
+ Parameters:
+ - model: dict - the deployment to index
+ - idx: int - the index in model_list
+ """
+ team_id = (model.get("model_info") or {}).get("team_id")
+ team_public_model_name = (model.get("model_info") or {}).get(
+ "team_public_model_name"
+ )
+ if team_id and team_public_model_name:
+ key = (team_id, team_public_model_name)
+ if key not in self.team_model_to_deployment_indices:
+ self.team_model_to_deployment_indices[key] = []
+ if idx not in self.team_model_to_deployment_indices[key]:
+ self.team_model_to_deployment_indices[key].append(idx)
+
def _add_model_to_list_and_index_map(
self, model: dict, model_id: Optional[str] = None
) -> None:
@@ -7178,6 +7294,9 @@ class Router:
self.model_name_to_deployment_indices[model_name] = []
self.model_name_to_deployment_indices[model_name].append(idx)
+ # Update team_model index for O(1) team-scoped lookup
+ self._update_team_model_index(model, idx)
+
def upsert_deployment(self, deployment: Deployment) -> Optional[Deployment]:
"""
Add or update deployment
@@ -7196,7 +7315,10 @@ class Router:
)
if _deployment_on_router is not None:
# deployment with this model_id exists on the router
- if deployment.litellm_params == _deployment_on_router.litellm_params:
+ if (
+ deployment.litellm_params == _deployment_on_router.litellm_params
+ and deployment.model_info == _deployment_on_router.model_info
+ ):
# No need to update
return None
@@ -8008,6 +8130,7 @@ class Router:
instead of O(n) linear scan through the entire model_list.
"""
self.model_name_to_deployment_indices.clear()
+ self.team_model_to_deployment_indices.clear()
for idx, model in enumerate(model_list):
model_name = model.get("model_name")
@@ -8016,6 +8139,8 @@ class Router:
self.model_name_to_deployment_indices[model_name] = []
self.model_name_to_deployment_indices[model_name].append(idx)
+ self._update_team_model_index(model, idx)
+
def _build_model_id_to_deployment_index_map(self, model_list: list):
"""
Build model index from model list to enable O(1) lookups immediately.
@@ -8148,20 +8273,25 @@ class Router:
def map_team_model(self, team_model_name: str, team_id: str) -> Optional[str]:
"""
- Map a team model name to a team-specific model name.
+ Check if team_model_name resolves to team-specific deployments.
+
+ Returns the public model name (unchanged) so the router can find all
+ sibling deployments via team_id filtering, instead of collapsing to a
+ single internal model_name.
Returns:
- - deployment id: str - the deployment id of the team-specific model
- - None: if no team-specific model name is found
+ - str: the team_model_name if team deployments exist for this team
+ - None: if no team-specific model is found
"""
models = self.get_model_list(model_name=team_model_name, team_id=team_id)
if not models:
return None
for model in models:
if model.get("model_info", {}).get("team_id") == team_id:
- return model.get("model_name")
+ return team_model_name
- ## wildcard models
+ # No team-scoped deployment found; wildcard/pattern routes are
+ # handled downstream by the pattern_router in _common_checks_available_deployment.
return None
def should_include_deployment(
@@ -8172,12 +8302,22 @@ class Router:
"""
if (
team_id is not None
- and model["model_info"].get("team_id") == team_id
- and model_name == model["model_info"].get("team_public_model_name")
+ and (model.get("model_info") or {}).get("team_id") == team_id
+ and model_name
+ == (model.get("model_info") or {}).get("team_public_model_name")
):
return True
elif model_name is not None and model["model_name"] == model_name:
- return True
+ # Fallback: check by internal model_name for non-team deployments
+ # or deployments that haven't been migrated to team_public_model_name yet
+ model_team_id = (model.get("model_info") or {}).get("team_id")
+ if (
+ team_id is None # requester has no team constraint
+ or model_team_id is None # global deployment - accessible to all teams
+ or model_team_id == team_id # deployment belongs to requester's team
+ ):
+ return True
+ # No match: deployment is for a different team or doesn't match the requested model
return False
def _get_all_deployments(
@@ -8194,9 +8334,36 @@ class Router:
if team_id specified, only return team-specific models
Optimized with O(1) index lookup instead of O(n) linear scan.
+
+ Note: when team_id is provided, O(1) lookup in
+ `team_model_to_deployment_indices` only applies when `model_name` is the
+ team public model name. If a caller passes an internal deployment model
+ name (for example, `model_name__`), this method falls back
+ to the standard model-name index / scan path.
"""
returned_models: List[DeploymentTypedDict] = []
+ # O(1) lookup in team_model index when team_id is provided
+ if team_id is not None:
+ key = (team_id, model_name)
+ if key in self.team_model_to_deployment_indices:
+ indices = self.team_model_to_deployment_indices[key]
+ # O(k) where k = team deployments for this model_name (typically 1-10)
+ for idx in indices:
+ model = self.model_list[idx]
+ if not self.should_include_deployment(
+ model_name=model_name, model=model, team_id=team_id
+ ):
+ continue
+ if model_alias is not None:
+ alias_model = model.copy()
+ alias_model["model_name"] = model_alias
+ returned_models.append(alias_model)
+ else:
+ returned_models.append(model)
+ if returned_models:
+ return returned_models
+
# O(1) lookup in model_name index
if model_name in self.model_name_to_deployment_indices:
indices = self.model_name_to_deployment_indices[model_name]
@@ -8791,12 +8958,6 @@ class Router:
if i not in invalid_model_indices
]
- ## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2)
- if len(_returned_deployments) > 0:
- _returned_deployments = litellm.utils._get_order_filtered_deployments(
- _returned_deployments
- )
-
return _returned_deployments
def _get_model_from_alias(self, model: str) -> Optional[str]:
@@ -8867,6 +9028,16 @@ class Router:
model = _model_from_alias
if model not in self.model_names:
+ # Check for team-specific deployments by team_public_model_name.
+ # This intentionally takes priority over team pattern routers below,
+ # so that named team deployments shadow wildcard/pattern routes.
+ if request_team_id is not None:
+ team_deployments = self._get_all_deployments(
+ model_name=model, team_id=request_team_id
+ )
+ if team_deployments:
+ return model, team_deployments
+
# check if provider/ specific wildcard routing use pattern matching
pattern_deployments = self.pattern_router.get_deployments_by_pattern(
model=model,
@@ -8997,6 +9168,14 @@ class Router:
if isinstance(healthy_deployments, dict):
return healthy_deployments
+ # Health-check-based filtering (before cooldown)
+ healthy_deployments = (
+ await self._async_filter_health_check_unhealthy_deployments(
+ healthy_deployments=healthy_deployments,
+ parent_otel_span=parent_otel_span,
+ )
+ )
+
cooldown_deployments = await _async_get_cooldown_deployments(
litellm_router_instance=self, parent_otel_span=parent_otel_span
)
@@ -9035,6 +9214,12 @@ class Router:
),
)
+ ## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2)
+ _target_order = (request_kwargs or {}).pop("_target_order", None)
+ healthy_deployments = litellm.utils._get_order_filtered_deployments(
+ cast(List[Dict], healthy_deployments), target_order=_target_order
+ )
+
if len(healthy_deployments) == 0:
exception = await async_raise_no_deployment_exception(
litellm_router_instance=self,
@@ -9422,6 +9607,13 @@ class Router:
parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs(
request_kwargs
)
+
+ # Health-check-based filtering (before cooldown)
+ healthy_deployments = self._filter_health_check_unhealthy_deployments(
+ healthy_deployments=healthy_deployments,
+ parent_otel_span=parent_otel_span,
+ )
+
cooldown_deployments = _get_cooldown_deployments(
litellm_router_instance=self, parent_otel_span=parent_otel_span
)
@@ -9439,6 +9631,12 @@ class Router:
request_kwargs=request_kwargs,
)
+ ## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2)
+ _target_order = (request_kwargs or {}).pop("_target_order", None)
+ healthy_deployments = litellm.utils._get_order_filtered_deployments(
+ healthy_deployments, target_order=_target_order
+ )
+
if len(healthy_deployments) == 0:
model_ids = self.get_model_ids(model_name=model)
_cooldown_time = self.cooldown_cache.get_min_cooldown(
@@ -9581,10 +9779,14 @@ class Router:
llm_provider="",
)
- # 4. Apply cooldown filtering
+ # 4. Apply health-check and cooldown filtering
parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs(
request_kwargs
)
+ pass_through_deployments = self._filter_health_check_unhealthy_deployments(
+ healthy_deployments=pass_through_deployments,
+ parent_otel_span=parent_otel_span,
+ )
cooldown_deployments = _get_cooldown_deployments(
litellm_router_instance=self, parent_otel_span=parent_otel_span
)
@@ -9706,6 +9908,67 @@ class Router:
if deployment["model_info"]["id"] not in cooldown_set
]
+ async def _async_filter_health_check_unhealthy_deployments(
+ self,
+ healthy_deployments: List[Dict],
+ parent_otel_span: Optional[Span] = None,
+ ) -> List[Dict]:
+ """
+ Filter out deployments marked unhealthy by background health checks.
+ No-op when enable_health_check_routing is False.
+ Returns all deployments if health state is unavailable, stale, or would
+ exclude every candidate (safety net).
+ """
+ if not self.enable_health_check_routing:
+ return healthy_deployments
+
+ unhealthy_ids = (
+ await self.health_state_cache.async_get_unhealthy_deployment_ids(
+ parent_otel_span=parent_otel_span
+ )
+ )
+ if not unhealthy_ids:
+ return healthy_deployments
+
+ filtered = [
+ d for d in healthy_deployments if d["model_info"]["id"] not in unhealthy_ids
+ ]
+
+ if not filtered:
+ verbose_router_logger.warning(
+ "All deployments marked unhealthy by health checks, bypassing health filter"
+ )
+ return healthy_deployments
+
+ return filtered
+
+ def _filter_health_check_unhealthy_deployments(
+ self,
+ healthy_deployments: List[Dict],
+ parent_otel_span: Optional[Span] = None,
+ ) -> List[Dict]:
+ """Sync version of _async_filter_health_check_unhealthy_deployments."""
+ if not self.enable_health_check_routing:
+ return healthy_deployments
+
+ unhealthy_ids = self.health_state_cache.get_unhealthy_deployment_ids(
+ parent_otel_span=parent_otel_span
+ )
+ if not unhealthy_ids:
+ return healthy_deployments
+
+ filtered = [
+ d for d in healthy_deployments if d["model_info"]["id"] not in unhealthy_ids
+ ]
+
+ if not filtered:
+ verbose_router_logger.warning(
+ "All deployments marked unhealthy by health checks, bypassing health filter"
+ )
+ return healthy_deployments
+
+ return filtered
+
def _filter_pass_through_deployments(
self, healthy_deployments: List[Dict]
) -> List[Dict]:
diff --git a/litellm/router_utils/health_state_cache.py b/litellm/router_utils/health_state_cache.py
new file mode 100644
index 00000000000..65b064f19d2
--- /dev/null
+++ b/litellm/router_utils/health_state_cache.py
@@ -0,0 +1,100 @@
+"""
+Wrapper around router cache for health-check-driven routing.
+
+Stores per-deployment health state from background health checks
+and exposes it for router candidate filtering.
+"""
+
+import time
+from typing import TYPE_CHECKING, Any, Dict, Optional, Set, Union
+
+from typing_extensions import TypedDict
+
+from litellm import verbose_logger
+from litellm.caching.caching import DualCache
+
+if TYPE_CHECKING:
+ from opentelemetry.trace import Span as _Span
+
+ Span = Union[_Span, Any]
+else:
+ Span = Any
+
+
+class DeploymentHealthStateValue(TypedDict):
+ is_healthy: bool
+ timestamp: float
+ reason: str
+
+
+class DeploymentHealthCache:
+ """
+ Cache for deployment health states produced by background health checks.
+
+ Stores a single dict mapping deployment_id -> DeploymentHealthStateValue.
+ Staleness is enforced at read time: entries older than staleness_threshold
+ are treated as healthy (unknown).
+ """
+
+ CACHE_KEY = "litellm:health_check:deployment_health_state"
+
+ def __init__(self, cache: DualCache, staleness_threshold: float):
+ self.cache = cache
+ self.staleness_threshold = staleness_threshold
+
+ def set_deployment_health_states(
+ self, states: Dict[str, DeploymentHealthStateValue]
+ ) -> None:
+ """Bulk-write all deployment health states as a single cache entry."""
+ try:
+ self.cache.set_cache(
+ key=self.CACHE_KEY,
+ value=states,
+ ttl=int(self.staleness_threshold * 1.5),
+ )
+ except Exception as e:
+ verbose_logger.error(
+ "DeploymentHealthCache::set_deployment_health_states - Exception: %s",
+ str(e),
+ )
+
+ def _extract_unhealthy_ids(self, raw: Any) -> Set[str]:
+ """Given raw cache value, return set of non-stale unhealthy deployment IDs."""
+ if not raw or not isinstance(raw, dict):
+ return set()
+ now = time.time()
+ return {
+ model_id
+ for model_id, state in raw.items()
+ if isinstance(state, dict)
+ and not state.get("is_healthy", True)
+ and (now - state.get("timestamp", 0)) < self.staleness_threshold
+ }
+
+ async def async_get_unhealthy_deployment_ids(
+ self, parent_otel_span: Optional[Span] = None
+ ) -> Set[str]:
+ """Return set of deployment IDs currently marked unhealthy and not stale."""
+ try:
+ raw = await self.cache.async_get_cache(key=self.CACHE_KEY)
+ return self._extract_unhealthy_ids(raw)
+ except Exception as e:
+ verbose_logger.debug(
+ "DeploymentHealthCache::async_get_unhealthy_deployment_ids - Exception: %s",
+ str(e),
+ )
+ return set()
+
+ def get_unhealthy_deployment_ids(
+ self, parent_otel_span: Optional[Span] = None
+ ) -> Set[str]:
+ """Sync version: return set of deployment IDs currently marked unhealthy and not stale."""
+ try:
+ raw = self.cache.get_cache(key=self.CACHE_KEY)
+ return self._extract_unhealthy_ids(raw)
+ except Exception as e:
+ verbose_logger.debug(
+ "DeploymentHealthCache::get_unhealthy_deployment_ids - Exception: %s",
+ str(e),
+ )
+ return set()
diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py
index 0856d8a6f9b..0d1501664b9 100644
--- a/litellm/types/integrations/prometheus.py
+++ b/litellm/types/integrations/prometheus.py
@@ -159,6 +159,23 @@ LATENCY_BUCKETS = (
float("inf"),
)
+# Batch jobs can run for minutes to hours; buckets span 1 min → 24 h.
+BATCH_DURATION_BUCKETS = (
+ 60.0,
+ 120.0,
+ 300.0,
+ 600.0,
+ 900.0,
+ 1800.0,
+ 3600.0,
+ 7200.0,
+ 14400.0,
+ 28800.0,
+ 43200.0,
+ 86400.0,
+ float("inf"),
+)
+
class UserAPIKeyLabelNames(Enum):
END_USER = "end_user"
@@ -185,6 +202,8 @@ class UserAPIKeyLabelNames(Enum):
USER_AGENT = "user_agent"
CALLBACK_NAME = "callback_name"
STREAM = "stream"
+ ORG_ID = "org_id"
+ ORG_ALIAS = "org_alias"
DEFINED_PROMETHEUS_METRICS = Literal[
@@ -207,6 +226,9 @@ DEFINED_PROMETHEUS_METRICS = Literal[
"litellm_remaining_team_budget_metric",
"litellm_team_max_budget_metric",
"litellm_team_budget_remaining_hours_metric",
+ "litellm_remaining_org_budget_metric",
+ "litellm_org_max_budget_metric",
+ "litellm_org_budget_remaining_hours_metric",
"litellm_remaining_api_key_budget_metric",
"litellm_api_key_max_budget_metric",
"litellm_api_key_budget_remaining_hours_metric",
@@ -238,6 +260,16 @@ DEFINED_PROMETHEUS_METRICS = Literal[
"litellm_llm_api_failed_requests_metric",
"litellm_callback_logging_failures_metric",
"litellm_in_flight_requests",
+ # Managed batch metrics
+ "litellm_managed_batch_created_total",
+ "litellm_managed_file_size_bytes",
+ "litellm_managed_batch_duration_seconds",
+ "litellm_managed_file_created_total",
+ "litellm_managed_file_deleted_total",
+ "litellm_check_batch_cost_jobs_polled",
+ "litellm_check_batch_cost_jobs_processed_total",
+ "litellm_check_batch_cost_errors_total",
+ "litellm_check_batch_cost_last_run_timestamp",
]
@@ -490,6 +522,21 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.TEAM_ALIAS.value,
]
+ litellm_remaining_org_budget_metric = [
+ UserAPIKeyLabelNames.ORG_ID.value,
+ UserAPIKeyLabelNames.ORG_ALIAS.value,
+ ]
+
+ litellm_org_max_budget_metric = [
+ UserAPIKeyLabelNames.ORG_ID.value,
+ UserAPIKeyLabelNames.ORG_ALIAS.value,
+ ]
+
+ litellm_org_budget_remaining_hours_metric = [
+ UserAPIKeyLabelNames.ORG_ID.value,
+ UserAPIKeyLabelNames.ORG_ALIAS.value,
+ ]
+
litellm_remaining_api_key_budget_metric = [
UserAPIKeyLabelNames.API_KEY_HASH.value,
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
@@ -618,6 +665,43 @@ class PrometheusMetricLabels:
litellm_cache_misses_metric = _cache_metric_labels
litellm_cached_tokens_metric = _cache_metric_labels
+ # Managed batch metrics
+ _batch_user_labels = [
+ UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
+ UserAPIKeyLabelNames.API_PROVIDER.value,
+ UserAPIKeyLabelNames.USER.value,
+ UserAPIKeyLabelNames.USER_EMAIL.value,
+ UserAPIKeyLabelNames.API_KEY_ALIAS.value,
+ ]
+
+ litellm_managed_batch_created_total = _batch_user_labels
+
+ litellm_managed_file_size_bytes: List[
+ str
+ ] = [] # labels: purpose, file_type, model, api_provider, user (custom)
+
+ litellm_managed_batch_duration_seconds = [
+ UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
+ UserAPIKeyLabelNames.API_PROVIDER.value,
+ ]
+
+ litellm_managed_file_created_total = _batch_user_labels
+
+ litellm_managed_file_deleted_total: List[
+ str
+ ] = [] # only "result" label, added at metric creation
+
+ litellm_check_batch_cost_jobs_polled: List[str] = []
+
+ litellm_check_batch_cost_jobs_processed_total = [
+ UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
+ UserAPIKeyLabelNames.API_PROVIDER.value,
+ ]
+
+ litellm_check_batch_cost_errors_total: List[str] = [] # label: error_type (custom)
+
+ litellm_check_batch_cost_last_run_timestamp: List[str] = []
+
@staticmethod
def get_labels(label_name: DEFINED_PROMETHEUS_METRICS) -> List[str]:
default_labels = getattr(PrometheusMetricLabels, label_name)
@@ -721,6 +805,12 @@ class UserAPIKeyLabelValues(BaseModel):
stream: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.STREAM.value)
] = None
+ org_id: Annotated[
+ Optional[str], Field(..., alias=UserAPIKeyLabelNames.ORG_ID.value)
+ ] = None
+ org_alias: Annotated[
+ Optional[str], Field(..., alias=UserAPIKeyLabelNames.ORG_ALIAS.value)
+ ] = None
@field_validator("stream", mode="before")
@classmethod
diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py
index 5a80b40d61f..9e669655761 100644
--- a/litellm/types/llms/openai.py
+++ b/litellm/types/llms/openai.py
@@ -315,11 +315,11 @@ class OpenAIFileObject(BaseModel):
`fine-tune`, `fine-tune-results`, `vision`, and `user_data`.
"""
- status: Optional[Literal["uploaded", "processed", "error"]] = None
+ status: Optional[Literal["uploaded", "processed", "error", "pending"]] = None
"""Deprecated.
- The current status of the file, which can be either `uploaded`, `processed`, or
- `error`.
+ The current status of the file, which can be either `uploaded`, `processed`,
+ `error`, or `pending` (Azure may return `pending` immediately after upload).
"""
expires_at: Optional[int] = None
@@ -536,6 +536,20 @@ class ChatCompletionRedactedThinkingBlock(TypedDict, total=False):
cache_control: Optional[Union[dict, ChatCompletionCachedContent]]
+class ChatCompletionReasoningSummaryTextBlock(TypedDict, total=False):
+ type: Required[Literal["summary_text"]]
+ text: str
+
+
+class ChatCompletionReasoningItem(TypedDict, total=False):
+ """Represents an OpenAI Responses API reasoning item for round-tripping in conversation history."""
+
+ type: Required[Literal["reasoning"]]
+ id: str
+ encrypted_content: Optional[str]
+ summary: List["ChatCompletionReasoningSummaryTextBlock"]
+
+
class WebSearchOptionsUserLocationApproximate(TypedDict, total=False):
city: str
"""Free text input for the city of the user, e.g. `San Francisco`."""
diff --git a/litellm/types/utils.py b/litellm/types/utils.py
index bd673da8bed..82557513a8a 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -58,6 +58,7 @@ from .llms.openai import (
AllMessageValues,
Batch,
ChatCompletionAnnotation,
+ ChatCompletionReasoningItem,
ChatCompletionRedactedThinkingBlock,
ChatCompletionThinkingBlock,
ChatCompletionToolCallChunk,
@@ -132,6 +133,7 @@ class ProviderSpecificModelInfo(TypedDict, total=False):
supports_audio_output: Optional[bool]
supports_pdf_input: Optional[bool]
supports_native_streaming: Optional[bool]
+ supports_native_structured_output: Optional[bool]
supports_parallel_function_calling: Optional[bool]
supports_web_search: Optional[bool]
supports_reasoning: Optional[bool]
@@ -1132,6 +1134,7 @@ class Message(SafeAttributeModel, OpenAIObject):
thinking_blocks: Optional[
List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]
] = None
+ reasoning_items: Optional[List[ChatCompletionReasoningItem]] = None
provider_specific_fields: Optional[Dict[str, Any]] = Field(default=None)
annotations: Optional[List[ChatCompletionAnnotation]] = None
@@ -1150,6 +1153,7 @@ class Message(SafeAttributeModel, OpenAIObject):
Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]
]
] = None,
+ reasoning_items: Optional[List[ChatCompletionReasoningItem]] = None,
annotations: Optional[List[ChatCompletionAnnotation]] = None,
**params,
):
@@ -1182,6 +1186,9 @@ class Message(SafeAttributeModel, OpenAIObject):
if thinking_blocks is not None:
init_values["thinking_blocks"] = thinking_blocks
+ if reasoning_items is not None:
+ init_values["reasoning_items"] = reasoning_items
+
if annotations is not None:
init_values["annotations"] = annotations
@@ -1219,6 +1226,11 @@ class Message(SafeAttributeModel, OpenAIObject):
if hasattr(self, "thinking_blocks"):
del self.thinking_blocks
+ if reasoning_items is None:
+ # ensure default response matches OpenAI spec
+ if hasattr(self, "reasoning_items"):
+ del self.reasoning_items
+
add_provider_specific_fields(self, provider_specific_fields)
def get(self, key, default=None):
@@ -1246,6 +1258,7 @@ class Delta(SafeAttributeModel, OpenAIObject):
thinking_blocks: Optional[
List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]
] = None
+ reasoning_items: Optional[List[ChatCompletionReasoningItem]] = None
provider_specific_fields: Optional[Dict[str, Any]] = Field(default=None)
def __init__(
@@ -1262,6 +1275,7 @@ class Delta(SafeAttributeModel, OpenAIObject):
Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]
]
] = None,
+ reasoning_items: Optional[List[ChatCompletionReasoningItem]] = None,
annotations: Optional[List[ChatCompletionAnnotation]] = None,
**params,
):
@@ -1295,6 +1309,13 @@ class Delta(SafeAttributeModel, OpenAIObject):
# ensure default response matches OpenAI spec
del self.thinking_blocks
+ if reasoning_items is not None:
+ self.reasoning_items = reasoning_items
+ else:
+ # ensure default response matches OpenAI spec
+ if hasattr(self, "reasoning_items"):
+ del self.reasoning_items
+
# Add annotations to the delta, ensure they are only on Delta if they exist (Match OpenAI spec)
if annotations is not None:
self.annotations = annotations
diff --git a/litellm/utils.py b/litellm/utils.py
index 088ee07d630..83d1242f3fd 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -2735,6 +2735,19 @@ def supports_reasoning(model: str, custom_llm_provider: Optional[str] = None) ->
)
+def supports_native_structured_output(
+ model: str, custom_llm_provider: Optional[str] = None
+) -> bool:
+ """
+ Check if the given model supports native structured outputs and return a boolean value.
+ """
+ return _supports_factory(
+ model=model,
+ custom_llm_provider=custom_llm_provider,
+ key="supports_native_structured_output",
+ )
+
+
def get_supported_regions(
model: str, custom_llm_provider: Optional[str] = None
) -> Optional[List[str]]:
@@ -4866,7 +4879,21 @@ def calculate_max_parallel_requests(
return None
-def _get_order_filtered_deployments(healthy_deployments: List[Dict]) -> List:
+def _get_order_filtered_deployments(
+ healthy_deployments: List[Dict], target_order: Optional[int] = None
+) -> List:
+ if target_order is not None:
+ filtered = [
+ d
+ for d in healthy_deployments
+ if d["litellm_params"].get("order") == target_order
+ ]
+ if filtered:
+ return filtered
+ # target_order doesn't match any deployment (e.g., external fallback model) — return all
+ return healthy_deployments
+
+ # Default: pick min order group
min_order = min(
(
deployment["litellm_params"]["order"]
@@ -5831,6 +5858,9 @@ def _get_model_info_helper( # noqa: PLR0915
supports_native_streaming=_model_info.get(
"supports_native_streaming", None
),
+ supports_native_structured_output=_model_info.get(
+ "supports_native_structured_output", None
+ ),
supports_web_search=_model_info.get("supports_web_search", None),
supports_url_context=_model_info.get("supports_url_context", None),
supports_reasoning=_model_info.get("supports_reasoning", None),
diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json
index c53ee943c58..d4f986edd9b 100644
--- a/model_prices_and_context_window.json
+++ b/model_prices_and_context_window.json
@@ -722,7 +722,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"anthropic.claude-haiku-4-5@20251001": {
"cache_creation_input_token_cost": 1.25e-06,
@@ -745,7 +746,8 @@
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
- "supports_native_streaming": true
+ "supports_native_streaming": true,
+ "supports_native_structured_output": true
},
"anthropic.claude-3-5-sonnet-20240620-v1:0": {
"input_cost_per_token": 3e-06,
@@ -967,22 +969,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_structured_output": true
},
"anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -997,22 +996,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"global.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1027,22 +1023,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"us.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.875e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
"cache_read_input_token_cost": 5.5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
"input_cost_per_token": 5.5e-06,
- "input_cost_per_token_above_200k_tokens": 1.1e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.75e-05,
- "output_cost_per_token_above_200k_tokens": 4.125e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1057,22 +1050,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"eu.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.875e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
"cache_read_input_token_cost": 5.5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
"input_cost_per_token": 5.5e-06,
- "input_cost_per_token_above_200k_tokens": 1.1e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.75e-05,
- "output_cost_per_token_above_200k_tokens": 4.125e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1087,22 +1077,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"au.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.875e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
"cache_read_input_token_cost": 5.5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
"input_cost_per_token": 5.5e-06,
- "input_cost_per_token_above_200k_tokens": 1.1e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.75e-05,
- "output_cost_per_token_above_200k_tokens": 4.125e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1117,22 +1104,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token": 3e-06,
- "input_cost_per_token_above_200k_tokens": 6e-06,
"litellm_provider": "bedrock_converse",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
- "output_cost_per_token_above_200k_tokens": 2.25e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1147,22 +1131,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"global.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token": 3e-06,
- "input_cost_per_token_above_200k_tokens": 6e-06,
"litellm_provider": "bedrock_converse",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
- "output_cost_per_token_above_200k_tokens": 2.25e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1177,22 +1158,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"us.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 4.125e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_read_input_token_cost": 3.3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"input_cost_per_token": 3.3e-06,
- "input_cost_per_token_above_200k_tokens": 6.6e-06,
"litellm_provider": "bedrock_converse",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
- "output_cost_per_token_above_200k_tokens": 2.475e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1207,22 +1185,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"eu.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 4.125e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_read_input_token_cost": 3.3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"input_cost_per_token": 3.3e-06,
- "input_cost_per_token_above_200k_tokens": 6.6e-06,
"litellm_provider": "bedrock_converse",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
- "output_cost_per_token_above_200k_tokens": 2.475e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1237,22 +1212,19 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"au.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 4.125e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_read_input_token_cost": 3.3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"input_cost_per_token": 3.3e-06,
- "input_cost_per_token_above_200k_tokens": 6.6e-06,
"litellm_provider": "bedrock_converse",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
- "output_cost_per_token_above_200k_tokens": 2.475e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -1267,7 +1239,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"anthropic.claude-sonnet-4-20250514-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -1327,7 +1300,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_structured_output": true
},
"anthropic.claude-v1": {
"input_cost_per_token": 8e-06,
@@ -1577,7 +1551,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"apac.anthropic.claude-3-sonnet-20240229-v1:0": {
"input_cost_per_token": 3e-06,
@@ -1665,7 +1640,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"azure/ada": {
"input_cost_per_token": 1e-07,
@@ -1831,7 +1807,7 @@
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "azure_ai",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
@@ -8503,18 +8479,14 @@
},
"claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token": 3e-06,
- "input_cost_per_token_above_200k_tokens": 6e-06,
"litellm_provider": "anthropic",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
- "output_cost_per_token_above_200k_tokens": 2.25e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -8695,19 +8667,15 @@
},
"claude-opus-4-6": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "anthropic",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -8730,19 +8698,15 @@
},
"claude-opus-4-6-20260205": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "anthropic",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -11742,7 +11706,8 @@
"output_cost_per_token": 1.68e-06,
"supports_function_calling": true,
"supports_reasoning": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_structured_output": true
},
"deepseek.v3.2": {
"input_cost_per_token": 6.2e-07,
@@ -12187,7 +12152,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"eu.anthropic.claude-3-5-sonnet-20240620-v1:0": {
"input_cost_per_token": 3e-06,
@@ -12401,7 +12367,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"eu.meta.llama3-2-1b-instruct-v1:0": {
"input_cost_per_token": 1.3e-07,
@@ -14631,18 +14598,6 @@
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"uses_embed_content": true
},
- "vertex_ai/gemini-embedding-2-preview": {
- "input_cost_per_token": 1.5e-07,
- "litellm_provider": "vertex_ai",
- "max_input_tokens": 8192,
- "max_tokens": 8192,
- "mode": "embedding",
- "output_cost_per_token": 0,
- "output_vector_size": 3072,
- "source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal",
- "supports_multimodal": true,
- "uses_embed_content": true
- },
"gemini/gemini-embedding-001": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "gemini",
@@ -15931,6 +15886,55 @@
"supports_tool_choice": true,
"supports_vision": true
},
+ "gemini/lyria-3-clip-preview": {
+ "input_cost_per_token": 0,
+ "litellm_provider": "gemini",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_image": 0.04,
+ "output_cost_per_token": 0,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "audio"
+ ],
+ "supports_audio_input": false,
+ "supports_audio_output": true,
+ "supports_function_calling": false,
+ "supports_prompt_caching": false,
+ "supports_response_schema": false,
+ "supports_system_messages": false,
+ "supports_vision": false,
+ "supports_web_search": false
+ },
+ "gemini/lyria-3-pro-preview": {
+ "input_cost_per_token": 0,
+ "litellm_provider": "gemini",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "audio"
+ ],
+ "supports_audio_input": false,
+ "supports_audio_output": true,
+ "supports_function_calling": false,
+ "supports_prompt_caching": false,
+ "supports_response_schema": false,
+ "supports_system_messages": false,
+ "supports_vision": false,
+ "supports_web_search": false
+ },
"gemini/veo-2.0-generate-001": {
"litellm_provider": "gemini",
"max_input_tokens": 1024,
@@ -16770,7 +16774,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"global.anthropic.claude-sonnet-4-20250514-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -16822,7 +16827,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"global.amazon.nova-2-lite-v1:0": {
"cache_read_input_token_cost": 7.5e-08,
@@ -20551,7 +20557,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"jp.anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.375e-06,
@@ -20573,7 +20580,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"lambda_ai/deepseek-llama3.3-70b": {
"input_cost_per_token": 2e-07,
@@ -21268,7 +21276,8 @@
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.2e-06,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"minimax.minimax-m2.1": {
"input_cost_per_token": 3e-07,
@@ -21424,7 +21433,8 @@
"mode": "chat",
"output_cost_per_token": 2e-07,
"supports_function_calling": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral.ministral-3-3b-instruct": {
"input_cost_per_token": 1e-07,
@@ -21435,7 +21445,8 @@
"mode": "chat",
"output_cost_per_token": 1e-07,
"supports_function_calling": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral.ministral-3-8b-instruct": {
"input_cost_per_token": 1.5e-07,
@@ -21446,7 +21457,8 @@
"mode": "chat",
"output_cost_per_token": 1.5e-07,
"supports_function_calling": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral.mistral-7b-instruct-v0:2": {
"input_cost_per_token": 1.5e-07,
@@ -21488,7 +21500,8 @@
"mode": "chat",
"output_cost_per_token": 1.5e-06,
"supports_function_calling": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral.mistral-small-2402-v1:0": {
"input_cost_per_token": 1e-06,
@@ -21519,7 +21532,8 @@
"mode": "chat",
"output_cost_per_token": 4e-08,
"supports_audio_input": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral.voxtral-small-24b-2507": {
"input_cost_per_token": 1e-07,
@@ -21530,7 +21544,8 @@
"mode": "chat",
"output_cost_per_token": 3e-07,
"supports_audio_input": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"mistral/codestral-2405": {
"input_cost_per_token": 1e-06,
@@ -22217,7 +22232,8 @@
"mode": "chat",
"output_cost_per_token": 2.5e-06,
"supports_reasoning": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"moonshotai.kimi-k2.5": {
"input_cost_per_token": 6e-07,
@@ -23092,7 +23108,8 @@
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "source": "https://aws.amazon.com/bedrock/pricing/"
+ "source": "https://aws.amazon.com/bedrock/pricing/",
+ "supports_native_structured_output": true
},
"o1": {
"cache_read_input_token_cost": 7.5e-06,
@@ -26176,7 +26193,8 @@
"output_cost_per_token": 1.8e-06,
"supports_function_calling": true,
"supports_reasoning": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-235b-a22b-2507-v1:0": {
"input_cost_per_token": 2.2e-07,
@@ -26188,7 +26206,8 @@
"output_cost_per_token": 8.8e-07,
"supports_function_calling": true,
"supports_reasoning": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-coder-30b-a3b-v1:0": {
"input_cost_per_token": 1.5e-07,
@@ -26200,7 +26219,8 @@
"output_cost_per_token": 6e-07,
"supports_function_calling": true,
"supports_reasoning": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-32b-v1:0": {
"input_cost_per_token": 1.5e-07,
@@ -26212,7 +26232,8 @@
"output_cost_per_token": 6e-07,
"supports_function_calling": true,
"supports_reasoning": true,
- "supports_tool_choice": true
+ "supports_tool_choice": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-next-80b-a3b": {
"input_cost_per_token": 1.5e-07,
@@ -26223,7 +26244,8 @@
"mode": "chat",
"output_cost_per_token": 1.2e-06,
"supports_function_calling": true,
- "supports_system_messages": true
+ "supports_system_messages": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-vl-235b-a22b": {
"input_cost_per_token": 5.3e-07,
@@ -26235,7 +26257,8 @@
"output_cost_per_token": 2.66e-06,
"supports_function_calling": true,
"supports_system_messages": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_native_structured_output": true
},
"qwen.qwen3-coder-next": {
"input_cost_per_token": 5e-07,
@@ -28260,7 +28283,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"us.anthropic.claude-3-5-sonnet-20240620-v1:0": {
"input_cost_per_token": 3e-06,
@@ -28418,7 +28442,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "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,
@@ -28439,7 +28464,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 346
+ "tool_use_system_prompt_tokens": 346,
+ "supports_native_structured_output": true
},
"us.anthropic.claude-opus-4-20250514-v1:0": {
"cache_creation_input_token_cost": 1.875e-05,
@@ -28491,7 +28517,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_structured_output": true
},
"global.anthropic.claude-opus-4-5-20251101-v1:0": {
"cache_creation_input_token_cost": 6.25e-06,
@@ -28517,7 +28544,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_structured_output": true
},
"eu.anthropic.claude-opus-4-5-20251101-v1:0": {
"cache_creation_input_token_cost": 6.25e-06,
@@ -28543,7 +28571,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
- "tool_use_system_prompt_tokens": 159
+ "tool_use_system_prompt_tokens": 159,
+ "supports_native_structured_output": true
},
"us.anthropic.claude-sonnet-4-20250514-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
@@ -30323,18 +30352,14 @@
},
"vertex_ai/claude-opus-4-6": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -30353,18 +30378,14 @@
},
"vertex_ai/claude-opus-4-6@default": {
"cache_creation_input_token_cost": 6.25e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@@ -30409,18 +30430,14 @@
},
"vertex_ai/claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token": 3e-06,
- "input_cost_per_token_above_200k_tokens": 6e-06,
"litellm_provider": "vertex_ai-anthropic_models",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
- "output_cost_per_token_above_200k_tokens": 2.25e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@@ -36830,6 +36847,38 @@
"supports_audio_input": true,
"supports_audio_output": true
},
+ "gemini-3.1-flash-live-preview": {
+ "input_cost_per_audio_token": 3e-06,
+ "input_cost_per_image_token": 1e-06,
+ "input_cost_per_token": 7.5e-07,
+ "input_cost_per_video_per_second": 3.3333333333333335e-05,
+ "litellm_provider": "gemini",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 65536,
+ "max_tokens": 65536,
+ "mode": "chat",
+ "output_cost_per_audio_token": 1.2e-05,
+ "output_cost_per_token": 4.5e-06,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_endpoints": [
+ "/v1/realtime"
+ ],
+ "supported_modalities": [
+ "text",
+ "image",
+ "audio",
+ "video"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "audio"
+ ],
+ "supports_audio_input": true,
+ "supports_audio_output": true,
+ "supports_function_calling": true,
+ "supports_vision": true,
+ "supports_web_search": true
+ },
"gemini/gemini-2.5-flash-native-audio-latest": {
"input_cost_per_audio_token": 1e-06,
"input_cost_per_token": 3e-07,
@@ -36908,6 +36957,40 @@
"tpm": 250000,
"rpm": 10
},
+ "gemini/gemini-3.1-flash-live-preview": {
+ "input_cost_per_audio_token": 3e-06,
+ "input_cost_per_image_token": 1e-06,
+ "input_cost_per_token": 7.5e-07,
+ "input_cost_per_video_per_second": 3.3333333333333335e-05,
+ "litellm_provider": "gemini",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 65536,
+ "max_tokens": 65536,
+ "mode": "chat",
+ "output_cost_per_audio_token": 1.2e-05,
+ "output_cost_per_token": 4.5e-06,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_endpoints": [
+ "/v1/realtime"
+ ],
+ "supported_modalities": [
+ "text",
+ "image",
+ "audio",
+ "video"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "audio"
+ ],
+ "supports_audio_input": true,
+ "supports_audio_output": true,
+ "supports_function_calling": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "tpm": 250000,
+ "rpm": 10
+ },
"gemini-2.5-flash-preview-tts": {
"input_cost_per_token": 3e-07,
"litellm_provider": "gemini",
@@ -37153,18 +37236,14 @@
},
"vertex_ai/claude-sonnet-4-6@default": {
"cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token": 3e-06,
- "input_cost_per_token_above_200k_tokens": 6e-06,
"litellm_provider": "vertex_ai-anthropic_models",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
- "output_cost_per_token_above_200k_tokens": 2.25e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
diff --git a/poetry.lock b/poetry.lock
index b9eba6e0a30..958e0ac65b8 100644
--- a/poetry.lock
+++ b/poetry.lock
@@ -3219,15 +3219,15 @@ files = [
[[package]]
name = "litellm-proxy-extras"
-version = "0.4.60"
+version = "0.4.61"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
optional = true
python-versions = "!=2.7.*,!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,!=3.7.*,>=3.8"
groups = ["main"]
markers = "extra == \"proxy\""
files = [
- {file = "litellm_proxy_extras-0.4.60-py3-none-any.whl", hash = "sha256:7abcc811f7430e4b24e7a8ba7186219a4845a955ae7a71d8822bd03fd9fc3393"},
- {file = "litellm_proxy_extras-0.4.60.tar.gz", hash = "sha256:1c122f2a7e0eb58fa4c6d8da9da82ac1fe2869de3510bcfade5c2932af202328"},
+ {file = "litellm_proxy_extras-0.4.61-py3-none-any.whl", hash = "sha256:9bd1e57ef51972cacff52172ef5d70b0ff689f57f3d240877667301ab8f8590e"},
+ {file = "litellm_proxy_extras-0.4.61.tar.gz", hash = "sha256:dce8e39b1547abf90d912ddd0f2a876beadf789d700ef04c165362d78ad56aee"},
]
[[package]]
@@ -8009,4 +8009,4 @@ utils = ["numpydoc"]
[metadata]
lock-version = "2.1"
python-versions = ">=3.9,<4.0"
-content-hash = "b4e3ee072f600fab9810024afdd550407d25733a9b6752476aa61826e33bc08e"
+content-hash = "8dad0e86d75e574f12c57c9f32614b7b4ea2181e931874046d43a398aef1e998"
diff --git a/pyproject.toml b/pyproject.toml
index b2a446149e5..ced8c4eb712 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -61,7 +61,7 @@ boto3 = { version = "^1.40.76", optional = true }
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
mcp = {version = ">=1.25.0,<2.0.0", optional = true, python = ">=3.10"}
a2a-sdk = {version = "^0.3.22", optional = true, python = ">=3.10"}
-litellm-proxy-extras = {version = "^0.4.60", optional = true}
+litellm-proxy-extras = {version = "^0.4.61", optional = true}
rich = {version = "^13.7.1", optional = true}
litellm-enterprise = {version = "0.1.35", optional = true}
diskcache = {version = "^5.6.1", optional = true}
diff --git a/requirements.txt b/requirements.txt
index 7ce9ab04d2a..d9e188c6e64 100644
--- a/requirements.txt
+++ b/requirements.txt
@@ -57,7 +57,7 @@ grpcio>=1.75.0; python_version >= "3.14"
sentry_sdk==2.21.0 # for sentry error handling
detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests
tzdata==2025.1 # IANA time zone database
-litellm-proxy-extras==0.4.60 # for proxy extras - e.g. prisma migrations
+litellm-proxy-extras==0.4.61 # for proxy extras - e.g. prisma migrations
llm-sandbox==0.3.31 # for skill execution in sandbox
### LITELLM PACKAGE DEPENDENCIES
python-dotenv==1.0.1 # for env
diff --git a/schema.prisma b/schema.prisma
index fde9a466a28..46be6b31e1f 100644
--- a/schema.prisma
+++ b/schema.prisma
@@ -320,6 +320,15 @@ model LiteLLM_MCPServerTable {
is_byok Boolean @default(false)
byok_description String[] @default([])
byok_api_key_help_url String?
+ source_url String?
+ // BYOM submission lifecycle
+ approval_status String? @default("active")
+ submitted_by String?
+ submitted_at DateTime?
+ reviewed_at DateTime?
+ review_notes String?
+
+ @@index([approval_status])
}
// Per-user BYOK credentials for MCP servers
diff --git a/tests/agent_tests/test_a2a_agent.py b/tests/agent_tests/test_a2a_agent.py
index ace1f8c54c5..ed7a5ab9823 100644
--- a/tests/agent_tests/test_a2a_agent.py
+++ b/tests/agent_tests/test_a2a_agent.py
@@ -1,38 +1,72 @@
"""
Simple A2A agent tests - non-streaming and streaming.
-These tests validate the localhost URL retry logic: if an A2A agent's card
-contains a localhost/internal URL (e.g., http://0.0.0.0:8001/), the request
-will fail with a connection error. LiteLLM detects this and automatically
-retries using the original api_base URL instead.
-
-Requires A2A_AGENT_URL environment variable to be set.
-
-Run with:
- A2A_AGENT_URL=https://your-agent.example.com pytest tests/agent_tests/test_a2a_agent.py -v -s
+These tests use a mocked A2A client to avoid network/env dependencies.
"""
-import os
-
-import pytest
+from types import SimpleNamespace
from uuid import uuid4
+import pytest
-def get_a2a_agent_url():
- """Get A2A agent URL from environment, skip test if not set."""
- url = os.environ.get("A2A_AGENT_URL")
- return url
+
+class MockA2AResponse:
+ def __init__(self, text: str):
+ self._payload = {
+ "id": str(uuid4()),
+ "jsonrpc": "2.0",
+ "result": {
+ "message": {
+ "role": "agent",
+ "parts": [{"kind": "text", "text": text}],
+ "messageId": uuid4().hex,
+ }
+ },
+ }
+
+ def model_dump(self, mode="json", exclude_none=True):
+ return self._payload
+
+
+class MockA2AStreamingChunk(MockA2AResponse):
+ def __init__(self, text: str, state: str):
+ super().__init__(text=text)
+ self._payload["result"]["status"] = {"state": state}
+
+
+class MockA2AClient:
+ def __init__(self):
+ self._litellm_agent_card = SimpleNamespace(
+ name="mock-agent", url="http://mock-agent.local"
+ )
+
+ async def send_message(self, request):
+ return MockA2AResponse(text="hello")
+
+ def send_message_streaming(self, request):
+ async def _stream():
+ yield MockA2AStreamingChunk(text="hel", state="in_progress")
+ yield MockA2AStreamingChunk(text="hello", state="completed")
+
+ return _stream()
+
+
+@pytest.fixture
+def mock_a2a_client(monkeypatch):
+ import litellm.a2a_protocol.main as a2a_main
+
+ async def _fake_create_a2a_client(base_url, timeout=60.0, extra_headers=None):
+ return MockA2AClient()
+
+ monkeypatch.setattr(a2a_main, "create_a2a_client", _fake_create_a2a_client)
@pytest.mark.asyncio
-@pytest.mark.flaky(retries=3, delay=5)
-async def test_a2a_non_streaming():
+async def test_a2a_non_streaming(mock_a2a_client):
"""Test non-streaming A2A request."""
from a2a.types import MessageSendParams, SendMessageRequest
from litellm.a2a_protocol import asend_message
- api_base = get_a2a_agent_url()
-
request = SendMessageRequest(
id=str(uuid4()),
params=MessageSendParams(
@@ -46,7 +80,7 @@ async def test_a2a_non_streaming():
response = await asend_message(
request=request,
- api_base=api_base,
+ api_base="http://mock",
)
assert response is not None
@@ -54,13 +88,11 @@ async def test_a2a_non_streaming():
@pytest.mark.asyncio
-async def test_a2a_streaming():
+async def test_a2a_streaming(mock_a2a_client):
"""Test streaming A2A request."""
from a2a.types import MessageSendParams, SendStreamingMessageRequest
from litellm.a2a_protocol import asend_message_streaming
- api_base = get_a2a_agent_url()
-
request = SendStreamingMessageRequest(
id=str(uuid4()),
params=MessageSendParams(
@@ -75,7 +107,7 @@ async def test_a2a_streaming():
chunks = []
async for chunk in asend_message_streaming(
request=request,
- api_base=api_base,
+ api_base="http://mock",
):
chunks.append(chunk)
print(f"\nStreaming chunk: {chunk}")
diff --git a/tests/audio_tests/test_audio_speech.py b/tests/audio_tests/test_audio_speech.py
index 67e0dbffa61..abf0df22a6a 100644
--- a/tests/audio_tests/test_audio_speech.py
+++ b/tests/audio_tests/test_audio_speech.py
@@ -35,8 +35,8 @@ import litellm
[
(
"azure/tts",
- os.getenv("AZURE_SWEDEN_API_KEY"),
- os.getenv("AZURE_SWEDEN_API_BASE"),
+ os.getenv("AZURE_TTS_API_KEY"),
+ os.getenv("AZURE_TTS_API_BASE"),
),
("openai/tts-1", os.getenv("OPENAI_API_KEY"), None),
],
@@ -286,9 +286,9 @@ async def test_speech_litellm_vertex_async_with_voice_ssml():
def test_audio_speech_cost_calc():
from litellm.integrations.custom_logger import CustomLogger
- model = "azure/azure-tts"
- api_base = os.getenv("AZURE_SWEDEN_API_BASE")
- api_key = os.getenv("AZURE_SWEDEN_API_KEY")
+ model = "azure/tts"
+ api_base = os.getenv("AZURE_TTS_API_BASE")
+ api_key = os.getenv("AZURE_TTS_API_KEY")
custom_logger = CustomLogger()
litellm.set_verbose = True
@@ -301,7 +301,7 @@ def test_audio_speech_cost_calc():
input="the quick brown fox jumped over the lazy dogs",
api_base=api_base,
api_key=api_key,
- base_model="azure/tts-1",
+ base_model="azure/tts",
)
time.sleep(1)
@@ -337,10 +337,9 @@ async def test_azure_ava_tts_async():
litellm._turn_on_debug()
api_key = os.getenv("AZURE_TTS_API_KEY")
api_base = os.getenv("AZURE_TTS_API_BASE")
-
speech_file_path = Path(__file__).parent / "azure_speech.mp3"
-
+
try:
response = await litellm.aspeech(
model="azure/speech/azure-tts",
@@ -354,30 +353,34 @@ async def test_azure_ava_tts_async():
# Assert the response is HttpxBinaryResponseContent
from litellm.types.llms.openai import HttpxBinaryResponseContent
-
+
assert isinstance(response, HttpxBinaryResponseContent)
-
+
# Get the binary content
binary_content = response.content
assert len(binary_content) > 0
-
+
# MP3 files start with these magic bytes
# ID3 tag or MPEG sync word
- assert binary_content[:3] == b"ID3" or binary_content[:2] == b"\xff\xfb" or binary_content[:2] == b"\xff\xf3"
-
+ assert (
+ binary_content[:3] == b"ID3"
+ or binary_content[:2] == b"\xff\xfb"
+ or binary_content[:2] == b"\xff\xf3"
+ )
+
# Write to file
response.stream_to_file(speech_file_path)
-
+
# Verify file was created and has content
assert speech_file_path.exists()
assert speech_file_path.stat().st_size > 0
-
+
print(f"Azure TTS audio saved to: {speech_file_path}")
# assert response cost is greater than 0
print("Response cost: ", response._hidden_params["response_cost"])
assert response._hidden_params["response_cost"] > 0
-
+
except Exception as e:
pytest.fail(f"Test failed with exception: {str(e)}")
@@ -392,10 +395,9 @@ async def test_runwayml_tts_async():
litellm._turn_on_debug()
api_key = os.getenv("RUNWAYML_API_KEY")
api_base = os.getenv("RUNWAYML_API_BASE")
-
speech_file_path = Path(__file__).parent / "runwayml_speech.mp3"
-
+
try:
response = await litellm.aspeech(
model="runwayml/eleven_multilingual_v2",
@@ -409,30 +411,34 @@ async def test_runwayml_tts_async():
# Assert the response is HttpxBinaryResponseContent
from litellm.types.llms.openai import HttpxBinaryResponseContent
-
+
assert isinstance(response, HttpxBinaryResponseContent)
-
+
# Get the binary content
binary_content = response.content
assert len(binary_content) > 0
-
+
# MP3 files start with these magic bytes
# ID3 tag or MPEG sync word
- assert binary_content[:3] == b"ID3" or binary_content[:2] == b"\xff\xfb" or binary_content[:2] == b"\xff\xf3"
-
+ assert (
+ binary_content[:3] == b"ID3"
+ or binary_content[:2] == b"\xff\xfb"
+ or binary_content[:2] == b"\xff\xf3"
+ )
+
# Write to file
response.stream_to_file(speech_file_path)
-
+
# Verify file was created and has content
assert speech_file_path.exists()
assert speech_file_path.stat().st_size > 0
-
+
print(f"RunwayML TTS audio saved to: {speech_file_path}")
# assert response cost is greater than 0
print("Response cost: ", response._hidden_params["response_cost"])
assert response._hidden_params["response_cost"] > 0
-
+
except Exception as e:
pytest.fail(f"Test failed with exception: {str(e)}")
@@ -445,17 +451,19 @@ async def test_azure_ava_tts_with_custom_voice():
"""
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
-
+
# Mock response
mock_response_content = b"fake_audio_data"
mock_httpx_response = MagicMock(spec=httpx.Response)
mock_httpx_response.content = mock_response_content
mock_httpx_response.status_code = 200
mock_httpx_response.headers = {"content-type": "audio/mpeg"}
-
- with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post") as mock_post:
+
+ with patch(
+ "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post"
+ ) as mock_post:
mock_post.return_value = mock_httpx_response
-
+
response = await litellm.aspeech(
model="azure/speech/azure-tts",
voice="en-US-AndrewNeural",
@@ -464,14 +472,14 @@ async def test_azure_ava_tts_with_custom_voice():
api_key="fake-key",
response_format="mp3",
)
-
+
# Verify the mock was called
assert mock_post.called
-
+
# Get the call arguments
call_args = mock_post.call_args
ssml_body = call_args.kwargs.get("data")
-
+
# Verify the SSML contains the custom voice
assert ssml_body is not None
assert "en-US-AndrewNeural" in ssml_body
@@ -488,17 +496,19 @@ async def test_azure_ava_tts_fable_voice_mapping():
"""
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
-
+
# Mock response
mock_response_content = b"fake_audio_data"
mock_httpx_response = MagicMock(spec=httpx.Response)
mock_httpx_response.content = mock_response_content
mock_httpx_response.status_code = 200
mock_httpx_response.headers = {"content-type": "audio/mpeg"}
-
- with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post") as mock_post:
+
+ with patch(
+ "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post"
+ ) as mock_post:
mock_post.return_value = mock_httpx_response
-
+
response = await litellm.aspeech(
model="azure/speech/azure-tts",
voice="fable",
@@ -507,14 +517,14 @@ async def test_azure_ava_tts_fable_voice_mapping():
api_key="fake-key",
response_format="mp3",
)
-
+
# Verify the mock was called
assert mock_post.called
-
+
# Get the call arguments
call_args = mock_post.call_args
ssml_body = call_args.kwargs.get("data")
-
+
# Verify the SSML contains the mapped voice (en-GB-RyanNeural, not 'fable')
assert ssml_body is not None
assert "en-GB-RyanNeural" in ssml_body
@@ -541,7 +551,9 @@ async def test_aws_polly_tts_with_native_voice():
mock_httpx_response.status_code = 200
mock_httpx_response.headers = {"content-type": "audio/mpeg"}
- with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post") as mock_post:
+ with patch(
+ "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post"
+ ) as mock_post:
mock_post.return_value = mock_httpx_response
response = await litellm.aspeech(
@@ -586,7 +598,9 @@ async def test_aws_polly_tts_with_openai_voice_mapping():
mock_httpx_response.status_code = 200
mock_httpx_response.headers = {"content-type": "audio/mpeg"}
- with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post") as mock_post:
+ with patch(
+ "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post"
+ ) as mock_post:
mock_post.return_value = mock_httpx_response
response = await litellm.aspeech(
@@ -628,7 +642,9 @@ async def test_aws_polly_tts_with_ssml():
ssml_input = 'Hello, this is SSML.'
- with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post") as mock_post:
+ with patch(
+ "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post"
+ ) as mock_post:
mock_post.return_value = mock_httpx_response
response = await litellm.aspeech(
@@ -676,7 +692,11 @@ async def test_aws_polly_tts_real_api():
assert len(binary_content) > 0
# MP3 files start with ID3 tag or MPEG sync word
- assert binary_content[:3] == b"ID3" or binary_content[:2] == b"\xff\xfb" or binary_content[:2] == b"\xff\xf3"
+ assert (
+ binary_content[:3] == b"ID3"
+ or binary_content[:2] == b"\xff\xfb"
+ or binary_content[:2] == b"\xff\xf3"
+ )
response.stream_to_file(speech_file_path)
diff --git a/tests/batches_tests/test_fine_tuning_api.py b/tests/batches_tests/test_fine_tuning_api.py
index 7e238173480..034f4e1b495 100644
--- a/tests/batches_tests/test_fine_tuning_api.py
+++ b/tests/batches_tests/test_fine_tuning_api.py
@@ -123,63 +123,6 @@ async def test_create_fine_tune_jobs_async():
pass
-@pytest.mark.asyncio
-async def test_azure_create_fine_tune_jobs_async():
- try:
- verbose_logger.setLevel(logging.DEBUG)
- file_name = "azure_fine_tune.jsonl"
- _current_dir = os.path.dirname(os.path.abspath(__file__))
- file_path = os.path.join(_current_dir, file_name)
-
- file_id = "file-5e4b20ecbd724182b9964f3cd2ab7212"
-
- create_fine_tuning_response = await litellm.acreate_fine_tuning_job(
- model="gpt-35-turbo-1106",
- training_file=file_id,
- custom_llm_provider="azure",
- api_base="https://exampleopenaiendpoint-production.up.railway.app",
- )
-
- print(
- "response from litellm.create_fine_tuning_job=", create_fine_tuning_response
- )
-
- assert create_fine_tuning_response.id is not None
-
- # response from Example/mocked endpoint
- assert create_fine_tuning_response.model == "davinci-002"
-
- # list fine tuning jobs
- print("listing ft jobs")
- ft_jobs = await litellm.alist_fine_tuning_jobs(
- limit=2,
- custom_llm_provider="azure",
- api_base="https://exampleopenaiendpoint-production.up.railway.app",
- )
- print("response from litellm.list_fine_tuning_jobs=", ft_jobs)
-
- # cancel ft job
- response = await litellm.acancel_fine_tuning_job(
- fine_tuning_job_id=create_fine_tuning_response.id,
- custom_llm_provider="azure",
- api_key=os.getenv("AZURE_SWEDEN_API_KEY"),
- api_base="https://exampleopenaiendpoint-production.up.railway.app",
- )
-
- print("response from litellm.cancel_fine_tuning_job=", response)
-
- assert response.status == "cancelled"
- assert response.id == create_fine_tuning_response.id
- except openai.RateLimitError:
- pass
- except Exception as e:
- if "Job has already completed" in str(e):
- pass
- else:
- pytest.fail(f"Error occurred: {e}")
- pass
-
-
@pytest.mark.asyncio()
async def test_create_vertex_fine_tune_jobs_mocked():
load_vertex_ai_credentials()
@@ -601,11 +544,10 @@ async def test_mock_openai_retrieve_fine_tune_job():
@pytest.mark.asyncio
async def test_mock_azure_create_fine_tune_job_with_azure_specific_params():
"""Test that Azure-specific parameters are passed through extra_body"""
- from openai import AsyncAzureOpenAI
- from openai.types.fine_tuning.fine_tuning_job import FineTuningJob
from openai.types.fine_tuning.fine_tuning_job import Hyperparameters as OAIHyperparameters
+ from litellm.types.utils import LiteLLMFineTuningJob
- mock_response = FineTuningJob(
+ mock_response = LiteLLMFineTuningJob(
id="ft-azure-123",
model="gpt-4.1-mini-2025-04-14",
created_at=1677610602,
@@ -619,8 +561,11 @@ async def test_mock_azure_create_fine_tune_job_with_azure_specific_params():
result_files=[],
)
+ async def mock_async_create(*args, **kwargs):
+ return mock_response
+
with patch("litellm.llms.azure.fine_tuning.handler.AzureOpenAIFineTuningAPI.create_fine_tuning_job") as mock_create:
- mock_create.return_value = mock_response
+ mock_create.return_value = mock_async_create()
response = await litellm.acreate_fine_tuning_job(
model="gpt-4.1-mini-2025-04-14",
diff --git a/tests/batches_tests/test_openai_batches_and_files.py b/tests/batches_tests/test_openai_batches_and_files.py
index e1165812e24..e4e1180eb13 100644
--- a/tests/batches_tests/test_openai_batches_and_files.py
+++ b/tests/batches_tests/test_openai_batches_and_files.py
@@ -75,8 +75,8 @@ def load_vertex_ai_credentials():
service_account_key_data = {}
# Update the service_account_key_data with environment variables
- private_key_id = os.environ.get("GCS_PRIVATE_KEY_ID", "")
- private_key = os.environ.get("GCS_PRIVATE_KEY", "")
+ private_key_id = os.environ.get("VERTEX_AI_PRIVATE_KEY_ID", "")
+ private_key = os.environ.get("VERTEX_AI_PRIVATE_KEY", "")
private_key = private_key.replace("\\n", "\n")
service_account_key_data["private_key_id"] = private_key_id
service_account_key_data["private_key"] = private_key
@@ -234,9 +234,9 @@ def cleanup_azure_ft_models():
import requests
client = AzureOpenAI(
- api_key=os.getenv("AZURE_FT_API_KEY"),
- azure_endpoint=os.getenv("AZURE_FT_API_BASE"),
- api_version=os.getenv("AZURE_API_VERSION"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
+ azure_endpoint=os.getenv("AZURE_AI_API_BASE"),
+ api_version=os.getenv("AZURE_AI_API_VERSION"),
)
_list_ft_jobs = client.fine_tuning.jobs.list()
@@ -577,7 +577,10 @@ async def test_vertex_list_batches(monkeypatch):
monkeypatch.setattr(
"litellm.llms.vertex_ai.batches.handler.VertexAIBatchPrediction._ensure_access_token",
- lambda self, credentials, project_id, custom_llm_provider: ("mock-token", "litellm-test-project"),
+ lambda self, credentials, project_id, custom_llm_provider: (
+ "mock-token",
+ "litellm-test-project",
+ ),
)
with patch(
@@ -648,7 +651,7 @@ async def test_vertex_async_create_batch_logs_error_body_on_http_error():
async def test_delete_batch_output_file():
"""
Test that deleting a batch output file works correctly.
-
+
This test verifies the fix for:
- When a batch is retrieved and has an output_file_id, the file object is properly stored
- The output file can be deleted without validation errors
@@ -656,11 +659,11 @@ async def test_delete_batch_output_file():
"""
litellm._turn_on_debug()
print("Testing delete batch output file")
-
+
file_name = "openai_batch_completions.jsonl"
_current_dir = os.path.dirname(os.path.abspath(__file__))
file_path = os.path.join(_current_dir, file_name)
-
+
# Create file for batch
file_obj = await litellm.acreate_file(
file=open(file_path, "rb"),
@@ -669,7 +672,7 @@ async def test_delete_batch_output_file():
)
print("Response from creating file=", file_obj)
batch_input_file_id = file_obj.id
-
+
# Create batch
create_batch_response = await litellm.acreate_batch(
completion_window="24h",
@@ -678,36 +681,37 @@ async def test_delete_batch_output_file():
custom_llm_provider="openai",
)
print("Batch created with ID=", create_batch_response.id)
-
+
# Retrieve batch to get output_file_id
retrieved_batch = await litellm.aretrieve_batch(
- batch_id=create_batch_response.id,
- custom_llm_provider="openai"
+ batch_id=create_batch_response.id, custom_llm_provider="openai"
)
print("Retrieved batch=", retrieved_batch)
-
+
# If batch has completed and has output file, test deleting it
if retrieved_batch.output_file_id:
print(f"Testing deletion of output file: {retrieved_batch.output_file_id}")
-
+
# This is the key test - deleting the output file should work
# without validation errors (file_object should not be None)
delete_output_file_response = await litellm.afile_delete(
- file_id=retrieved_batch.output_file_id,
- custom_llm_provider="openai"
+ file_id=retrieved_batch.output_file_id, custom_llm_provider="openai"
)
-
+
print("Delete output file response=", delete_output_file_response)
assert delete_output_file_response.id == retrieved_batch.output_file_id
- assert delete_output_file_response.deleted is True or hasattr(delete_output_file_response, 'id')
+ assert delete_output_file_response.deleted is True or hasattr(
+ delete_output_file_response, "id"
+ )
print("✓ Successfully deleted batch output file")
else:
- print("⚠ Batch has not completed yet or no output file available, skipping output file deletion test")
-
+ print(
+ "⚠ Batch has not completed yet or no output file available, skipping output file deletion test"
+ )
+
# Clean up - delete the input file
delete_input_file_response = await litellm.afile_delete(
- file_id=batch_input_file_id,
- custom_llm_provider="openai"
+ file_id=batch_input_file_id, custom_llm_provider="openai"
)
print("Delete input file response=", delete_input_file_response)
assert delete_input_file_response.id == batch_input_file_id
diff --git a/tests/documentation_tests/test_router_settings.py b/tests/documentation_tests/test_router_settings.py
index c66a02d6849..290aa283af4 100644
--- a/tests/documentation_tests/test_router_settings.py
+++ b/tests/documentation_tests/test_router_settings.py
@@ -37,14 +37,12 @@ print(router_init_params)
router_init_params.remove("model_list")
# Parse the documentation to extract documented keys
-repo_base = "./"
-print(os.listdir(repo_base))
-docs_path = (
- "./docs/my-website/docs/proxy/config_settings.md" # Path to the documentation
+_test_dir = os.path.dirname(os.path.abspath(__file__))
+_repo_root = os.path.abspath(os.path.join(_test_dir, "..", ".."))
+print(os.listdir(_repo_root))
+docs_path = os.path.join(
+ _repo_root, "docs", "my-website", "docs", "proxy", "config_settings.md"
)
-# docs_path = (
-# "../../docs/my-website/docs/proxy/config_settings.md" # Path to the documentation
-# )
documented_keys = set()
try:
with open(docs_path, "r", encoding="utf-8") as docs_file:
diff --git a/tests/enterprise/litellm_enterprise/integrations/test_prometheus_unit_tests.py b/tests/enterprise/litellm_enterprise/integrations/test_prometheus_unit_tests.py
index ab37a0a84cd..76a57783472 100644
--- a/tests/enterprise/litellm_enterprise/integrations/test_prometheus_unit_tests.py
+++ b/tests/enterprise/litellm_enterprise/integrations/test_prometheus_unit_tests.py
@@ -169,9 +169,9 @@ async def test_prometheus_metric_tracking():
"model_name": "gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
- "api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
+ "api_version": os.getenv("AZURE_AI_API_VERSION"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"model_info": {"id": "azure-model-id"},
},
diff --git a/tests/enterprise/litellm_enterprise/proxy/auth/test_user_api_key_auth.py b/tests/enterprise/litellm_enterprise/proxy/auth/test_user_api_key_auth.py
index a45df5df008..0a2930aebb0 100644
--- a/tests/enterprise/litellm_enterprise/proxy/auth/test_user_api_key_auth.py
+++ b/tests/enterprise/litellm_enterprise/proxy/auth/test_user_api_key_auth.py
@@ -72,7 +72,7 @@ async def test_enterprise_custom_auth_returns_string():
auth_obj = await _user_api_key_auth_builder(
request=request,
api_key="my-custom-key",
- azure_api_key_header="",
+ AZURE_AI_API_KEY_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
diff --git a/tests/image_gen_tests/test_image_edits.py b/tests/image_gen_tests/test_image_edits.py
index 393b4cb67a1..8fe2e25a994 100644
--- a/tests/image_gen_tests/test_image_edits.py
+++ b/tests/image_gen_tests/test_image_edits.py
@@ -23,10 +23,11 @@ from litellm.types.utils import StandardLoggingPayload
# Configure pytest marks to avoid warnings
pytestmark = pytest.mark.asyncio
+
class TestCustomLogger(CustomLogger):
def __init__(self):
self.standard_logging_payload: Optional[StandardLoggingPayload] = None
-
+
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
self.standard_logging_payload = kwargs.get("standard_logging_object", None)
pass
@@ -80,12 +81,12 @@ class BaseLLMImageEditTest(ABC):
result = self.image_edit_function(**call_args)
else:
result = await self.async_image_edit_function(**call_args)
-
+
print("result from image edit", result)
# Validate the response meets expected schema
ImageResponse.model_validate(result)
-
+
if isinstance(result, ImageResponse) and result.data:
image_base64 = result.data[0].b64_json
if image_base64:
@@ -97,6 +98,7 @@ class BaseLLMImageEditTest(ABC):
except litellm.ContentPolicyViolationError as e:
pass
+
# Get the current directory of the file being run
pwd = os.path.dirname(os.path.realpath(__file__))
@@ -107,6 +109,7 @@ TEST_IMAGES = [
SINGLE_TEST_IMAGE = open(os.path.join(pwd, "ishaan_github.png"), "rb")
+
def get_test_images_as_bytesio():
"""Helper function to get test images as BytesIO objects"""
bytesio_images = []
@@ -129,6 +132,7 @@ class TestOpenAIImageEditGPTImage1(BaseLLMImageEditTest):
"image": TEST_IMAGES,
}
+
class TestOpenAIImageEditDallE2(BaseLLMImageEditTest):
"""
Concrete implementation of BaseLLMImageEditTest for OpenAI DALL-E-2 image edits.
@@ -155,7 +159,7 @@ class TestAzureAIFlux2ImageEdit(BaseLLMImageEditTest):
"model": "azure_ai/flux.2-pro",
"image": SINGLE_TEST_IMAGE,
"api_base": "https://litellm-ci-cd-prod.services.ai.azure.com",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": "preview",
}
@@ -187,7 +191,7 @@ async def test_openai_image_edit_litellm_router():
# Validate the response meets expected schema
ImageResponse.model_validate(result)
-
+
if isinstance(result, ImageResponse) and result.data:
image_base64 = result.data[0].b64_json
if image_base64:
@@ -199,17 +203,19 @@ async def test_openai_image_edit_litellm_router():
except litellm.ContentPolicyViolationError as e:
pass
+
@pytest.mark.flaky(retries=3, delay=2)
@pytest.mark.asyncio
async def test_openai_image_edit_with_bytesio():
"""Test image editing using BytesIO objects instead of file readers"""
from litellm import image_edit, aimage_edit
+
litellm._turn_on_debug()
try:
prompt = """
Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO.
"""
-
+
# Get images as BytesIO objects
bytesio_images = get_test_images_as_bytesio()
@@ -222,7 +228,7 @@ async def test_openai_image_edit_with_bytesio():
# Validate the response meets expected schema
ImageResponse.model_validate(result)
-
+
if isinstance(result, ImageResponse) and result.data:
image_base64 = result.data[0].b64_json
if image_base64:
@@ -239,7 +245,7 @@ async def test_openai_image_edit_with_bytesio():
async def test_azure_image_edit_litellm_sdk():
"""Test Azure image edit with mocked httpx request to validate request body and URL"""
from litellm import image_edit, aimage_edit
-
+
# Mock response for Azure image edit
mock_response = {
"created": 1589478378,
@@ -247,7 +253,7 @@ async def test_azure_image_edit_litellm_sdk():
{
"b64_json": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg=="
}
- ]
+ ],
}
class MockResponse:
@@ -267,16 +273,16 @@ async def test_azure_image_edit_litellm_sdk():
mock_post.return_value = MockResponse(mock_response, 200)
litellm._turn_on_debug()
-
+
prompt = """
Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO.
"""
-
+
# Set up test environment variables
test_api_base = "https://ai-api-gw-uae-north.openai.azure.com"
test_api_key = "test-api-key"
test_api_version = "2025-04-01-preview"
-
+
result = await aimage_edit(
prompt=prompt,
model="azure/gpt-image-1",
@@ -285,41 +291,54 @@ async def test_azure_image_edit_litellm_sdk():
api_version=test_api_version,
image=TEST_IMAGES,
)
-
+
# Verify the request was made correctly
mock_post.assert_called_once()
-
+
# Check the URL
call_args = mock_post.call_args
expected_url = f"{test_api_base}/openai/deployments/gpt-image-1/images/edits?api-version={test_api_version}"
- actual_url = call_args.args[0] if call_args.args else call_args.kwargs.get('url')
+ actual_url = (
+ call_args.args[0] if call_args.args else call_args.kwargs.get("url")
+ )
print(f"Expected URL: {expected_url}")
print(f"Actual URL: {actual_url}")
- assert actual_url == expected_url, f"URL mismatch. Expected: {expected_url}, Got: {actual_url}"
-
+ assert (
+ actual_url == expected_url
+ ), f"URL mismatch. Expected: {expected_url}, Got: {actual_url}"
+
# Check the request body
- if 'data' in call_args.kwargs:
+ if "data" in call_args.kwargs:
# For multipart form data, check the data parameter
- form_data = call_args.kwargs['data']
- print("Form data keys:", list(form_data.keys()) if hasattr(form_data, 'keys') else "Not a dict")
-
+ form_data = call_args.kwargs["data"]
+ print(
+ "Form data keys:",
+ list(form_data.keys()) if hasattr(form_data, "keys") else "Not a dict",
+ )
+
# Validate that model and prompt are in the form data
- assert 'model' in form_data, "model should be in form data"
- assert 'prompt' in form_data, "prompt should be in form data"
- assert form_data['model'] == 'gpt-image-1', f"Expected model 'gpt-image-1', got {form_data['model']}"
- assert prompt.strip() in form_data['prompt'], f"Expected prompt to contain '{prompt.strip()}'"
-
+ assert "model" in form_data, "model should be in form data"
+ assert "prompt" in form_data, "prompt should be in form data"
+ assert (
+ form_data["model"] == "gpt-image-1"
+ ), f"Expected model 'gpt-image-1', got {form_data['model']}"
+ assert (
+ prompt.strip() in form_data["prompt"]
+ ), f"Expected prompt to contain '{prompt.strip()}'"
+
# Check headers
- headers = call_args.kwargs.get('headers', {})
+ headers = call_args.kwargs.get("headers", {})
print("Request headers:", headers)
- assert 'Authorization' in headers, "Authorization header should be present"
- assert headers['Authorization'].startswith('Bearer '), "Authorization should be Bearer token"
-
+ assert "Authorization" in headers, "Authorization header should be present"
+ assert headers["Authorization"].startswith(
+ "Bearer "
+ ), "Authorization should be Bearer token"
+
print("result from image edit", result)
# Validate the response meets expected schema
ImageResponse.model_validate(result)
-
+
if isinstance(result, ImageResponse) and result.data:
image_base64 = result.data[0].b64_json
if image_base64:
@@ -330,15 +349,15 @@ async def test_azure_image_edit_litellm_sdk():
f.write(image_bytes)
-
@pytest.mark.asyncio
async def test_openai_image_edit_cost_tracking():
"""Test OpenAI image edit cost tracking with custom logger"""
from litellm import image_edit, aimage_edit
+
test_custom_logger = TestCustomLogger()
litellm.logging_callback_manager._reset_all_callbacks()
litellm.callbacks = [test_custom_logger]
-
+
# Mock response for Azure image edit with usage data for cost tracking
mock_response = {
"created": 1589478378,
@@ -350,12 +369,9 @@ async def test_openai_image_edit_cost_tracking():
"usage": {
"total_tokens": 1100,
"input_tokens": 100,
- "input_tokens_details": {
- "image_tokens": 50,
- "text_tokens": 50
- },
- "output_tokens": 1000
- }
+ "input_tokens_details": {"image_tokens": 50, "text_tokens": 50},
+ "output_tokens": 1000,
+ },
}
class MockResponse:
@@ -375,26 +391,25 @@ async def test_openai_image_edit_cost_tracking():
mock_post.return_value = MockResponse(mock_response, 200)
litellm._turn_on_debug()
-
+
prompt = """
Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO.
"""
-
+
# Set up test environment variables
-
+
result = await aimage_edit(
prompt=prompt,
model="openai/gpt-image-1",
image=TEST_IMAGES,
)
-
+
# Verify the request was made correctly
mock_post.assert_called_once()
-
# Validate the response meets expected schema
ImageResponse.model_validate(result)
-
+
if isinstance(result, ImageResponse) and result.data:
image_base64 = result.data[0].b64_json
if image_base64:
@@ -403,30 +418,36 @@ async def test_openai_image_edit_cost_tracking():
# Save the image to a file
with open("test_image_edit.png", "wb") as f:
f.write(image_bytes)
-
await asyncio.sleep(5)
- print("standard logging payload", json.dumps(test_custom_logger.standard_logging_payload, indent=4, default=str))
+ print(
+ "standard logging payload",
+ json.dumps(
+ test_custom_logger.standard_logging_payload, indent=4, default=str
+ ),
+ )
# check model
assert test_custom_logger.standard_logging_payload["model"] == "gpt-image-1"
- assert test_custom_logger.standard_logging_payload["custom_llm_provider"] == "openai"
+ assert (
+ test_custom_logger.standard_logging_payload["custom_llm_provider"]
+ == "openai"
+ )
# check response_cost
assert test_custom_logger.standard_logging_payload["response_cost"] is not None
assert test_custom_logger.standard_logging_payload["response_cost"] > 0
-
-
@pytest.mark.asyncio
async def test_azure_image_edit_cost_tracking():
"""Test Azure image edit cost tracking with custom logger"""
from litellm import image_edit, aimage_edit
+
test_custom_logger = TestCustomLogger()
litellm.logging_callback_manager._reset_all_callbacks()
litellm.callbacks = [test_custom_logger]
-
+
# Mock response for Azure image edit with usage data for cost tracking
mock_response = {
"created": 1589478378,
@@ -438,12 +459,9 @@ async def test_azure_image_edit_cost_tracking():
"usage": {
"total_tokens": 1100,
"input_tokens": 100,
- "input_tokens_details": {
- "image_tokens": 50,
- "text_tokens": 50
- },
- "output_tokens": 1000
- }
+ "input_tokens_details": {"image_tokens": 50, "text_tokens": 50},
+ "output_tokens": 1000,
+ },
}
class MockResponse:
@@ -463,27 +481,26 @@ async def test_azure_image_edit_cost_tracking():
mock_post.return_value = MockResponse(mock_response, 200)
litellm._turn_on_debug()
-
+
prompt = """
Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO.
"""
-
+
# Set up test environment variables
-
+
result = await aimage_edit(
prompt=prompt,
model="azure/CUSTOM_AZURE_DEPLOYMENT_NAME",
base_model="azure/gpt-image-1",
image=TEST_IMAGES,
)
-
+
# Verify the request was made correctly
mock_post.assert_called_once()
-
# Validate the response meets expected schema
ImageResponse.model_validate(result)
-
+
if isinstance(result, ImageResponse) and result.data:
image_base64 = result.data[0].b64_json
if image_base64:
@@ -492,14 +509,24 @@ async def test_azure_image_edit_cost_tracking():
# Save the image to a file
with open("test_image_edit.png", "wb") as f:
f.write(image_bytes)
-
await asyncio.sleep(5)
- print("standard logging payload", json.dumps(test_custom_logger.standard_logging_payload, indent=4, default=str))
+ print(
+ "standard logging payload",
+ json.dumps(
+ test_custom_logger.standard_logging_payload, indent=4, default=str
+ ),
+ )
# check model
- assert test_custom_logger.standard_logging_payload["model"] == "CUSTOM_AZURE_DEPLOYMENT_NAME"
- assert test_custom_logger.standard_logging_payload["custom_llm_provider"] == "azure"
+ assert (
+ test_custom_logger.standard_logging_payload["model"]
+ == "CUSTOM_AZURE_DEPLOYMENT_NAME"
+ )
+ assert (
+ test_custom_logger.standard_logging_payload["custom_llm_provider"]
+ == "azure"
+ )
# check response_cost
assert test_custom_logger.standard_logging_payload["response_cost"] is not None
@@ -511,6 +538,7 @@ async def test_azure_image_edit_cost_tracking():
async def test_recraft_image_edit_api():
from litellm import aimage_edit
import requests
+
litellm._turn_on_debug()
global TEST_IMAGES
try:
@@ -526,10 +554,10 @@ async def test_recraft_image_edit_api():
# Validate the response meets expected schema
ImageResponse.model_validate(result)
-
+
if isinstance(result, ImageResponse) and result.data:
image_url = result.data[0].url
-
+
# download the image
image_bytes = requests.get(image_url).content
with open("test_image_edit.png", "wb") as f:
@@ -545,51 +573,55 @@ def test_recraft_image_edit_config():
from litellm.llms.recraft.image_edit.transformation import RecraftImageEditConfig
from litellm.types.images.main import ImageEditOptionalRequestParams
from litellm.types.router import GenericLiteLLMParams
-
+
config = RecraftImageEditConfig()
-
+
# Test supported OpenAI params
supported_params = config.get_supported_openai_params("recraftv3")
expected_params = ["n", "response_format", "style"]
assert supported_params == expected_params
-
+
# Test parameter mapping (reuses OpenAI logic with filtering)
- image_edit_params = ImageEditOptionalRequestParams({
- "n": 2,
- "response_format": "b64_json",
- "style": "realistic_image",
- "size": "1024x1024", # Should be dropped
- "quality": "high" # Should be dropped
- })
-
- mapped_params = config.map_openai_params(image_edit_params, "recraftv3", drop_params=True)
-
+ image_edit_params = ImageEditOptionalRequestParams(
+ {
+ "n": 2,
+ "response_format": "b64_json",
+ "style": "realistic_image",
+ "size": "1024x1024", # Should be dropped
+ "quality": "high", # Should be dropped
+ }
+ )
+
+ mapped_params = config.map_openai_params(
+ image_edit_params, "recraftv3", drop_params=True
+ )
+
# Should only contain supported params
assert mapped_params["n"] == 2
assert mapped_params["response_format"] == "b64_json"
assert mapped_params["style"] == "realistic_image"
assert "size" not in mapped_params # Should be dropped
assert "quality" not in mapped_params # Should be dropped
-
+
# Test request transformation (reuses OpenAI file handling)
mock_image = b"fake_image_data"
prompt = "winter landscape"
litellm_params = GenericLiteLLMParams(api_key="test_key")
-
+
data, files = config.transform_image_edit_request(
model="recraftv3",
prompt=prompt,
image=mock_image,
image_edit_optional_request_params={"strength": 0.7, "n": 1},
litellm_params=litellm_params,
- headers={}
+ headers={},
)
-
+
# Check data structure (like OpenAI but with Recraft additions)
assert data["prompt"] == prompt
assert data["strength"] == 0.7 # Recraft-specific parameter
assert data["model"] == "recraftv3"
-
+
# Check file structure (reuses OpenAI logic)
assert len(files) == 1
assert files[0][0] == "image" # Field name (not image[] like OpenAI)
@@ -603,11 +635,12 @@ def test_recraft_image_edit_config():
async def test_multiple_vs_single_image_edit(sync_mode):
"""Test that both single and multiple image editing work correctly"""
from litellm import image_edit, aimage_edit
+
litellm._turn_on_debug()
-
+
try:
prompt = "Add a soft blue tint to the image(s)"
-
+
# Test single image
if sync_mode:
single_result = image_edit(
@@ -621,10 +654,10 @@ async def test_multiple_vs_single_image_edit(sync_mode):
model="gpt-image-1",
image=SINGLE_TEST_IMAGE,
)
-
+
print("Single image result:", single_result)
ImageResponse.model_validate(single_result)
-
+
# Test multiple images
if sync_mode:
multiple_result = image_edit(
@@ -638,10 +671,10 @@ async def test_multiple_vs_single_image_edit(sync_mode):
model="gpt-image-1",
image=TEST_IMAGES,
)
-
+
print("Multiple images result:", multiple_result)
ImageResponse.model_validate(multiple_result)
-
+
# Both should return valid responses
assert single_result is not None
assert multiple_result is not None
@@ -649,7 +682,7 @@ async def test_multiple_vs_single_image_edit(sync_mode):
assert multiple_result.data is not None
assert len(single_result.data) > 0
assert len(multiple_result.data) > 0
-
+
except litellm.ContentPolicyViolationError as e:
pytest.skip(f"Content policy violation: {e}")
@@ -659,36 +692,37 @@ async def test_multiple_vs_single_image_edit(sync_mode):
async def test_multiple_image_edit_with_different_formats():
"""Test multiple images editing with different file formats and types"""
from litellm import aimage_edit
+
litellm._turn_on_debug()
-
+
try:
prompt = "Create a cohesive artistic style across all images"
-
+
# Test with mixed BytesIO and file objects
mixed_images = [
SINGLE_TEST_IMAGE, # File object
- get_test_images_as_bytesio()[1] # BytesIO object
+ get_test_images_as_bytesio()[1], # BytesIO object
]
-
+
result = await aimage_edit(
prompt=prompt,
model="gpt-image-1",
image=mixed_images,
)
-
+
print("Mixed format images result:", result)
ImageResponse.model_validate(result)
-
+
assert result is not None
assert result.data is not None
assert len(result.data) > 0
-
+
# Save result if available
if result.data and result.data[0].b64_json:
image_bytes = base64.b64decode(result.data[0].b64_json)
with open("test_multiple_image_edit_mixed.png", "wb") as f:
f.write(image_bytes)
-
+
except litellm.ContentPolicyViolationError as e:
pytest.skip(f"Content policy violation: {e}")
@@ -698,7 +732,7 @@ async def test_multiple_image_edit_with_different_formats():
async def test_image_edit_array_handling():
"""Test that the image parameter correctly handles both single items and arrays"""
from litellm import aimage_edit
-
+
# Mock response
mock_response = {
"created": 1589478378,
@@ -706,7 +740,7 @@ async def test_image_edit_array_handling():
{
"b64_json": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg=="
}
- ]
+ ],
}
class MockResponse:
@@ -723,29 +757,26 @@ async def test_image_edit_array_handling():
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = MockResponse(mock_response, 200)
-
+
prompt = "Test prompt"
-
+
# Test 1: Single image (should be converted to list internally)
result1 = await aimage_edit(
prompt=prompt,
model="gpt-image-1",
image=SINGLE_TEST_IMAGE,
)
-
+
# Test 2: Multiple images (already a list)
result2 = await aimage_edit(
prompt=prompt,
model="gpt-image-1",
image=TEST_IMAGES,
)
-
# Both valid calls should succeed
ImageResponse.model_validate(result1)
ImageResponse.model_validate(result2)
-
+
# Verify that both calls were made to the API
assert mock_post.call_count == 2
-
-
diff --git a/tests/image_gen_tests/test_image_generation.py b/tests/image_gen_tests/test_image_generation.py
index 3b4abeeb82f..aca5c41de38 100644
--- a/tests/image_gen_tests/test_image_generation.py
+++ b/tests/image_gen_tests/test_image_generation.py
@@ -121,6 +121,7 @@ class TestVertexImageGeneration(BaseImageGenTest):
class TestVertexAIGeminiImageGeneration(BaseImageGenTest):
"""Test Gemini image generation models (Nano Banana)"""
+
def get_base_image_generation_call_args(self) -> dict:
# comment this when running locally
load_vertex_ai_credentials()
@@ -212,7 +213,9 @@ class TestAimlImageGeneration(BaseImageGenTest):
custom_logger = TestCustomLogger()
litellm.logging_callback_manager._reset_all_callbacks()
litellm.callbacks = [custom_logger]
- base_image_generation_call_args = self.get_base_image_generation_call_args()
+ base_image_generation_call_args = (
+ self.get_base_image_generation_call_args()
+ )
litellm.set_verbose = True
# Pass dummy api_key so validate_environment passes; HTTP is mocked
response = await litellm.aimage_generation(
@@ -229,7 +232,9 @@ class TestAimlImageGeneration(BaseImageGenTest):
# print("response_cost", response._hidden_params["response_cost"])
logged_standard_logging_payload = custom_logger.standard_logging_payload
- print("logged_standard_logging_payload", logged_standard_logging_payload)
+ print(
+ "logged_standard_logging_payload", logged_standard_logging_payload
+ )
assert logged_standard_logging_payload is not None
assert logged_standard_logging_payload["response_cost"] is not None
assert logged_standard_logging_payload["response_cost"] > 0
@@ -244,7 +249,9 @@ class TestAimlImageGeneration(BaseImageGenTest):
response_dict["usage"] = dict(response_dict["usage"])
print("response usage=", response_dict.get("usage"))
- assert response.data is not None # type guard for iteration (base fails here if None)
+ assert (
+ response.data is not None
+ ) # type guard for iteration (base fails here if None)
for d in response.data:
assert isinstance(d, Image)
print("data in response.data", d)
@@ -266,25 +273,27 @@ class TestGoogleImageGen(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
return {"model": "gemini/imagen-4.0-generate-001"}
+
@pytest.mark.skip(reason="Runwayml image generation API only tested locally")
class TestRunwaymlImageGeneration(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
return {"model": "runwayml/gen4_image"}
-class TestAzureOpenAIDalle3(BaseImageGenTest):
- def get_base_image_generation_call_args(self) -> dict:
- return {
- "model": "azure/dall-e-3",
- "api_version": "2024-02-01",
- "api_base": os.getenv("AZURE_API_BASE"),
- "api_key": os.getenv("AZURE_API_KEY"),
- "metadata": {
- "model_info": {
- "base_model": "azure/dall-e-3",
- }
- },
- }
+## AZURE AI DALL-E 3 is deprecated and new deployments cannot be made
+# class TestAzureOpenAIDalle3(BaseImageGenTest):
+# def get_base_image_generation_call_args(self) -> dict:
+# return {
+# "model": "azure/dall-e-3",
+# "api_version": "2024-02-01",
+# "api_base": os.getenv("AZURE_AI_API_BASE"),
+# "api_key": os.getenv("AZURE_AI_API_KEY"),
+# "metadata": {
+# "model_info": {
+# "base_model": "azure/dall-e-3",
+# }
+# },
+# }
@pytest.mark.skip(reason="model EOL")
diff --git a/tests/image_gen_tests/vertex_key.json b/tests/image_gen_tests/vertex_key.json
index 45ca6acc010..800969fb305 100644
--- a/tests/image_gen_tests/vertex_key.json
+++ b/tests/image_gen_tests/vertex_key.json
@@ -1,13 +1,13 @@
{
"type": "service_account",
- "project_id": "pathrise-convert-1606954137718",
+ "project_id": "litellm-ci-cd",
"private_key_id": "",
"private_key": "",
- "client_email": "ci-cd-723@pathrise-convert-1606954137718.iam.gserviceaccount.com",
- "client_id": "109577393201924326488",
+ "client_email": "test-litellm-ci-cd@litellm-ci-cd.iam.gserviceaccount.com",
+ "client_id": "116563532503305622785",
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
"token_uri": "https://oauth2.googleapis.com/token",
"auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs",
- "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/ci-cd-723%40pathrise-convert-1606954137718.iam.gserviceaccount.com",
+ "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/test-litellm-ci-cd%40litellm-ci-cd.iam.gserviceaccount.com",
"universe_domain": "googleapis.com"
}
diff --git a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py
index c1e3fd5072f..062748f3387 100644
--- a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py
+++ b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py
@@ -222,6 +222,27 @@ class TestMigrationSQLIdempotency:
+ "\n".join(violations)
)
+ _DROP_COLUMN_ALLOWLIST = {
+ "20250918083359_drop_spec_version_column_from_mcp_table",
+ "20260213170952_access_group_change_to_model_name",
+ "20260224203854_add_agent_object_permissions_table",
+ }
+
+ def test_no_drop_column_statements(self, all_migrations):
+ """Migrations must not drop columns — dropping columns is destructive
+ and can break running application instances during rolling deploys."""
+ violations = []
+ for migration_name, sql in all_migrations:
+ if migration_name in self._DROP_COLUMN_ALLOWLIST:
+ continue
+ for line_num, line in enumerate(sql.splitlines(), 1):
+ if re.search(r"DROP\s+COLUMN", line, re.IGNORECASE):
+ violations.append(f" {migration_name}:{line_num}: {line.strip()}")
+ assert not violations, (
+ "DROP COLUMN found in migrations (destructive, not allowed):\n"
+ + "\n".join(violations)
+ )
+
def test_drop_index_uses_if_exists(self, all_migrations):
"""DROP INDEX statements must use IF EXISTS"""
violations = []
diff --git a/tests/litellm_utils_tests/test_health_check.py b/tests/litellm_utils_tests/test_health_check.py
index b048590d51a..d014b3f72af 100644
--- a/tests/litellm_utils_tests/test_health_check.py
+++ b/tests/litellm_utils_tests/test_health_check.py
@@ -21,9 +21,9 @@ async def test_azure_health_check():
model_params={
"model": "azure/gpt-4.1-mini",
"messages": [{"role": "user", "content": "Hey, how's it going?"}],
- "api_key": os.getenv("AZURE_API_KEY"),
- "api_base": os.getenv("AZURE_API_BASE"),
- "api_version": os.getenv("AZURE_API_VERSION"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
+ "api_version": os.getenv("AZURE_AI_API_VERSION"),
}
)
print(f"response: {response}")
@@ -51,9 +51,9 @@ async def test_azure_embedding_health_check():
response = await litellm.ahealth_check(
model_params={
"model": "azure/text-embedding-ada-002",
- "api_key": os.getenv("AZURE_API_KEY"),
- "api_base": os.getenv("AZURE_API_BASE"),
- "api_version": os.getenv("AZURE_API_VERSION"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
+ "api_version": os.getenv("AZURE_AI_API_VERSION"),
},
input=["test for litellm"],
mode="embedding",
@@ -83,7 +83,9 @@ async def test_openai_img_gen_health_check():
# asyncio.run(test_openai_img_gen_health_check())
-@pytest.mark.skip(reason="Azure DALL-E 3 model deployment is deprecated (410 ModelDeprecated)")
+@pytest.mark.skip(
+ reason="Azure DALL-E 3 model deployment is deprecated (410 ModelDeprecated)"
+)
@pytest.mark.asyncio
async def test_azure_img_gen_health_check():
"""
@@ -98,8 +100,8 @@ async def test_azure_img_gen_health_check():
response = await litellm.ahealth_check(
model_params={
"model": "azure/dall-e-3",
- "api_base": os.getenv("AZURE_API_BASE"),
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
},
mode="image_generation",
prompt="cute baby sea otter",
@@ -244,35 +246,6 @@ async def test_audio_transcription_health_check():
print(response)
-@pytest.mark.asyncio
-@pytest.mark.parametrize(
- "model", ["azure/gpt-4o-realtime-preview", "openai/gpt-4o-realtime-preview"]
-)
-async def test_async_realtime_health_check(model, mocker):
- """
- Test Health Check with Valid models passes
-
- """
- mock_websocket = AsyncMock()
- mock_connect = AsyncMock().__aenter__.return_value = mock_websocket
- mocker.patch("websockets.connect", return_value=mock_connect)
-
- litellm.set_verbose = True
- model_params = {
- "model": model,
- }
- if model == "azure/gpt-4o-realtime-preview":
- model_params["api_base"] = os.getenv("AZURE_REALTIME_API_BASE")
- model_params["api_key"] = os.getenv("AZURE_REALTIME_API_KEY")
- model_params["api_version"] = os.getenv("AZURE_REALTIME_API_VERSION")
- response = await litellm.ahealth_check(
- model_params=model_params,
- mode="realtime",
- )
- print(response)
- assert response == {}
-
-
def test_update_litellm_params_for_health_check():
"""
Test if _update_litellm_params_for_health_check correctly:
@@ -500,7 +473,9 @@ async def test_perform_health_check_filters_by_model_id():
async def mock_perform_health_check(m_list, details=True, **kwargs):
captured_list.append(m_list)
- return [{"model": "gpt-4", "api_key": m_list[0]["litellm_params"]["api_key"]}], []
+ return [
+ {"model": "gpt-4", "api_key": m_list[0]["litellm_params"]["api_key"]}
+ ], []
with patch(
"litellm.proxy.health_check._perform_health_check",
@@ -657,7 +632,8 @@ async def test_health_check_creates_only_bounded_initial_tasks():
return real_create_task(coro)
with patch("litellm.ahealth_check", side_effect=mock_health_check), patch(
- "litellm.proxy.health_check.asyncio.create_task", side_effect=tracked_create_task
+ "litellm.proxy.health_check.asyncio.create_task",
+ side_effect=tracked_create_task,
):
perform_task = real_create_task(
_perform_health_check(model_list, max_concurrency=2)
diff --git a/tests/litellm_utils_tests/vertex_key.json b/tests/litellm_utils_tests/vertex_key.json
index 45ca6acc010..800969fb305 100644
--- a/tests/litellm_utils_tests/vertex_key.json
+++ b/tests/litellm_utils_tests/vertex_key.json
@@ -1,13 +1,13 @@
{
"type": "service_account",
- "project_id": "pathrise-convert-1606954137718",
+ "project_id": "litellm-ci-cd",
"private_key_id": "",
"private_key": "",
- "client_email": "ci-cd-723@pathrise-convert-1606954137718.iam.gserviceaccount.com",
- "client_id": "109577393201924326488",
+ "client_email": "test-litellm-ci-cd@litellm-ci-cd.iam.gserviceaccount.com",
+ "client_id": "116563532503305622785",
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
"token_uri": "https://oauth2.googleapis.com/token",
"auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs",
- "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/ci-cd-723%40pathrise-convert-1606954137718.iam.gserviceaccount.com",
+ "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/test-litellm-ci-cd%40litellm-ci-cd.iam.gserviceaccount.com",
"universe_domain": "googleapis.com"
}
diff --git a/tests/llm_responses_api_testing/test_azure_responses_api.py b/tests/llm_responses_api_testing/test_azure_responses_api.py
index 5e876ce0848..fed9e9e11f0 100644
--- a/tests/llm_responses_api_testing/test_azure_responses_api.py
+++ b/tests/llm_responses_api_testing/test_azure_responses_api.py
@@ -25,11 +25,11 @@ class TestAzureResponsesAPITest(BaseResponsesAPITest):
return {
"model": "azure/gpt-4.1-mini",
"truncation": "auto",
- "api_base": os.getenv("AZURE_API_BASE"),
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": "2025-03-01-preview",
}
-
+
def get_advanced_model_for_shell_tool(self) -> Optional[str]:
"""If specified, overrides the model used by test_responses_api_shell_tool_streaming_sees_shell_output (e.g. openai/gpt-5.2 for shell support)."""
return "azure/gpt-5-mini"
@@ -45,8 +45,8 @@ async def test_azure_responses_api_preview_api_version():
model="azure/gpt-5-mini",
truncation="auto",
api_version="preview",
- api_base=os.getenv("AZURE_API_BASE"),
- api_key=os.getenv("AZURE_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
input="Hello, can you tell me a short joke?",
)
@@ -108,7 +108,9 @@ async def test_azure_responses_api_status_error():
"role": "assistant",
"type": "message",
"status": "completed",
- "content": [{"type": "output_text", "text": "Here's an interesting fact."}],
+ "content": [
+ {"type": "output_text", "text": "Here's an interesting fact."}
+ ],
}
],
}
@@ -124,7 +126,7 @@ async def test_azure_responses_api_status_error():
captured_request_body = json.loads(kwargs["data"])
import httpx
-
+
# Create a proper httpx Response object
response_content = json.dumps(mock_response_data).encode("utf-8")
response = httpx.Response(
@@ -149,18 +151,17 @@ async def test_azure_responses_api_status_error():
)
# Verify that 'status' field is not present in any of the input messages
- print("Final request body:", json.dumps(captured_request_body, indent=4, default=str))
+ print(
+ "Final request body:", json.dumps(captured_request_body, indent=4, default=str)
+ )
assert "input" in captured_request_body, "Request body should contain 'input' field"
-
+
expected_input = [
- {
- "content": "tell me an interesting fact",
- "role": "user"
- },
+ {"content": "tell me an interesting fact", "role": "user"},
{
"id": "rs_0ab687487834d9df0068e462a1b2d88197aabbc832c9ba5316",
"summary": [],
- "type": "reasoning"
+ "type": "reasoning",
},
{
"id": "msg_0ab687487834d9df0068e462a1df188197b74b1eef05102c18",
@@ -169,18 +170,15 @@ async def test_azure_responses_api_status_error():
"annotations": [],
"text": "very good morning",
"type": "output_text",
- "logprobs": []
+ "logprobs": [],
}
],
"role": "assistant",
- "type": "message"
+ "type": "message",
},
- {
- "role": "user",
- "content": "tell me another"
- }
+ {"role": "user", "content": "tell me another"},
]
-
+
assert captured_request_body["input"] == expected_input, (
f"Request body input should match expected format without 'status' field.\n"
f"Expected: {json.dumps(expected_input, indent=2)}\n"
@@ -193,9 +191,9 @@ async def test_azure_responses_api_headers_with_llm_provider_prefix():
"""
Test that Azure-specific headers like 'x-request-id' and 'apim-request-id'
are properly forwarded with 'llm_provider-' prefix in response._hidden_params["headers"].
-
+
Issue: https://github.com/BerriAI/litellm/issues/16538
-
+
The fix ensures that processed headers (with llm_provider- prefix) are stored
in response._hidden_params["headers"] instead of additional_headers, making them
accessible via completion.headers in the same way as the completion API.
@@ -253,12 +251,12 @@ async def test_azure_responses_api_headers_with_llm_provider_prefix():
# Check that the response has the expected headers structure
assert hasattr(response, "_hidden_params"), "Response should have _hidden_params"
- assert "additional_headers" in response._hidden_params, (
- "Response _hidden_params should contain 'additional_headers' with the LLM provider headers"
- )
+ assert (
+ "additional_headers" in response._hidden_params
+ ), "Response _hidden_params should contain 'additional_headers' with the LLM provider headers"
headers = response._hidden_params["additional_headers"]
-
+
# Verify that Azure-specific headers are present with llm_provider- prefix
assert "llm_provider-x-request-id" in headers, (
f"Response should contain 'llm_provider-x-request-id' header. "
@@ -268,12 +266,17 @@ async def test_azure_responses_api_headers_with_llm_provider_prefix():
f"Response should contain 'llm_provider-apim-request-id' header. "
f"Headers: {list(headers.keys())}"
)
-
+
# Verify the header values match
- assert headers["llm_provider-x-request-id"] == "12086715-aca3-4006-a29f-2f1e1d552043"
- assert headers["llm_provider-apim-request-id"] == "25664b0d-cf4b-4e10-8d27-c7272e7efd49"
+ assert (
+ headers["llm_provider-x-request-id"] == "12086715-aca3-4006-a29f-2f1e1d552043"
+ )
+ assert (
+ headers["llm_provider-apim-request-id"]
+ == "25664b0d-cf4b-4e10-8d27-c7272e7efd49"
+ )
assert headers["llm_provider-x-ms-region"] == "Sweden Central"
-
+
# Also verify openai-compatible headers are included
assert "x-ratelimit-limit-tokens" in headers
assert "x-ratelimit-remaining-tokens" in headers
diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py
index 8630ba65610..b6a98e4b2c7 100644
--- a/tests/llm_translation/test_anthropic_completion.py
+++ b/tests/llm_translation/test_anthropic_completion.py
@@ -1350,6 +1350,9 @@ def test_anthropic_text_editor():
@pytest.mark.parametrize("spec", ["anthropic", "openai"])
+@pytest.mark.skipif(
+ os.getenv("ZAPIER_CI_CD_MCP_TOKEN") is None, reason="ZAPIER_CI_CD_MCP_TOKEN not set"
+)
def test_anthropic_mcp_server_tool_use(spec: str):
litellm._turn_on_debug()
@@ -1391,6 +1394,9 @@ def test_anthropic_mcp_server_tool_use(spec: str):
@pytest.mark.parametrize(
"model", ["openai/gpt-4.1", "anthropic/claude-sonnet-4-20250514"]
)
+@pytest.mark.skipif(
+ os.getenv("ZAPIER_CI_CD_MCP_TOKEN") is None, reason="ZAPIER_CI_CD_MCP_TOKEN not set"
+)
def test_anthropic_mcp_server_responses_api(model: str):
from litellm import responses
@@ -1800,3 +1806,81 @@ def test_anthropic_structured_output_chat_completion_api():
)
assert response is not None
print(f"response: {response}")
+
+
+def _make_transform_request(optional_params: dict, litellm_params: dict) -> dict:
+ from litellm.llms.anthropic.chat.transformation import AnthropicConfig
+
+ return AnthropicConfig().transform_request(
+ model="claude-3-5-sonnet-20241022",
+ messages=[{"role": "user", "content": "hi"}],
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ headers={},
+ )
+
+
+def test_metadata_only_user_id_passes_through():
+ """metadata with only user_id is forwarded as-is."""
+ data = _make_transform_request(
+ optional_params={"metadata": {"user_id": "abc123"}},
+ litellm_params={},
+ )
+ assert data.get("metadata") == {"user_id": "abc123"}
+
+
+def test_metadata_extra_keys_are_stripped():
+ """Extra keys in metadata are removed; only user_id is sent."""
+ data = _make_transform_request(
+ optional_params={"metadata": {"user_id": "abc123", "extra_key": "val"}},
+ litellm_params={},
+ )
+ assert data.get("metadata") == {"user_id": "abc123"}
+
+
+def test_metadata_without_user_id_is_dropped():
+ """metadata with no user_id is removed entirely."""
+ data = _make_transform_request(
+ optional_params={"metadata": {"only_other_key": "val"}},
+ litellm_params={},
+ )
+ assert "metadata" not in data
+
+
+def test_metadata_user_id_from_litellm_params_strips_extras():
+ """user_id from litellm_params metadata is extracted; extra keys are not forwarded."""
+ data = _make_transform_request(
+ optional_params={},
+ litellm_params={"metadata": {"user_id": "abc123", "trace_id": "xyz"}},
+ )
+ assert data.get("metadata") == {"user_id": "abc123"}
+
+
+def test_metadata_filter_applies_to_vertex_anthropic():
+ """VertexAIAnthropicConfig inherits the metadata filter."""
+ from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.transformation import (
+ VertexAIAnthropicConfig,
+ )
+
+ data = VertexAIAnthropicConfig().transform_request(
+ model="claude-3-5-sonnet-20241022",
+ messages=[{"role": "user", "content": "hi"}],
+ optional_params={"metadata": {"user_id": "u1", "extra": "drop_me"}},
+ litellm_params={},
+ headers={},
+ )
+ assert data.get("metadata") == {"user_id": "u1"}
+
+
+def test_metadata_filter_applies_to_azure_anthropic():
+ """AzureAnthropicConfig inherits the metadata filter."""
+ from litellm.llms.azure_ai.anthropic.transformation import AzureAnthropicConfig
+
+ data = AzureAnthropicConfig().transform_request(
+ model="claude-3-5-sonnet-20241022",
+ messages=[{"role": "user", "content": "hi"}],
+ optional_params={"metadata": {"user_id": "u2", "extra": "drop_me"}},
+ litellm_params={},
+ headers={},
+ )
+ assert data.get("metadata") == {"user_id": "u2"}
diff --git a/tests/llm_translation/test_azure_ai.py b/tests/llm_translation/test_azure_ai.py
index 972ba34a179..b6f53e4e955 100644
--- a/tests/llm_translation/test_azure_ai.py
+++ b/tests/llm_translation/test_azure_ai.py
@@ -188,35 +188,6 @@ def test_azure_ai_services_with_api_version():
)
-@pytest.mark.skip(reason="Skipping due to cohere ssl issues")
-def test_completion_azure_ai_command_r():
- try:
- import os
-
- litellm.set_verbose = True
-
- os.environ["AZURE_AI_API_BASE"] = os.getenv("AZURE_COHERE_API_BASE", "")
- os.environ["AZURE_AI_API_KEY"] = os.getenv("AZURE_COHERE_API_KEY", "")
-
- response = completion(
- model="azure_ai/command-r-plus",
- messages=[
- {
- "role": "user",
- "content": [
- {"type": "text", "text": "What is the meaning of life?"}
- ],
- }
- ],
- ) # type: ignore
-
- assert "azure_ai" in response.model
- except litellm.Timeout as e:
- pass
- except Exception as e:
- pytest.fail(f"Error occurred: {e}")
-
-
def test_azure_deepseek_reasoning_content():
import json
@@ -283,8 +254,8 @@ async def test_azure_ai_request_format():
litellm._turn_on_debug()
# Set up the test parameters
- api_key = os.getenv("AZURE_API_KEY")
- api_base = os.getenv("AZURE_API_BASE")
+ api_key = os.getenv("AZURE_AI_API_KEY")
+ api_base = os.getenv("AZURE_AI_API_BASE")
model = "azure_ai/gpt-4.1-mini"
messages = [
{"role": "user", "content": "hi"},
@@ -310,17 +281,17 @@ async def test_azure_gpt5_reasoning(model):
messages=[{"role": "user", "content": "What is the capital of France?"}],
reasoning_effort="minimal",
max_tokens=10,
- api_base=os.getenv("AZURE_API_BASE"),
- api_key=os.getenv("AZURE_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
)
print("response: ", response)
assert response.choices[0].message.content is not None
-
def test_completion_azure():
try:
from litellm import completion_cost
+
litellm.set_verbose = False
## Test azure call
response = completion(
@@ -331,7 +302,7 @@ def test_completion_azure():
"content": "Hello, how are you?",
}
],
- api_key="os.environ/AZURE_API_KEY",
+ api_key="os.environ/AZURE_AI_API_KEY",
)
print(f"response: {response}")
print(f"response hidden params: {response._hidden_params}")
@@ -358,7 +329,7 @@ def test_completion_azure_ai_gpt_4o_with_flexible_api_base(api_base):
response = completion(
model="azure_ai/gpt-4.1-mini",
api_base=api_base,
- api_key=os.getenv("AZURE_API_KEY"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
messages=[{"role": "user", "content": "What is the meaning of life?"}],
)
@@ -374,18 +345,20 @@ async def test_azure_ai_model_router():
"""
Test Azure AI model router non-streaming response cost tracking.
Verifies that the flat cost of $0.14 per M input tokens is applied.
-
+
Tests the pattern: azure_ai/model_router/
Where deployment-name is the Azure deployment (e.g., "azure-model-router").
The model_router prefix is stripped before sending to Azure API.
"""
- from litellm.llms.azure_ai.cost_calculator import calculate_azure_model_router_flat_cost
-
+ from litellm.llms.azure_ai.cost_calculator import (
+ calculate_azure_model_router_flat_cost,
+ )
+
litellm._turn_on_debug()
response = await litellm.acompletion(
model="azure_ai/model_router/azure-model-router",
messages=[{"role": "user", "content": "hi who is this"}],
- api_base="https://ishaa-mh6uutut-swedencentral.cognitiveservices.azure.com/openai/v1/",
+ api_base=os.getenv("AZURE_MODEL_ROUTER_API_BASE"),
api_key=os.getenv("AZURE_MODEL_ROUTER_API_KEY"),
)
print("response: ", response)
@@ -394,23 +367,22 @@ async def test_azure_ai_model_router():
tracked_cost = response._hidden_params["response_cost"]
assert tracked_cost > 0
print("Tracked cost: ", tracked_cost)
-
+
# Verify flat cost is included using the helper function
usage = response.usage
if usage and usage.prompt_tokens:
expected_flat_cost = calculate_azure_model_router_flat_cost(
- model="model_router/azure-model-router",
- prompt_tokens=usage.prompt_tokens
+ model="model_router/azure-model-router", prompt_tokens=usage.prompt_tokens
)
print(f"Prompt tokens: {usage.prompt_tokens}")
print(f"Expected flat cost: ${expected_flat_cost:.9f}")
print(f"Total tracked cost: ${tracked_cost:.9f}")
-
+
# Total cost should be at least the flat cost
- assert tracked_cost >= expected_flat_cost, (
- f"Cost ${tracked_cost:.9f} should be >= flat cost ${expected_flat_cost:.9f}"
- )
-
+ assert (
+ tracked_cost >= expected_flat_cost
+ ), f"Cost ${tracked_cost:.9f} should be >= flat cost ${expected_flat_cost:.9f}"
+
# Verify the flat cost is non-zero
assert expected_flat_cost > 0, "Flat cost should be greater than 0"
@@ -445,15 +417,20 @@ async def test_azure_ai_model_router_streaming_model_in_chunk():
# The model should NOT be azure-model-router (the request model)
# It should be the actual model from the response (e.g., gpt-4.1-nano, gpt-5-nano, etc.)
for model in chunks_with_model:
- assert model != "azure-model-router", f"Chunk model should be actual model, not request model. Got: {model}"
+ assert (
+ model != "azure-model-router"
+ ), f"Chunk model should be actual model, not request model. Got: {model}"
# The actual model should be a real model name like gpt-4.1-nano, gpt-5-nano, etc.
print(f"Verified chunk has actual model: {model}")
-class AzureModelRouterStreamingCallback(litellm.integrations.custom_logger.CustomLogger):
+class AzureModelRouterStreamingCallback(
+ litellm.integrations.custom_logger.CustomLogger
+):
"""
Custom callback to capture streaming cost tracking for Azure Model Router.
"""
+
def __init__(self):
self.standard_logging_payload = None
self.response_cost = None
@@ -466,17 +443,21 @@ class AzureModelRouterStreamingCallback(litellm.integrations.custom_logger.Custo
self.async_success_called = True
self.standard_logging_payload = kwargs.get("standard_logging_object")
self.complete_streaming_response = kwargs.get("complete_streaming_response")
-
+
if self.standard_logging_payload:
self.response_cost = self.standard_logging_payload.get("response_cost")
- print(f"standard_logging_payload model: {self.standard_logging_payload.get('model')}")
+ print(
+ f"standard_logging_payload model: {self.standard_logging_payload.get('model')}"
+ )
print(f"standard_logging_payload response_cost: {self.response_cost}")
-
+
if self.complete_streaming_response:
- print(f"complete_streaming_response model: {self.complete_streaming_response.model}")
- print(f"complete_streaming_response usage: {self.complete_streaming_response.usage}")
-
-
+ print(
+ f"complete_streaming_response model: {self.complete_streaming_response.model}"
+ )
+ print(
+ f"complete_streaming_response usage: {self.complete_streaming_response.usage}"
+ )
@pytest.mark.asyncio
@@ -504,10 +485,16 @@ async def test_azure_ai_model_router_streaming_cost_with_stream_options():
full_response = ""
chunks_with_model = []
async for chunk in response:
- print(f"Chunk: model={chunk.model}, choices={len(chunk.choices) if chunk.choices else 0}")
+ print(
+ f"Chunk: model={chunk.model}, choices={len(chunk.choices) if chunk.choices else 0}"
+ )
if chunk.model:
chunks_with_model.append(chunk.model)
- if chunk.choices and chunk.choices[0].delta and chunk.choices[0].delta.content:
+ if (
+ chunk.choices
+ and chunk.choices[0].delta
+ and chunk.choices[0].delta.content
+ ):
full_response += chunk.choices[0].delta.content
print(f"Full streamed response: {full_response}")
@@ -515,27 +502,42 @@ async def test_azure_ai_model_router_streaming_cost_with_stream_options():
# Give async logging time to complete
import asyncio
+
await asyncio.sleep(1)
# Verify callback was called
- assert test_callback.async_success_called is True, "async_log_success_event was not called"
- assert test_callback.standard_logging_payload is not None, "standard_logging_payload is None"
+ assert (
+ test_callback.async_success_called is True
+ ), "async_log_success_event was not called"
+ assert (
+ test_callback.standard_logging_payload is not None
+ ), "standard_logging_payload is None"
# Check response cost
print(f"Final response_cost: {test_callback.response_cost}")
-
+
# The first chunk may have the request model (azure-model-router) because it's created
# before the API response is received. Subsequent chunks should have the actual model.
# At least some chunks should have the actual model (not azure-model-router)
- actual_model_chunks = [m for m in chunks_with_model if m != "azure-model-router"]
- assert len(actual_model_chunks) > 0, "No chunks had the actual model from the API response"
+ actual_model_chunks = [
+ m for m in chunks_with_model if m != "azure-model-router"
+ ]
+ assert (
+ len(actual_model_chunks) > 0
+ ), "No chunks had the actual model from the API response"
print(f"Chunks with actual model: {actual_model_chunks}")
# Verify response cost is tracked - this is the main goal of this test
- assert test_callback.response_cost is not None, "response_cost is None with stream_options"
- assert test_callback.response_cost > 0, f"response_cost should be > 0, got {test_callback.response_cost}"
- print(f"Streaming cost tracking with stream_options passed. Cost: {test_callback.response_cost}")
+ assert (
+ test_callback.response_cost is not None
+ ), "response_cost is None with stream_options"
+ assert (
+ test_callback.response_cost > 0
+ ), f"response_cost should be > 0, got {test_callback.response_cost}"
+ print(
+ f"Streaming cost tracking with stream_options passed. Cost: {test_callback.response_cost}"
+ )
finally:
litellm.logging_callback_manager._reset_all_callbacks()
- litellm.callbacks = []
\ No newline at end of file
+ litellm.callbacks = []
diff --git a/tests/llm_translation/test_azure_o_series.py b/tests/llm_translation/test_azure_o_series.py
index 925453e68c1..181d3b677d2 100644
--- a/tests/llm_translation/test_azure_o_series.py
+++ b/tests/llm_translation/test_azure_o_series.py
@@ -24,9 +24,9 @@ class TestAzureOpenAIO3Mini(BaseOSeriesModelsTest, BaseLLMChatTest):
litellm.in_memory_llm_clients_cache.flush_cache()
return {
"model": "azure/o3-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
- "api_base": os.getenv("AZURE_API_BASE"),
- "api_version": "2024-12-01-preview"
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
+ "api_version": "2024-12-01-preview",
}
def get_client(self):
@@ -187,13 +187,31 @@ async def test_azure_o1_series_response_format_extra_params():
litellm.set_verbose = True
client = AsyncAzureOpenAI(
- api_key="fake-api-key",
- base_url="https://openai-prod-test.openai.azure.com/openai/deployments/o1/chat/completions?api-version=2025-01-01-preview",
- api_version="2025-01-01-preview"
+ api_key="fake-api-key",
+ base_url="https://openai-prod-test.openai.azure.com/openai/deployments/o1/chat/completions?api-version=2025-01-01-preview",
+ api_version="2025-01-01-preview",
)
- tools = [{'type': 'function', 'function': {'name': 'get_current_time', 'description': 'Get the current time in a given location.', 'parameters': {'type': 'object', 'properties': {'location': {'type': 'string', 'description': 'The city name, e.g. San Francisco'}}, 'required': ['location']}}}]
- response_format = {'type': 'json_object'}
+ tools = [
+ {
+ "type": "function",
+ "function": {
+ "name": "get_current_time",
+ "description": "Get the current time in a given location.",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "location": {
+ "type": "string",
+ "description": "The city name, e.g. San Francisco",
+ }
+ },
+ "required": ["location"],
+ },
+ },
+ }
+ ]
+ response_format = {"type": "json_object"}
tool_choice = "auto"
with patch.object(
client.chat.completions.with_raw_response, "create"
@@ -208,7 +226,7 @@ async def test_azure_o1_series_response_format_extra_params():
messages=[{"role": "user", "content": "Hello! return a json object"}],
tools=tools,
response_format=response_format,
- tool_choice=tool_choice
+ tool_choice=tool_choice,
)
except Exception as e:
print(f"Error: {e}")
@@ -220,7 +238,3 @@ async def test_azure_o1_series_response_format_extra_params():
assert request_body["tools"] == tools
assert request_body["response_format"] == response_format
assert request_body["tool_choice"] == tool_choice
-
-
-
-
diff --git a/tests/llm_translation/test_azure_openai.py b/tests/llm_translation/test_azure_openai.py
index 1da380b57a2..f3dc954020c 100644
--- a/tests/llm_translation/test_azure_openai.py
+++ b/tests/llm_translation/test_azure_openai.py
@@ -208,8 +208,8 @@ class TestAzureEmbedding(BaseLLMEmbeddingTest):
def get_base_embedding_call_args(self) -> dict:
return {
"model": "azure/text-embedding-ada-002",
- "api_key": os.getenv("AZURE_API_KEY"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
}
def get_custom_llm_provider(self) -> litellm.LlmProviders:
@@ -618,8 +618,8 @@ def test_azure_safety_result():
response = completion(
model="azure/gpt-4.1-mini",
- api_key=os.getenv("AZURE_API_KEY"),
- api_base=os.getenv("AZURE_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
api_version="2024-12-01-preview",
messages=[{"role": "user", "content": "Hello world"}],
)
@@ -671,6 +671,8 @@ def test_completion_azure_deployment_id():
)
# Add any assertions here to check the response
print(response)
+
+
def test_azure_with_content_safety_error():
"""
Verify user can access innererror from the Azure OpenAI exception
@@ -679,55 +681,55 @@ def test_azure_with_content_safety_error():
from litellm.exceptions import ContentPolicyViolationError
from litellm.litellm_core_utils.exception_mapping_utils import exception_type
from unittest.mock import MagicMock
-
- mock_exception = Exception("The response was filtered due to the prompt triggering Azure OpenAI's content management policy")
+
+ mock_exception = Exception(
+ "The response was filtered due to the prompt triggering Azure OpenAI's content management policy"
+ )
mock_exception.body = {
"innererror": {
"code": "ResponsibleAIPolicyViolation",
"content_filter_result": {
- "hate": {
- "filtered": False,
- "severity": "safe"
- },
- "jailbreak": {
- "filtered": False,
- "detected": False
- },
- "self_harm": {
- "filtered": False,
- "severity": "safe"
- },
- "sexual": {
- "filtered": False,
- "severity": "safe"
- },
- "violence": {
- "filtered": True,
- "severity": "high"
- }
- }
+ "hate": {"filtered": False, "severity": "safe"},
+ "jailbreak": {"filtered": False, "detected": False},
+ "self_harm": {"filtered": False, "severity": "safe"},
+ "sexual": {"filtered": False, "severity": "safe"},
+ "violence": {"filtered": True, "severity": "high"},
+ },
}
}
-
+
mock_response = MagicMock()
mock_response.status_code = 400
mock_exception.response = mock_response
-
+
with pytest.raises(ContentPolicyViolationError) as exc_info:
exception_type(
model="azure/gpt-4o-new-test",
original_exception=mock_exception,
- custom_llm_provider="azure"
+ custom_llm_provider="azure",
)
-
+
e = exc_info.value
print("got exception=", e)
assert e.provider_specific_fields is not None
print("got provider_specific_fields=", e.provider_specific_fields)
assert e.provider_specific_fields.get("innererror") is not None
- assert e.provider_specific_fields["innererror"]["code"] == "ResponsibleAIPolicyViolation"
- assert e.provider_specific_fields["innererror"]["content_filter_result"]["violence"]["filtered"] is True
- assert e.provider_specific_fields["innererror"]["content_filter_result"]["violence"]["severity"] == "high"
+ assert (
+ e.provider_specific_fields["innererror"]["code"]
+ == "ResponsibleAIPolicyViolation"
+ )
+ assert (
+ e.provider_specific_fields["innererror"]["content_filter_result"]["violence"][
+ "filtered"
+ ]
+ is True
+ )
+ assert (
+ e.provider_specific_fields["innererror"]["content_filter_result"]["violence"][
+ "severity"
+ ]
+ == "high"
+ )
def test_azure_openai_with_prompt_cache_key():
@@ -737,9 +739,9 @@ def test_azure_openai_with_prompt_cache_key():
litellm._turn_on_debug()
response = litellm.completion(
model="azure/gpt-4.1-mini",
- api_key=os.getenv("AZURE_API_KEY"),
- api_base=os.getenv("AZURE_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
api_version="2024-12-01-preview",
messages=[{"role": "user", "content": "What is the weather in San Francisco?"}],
prompt_cache_key="test_streaming_azure_openai",
- )
\ No newline at end of file
+ )
diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py
index b71e4e51877..f57c77db54d 100644
--- a/tests/llm_translation/test_bedrock_completion.py
+++ b/tests/llm_translation/test_bedrock_completion.py
@@ -217,60 +217,6 @@ def test_completion_bedrock_claude_external_client_auth():
# test_completion_bedrock_claude_external_client_auth()
-@pytest.mark.skip(reason="Expired token, need to renew")
-def test_completion_bedrock_claude_sts_client_auth():
- print("\ncalling bedrock claude external client auth")
- import os
-
- aws_access_key_id = os.environ["AWS_TEMP_ACCESS_KEY_ID"]
- aws_secret_access_key = os.environ["AWS_TEMP_SECRET_ACCESS_KEY"]
- aws_region_name = os.environ["AWS_REGION_NAME"]
- aws_role_name = os.environ["AWS_TEMP_ROLE_NAME"]
-
- try:
- import boto3
-
- litellm.set_verbose = True
-
- response = completion(
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- messages=messages,
- max_tokens=10,
- temperature=0.1,
- aws_region_name=aws_region_name,
- aws_access_key_id=aws_access_key_id,
- aws_secret_access_key=aws_secret_access_key,
- aws_role_name=aws_role_name,
- aws_session_name="my-test-session",
- )
-
- response = embedding(
- model="cohere.embed-multilingual-v3",
- input=["hello world"],
- aws_region_name="us-east-1",
- aws_access_key_id=aws_access_key_id,
- aws_secret_access_key=aws_secret_access_key,
- aws_role_name=aws_role_name,
- aws_session_name="my-test-session",
- )
-
- response = completion(
- model="gpt-3.5-turbo",
- messages=messages,
- aws_region_name="us-east-1",
- aws_access_key_id=aws_access_key_id,
- aws_secret_access_key=aws_secret_access_key,
- aws_role_name=aws_role_name,
- aws_session_name="my-test-session",
- )
- # Add any assertions here to check the response
- print(response)
- except RateLimitError:
- pass
- except Exception as e:
- pytest.fail(f"Error occurred: {e}")
-
-
@pytest.fixture()
def bedrock_session_token_creds():
print("\ncalling oidc auto to get aws_session_token credentials")
@@ -3413,7 +3359,8 @@ def test_bedrock_openai_imported_model():
print(f"URL: {url}")
assert "bedrock-runtime.us-east-1.amazonaws.com" in url
assert (
- "arn:aws:bedrock:us-east-1:117159858402:imported-model%2Fm4gc1mrfuddy" in url
+ "arn:aws:bedrock:us-east-1:117159858402:imported-model%2Fm4gc1mrfuddy"
+ in url
)
assert "/invoke" in url
@@ -3850,10 +3797,12 @@ def test_bedrock_openai_error_handling():
assert exc_info.value.status_code == 422
print("✓ Error handling works correctly")
+
# ============================================================================
# Nova Grounding (web_search_options) Unit Tests (Mocked)
# ============================================================================
+
def test_bedrock_nova_grounding_web_search_options_non_streaming():
"""
Unit test for Nova grounding using web_search_options parameter (non-streaming).
@@ -3907,7 +3856,9 @@ def test_bedrock_nova_grounding_web_search_options_non_streaming():
break
assert system_tool_found, "systemTool with nova_grounding should be present"
- print(f"✓ web_search_options correctly transformed to systemTool (non-streaming)")
+ print(
+ f"✓ web_search_options correctly transformed to systemTool (non-streaming)"
+ )
def test_bedrock_nova_grounding_with_function_tools():
@@ -3987,7 +3938,9 @@ def test_bedrock_nova_grounding_with_function_tools():
assert tool["systemTool"]["name"] == "nova_grounding"
system_tool_found = True
- assert function_tool_found, "Function tool (get_stock_price) should be present"
+ assert (
+ function_tool_found
+ ), "Function tool (get_stock_price) should be present"
assert system_tool_found, "systemTool (nova_grounding) should be present"
print(f"✓ Both function tools and web_search_options correctly combined")
@@ -4092,10 +4045,12 @@ def test_bedrock_nova_grounding_request_transformation():
mock_post.return_value = MagicMock(
status_code=200,
json=lambda: {
- "output": {"message": {"role": "assistant", "content": [{"text": "Test"}]}},
+ "output": {
+ "message": {"role": "assistant", "content": [{"text": "Test"}]}
+ },
"stopReason": "end_turn",
- "usage": {"inputTokens": 10, "outputTokens": 5}
- }
+ "usage": {"inputTokens": 10, "outputTokens": 5},
+ },
)
try:
diff --git a/tests/llm_translation/test_clarifai_completion.py b/tests/llm_translation/test_clarifai_completion.py
deleted file mode 100644
index 5080413f2e5..00000000000
--- a/tests/llm_translation/test_clarifai_completion.py
+++ /dev/null
@@ -1,109 +0,0 @@
-import sys, os
-import traceback
-from dotenv import load_dotenv
-import asyncio, logging
-
-load_dotenv()
-import os, io
-
-sys.path.insert(
- 0, os.path.abspath("../..")
-) # Adds the parent directory to the system path
-import pytest
-import litellm
-from litellm import (
- embedding,
- completion,
- acompletion,
- acreate,
- completion_cost,
- Timeout,
- ModelResponse,
-)
-from litellm import RateLimitError
-
-# litellm.num_retries = 3
-litellm.cache = None
-litellm.success_callback = []
-user_message = "Write a short poem about the sky"
-messages = [{"content": user_message, "role": "user"}]
-
-
-@pytest.fixture(autouse=True)
-def reset_callbacks():
- print("\npytest fixture - resetting callbacks")
- litellm.success_callback = []
- litellm._async_success_callback = []
- litellm.failure_callback = []
- litellm.callbacks = []
-
-
-@pytest.mark.skip(reason="Account rate limited.")
-def test_completion_clarifai_claude_2_1():
- print("calling clarifai claude completion")
- import os
-
- clarifai_pat = os.environ["CLARIFAI_API_KEY"]
-
- try:
- response = completion(
- model="clarifai/anthropic.completion.claude-2_1",
- num_retries=3,
- messages=messages,
- max_tokens=10,
- temperature=0.1,
- )
- print(response)
-
- except RateLimitError:
- pass
-
- except Exception as e:
- pytest.fail(f"Error occured: {e}")
-
-
-@pytest.mark.skip(reason="Account rate limited")
-def test_completion_clarifai_mistral_large():
- try:
- litellm.set_verbose = True
- response: ModelResponse = completion(
- model="clarifai/mistralai.completion.mistral-small",
- messages=messages,
- num_retries=3,
- max_tokens=10,
- temperature=0.78,
- )
- # Add any assertions here to check the response
- assert len(response.choices) > 0
- assert len(response.choices[0].message.content) > 0
- except RateLimitError:
- pass
- except Exception as e:
- pytest.fail(f"Error occurred: {e}")
-
-
-@pytest.mark.skip(reason="Account rate limited")
-@pytest.mark.asyncio
-def test_async_completion_clarifai():
- import asyncio
-
- litellm.set_verbose = True
-
- async def test_get_response():
- user_message = "Hello, how are you?"
- messages = [{"content": user_message, "role": "user"}]
- try:
- response = await acompletion(
- model="clarifai/openai.chat-completion.GPT-4",
- messages=messages,
- num_retries=3,
- timeout=10,
- api_key=os.getenv("CLARIFAI_API_KEY"),
- )
- print(f"response: {response}")
- except litellm.Timeout as e:
- pass
- except Exception as e:
- pytest.fail(f"An exception occurred: {e}")
-
- asyncio.run(test_get_response())
diff --git a/tests/llm_translation/test_rerank.py b/tests/llm_translation/test_rerank.py
index edd5981c165..d784677060a 100644
--- a/tests/llm_translation/test_rerank.py
+++ b/tests/llm_translation/test_rerank.py
@@ -148,48 +148,6 @@ async def test_basic_rerank_together_ai(sync_mode):
raise e
-@pytest.mark.asyncio()
-@pytest.mark.parametrize("sync_mode", [True, False])
-@pytest.mark.skip(reason="Skipping test due to Cohere RBAC issues")
-async def test_basic_rerank_azure_ai(sync_mode):
- import os
-
- litellm.set_verbose = True
-
- if sync_mode is True:
- response = litellm.rerank(
- model="azure_ai/Cohere-rerank-v3-multilingual-ko",
- query="hello",
- documents=["hello", "world"],
- top_n=3,
- api_key=os.getenv("AZURE_AI_COHERE_API_KEY"),
- api_base=os.getenv("AZURE_AI_COHERE_API_BASE"),
- )
-
- print("re rank response: ", response)
-
- assert response.id is not None
- assert response.results is not None
-
- assert_response_shape(response, custom_llm_provider="together_ai")
- else:
- response = await litellm.arerank(
- model="azure_ai/Cohere-rerank-v3-multilingual-ko",
- query="hello",
- documents=["hello", "world"],
- top_n=3,
- api_key=os.getenv("AZURE_AI_COHERE_API_KEY"),
- api_base=os.getenv("AZURE_AI_COHERE_API_BASE"),
- )
-
- print("async re rank response: ", response)
-
- assert response.id is not None
- assert response.results is not None
-
- assert_response_shape(response, custom_llm_provider="together_ai")
-
-
@pytest.mark.asyncio()
@pytest.mark.parametrize("version", ["v1", "v2"])
async def test_rerank_custom_api_base(version):
diff --git a/tests/llm_translation/test_router_llm_translation_tests.py b/tests/llm_translation/test_router_llm_translation_tests.py
index 61446ce6136..26456ab0a35 100644
--- a/tests/llm_translation/test_router_llm_translation_tests.py
+++ b/tests/llm_translation/test_router_llm_translation_tests.py
@@ -66,8 +66,8 @@ def test_router_azure_acompletion():
print("Router Test Azure - Acompletion, Acompletion with stream")
# remove api key from env to repro how proxy passes key to router
- old_api_key = os.environ["AZURE_API_KEY"]
- os.environ.pop("AZURE_API_KEY", None)
+ old_api_key = os.environ["AZURE_AI_API_KEY"]
+ os.environ.pop("AZURE_AI_API_KEY", None)
model_list = [
{
@@ -75,8 +75,8 @@ def test_router_azure_acompletion():
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
"api_key": old_api_key,
- "api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_version": os.getenv("AZURE_AI_API_VERSION"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"rpm": 1800,
},
@@ -85,8 +85,8 @@ def test_router_azure_acompletion():
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
"api_key": old_api_key,
- "api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_version": os.getenv("AZURE_AI_API_VERSION"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"rpm": 1800,
},
@@ -126,9 +126,9 @@ def test_router_azure_acompletion():
asyncio.run(test2())
print("\n Passed Streaming")
- os.environ["AZURE_API_KEY"] = old_api_key
+ os.environ["AZURE_AI_API_KEY"] = old_api_key
router.reset()
except Exception as e:
- os.environ["AZURE_API_KEY"] = old_api_key
+ os.environ["AZURE_AI_API_KEY"] = old_api_key
print(f"FAILED TEST")
pytest.fail(f"Got unexpected exception on router! - {e}")
diff --git a/tests/llm_translation/test_snowflake.py b/tests/llm_translation/test_snowflake.py
index 83aa5635f4b..63878c8fc69 100644
--- a/tests/llm_translation/test_snowflake.py
+++ b/tests/llm_translation/test_snowflake.py
@@ -46,45 +46,51 @@ def mock_snowflake_streaming_response_chunks() -> List[str]:
Mock streaming response chunks for Snowflake.
"""
return [
- json.dumps({
- "id": "chatcmpl-snowflake-stream-123",
- "object": "chat.completion.chunk",
- "created": 1700000000,
- "model": "mistral-7b",
- "choices": [
- {
- "index": 0,
- "delta": {"role": "assistant", "content": "The"},
- "finish_reason": None,
- }
- ],
- }),
- json.dumps({
- "id": "chatcmpl-snowflake-stream-123",
- "object": "chat.completion.chunk",
- "created": 1700000000,
- "model": "mistral-7b",
- "choices": [
- {
- "index": 0,
- "delta": {"content": " sky"},
- "finish_reason": None,
- }
- ],
- }),
- json.dumps({
- "id": "chatcmpl-snowflake-stream-123",
- "object": "chat.completion.chunk",
- "created": 1700000000,
- "model": "mistral-7b",
- "choices": [
- {
- "index": 0,
- "delta": {"content": " is blue"},
- "finish_reason": "stop",
- }
- ],
- }),
+ json.dumps(
+ {
+ "id": "chatcmpl-snowflake-stream-123",
+ "object": "chat.completion.chunk",
+ "created": 1700000000,
+ "model": "mistral-7b",
+ "choices": [
+ {
+ "index": 0,
+ "delta": {"role": "assistant", "content": "The"},
+ "finish_reason": None,
+ }
+ ],
+ }
+ ),
+ json.dumps(
+ {
+ "id": "chatcmpl-snowflake-stream-123",
+ "object": "chat.completion.chunk",
+ "created": 1700000000,
+ "model": "mistral-7b",
+ "choices": [
+ {
+ "index": 0,
+ "delta": {"content": " sky"},
+ "finish_reason": None,
+ }
+ ],
+ }
+ ),
+ json.dumps(
+ {
+ "id": "chatcmpl-snowflake-stream-123",
+ "object": "chat.completion.chunk",
+ "created": 1700000000,
+ "model": "mistral-7b",
+ "choices": [
+ {
+ "index": 0,
+ "delta": {"content": " is blue"},
+ "finish_reason": "stop",
+ }
+ ],
+ }
+ ),
]
@@ -120,6 +126,7 @@ def test_chat_completion_snowflake(sync_mode):
async_handler = AsyncHTTPHandler()
with patch.object(AsyncHTTPHandler, "post", return_value=mock_response):
import asyncio
+
response = asyncio.run(
acompletion(
model="snowflake/mistral-7b",
@@ -148,16 +155,16 @@ def test_chat_completion_snowflake_stream(sync_mode):
if sync_mode:
sync_handler = HTTPHandler()
mock_chunks = mock_snowflake_streaming_response_chunks()
-
+
def mock_iter_lines():
for chunk in mock_chunks:
for line in [f"data: {chunk}", "data: [DONE]"]:
yield line
-
+
mock_response = MagicMock()
mock_response.iter_lines.side_effect = mock_iter_lines
mock_response.status_code = 200
-
+
with patch.object(HTTPHandler, "post", return_value=mock_response):
response = completion(
model="snowflake/mistral-7b",
@@ -167,28 +174,28 @@ def test_chat_completion_snowflake_stream(sync_mode):
api_base="https://exampleopenaiendpoint-production.up.railway.app/v1/chat/completions",
client=sync_handler,
)
-
+
chunks_received = []
for chunk in response:
chunks_received.append(chunk)
-
+
assert len(chunks_received) > 0
else:
async_handler = AsyncHTTPHandler()
mock_chunks = mock_snowflake_streaming_response_chunks()
-
+
async def mock_iter_lines():
for chunk in mock_chunks:
for line in [f"data: {chunk}", "data: [DONE]"]:
yield line
-
+
mock_response = MagicMock()
mock_response.iter_lines.side_effect = mock_iter_lines
mock_response.status_code = 200
-
+
with patch.object(AsyncHTTPHandler, "post", return_value=mock_response):
import asyncio
-
+
async def test_async_stream():
response = await acompletion(
model="snowflake/mistral-7b",
@@ -198,78 +205,11 @@ def test_chat_completion_snowflake_stream(sync_mode):
api_base="https://exampleopenaiendpoint-production.up.railway.app/v1/chat/completions",
client=async_handler,
)
-
+
chunks_received = []
async for chunk in response:
chunks_received.append(chunk)
-
+
assert len(chunks_received) > 0
-
+
asyncio.run(test_async_stream())
-
-
-@pytest.mark.skip(reason="Requires Snowflake credentials - run manually when needed")
-def test_snowflake_tool_calling_responses_api():
- """
- Test Snowflake tool calling with Responses API.
- Requires SNOWFLAKE_JWT and SNOWFLAKE_ACCOUNT_ID environment variables.
- """
- import litellm
-
- # Skip if credentials not available
- if not os.getenv("SNOWFLAKE_JWT") or not os.getenv("SNOWFLAKE_ACCOUNT_ID"):
- pytest.skip("Snowflake credentials not available")
-
- litellm.drop_params = False # We now support tools!
-
- tools = [
- {
- "type": "function",
- "name": "get_weather",
- "description": "Get the current weather in a given location",
- "parameters": {
- "type": "object",
- "properties": {
- "location": {
- "type": "string",
- "description": "The city and state, e.g. San Francisco, CA",
- }
- },
- "required": ["location"],
- },
- }
- ]
-
- try:
- # Test with tool_choice to force tool use
- response = responses(
- model="snowflake/claude-3-5-sonnet",
- input="What's the weather in Paris?",
- tools=tools,
- tool_choice={"type": "function", "function": {"name": "get_weather"}},
- max_output_tokens=200,
- )
-
- assert response is not None
- assert hasattr(response, "output")
- assert len(response.output) > 0
-
- # Verify tool call was made
- tool_call_found = False
- for item in response.output:
- if hasattr(item, "type") and item.type == "function_call":
- tool_call_found = True
- assert item.name == "get_weather"
- assert hasattr(item, "arguments")
- print(f"✅ Tool call detected: {item.name}({item.arguments})")
- break
-
- assert tool_call_found, "Expected tool call but none was found"
-
- except APIConnectionError as e:
- if "JWT token is invalid" in str(e):
- pytest.skip("Invalid Snowflake JWT token")
- elif "Application failed to respond" in str(e) or "502" in str(e):
- pytest.skip(f"Snowflake API unavailable: {e}")
- else:
- raise
diff --git a/tests/load_tests/vertex_key.json b/tests/load_tests/vertex_key.json
index 45ca6acc010..800969fb305 100644
--- a/tests/load_tests/vertex_key.json
+++ b/tests/load_tests/vertex_key.json
@@ -1,13 +1,13 @@
{
"type": "service_account",
- "project_id": "pathrise-convert-1606954137718",
+ "project_id": "litellm-ci-cd",
"private_key_id": "",
"private_key": "",
- "client_email": "ci-cd-723@pathrise-convert-1606954137718.iam.gserviceaccount.com",
- "client_id": "109577393201924326488",
+ "client_email": "test-litellm-ci-cd@litellm-ci-cd.iam.gserviceaccount.com",
+ "client_id": "116563532503305622785",
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
"token_uri": "https://oauth2.googleapis.com/token",
"auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs",
- "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/ci-cd-723%40pathrise-convert-1606954137718.iam.gserviceaccount.com",
+ "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/test-litellm-ci-cd%40litellm-ci-cd.iam.gserviceaccount.com",
"universe_domain": "googleapis.com"
}
diff --git a/tests/local_testing/example_config_yaml/azure_config.yaml b/tests/local_testing/example_config_yaml/azure_config.yaml
index 0a015aefde8..05ba0c9bf54 100644
--- a/tests/local_testing/example_config_yaml/azure_config.yaml
+++ b/tests/local_testing/example_config_yaml/azure_config.yaml
@@ -4,12 +4,12 @@ model_list:
model: azure/gpt-4.1-mini
api_base: https://openai-gpt-4-test-v-1.openai.azure.com/
api_version: "2023-05-15"
- api_key: os.environ/AZURE_API_KEY
+ api_key: os.environ/AZURE_AI_API_KEY
tpm: 20_000
- model_name: gpt-4-team2
litellm_params:
model: azure/gpt-4
- api_key: os.environ/AZURE_API_KEY
+ api_key: os.environ/AZURE_AI_API_KEY
api_base: https://openai-gpt-4-test-v-2.openai.azure.com/
tpm: 100_000
diff --git a/tests/local_testing/test_acooldowns_router.py b/tests/local_testing/test_acooldowns_router.py
index ff992102984..18dc26bda9a 100644
--- a/tests/local_testing/test_acooldowns_router.py
+++ b/tests/local_testing/test_acooldowns_router.py
@@ -31,7 +31,7 @@ def _make_model_list():
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
diff --git a/tests/local_testing/test_alangfuse.py b/tests/local_testing/test_alangfuse.py
index 306c7749f18..7c2ec7e9f64 100644
--- a/tests/local_testing/test_alangfuse.py
+++ b/tests/local_testing/test_alangfuse.py
@@ -599,233 +599,6 @@ def test_langfuse_logging_function_calling():
# test_langfuse_logging_function_calling()
-@pytest.mark.skip(reason="Need to address this on main")
-def test_aaalangfuse_existing_trace_id():
- """
- When existing trace id is passed, don't set trace params -> prevents overwriting the trace
-
- Pass 1 logging object with a trace
-
- Pass 2nd logging object with the trace id
-
- Assert no changes to the trace
- """
- # Test - if the logs were sent to the correct team on langfuse
- import datetime
-
- import litellm
- from litellm.integrations.langfuse.langfuse import LangFuseLogger
-
- langfuse_Logger = LangFuseLogger(
- langfuse_public_key=os.getenv("LANGFUSE_PROJECT2_PUBLIC"),
- langfuse_secret=os.getenv("LANGFUSE_PROJECT2_SECRET"),
- )
- litellm.success_callback = ["langfuse"]
-
- # langfuse_args = {'kwargs': { 'start_time': 'end_time': datetime.datetime(2024, 5, 1, 7, 31, 29, 903685), 'user_id': None, 'print_verbose': , 'level': 'DEFAULT', 'status_message': None}
- response_obj = litellm.ModelResponse(
- id="chatcmpl-9K5HUAbVRqFrMZKXL0WoC295xhguY",
- choices=[
- litellm.Choices(
- finish_reason="stop",
- index=0,
- message=litellm.Message(
- content="I'm sorry, I am an AI assistant and do not have real-time information. I recommend checking a reliable weather website or app for the most up-to-date weather information in Boston.",
- role="assistant",
- ),
- )
- ],
- created=1714573888,
- model="gpt-3.5-turbo-0125",
- object="chat.completion",
- system_fingerprint="fp_3b956da36b",
- usage=litellm.Usage(completion_tokens=37, prompt_tokens=14, total_tokens=51),
- )
-
- ### NEW TRACE ###
- message = [{"role": "user", "content": "what's the weather in boston"}]
- langfuse_args = {
- "response_obj": response_obj,
- "kwargs": {
- "model": "gpt-3.5-turbo",
- "litellm_params": {
- "acompletion": False,
- "api_key": None,
- "force_timeout": 600,
- "logger_fn": None,
- "verbose": False,
- "custom_llm_provider": "openai",
- "api_base": "https://api.openai.com/v1/",
- "litellm_call_id": None,
- "model_alias_map": {},
- "completion_call_id": None,
- "metadata": None,
- "model_info": None,
- "proxy_server_request": None,
- "preset_cache_key": None,
- "no-log": False,
- "stream_response": {},
- },
- "messages": message,
- "optional_params": {"temperature": 0.1, "extra_body": {}},
- "start_time": "2024-05-01 07:31:27.986164",
- "stream": False,
- "user": None,
- "call_type": "completion",
- "litellm_call_id": None,
- "completion_start_time": "2024-05-01 07:31:29.903685",
- "temperature": 0.1,
- "extra_body": {},
- "input": [{"role": "user", "content": "what's the weather in boston"}],
- "api_key": "my-api-key",
- "additional_args": {
- "complete_input_dict": {
- "model": "gpt-3.5-turbo",
- "messages": [
- {"role": "user", "content": "what's the weather in boston"}
- ],
- "temperature": 0.1,
- "extra_body": {},
- }
- },
- "log_event_type": "successful_api_call",
- "end_time": "2024-05-01 07:31:29.903685",
- "cache_hit": None,
- "response_cost": 6.25e-05,
- },
- "start_time": datetime.datetime(2024, 5, 1, 7, 31, 27, 986164),
- "end_time": datetime.datetime(2024, 5, 1, 7, 31, 29, 903685),
- "user_id": None,
- "print_verbose": litellm.print_verbose,
- "level": "DEFAULT",
- "status_message": None,
- }
-
- langfuse_response_object = langfuse_Logger.log_event(**langfuse_args)
-
- import langfuse
-
- langfuse_client = langfuse.Langfuse(
- public_key=os.getenv("LANGFUSE_PROJECT2_PUBLIC"),
- secret_key=os.getenv("LANGFUSE_PROJECT2_SECRET"),
- )
-
- trace_id = langfuse_response_object["trace_id"]
-
- assert trace_id is not None
-
- langfuse_client.flush()
-
- time.sleep(2)
-
- print(langfuse_client.get_trace(id=trace_id))
-
- initial_langfuse_trace = langfuse_client.get_trace(id=trace_id)
-
- ### EXISTING TRACE ###
-
- new_metadata = {"existing_trace_id": trace_id}
- new_messages = [{"role": "user", "content": "What do you know?"}]
- new_response_obj = litellm.ModelResponse(
- id="chatcmpl-9K5HUAbVRqFrMZKXL0WoC295xhguY",
- choices=[
- litellm.Choices(
- finish_reason="stop",
- index=0,
- message=litellm.Message(
- content="What do I know?",
- role="assistant",
- ),
- )
- ],
- created=1714573888,
- model="gpt-3.5-turbo-0125",
- object="chat.completion",
- system_fingerprint="fp_3b956da36b",
- usage=litellm.Usage(completion_tokens=37, prompt_tokens=14, total_tokens=51),
- )
- langfuse_args = {
- "response_obj": new_response_obj,
- "kwargs": {
- "model": "gpt-3.5-turbo",
- "litellm_params": {
- "acompletion": False,
- "api_key": None,
- "force_timeout": 600,
- "logger_fn": None,
- "verbose": False,
- "custom_llm_provider": "openai",
- "api_base": "https://api.openai.com/v1/",
- "litellm_call_id": "508113a1-c6f1-48ce-a3e1-01c6cce9330e",
- "model_alias_map": {},
- "completion_call_id": None,
- "metadata": new_metadata,
- "model_info": None,
- "proxy_server_request": None,
- "preset_cache_key": None,
- "no-log": False,
- "stream_response": {},
- },
- "messages": new_messages,
- "optional_params": {"temperature": 0.1, "extra_body": {}},
- "start_time": "2024-05-01 07:31:27.986164",
- "stream": False,
- "user": None,
- "call_type": "completion",
- "litellm_call_id": "508113a1-c6f1-48ce-a3e1-01c6cce9330e",
- "completion_start_time": "2024-05-01 07:31:29.903685",
- "temperature": 0.1,
- "extra_body": {},
- "input": [{"role": "user", "content": "what's the weather in boston"}],
- "api_key": "my-api-key",
- "additional_args": {
- "complete_input_dict": {
- "model": "gpt-3.5-turbo",
- "messages": [
- {"role": "user", "content": "what's the weather in boston"}
- ],
- "temperature": 0.1,
- "extra_body": {},
- }
- },
- "log_event_type": "successful_api_call",
- "end_time": "2024-05-01 07:31:29.903685",
- "cache_hit": None,
- "response_cost": 6.25e-05,
- },
- "start_time": datetime.datetime(2024, 5, 1, 7, 31, 27, 986164),
- "end_time": datetime.datetime(2024, 5, 1, 7, 31, 29, 903685),
- "user_id": None,
- "print_verbose": litellm.print_verbose,
- "level": "DEFAULT",
- "status_message": None,
- }
-
- langfuse_response_object = langfuse_Logger.log_event(**langfuse_args)
-
- new_trace_id = langfuse_response_object["trace_id"]
-
- assert new_trace_id == trace_id
-
- langfuse_client.flush()
-
- time.sleep(2)
-
- print(langfuse_client.get_trace(id=trace_id))
-
- new_langfuse_trace = langfuse_client.get_trace(id=trace_id)
-
- initial_langfuse_trace_dict = dict(initial_langfuse_trace)
- initial_langfuse_trace_dict.pop("updatedAt")
- initial_langfuse_trace_dict.pop("timestamp")
-
- new_langfuse_trace_dict = dict(new_langfuse_trace)
- new_langfuse_trace_dict.pop("updatedAt")
- new_langfuse_trace_dict.pop("timestamp")
-
- assert initial_langfuse_trace_dict == new_langfuse_trace_dict
-
-
@pytest.mark.skipif(
condition=not os.environ.get("OPENAI_API_KEY", False),
reason="Authentication missing for openai",
@@ -928,42 +701,6 @@ async def test_make_request():
)
-@pytest.mark.skip(
- reason="local only test, use this to verify if dynamic langfuse logging works as expected"
-)
-def test_aaalangfuse_dynamic_logging():
- """
- pass in langfuse credentials via completion call
-
- assert call is logged.
-
- Covers the team-logging scenario.
- """
- from litellm._uuid import uuid
-
- import langfuse
-
- trace_id = str(uuid.uuid4())
- _ = litellm.completion(
- model="gpt-3.5-turbo",
- messages=[{"role": "user", "content": "Hey"}],
- mock_response="Hey! how's it going?",
- langfuse_public_key=os.getenv("LANGFUSE_PROJECT2_PUBLIC"),
- langfuse_secret_key=os.getenv("LANGFUSE_PROJECT2_SECRET"),
- metadata={"trace_id": trace_id},
- success_callback=["langfuse"],
- )
-
- time.sleep(3)
-
- langfuse_client = langfuse.Langfuse(
- public_key=os.getenv("LANGFUSE_PROJECT2_PUBLIC"),
- secret_key=os.getenv("LANGFUSE_PROJECT2_SECRET"),
- )
-
- langfuse_client.get_trace(id=trace_id)
-
-
import datetime
generation_params = {
diff --git a/tests/local_testing/test_arize_ai.py b/tests/local_testing/test_arize_ai.py
index 3b497d638ae..138858cee03 100644
--- a/tests/local_testing/test_arize_ai.py
+++ b/tests/local_testing/test_arize_ai.py
@@ -48,8 +48,8 @@ async def test_async_dynamic_arize_config():
messages=[{"role": "user", "content": "hi test from arize dynamic config"}],
temperature=0.1,
user="OTEL_USER",
- arize_api_key=os.getenv("ARIZE_SPACE_2_API_KEY"),
- arize_space_key=os.getenv("ARIZE_SPACE_2_KEY"),
+ arize_api_key=os.getenv("ARIZE_SPACE_API_KEY"),
+ arize_space_key=os.getenv("ARIZE_SPACE_KEY"),
)
await asyncio.sleep(2)
diff --git a/tests/local_testing/test_assistants.py b/tests/local_testing/test_assistants.py
index df0df6b79cd..93ebba3c345 100644
--- a/tests/local_testing/test_assistants.py
+++ b/tests/local_testing/test_assistants.py
@@ -37,10 +37,11 @@ V0 Scope:
- Run Thread -> `/v1/threads/{thread_id}/run`
"""
+
def _add_azure_related_dynamic_params(data: dict) -> dict:
data["api_version"] = "2024-02-15-preview"
- data["api_base"] = os.getenv("AZURE_API_BASE")
- data["api_key"] = os.getenv("AZURE_API_KEY")
+ data["api_base"] = os.getenv("AZURE_AI_API_BASE")
+ data["api_key"] = os.getenv("AZURE_AI_API_KEY")
return data
@@ -236,8 +237,6 @@ async def test_aarun_thread_litellm(sync_mode, provider, is_streaming):
"""
import openai
-
-
try:
get_assistants_data = {
"custom_llm_provider": provider,
diff --git a/tests/local_testing/test_azure_openai.py b/tests/local_testing/test_azure_openai.py
index e95c1b6fcce..1b99140b6e6 100644
--- a/tests/local_testing/test_azure_openai.py
+++ b/tests/local_testing/test_azure_openai.py
@@ -40,11 +40,13 @@ async def test_aaaaazure_tenant_id_auth(respx_mock: MockRouter):
PROD Test
"""
- litellm.disable_aiohttp_transport = True # since this uses respx, we need to set use_aiohttp_transport to False
-
+ litellm.disable_aiohttp_transport = (
+ True # since this uses respx, we need to set use_aiohttp_transport to False
+ )
+
# Clear the HTTP client cache to ensure respx mocking works
# This is critical because respx only intercepts clients created AFTER mocking is active
- if hasattr(litellm, 'in_memory_llm_clients_cache'):
+ if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
router = Router(
@@ -53,7 +55,7 @@ async def test_aaaaazure_tenant_id_auth(respx_mock: MockRouter):
"model_name": "gpt-3.5-turbo",
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
"tenant_id": os.getenv("AZURE_TENANT_ID"),
"client_id": os.getenv("AZURE_CLIENT_ID"),
"client_secret": os.getenv("AZURE_CLIENT_SECRET"),
diff --git a/tests/local_testing/test_azure_perf.py b/tests/local_testing/test_azure_perf.py
index 1e2d5cc4f7b..57d56a24a15 100644
--- a/tests/local_testing/test_azure_perf.py
+++ b/tests/local_testing/test_azure_perf.py
@@ -9,8 +9,8 @@
# from openai import AsyncAzureOpenAI
# client = AsyncAzureOpenAI(
-# api_key=os.getenv("AZURE_API_KEY"),
-# azure_endpoint=os.getenv("AZURE_API_BASE"), # type: ignore
+# api_key=os.getenv("AZURE_AI_API_KEY"),
+# azure_endpoint=os.getenv("AZURE_AI_API_BASE"), # type: ignore
# api_version=os.getenv("AZURE_API_VERSION"),
# )
@@ -19,8 +19,8 @@
# "model_name": "azure-test",
# "litellm_params": {
# "model": "azure/gpt-4.1-mini",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "api_base": os.getenv("AZURE_API_BASE"),
+# "api_key": os.getenv("AZURE_AI_API_KEY"),
+# "api_base": os.getenv("AZURE_AI_API_BASE"),
# "api_version": os.getenv("AZURE_API_VERSION"),
# },
# }
diff --git a/tests/local_testing/test_caching.py b/tests/local_testing/test_caching.py
index 01004e4bfa0..0b06b7195bd 100644
--- a/tests/local_testing/test_caching.py
+++ b/tests/local_testing/test_caching.py
@@ -147,7 +147,12 @@ def test_caching_dynamic_args(): # test in memory cache
port=_redis_port_env,
password=_redis_password_env,
)
- response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
+ response1 = completion(
+ model="gpt-3.5-turbo",
+ messages=messages,
+ caching=True,
+ mock_response="Hello world from cache test",
+ )
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
print(f"response1: {response1}")
print(f"response2: {response2}")
@@ -173,7 +178,12 @@ def test_caching_v2(): # test in memory cache
try:
litellm.set_verbose = True
litellm.cache = Cache()
- response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
+ response1 = completion(
+ model="gpt-3.5-turbo",
+ messages=messages,
+ caching=True,
+ mock_response="Hello world from cache test",
+ )
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
print(f"response1: {response1}")
print(f"response2: {response2}")
@@ -200,9 +210,18 @@ def test_caching_with_ttl():
litellm.set_verbose = True
litellm.cache = Cache()
response1 = completion(
- model="gpt-3.5-turbo", messages=messages, caching=True, ttl=0, mock_response="Hello world from cache test 1"
+ model="gpt-3.5-turbo",
+ messages=messages,
+ caching=True,
+ ttl=0,
+ mock_response="Hello world from cache test 1",
+ )
+ response2 = completion(
+ model="gpt-3.5-turbo",
+ messages=messages,
+ caching=True,
+ mock_response="Hello world from cache test 2",
)
- response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test 2")
print(f"response1: {response1}")
print(f"response2: {response2}")
litellm.cache = None # disable cache
@@ -221,8 +240,18 @@ def test_caching_with_default_ttl():
try:
litellm.set_verbose = True
litellm.cache = Cache(ttl=0)
- response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
- response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
+ response1 = completion(
+ model="gpt-3.5-turbo",
+ messages=messages,
+ caching=True,
+ mock_response="Hello world from cache test",
+ )
+ response2 = completion(
+ model="gpt-3.5-turbo",
+ messages=messages,
+ caching=True,
+ mock_response="Hello world from cache test",
+ )
print(f"response1: {response1}")
print(f"response2: {response2}")
litellm.cache = None # disable cache
@@ -247,10 +276,16 @@ async def test_caching_with_cache_controls(sync_flag):
if sync_flag:
## TTL = 0
response1 = completion(
- model="gpt-3.5-turbo", messages=messages, cache={"ttl": 0}, mock_response="Hello world"
+ model="gpt-3.5-turbo",
+ messages=messages,
+ cache={"ttl": 0},
+ mock_response="Hello world",
)
response2 = completion(
- model="gpt-3.5-turbo", messages=messages, cache={"s-maxage": 10}, mock_response="Hello world"
+ model="gpt-3.5-turbo",
+ messages=messages,
+ cache={"s-maxage": 10},
+ mock_response="Hello world",
)
assert response2["id"] != response1["id"]
@@ -322,9 +357,19 @@ def test_caching_with_models_v2():
litellm.cache = Cache()
print("test2 for caching")
litellm.set_verbose = True
- response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
+ response1 = completion(
+ model="gpt-3.5-turbo",
+ messages=messages,
+ caching=True,
+ mock_response="Hello world from cache test",
+ )
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
- response3 = completion(model="gpt-4.1-nano", messages=messages, caching=True, mock_response="Different model response")
+ response3 = completion(
+ model="gpt-4.1-nano",
+ messages=messages,
+ caching=True,
+ mock_response="Different model response",
+ )
print(f"response1: {response1}")
print(f"response2: {response2}")
print(f"response3: {response3}")
@@ -423,7 +468,10 @@ def test_embedding_caching():
text_to_embed = [embedding_large_text]
start_time = time.time()
embedding1 = embedding(
- model="text-embedding-ada-002", input=text_to_embed, caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
+ model="text-embedding-ada-002",
+ input=text_to_embed,
+ caching=True,
+ mock_response="0.1,0.2,0.3,0.4,0.5",
)
end_time = time.time()
print(f"Embedding 1 response time: {end_time - start_time} seconds")
@@ -459,12 +507,18 @@ async def test_embedding_caching_individual_items_and_then_list():
"world",
]
embedding1 = await aembedding(
- model="text-embedding-ada-002", input=text_to_embed[0], caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
+ model="text-embedding-ada-002",
+ input=text_to_embed[0],
+ caching=True,
+ mock_response="0.1,0.2,0.3,0.4,0.5",
)
initial_prompt_tokens = embedding1.usage.prompt_tokens
await asyncio.sleep(1)
embedding2 = await aembedding(
- model="text-embedding-ada-002", input=text_to_embed[1], caching=True, mock_response="0.6,0.7,0.8,0.9,1.0"
+ model="text-embedding-ada-002",
+ input=text_to_embed[1],
+ caching=True,
+ mock_response="0.6,0.7,0.8,0.9,1.0",
)
await asyncio.sleep(1)
embedding3 = await aembedding(
@@ -480,7 +534,10 @@ async def test_embedding_caching_individual_items_and_then_list():
additional_text = "this is a new text"
text_to_embed.append(additional_text)
embedding4 = await aembedding(
- model="text-embedding-ada-002", input=text_to_embed, caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
+ model="text-embedding-ada-002",
+ input=text_to_embed,
+ caching=True,
+ mock_response="0.1,0.2,0.3,0.4,0.5",
)
assert embedding4.usage.prompt_tokens > embedding3.usage.prompt_tokens
@@ -490,7 +547,10 @@ async def test_embedding_caching_individual_items():
litellm.cache = Cache()
text_to_embed = "hello"
embedding1 = await aembedding(
- model="text-embedding-ada-002", input=text_to_embed, caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
+ model="text-embedding-ada-002",
+ input=text_to_embed,
+ caching=True,
+ mock_response="0.1,0.2,0.3,0.4,0.5",
)
await asyncio.sleep(1)
@@ -512,13 +572,13 @@ def test_embedding_caching_azure():
litellm.cache = Cache()
text_to_embed = [embedding_large_text]
- api_key = os.environ["AZURE_API_KEY"]
- api_base = os.environ["AZURE_API_BASE"]
+ api_key = os.environ["AZURE_AI_API_KEY"]
+ api_base = os.environ["AZURE_AI_API_BASE"]
api_version = os.environ["AZURE_API_VERSION"]
os.environ["AZURE_API_VERSION"] = ""
- os.environ["AZURE_API_BASE"] = ""
- os.environ["AZURE_API_KEY"] = ""
+ os.environ["AZURE_AI_API_BASE"] = ""
+ os.environ["AZURE_AI_API_KEY"] = ""
start_time = time.time()
print("AZURE CONFIGS")
@@ -560,8 +620,8 @@ def test_embedding_caching_azure():
pytest.fail("Error occurred: Embedding caching failed")
os.environ["AZURE_API_VERSION"] = api_version
- os.environ["AZURE_API_BASE"] = api_base
- os.environ["AZURE_API_KEY"] = api_key
+ os.environ["AZURE_AI_API_BASE"] = api_base
+ os.environ["AZURE_AI_API_KEY"] = api_key
# test_embedding_caching_azure()
@@ -851,9 +911,18 @@ def test_redis_cache_completion():
model="gpt-3.5-turbo", messages=messages, caching=True, max_tokens=20
)
response3 = completion(
- model="gpt-3.5-turbo", messages=messages, caching=True, temperature=0.5, mock_response="Different params response"
+ model="gpt-3.5-turbo",
+ messages=messages,
+ caching=True,
+ temperature=0.5,
+ mock_response="Different params response",
+ )
+ response4 = completion(
+ model="gpt-4o-mini",
+ messages=messages,
+ caching=True,
+ mock_response="Different model response",
)
- response4 = completion(model="gpt-4o-mini", messages=messages, caching=True, mock_response="Different model response")
print("\nresponse 1", response1)
print("\nresponse 2", response2)
@@ -1127,7 +1196,11 @@ async def test_redis_cache_atext_completion():
print("test for caching, atext_completion")
response1 = await litellm.atext_completion(
- model="gpt-3.5-turbo-instruct", prompt=prompt, max_tokens=40, temperature=1, mock_response="Hello world from cache test"
+ model="gpt-3.5-turbo-instruct",
+ prompt=prompt,
+ max_tokens=40,
+ temperature=1,
+ mock_response="Hello world from cache test",
)
await asyncio.sleep(0.5)
@@ -1458,11 +1531,17 @@ def test_cache_override():
# test embedding
response1 = embedding(
- model="text-embedding-ada-002", input=["hello who are you"], caching=False, mock_response="0.1,0.2,0.3,0.4,0.5"
+ model="text-embedding-ada-002",
+ input=["hello who are you"],
+ caching=False,
+ mock_response="0.1,0.2,0.3,0.4,0.5",
)
response2 = embedding(
- model="text-embedding-ada-002", input=["hello who are you"], caching=False, mock_response="0.6,0.7,0.8,0.9,1.0"
+ model="text-embedding-ada-002",
+ input=["hello who are you"],
+ caching=False,
+ mock_response="0.6,0.7,0.8,0.9,1.0",
)
# When caching=False, responses should have different IDs
@@ -2787,7 +2866,7 @@ def test_caching_thinking_args_hit(): # test in memory cache
async def test_cache_key_in_hidden_params_acompletion():
"""
Test that cache_key is present in _hidden_params on cache hits for acompletion.
-
+
Validates fix for missing x-litellm-cache-key header on proxy cache hits.
"""
litellm.cache = Cache(
@@ -2796,10 +2875,10 @@ async def test_cache_key_in_hidden_params_acompletion():
port=os.environ["REDIS_PORT"],
password=os.environ["REDIS_PASSWORD"],
)
-
+
unique_content = f"test cache key hidden params {uuid.uuid4()}"
messages = [{"role": "user", "content": unique_content}]
-
+
# First call - cache miss
response1 = await litellm.acompletion(
model="gpt-3.5-turbo",
@@ -2807,12 +2886,12 @@ async def test_cache_key_in_hidden_params_acompletion():
mock_response="test response",
caching=True,
)
-
+
print(f"Response 1 _hidden_params: {response1._hidden_params}")
assert response1._hidden_params.get("cache_hit") is not True
-
+
await asyncio.sleep(0.5)
-
+
# Second call - cache hit
response2 = await litellm.acompletion(
model="gpt-3.5-turbo",
@@ -2820,17 +2899,17 @@ async def test_cache_key_in_hidden_params_acompletion():
mock_response="test response",
caching=True,
)
-
+
print(f"Response 2 _hidden_params: {response2._hidden_params}")
-
+
# Verify cache hit occurred
assert response2._hidden_params.get("cache_hit") is True
-
+
# Verify cache_key is present in _hidden_params
assert "cache_key" in response2._hidden_params
assert response2._hidden_params["cache_key"] is not None
-
+
# Verify both responses have same ID (cache hit)
assert response1.id == response2.id
-
+
litellm.cache = None
diff --git a/tests/local_testing/test_caching_ssl.py b/tests/local_testing/test_caching_ssl.py
index 523976f1237..a8a38e747a7 100644
--- a/tests/local_testing/test_caching_ssl.py
+++ b/tests/local_testing/test_caching_ssl.py
@@ -59,9 +59,9 @@ def test_caching_router():
"model_name": "gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
diff --git a/tests/local_testing/test_class.py b/tests/local_testing/test_class.py
index e02f59a2941..b4b4f85a9d0 100644
--- a/tests/local_testing/test_class.py
+++ b/tests/local_testing/test_class.py
@@ -56,9 +56,9 @@
# # "model_name": "gpt-3.5-turbo", # openai model name
# # "litellm_params": { # params for litellm completion/embedding call
# # "model": "azure/gpt-4.1-mini",
-# # "api_key": os.getenv("AZURE_API_KEY"),
+# # "api_key": os.getenv("AZURE_AI_API_KEY"),
# # "api_version": os.getenv("AZURE_API_VERSION"),
-# # "api_base": os.getenv("AZURE_API_BASE"),
+# # "api_base": os.getenv("AZURE_AI_API_BASE"),
# # },
# # }
# # ]
@@ -94,9 +94,9 @@
# # "model_name": "gpt-3.5-turbo", # openai model name
# # "litellm_params": { # params for litellm completion/embedding call
# # "model": "azure/gpt-4.1-mini",
-# # "api_key": os.getenv("AZURE_API_KEY"),
+# # "api_key": os.getenv("AZURE_AI_API_KEY"),
# # "api_version": os.getenv("AZURE_API_VERSION"),
-# # "api_base": os.getenv("AZURE_API_BASE"),
+# # "api_base": os.getenv("AZURE_AI_API_BASE"),
# # },
# # }
# # ],
diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py
index e6f5cd86517..d5da0719e0e 100644
--- a/tests/local_testing/test_completion.py
+++ b/tests/local_testing/test_completion.py
@@ -132,7 +132,6 @@ def test_null_role_response():
assert response.choices[0].message.role == "assistant"
-
def predibase_mock_post(url, data=None, json=None, headers=None, timeout=None):
mock_response = MagicMock()
mock_response.status_code = 200
@@ -177,34 +176,6 @@ def predibase_mock_post(url, data=None, json=None, headers=None, timeout=None):
return mock_response
-# @pytest.mark.skip(reason="local-only test")
-@pytest.mark.asyncio
-async def test_completion_predibase():
- try:
- litellm.set_verbose = True
-
- # with patch("requests.post", side_effect=predibase_mock_post):
- response = await litellm.acompletion(
- model="predibase/llama-3-8b-instruct",
- tenant_id="c4768f95",
- api_key=os.getenv("PREDIBASE_API_KEY"),
- messages=[{"role": "user", "content": "who are u?"}],
- max_tokens=10,
- timeout=5,
- )
-
- print(response)
- except litellm.Timeout as e:
- print("got a timeout error from predibase")
- pass
- except litellm.ServiceUnavailableError as e:
- pass
- except litellm.InternalServerError:
- pass
- except Exception as e:
- pytest.fail(f"Error occurred: {e}")
-
-
# test_completion_predibase()
@@ -286,7 +257,9 @@ def test_completion_claude_3_empty_response():
},
]
try:
- response = litellm.completion(model="claude-sonnet-4-5-20250929", messages=messages)
+ response = litellm.completion(
+ model="claude-sonnet-4-5-20250929", messages=messages
+ )
print(response)
except litellm.InternalServerError as e:
pytest.skip(f"InternalServerError - {str(e)}")
@@ -849,8 +822,8 @@ def test_completion_mistral_azure():
litellm.set_verbose = True
response = completion(
model="mistral/Mistral-large-nmefg",
- api_key=os.environ["MISTRAL_AZURE_API_KEY"],
- api_base=os.environ["MISTRAL_AZURE_API_BASE"],
+ api_key=os.environ["MISTRAL_AZURE_AI_API_KEY"],
+ api_base=os.environ["MISTRAL_AZURE_AI_API_BASE"],
max_tokens=5,
messages=[
{
@@ -996,59 +969,6 @@ def test_completion_gpt4_vision():
pytest.fail(f"Error occurred: {e}")
-# test_completion_gpt4_vision()
-
-
-def test_completion_azure_gpt4_vision():
- # azure/gpt-4, vision takes 5-seconds to respond
- try:
- litellm.set_verbose = True
- response = completion(
- model="azure/gpt-4-vision",
- timeout=5,
- messages=[
- {
- "role": "user",
- "content": [
- {"type": "text", "text": "Whats in this image?"},
- {
- "type": "image_url",
- "image_url": {
- "url": "https://avatars.githubusercontent.com/u/29436595?v=4"
- },
- },
- ],
- }
- ],
- base_url="https://gpt-4-vision-resource.openai.azure.com/openai/deployments/gpt-4-vision/extensions",
- api_key=os.getenv("AZURE_VISION_API_KEY"),
- enhancements={"ocr": {"enabled": True}, "grounding": {"enabled": True}},
- dataSources=[
- {
- "type": "AzureComputerVision",
- "parameters": {
- "endpoint": "https://gpt-4-vision-enhancement.cognitiveservices.azure.com/",
- "key": os.environ["AZURE_VISION_ENHANCE_KEY"],
- },
- }
- ],
- )
- print(response)
- except openai.APIError as e:
- pass
- except openai.APITimeoutError:
- print("got a timeout error")
- pass
- except openai.RateLimitError as e:
- print("got a rate liimt error", e)
- pass
- except openai.APIStatusError as e:
- print("got an api status error", e)
- pass
- except Exception as e:
- pytest.fail(f"Error occurred: {e}")
-
-
# test_completion_azure_gpt4_vision()
@@ -1751,7 +1671,6 @@ def test_completion_openai_pydantic(model, api_version):
pytest.fail(f"Error occurred: {e}")
-
def test_completion_text_openai():
try:
# litellm.set_verbose =True
@@ -2341,9 +2260,9 @@ def test_completion_azure_extra_headers():
response = completion(
model="azure/gpt-4.1-mini",
messages=messages,
- api_base=os.getenv("AZURE_API_BASE"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
api_version="2023-07-01-preview",
- api_key=os.getenv("AZURE_API_KEY"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
extra_headers={
"Authorization": "my-bad-key",
"Ocp-Apim-Subscription-Key": "hello-world-testing",
@@ -2379,8 +2298,8 @@ def test_completion_azure_ad_token():
litellm.set_verbose = True
- old_key = os.environ["AZURE_API_KEY"]
- os.environ.pop("AZURE_API_KEY", None)
+ old_key = os.environ["AZURE_AI_API_KEY"]
+ os.environ.pop("AZURE_AI_API_KEY", None)
http_client = Client()
@@ -2396,7 +2315,7 @@ def test_completion_azure_ad_token():
except Exception as e:
pass
finally:
- os.environ["AZURE_API_KEY"] = old_key
+ os.environ["AZURE_AI_API_KEY"] = old_key
mock_client.assert_called_once()
request = mock_client.call_args[0][0]
@@ -2412,8 +2331,8 @@ def test_completion_azure_key_completion_arg():
# DO NOT REMOVE THIS TEST. No MATTER WHAT Happens!
# If you want to remove it, speak to Ishaan!
# Ishaan will be very disappointed if this test is removed -> this is a standard way to pass api_key + the router + proxy use this
- old_key = os.environ["AZURE_API_KEY"]
- os.environ.pop("AZURE_API_KEY", None)
+ old_key = os.environ["AZURE_AI_API_KEY"]
+ os.environ.pop("AZURE_AI_API_KEY", None)
try:
print("azure gpt-3.5 test\n\n")
litellm.set_verbose = True
@@ -2430,9 +2349,9 @@ def test_completion_azure_key_completion_arg():
print("Hidden Params", response._hidden_params)
assert response._hidden_params["custom_llm_provider"] == "azure"
- os.environ["AZURE_API_KEY"] = old_key
+ os.environ["AZURE_AI_API_KEY"] = old_key
except Exception as e:
- os.environ["AZURE_API_KEY"] = old_key
+ os.environ["AZURE_AI_API_KEY"] = old_key
pytest.fail(f"Error occurred: {e}")
@@ -2443,8 +2362,8 @@ async def test_re_use_azure_async_client():
import openai
client = openai.AsyncAzureOpenAI(
- azure_endpoint=os.environ["AZURE_API_BASE"],
- api_key=os.environ["AZURE_API_KEY"],
+ azure_endpoint=os.environ["AZURE_AI_API_BASE"],
+ api_key=os.environ["AZURE_AI_API_KEY"],
api_version="2023-07-01-preview",
)
## Test azure call
@@ -2525,13 +2444,13 @@ def test_completion_azure2():
try:
print("azure gpt-3.5 test\n\n")
litellm.set_verbose = False
- api_base = os.environ["AZURE_API_BASE"]
- api_key = os.environ["AZURE_API_KEY"]
+ api_base = os.environ["AZURE_AI_API_BASE"]
+ api_key = os.environ["AZURE_AI_API_KEY"]
api_version = os.environ["AZURE_API_VERSION"]
- os.environ["AZURE_API_BASE"] = ""
+ os.environ["AZURE_AI_API_BASE"] = ""
os.environ["AZURE_API_VERSION"] = ""
- os.environ["AZURE_API_KEY"] = ""
+ os.environ["AZURE_AI_API_KEY"] = ""
## Test azure call
response = completion(
@@ -2546,9 +2465,9 @@ def test_completion_azure2():
# Add any assertions here to check the response
print(response)
- os.environ["AZURE_API_BASE"] = api_base
+ os.environ["AZURE_AI_API_BASE"] = api_base
os.environ["AZURE_API_VERSION"] = api_version
- os.environ["AZURE_API_KEY"] = api_key
+ os.environ["AZURE_AI_API_KEY"] = api_key
except Exception as e:
pytest.fail(f"Error occurred: {e}")
@@ -2562,13 +2481,13 @@ def test_completion_azure3():
try:
print("azure gpt-3.5 test\n\n")
litellm.set_verbose = True
- litellm.api_base = os.environ["AZURE_API_BASE"]
- litellm.api_key = os.environ["AZURE_API_KEY"]
+ litellm.api_base = os.environ["AZURE_AI_API_BASE"]
+ litellm.api_key = os.environ["AZURE_AI_API_KEY"]
litellm.api_version = os.environ["AZURE_API_VERSION"]
- os.environ["AZURE_API_BASE"] = ""
+ os.environ["AZURE_AI_API_BASE"] = ""
os.environ["AZURE_API_VERSION"] = ""
- os.environ["AZURE_API_KEY"] = ""
+ os.environ["AZURE_AI_API_KEY"] = ""
## Test azure call
response = completion(
@@ -2580,9 +2499,9 @@ def test_completion_azure3():
# Add any assertions here to check the response
print(response)
- os.environ["AZURE_API_BASE"] = litellm.api_base
+ os.environ["AZURE_AI_API_BASE"] = litellm.api_base
os.environ["AZURE_API_VERSION"] = litellm.api_version
- os.environ["AZURE_API_KEY"] = litellm.api_key
+ os.environ["AZURE_AI_API_KEY"] = litellm.api_key
except Exception as e:
pytest.fail(f"Error occurred: {e}")
@@ -2594,7 +2513,7 @@ def test_completion_azure3():
# new azure test for using litellm. vars,
# use the following vars in this test and make an azure_api_call
# litellm.api_type = self.azure_api_type
-# litellm.api_base = self.azure_api_base
+# litellm.api_base = self.AZURE_AI_API_BASE
# litellm.api_version = self.azure_api_version
# litellm.api_key = self.api_key
def test_completion_azure_with_litellm_key():
@@ -2604,14 +2523,14 @@ def test_completion_azure_with_litellm_key():
#### set litellm vars
litellm.api_type = "azure"
- litellm.api_base = os.environ["AZURE_API_BASE"]
+ litellm.api_base = os.environ["AZURE_AI_API_BASE"]
litellm.api_version = os.environ["AZURE_API_VERSION"]
- litellm.api_key = os.environ["AZURE_API_KEY"]
+ litellm.api_key = os.environ["AZURE_AI_API_KEY"]
######### UNSET ENV VARs for this ################
- os.environ["AZURE_API_BASE"] = ""
+ os.environ["AZURE_AI_API_BASE"] = ""
os.environ["AZURE_API_VERSION"] = ""
- os.environ["AZURE_API_KEY"] = ""
+ os.environ["AZURE_AI_API_KEY"] = ""
######### UNSET OpenAI vars for this ##############
openai.api_type = ""
@@ -2627,9 +2546,9 @@ def test_completion_azure_with_litellm_key():
print(response)
######### RESET ENV VARs for this ################
- os.environ["AZURE_API_BASE"] = litellm.api_base
+ os.environ["AZURE_AI_API_BASE"] = litellm.api_base
os.environ["AZURE_API_VERSION"] = litellm.api_version
- os.environ["AZURE_API_KEY"] = litellm.api_key
+ os.environ["AZURE_AI_API_KEY"] = litellm.api_key
######### UNSET litellm vars
litellm.api_type = None
@@ -3081,7 +3000,6 @@ async def test_completion_bedrock_httpx_models(sync_mode, model):
pytest.fail(f"An error occurred - {str(e)}")
-
# test_completion_bedrock_titan()
@@ -3256,7 +3174,6 @@ def test_completion_anyscale_api():
pytest.fail(f"Error occurred: {e}")
-
@pytest.mark.skip(reason="anyscale stopped serving public api endpoints")
def test_completion_anyscale_2():
try:
@@ -3293,23 +3210,6 @@ def test_mistral_anyscale_stream():
print(chunk["choices"][0]["delta"].get("content", ""), end="")
-# test_completion_anyscale_2()
-# def test_completion_with_fallbacks_multiple_keys():
-# print(f"backup key 1: {os.getenv('BACKUP_OPENAI_API_KEY_1')}")
-# print(f"backup key 2: {os.getenv('BACKUP_OPENAI_API_KEY_2')}")
-# backup_keys = [{"api_key": os.getenv("BACKUP_OPENAI_API_KEY_1")}, {"api_key": os.getenv("BACKUP_OPENAI_API_KEY_2")}]
-# try:
-# api_key = "bad-key"
-# response = completion(
-# model="gpt-3.5-turbo", messages=messages, force_timeout=120, fallbacks=backup_keys, api_key=api_key
-# )
-# # Add any assertions here to check the response
-# print(response)
-# except Exception as e:
-# error_str = traceback.format_exc()
-# pytest.fail(f"Error occurred: {error_str}")
-
-
# test_completion_with_fallbacks_multiple_keys()
def test_petals():
try:
@@ -3871,9 +3771,6 @@ async def test_dynamic_azure_params(stream, sync_mode):
raise e
-
-
-
@pytest.mark.parametrize(
"model",
["gpt-4o", "azure/gpt-4.1-mini"],
diff --git a/tests/local_testing/test_config.py b/tests/local_testing/test_config.py
index e74f92ca7a5..2c5d04d3815 100644
--- a/tests/local_testing/test_config.py
+++ b/tests/local_testing/test_config.py
@@ -47,8 +47,8 @@ async def test_delete_deployment():
litellm_params = LiteLLM_Params(
model="azure/gpt-4.1-mini",
- api_key=os.getenv("AZURE_API_KEY"),
- api_base=os.getenv("AZURE_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
api_version=os.getenv("AZURE_API_VERSION"),
)
encrypted_litellm_params = litellm_params.dict(exclude_none=True)
@@ -131,8 +131,8 @@ async def test_add_existing_deployment():
litellm_params = LiteLLM_Params(
model="gpt-3.5-turbo",
- api_key=os.getenv("AZURE_API_KEY"),
- api_base=os.getenv("AZURE_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
api_version=os.getenv("AZURE_API_VERSION"),
)
deployment = Deployment(model_name="gpt-3.5-turbo", litellm_params=litellm_params)
@@ -186,8 +186,8 @@ async def test_db_error_new_model_check():
litellm_params = LiteLLM_Params(
model="gpt-3.5-turbo",
- api_key=os.getenv("AZURE_API_KEY"),
- api_base=os.getenv("AZURE_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
api_version=os.getenv("AZURE_API_VERSION"),
)
deployment = Deployment(model_name="gpt-3.5-turbo", litellm_params=litellm_params)
@@ -233,8 +233,8 @@ async def test_db_error_new_model_check():
litellm_params = LiteLLM_Params(
model="azure/gpt-4.1-mini",
- api_key=os.getenv("AZURE_API_KEY"),
- api_base=os.getenv("AZURE_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
api_version=os.getenv("AZURE_API_VERSION"),
)
@@ -251,8 +251,8 @@ def _create_model_list(flag_value: Literal[0, 1], master_key: str):
new_litellm_params = LiteLLM_Params(
model="azure/gpt-4.1-mini-3",
- api_key=os.getenv("AZURE_API_KEY"),
- api_base=os.getenv("AZURE_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
api_version=os.getenv("AZURE_API_VERSION"),
)
@@ -421,4 +421,3 @@ def test_litellm_proxy_responses_api_config():
assert (
config.custom_llm_provider == LlmProviders.LITELLM_PROXY
), "custom_llm_provider should be LITELLM_PROXY"
-
diff --git a/tests/local_testing/test_configs/test_bad_config.yaml b/tests/local_testing/test_configs/test_bad_config.yaml
index 4a70886a93b..4bc4c7cc541 100644
--- a/tests/local_testing/test_configs/test_bad_config.yaml
+++ b/tests/local_testing/test_configs/test_bad_config.yaml
@@ -6,16 +6,16 @@ model_list:
- model_name: working-azure-gpt-3.5-turbo
litellm_params:
model: azure/gpt-4.1-mini
- api_base: os.environ/AZURE_API_BASE
- api_key: os.environ/AZURE_API_KEY
+ api_base: os.environ/AZURE_AI_API_BASE
+ api_key: os.environ/AZURE_AI_API_KEY
- model_name: azure-gpt-3.5-turbo
litellm_params:
model: azure/gpt-4.1-mini
- api_base: os.environ/AZURE_API_BASE
+ api_base: os.environ/AZURE_AI_API_BASE
api_key: bad-key
- model_name: azure-embedding
litellm_params:
model: azure/text-embedding-ada-002
- api_base: os.environ/AZURE_API_BASE
+ api_base: os.environ/AZURE_AI_API_BASE
api_key: bad-key
\ No newline at end of file
diff --git a/tests/local_testing/test_configs/test_cloudflare_azure_with_cache_config.yaml b/tests/local_testing/test_configs/test_cloudflare_azure_with_cache_config.yaml
index 99028356183..24240008fe2 100644
--- a/tests/local_testing/test_configs/test_cloudflare_azure_with_cache_config.yaml
+++ b/tests/local_testing/test_configs/test_cloudflare_azure_with_cache_config.yaml
@@ -3,7 +3,7 @@ model_list:
litellm_params:
model: azure/gpt-4.1-mini
api_base: https://gateway.ai.cloudflare.com/v1/0399b10e77ac6668c80404a5ff49eb37/litellm-test/azure-openai/openai-gpt-4-test-v-1
- api_key: os.environ/AZURE_API_KEY
+ api_key: os.environ/AZURE_AI_API_KEY
api_version: 2023-07-01-preview
litellm_settings:
diff --git a/tests/local_testing/test_configs/test_config_no_auth.yaml b/tests/local_testing/test_configs/test_config_no_auth.yaml
index cdc447a5ee6..f4896217049 100644
--- a/tests/local_testing/test_configs/test_config_no_auth.yaml
+++ b/tests/local_testing/test_configs/test_config_no_auth.yaml
@@ -11,7 +11,7 @@ model_list:
model_name: azure-model
- litellm_params:
api_base: https://gateway.ai.cloudflare.com/v1/0399b10e77ac6668c80404a5ff49eb37/litellm-test/azure-openai/openai-gpt-4-test-v-1
- api_key: os.environ/AZURE_API_KEY
+ api_key: os.environ/AZURE_AI_API_KEY
model: azure/gpt-4.1-mini
model_name: azure-cloudflare-model
- litellm_params:
@@ -49,8 +49,8 @@ model_list:
id: 79fc75bf-8e1b-47d5-8d24-9365a854af03
model_name: test_openai_models
- litellm_params:
- api_base: os.environ/AZURE_API_BASE
- api_key: os.environ/AZURE_API_KEY
+ api_base: os.environ/AZURE_AI_API_BASE
+ api_key: os.environ/AZURE_AI_API_KEY
api_version: 2023-07-01-preview
model: azure/text-embedding-ada-002
model_info:
@@ -94,16 +94,16 @@ model_list:
mode: image_generation
model_name: dall-e-3
- litellm_params:
- api_base: os.environ/AZURE_API_BASE
- api_key: os.environ/AZURE_API_KEY
+ api_base: os.environ/AZURE_AI_API_BASE
+ api_key: os.environ/AZURE_AI_API_KEY
api_version: 2023-06-01-preview
model: azure/
model_info:
mode: image_generation
model_name: dall-e-2
- litellm_params:
- api_base: os.environ/AZURE_API_BASE
- api_key: os.environ/AZURE_API_KEY
+ api_base: os.environ/AZURE_AI_API_BASE
+ api_key: os.environ/AZURE_AI_API_KEY
api_version: 2023-07-01-preview
model: azure/text-embedding-ada-002
model_info:
diff --git a/tests/local_testing/test_configs/test_custom_logger.yaml b/tests/local_testing/test_configs/test_custom_logger.yaml
index 22bbfe42be7..464d66bd783 100644
--- a/tests/local_testing/test_configs/test_custom_logger.yaml
+++ b/tests/local_testing/test_configs/test_custom_logger.yaml
@@ -2,8 +2,8 @@ model_list:
- model_name: Azure OpenAI GPT-4 Canada
litellm_params:
model: azure/gpt-4.1-mini
- api_base: os.environ/AZURE_API_BASE
- api_key: os.environ/AZURE_API_KEY
+ api_base: os.environ/AZURE_AI_API_BASE
+ api_key: os.environ/AZURE_AI_API_KEY
api_version: "2023-07-01-preview"
model_info:
mode: chat
@@ -12,8 +12,8 @@ model_list:
- model_name: azure-embedding-model
litellm_params:
model: azure/text-embedding-ada-002
- api_base: os.environ/AZURE_API_BASE
- api_key: os.environ/AZURE_API_KEY
+ api_base: os.environ/AZURE_AI_API_BASE
+ api_key: os.environ/AZURE_AI_API_KEY
api_version: "2023-07-01-preview"
model_info:
mode: embedding
diff --git a/tests/local_testing/test_embedding.py b/tests/local_testing/test_embedding.py
index 48d535fe450..e6776e3329b 100644
--- a/tests/local_testing/test_embedding.py
+++ b/tests/local_testing/test_embedding.py
@@ -108,7 +108,11 @@ def test_openai_embedding_3():
"model, api_base, api_key",
[
# ("azure/text-embedding-ada-002", None, None),
- ("together_ai/BAAI/bge-base-en-v1.5", None, None), # Updated to current Together AI embedding model
+ (
+ "together_ai/BAAI/bge-base-en-v1.5",
+ None,
+ None,
+ ), # Updated to current Together AI embedding model
],
)
@pytest.mark.parametrize("sync_mode", [True, False])
@@ -193,9 +197,9 @@ def _azure_ai_image_mock_response(*args, **kwargs):
"model, api_base, api_key",
[
(
- "azure_ai/Cohere-embed-v3-multilingual-jzu",
- "https://Cohere-embed-v3-multilingual-jzu.eastus2.models.ai.azure.com",
- os.getenv("AZURE_AI_COHERE_API_KEY_2"),
+ "azure_ai/Cohere-embed-v3-multilingual-2",
+ os.getenv("AZURE_AI_API_BASE"),
+ os.getenv("AZURE_AI_API_KEY"),
)
],
)
@@ -292,13 +296,13 @@ def test_openai_embedding_timeouts():
def test_openai_azure_embedding():
try:
- api_key = os.environ["AZURE_API_KEY"]
- api_base = os.environ["AZURE_API_BASE"]
+ api_key = os.environ["AZURE_AI_API_KEY"]
+ api_base = os.environ["AZURE_AI_API_BASE"]
api_version = os.environ["AZURE_API_VERSION"]
os.environ["AZURE_API_VERSION"] = ""
- os.environ["AZURE_API_BASE"] = ""
- os.environ["AZURE_API_KEY"] = ""
+ os.environ["AZURE_AI_API_BASE"] = ""
+ os.environ["AZURE_AI_API_KEY"] = ""
response = embedding(
model="azure/text-embedding-ada-002",
@@ -310,8 +314,8 @@ def test_openai_azure_embedding():
print(response)
os.environ["AZURE_API_VERSION"] = api_version
- os.environ["AZURE_API_BASE"] = api_base
- os.environ["AZURE_API_KEY"] = api_key
+ os.environ["AZURE_AI_API_BASE"] = api_base
+ os.environ["AZURE_AI_API_KEY"] = api_key
except Exception as e:
pytest.fail(f"Error occurred: {e}")
@@ -353,11 +357,11 @@ def test_openai_azure_embedding_optional_arg():
)
mock_client.assert_called_once_with(
- model="test",
- input=["test"],
- extra_body={"azure_ad_token": "test"},
- timeout=600,
- extra_headers={"X-Stainless-Raw-Response": "true"}
+ model="test",
+ input=["test"],
+ extra_body={"azure_ad_token": "test"},
+ timeout=600,
+ extra_headers={"X-Stainless-Raw-Response": "true"},
)
# Verify azure_ad_token is passed in extra_body, not as a direct parameter
assert "azure_ad_token" not in mock_client.call_args.kwargs
@@ -545,7 +549,7 @@ def test_bedrock_embedding_cohere():
"good morning from litellm, attempting to embed data",
"lets test a second string for good measure",
],
- aws_region_name="os.environ/AWS_REGION_NAME_2",
+ aws_region_name="os.environ/AWS_REGION_NAME",
)
assert isinstance(
response["data"][0]["embedding"], list
@@ -811,7 +815,7 @@ def test_watsonx_embeddings(monkeypatch):
monkeypatch.setenv("WATSONX_PROJECT_ID", "mock-project-id")
client = HTTPHandler()
-
+
# Track the actual request made
captured_request = {}
@@ -820,7 +824,7 @@ def test_watsonx_embeddings(monkeypatch):
captured_request["url"] = url
captured_request["headers"] = kwargs.get("headers", {})
captured_request["data"] = kwargs.get("data")
-
+
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {"Content-Type": "application/json"}
@@ -843,10 +847,12 @@ def test_watsonx_embeddings(monkeypatch):
print(f"response: {response}")
assert isinstance(response.usage, litellm.Usage)
-
+
# Verify the request was made correctly
assert "Authorization" in captured_request["headers"]
- assert captured_request["headers"]["Authorization"] == "Bearer mock-watsonx-token"
+ assert (
+ captured_request["headers"]["Authorization"] == "Bearer mock-watsonx-token"
+ )
assert "us-south.ml.cloud.ibm.com" in captured_request["url"]
except litellm.RateLimitError as e:
pass
@@ -1256,9 +1262,7 @@ def test_jina_ai_img_embeddings(input_data, expected_payload_input):
# Call the function we want to test
try:
- litellm.embedding(
- model="jina_ai/jina-embeddings-v4", input=input_data
- )
+ litellm.embedding(model="jina_ai/jina-embeddings-v4", input=input_data)
except Exception as e:
pytest.fail(
f"litellm.embedding call failed with an unexpected exception: {e}"
@@ -1285,105 +1289,113 @@ def test_jina_ai_img_embeddings(input_data, expected_payload_input):
def test_encoding_format_none_not_omitted_from_openai_sdk():
"""
Test that encoding_format=None is explicitly sent to OpenAI SDK.
-
+
This test verifies that when encoding_format is not provided by the user,
liteLLM explicitly sets it to None rather than omitting it. This prevents
the OpenAI SDK from adding its default value of 'base64'.
-
+
Without this fix:
- OpenAI SDK adds encoding_format='base64' as default when parameter is missing
- This causes issues with providers that don't support encoding_format (like Gemini)
-
+
With this fix:
- encoding_format=None is explicitly passed
- OpenAI SDK respects the explicit None and doesn't add defaults
"""
- with patch("litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client") as mock_get_client:
+ with patch(
+ "litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client"
+ ) as mock_get_client:
# Create a mock client instance
mock_client_instance = MagicMock()
mock_get_client.return_value = mock_client_instance
-
+
# Mock the embeddings.with_raw_response.create method
mock_response = MagicMock()
mock_response.parse.return_value = MagicMock(
model_dump=lambda: {
- 'data': [{'embedding': [0.1, 0.2, 0.3], 'index': 0}],
- 'model': 'text-embedding-ada-002',
- 'object': 'list',
- 'usage': {'prompt_tokens': 1, 'total_tokens': 1}
+ "data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}],
+ "model": "text-embedding-ada-002",
+ "object": "list",
+ "usage": {"prompt_tokens": 1, "total_tokens": 1},
}
)
mock_response.headers = {}
-
- mock_client_instance.embeddings.with_raw_response.create.return_value = mock_response
-
+
+ mock_client_instance.embeddings.with_raw_response.create.return_value = (
+ mock_response
+ )
+
# Call the embedding function without encoding_format
response = embedding(
model="text-embedding-ada-002",
input="Hello world",
)
-
+
# Get the call arguments to verify what was sent to OpenAI SDK
call_args = mock_client_instance.embeddings.with_raw_response.create.call_args
- assert call_args is not None, "OpenAI SDK embeddings.create should have been called"
-
+ assert (
+ call_args is not None
+ ), "OpenAI SDK embeddings.create should have been called"
+
call_kwargs = call_args[1] # Get kwargs
-
+
# The key assertion: encoding_format should be in the request with value None
# This prevents OpenAI SDK from adding its default 'base64' value
- assert 'encoding_format' in call_kwargs, (
+ assert "encoding_format" in call_kwargs, (
"encoding_format should be explicitly passed to OpenAI SDK "
"(even if None) to prevent SDK from adding default value"
)
- assert call_kwargs['encoding_format'] is None, (
- "encoding_format should be None when not provided by user"
- )
-
+ assert (
+ call_kwargs["encoding_format"] is None
+ ), "encoding_format should be None when not provided by user"
+
print("✅ PASS: encoding_format=None is correctly passed to OpenAI SDK")
def test_encoding_format_explicit_value_preserved():
"""
Test that explicitly provided encoding_format values are preserved.
-
- When user provides encoding_format='float' or 'base64', it should be
+
+ When user provides encoding_format='float' or 'base64', it should be
sent as-is to the OpenAI SDK.
"""
- with patch("litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client") as mock_get_client:
+ with patch(
+ "litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client"
+ ) as mock_get_client:
# Create a mock client instance
mock_client_instance = MagicMock()
mock_get_client.return_value = mock_client_instance
-
+
# Mock the embeddings.with_raw_response.create method
mock_response = MagicMock()
mock_response.parse.return_value = MagicMock(
model_dump=lambda: {
- 'data': [{'embedding': [0.1, 0.2, 0.3], 'index': 0}],
- 'model': 'text-embedding-ada-002',
- 'object': 'list',
- 'usage': {'prompt_tokens': 1, 'total_tokens': 1}
+ "data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}],
+ "model": "text-embedding-ada-002",
+ "object": "list",
+ "usage": {"prompt_tokens": 1, "total_tokens": 1},
}
)
mock_response.headers = {}
-
- mock_client_instance.embeddings.with_raw_response.create.return_value = mock_response
-
+
+ mock_client_instance.embeddings.with_raw_response.create.return_value = (
+ mock_response
+ )
+
# Test with explicit encoding_format='float'
response = embedding(
- model="text-embedding-ada-002",
- input="Hello world",
- encoding_format="float"
+ model="text-embedding-ada-002", input="Hello world", encoding_format="float"
)
-
+
# Verify the encoding_format was passed correctly
call_args = mock_client_instance.embeddings.with_raw_response.create.call_args
call_kwargs = call_args[1]
-
- assert 'encoding_format' in call_kwargs, (
- "encoding_format should be in the request"
- )
- assert call_kwargs['encoding_format'] == 'float', (
- "encoding_format should be 'float' when explicitly provided"
- )
-
+
+ assert (
+ "encoding_format" in call_kwargs
+ ), "encoding_format should be in the request"
+ assert (
+ call_kwargs["encoding_format"] == "float"
+ ), "encoding_format should be 'float' when explicitly provided"
+
print("✅ PASS: encoding_format='float' is correctly preserved")
diff --git a/tests/local_testing/test_exceptions.py b/tests/local_testing/test_exceptions.py
index 2c950d79067..123442d1cf6 100644
--- a/tests/local_testing/test_exceptions.py
+++ b/tests/local_testing/test_exceptions.py
@@ -162,8 +162,8 @@ def invalid_auth(model): # set the model key to an invalid key, depending on th
temporary_secret_key = os.environ["AWS_SECRET_ACCESS_KEY"]
os.environ["AWS_SECRET_ACCESS_KEY"] = "bad-key"
elif model == "azure/gpt-4.1-mini":
- temporary_key = os.environ["AZURE_API_KEY"]
- os.environ["AZURE_API_KEY"] = "bad-key"
+ temporary_key = os.environ["AZURE_AI_API_KEY"]
+ os.environ["AZURE_AI_API_KEY"] = "bad-key"
elif model == "claude-3-5-haiku-20241022":
temporary_key = os.environ["ANTHROPIC_API_KEY"]
os.environ["ANTHROPIC_API_KEY"] = "bad-key"
@@ -175,9 +175,7 @@ def invalid_auth(model): # set the model key to an invalid key, depending on th
os.environ["AI21_API_KEY"] = "bad-key"
elif "togethercomputer" in model:
temporary_key = os.environ["TOGETHERAI_API_KEY"]
- os.environ["TOGETHERAI_API_KEY"] = (
- "sk-test-togetherai-key-808"
- )
+ os.environ["TOGETHERAI_API_KEY"] = "sk-test-togetherai-key-808"
elif model in litellm.openrouter_models:
temporary_key = os.environ["OPENROUTER_API_KEY"]
os.environ["OPENROUTER_API_KEY"] = "bad-key"
@@ -185,7 +183,6 @@ def invalid_auth(model): # set the model key to an invalid key, depending on th
temporary_key = os.environ["ALEPH_ALPHA_API_KEY"]
os.environ["ALEPH_ALPHA_API_KEY"] = "bad-key"
elif model in litellm.nlp_cloud_models:
- temporary_key = os.environ["NLP_CLOUD_API_KEY"]
os.environ["NLP_CLOUD_API_KEY"] = "bad-key"
elif (
model
@@ -212,7 +209,7 @@ def invalid_auth(model): # set the model key to an invalid key, depending on th
if model == "gpt-3.5-turbo":
os.environ["OPENAI_API_KEY"] = temporary_key
elif model == "chatgpt-test":
- os.environ["AZURE_API_KEY"] = temporary_key
+ os.environ["AZURE_AI_API_KEY"] = temporary_key
azure = True
elif model == "claude-3-5-haiku-20241022":
os.environ["ANTHROPIC_API_KEY"] = temporary_key
@@ -230,7 +227,7 @@ def invalid_auth(model): # set the model key to an invalid key, depending on th
elif model in litellm.aleph_alpha_models:
os.environ["ALEPH_ALPHA_API_KEY"] = temporary_key
elif model in litellm.nlp_cloud_models:
- os.environ["NLP_CLOUD_API_KEY"] = temporary_key
+ os.environ.pop("NLP_CLOUD_API_KEY", None)
elif "bedrock" in model:
os.environ["AWS_ACCESS_KEY_ID"] = temporary_aws_access_key
os.environ["AWS_REGION_NAME"] = temporary_aws_region_name
@@ -259,17 +256,17 @@ def test_completion_azure_exception():
print("azure gpt-3.5 test\n\n")
litellm.set_verbose = True
## Test azure call
- old_azure_key = os.environ["AZURE_API_KEY"]
- os.environ["AZURE_API_KEY"] = "good morning"
+ old_azure_key = os.environ["AZURE_AI_API_KEY"]
+ os.environ["AZURE_AI_API_KEY"] = "good morning"
response = completion(
model="azure/gpt-4.1-mini",
messages=[{"role": "user", "content": "hello"}],
)
- os.environ["AZURE_API_KEY"] = old_azure_key
+ os.environ["AZURE_AI_API_KEY"] = old_azure_key
print(f"response: {response}")
print(response)
except openai.AuthenticationError as e:
- os.environ["AZURE_API_KEY"] = old_azure_key
+ os.environ["AZURE_AI_API_KEY"] = old_azure_key
print("good job got the correct error for azure when key not set")
except Exception as e:
pytest.fail(f"Error occurred: {e}")
@@ -303,8 +300,8 @@ async def asynctest_completion_azure_exception():
print("azure gpt-3.5 test\n\n")
litellm.set_verbose = True
## Test azure call
- old_azure_key = os.environ["AZURE_API_KEY"]
- os.environ["AZURE_API_KEY"] = "good morning"
+ old_azure_key = os.environ["AZURE_AI_API_KEY"]
+ os.environ["AZURE_AI_API_KEY"] = "good morning"
response = await litellm.acompletion(
model="azure/gpt-4.1-mini",
messages=[{"role": "user", "content": "hello"}],
@@ -312,7 +309,7 @@ async def asynctest_completion_azure_exception():
print(f"response: {response}")
print(response)
except openai.AuthenticationError as e:
- os.environ["AZURE_API_KEY"] = old_azure_key
+ os.environ["AZURE_AI_API_KEY"] = old_azure_key
print("good job got the correct error for azure when key not set")
print(e)
except Exception as e:
@@ -495,6 +492,7 @@ def test_completion_bedrock_invalid_role_exception():
== "litellm.BadRequestError: Invalid Message passed in {'role': 'very-bad-role', 'content': 'hello'}"
)
+
@pytest.mark.skip(reason="OpenAI exception changed to a generic error")
def test_content_policy_exceptionimage_generation_openai():
try:
@@ -773,7 +771,15 @@ def test_litellm_predibase_exception():
@pytest.mark.parametrize(
- "provider", ["predibase", "vertex_ai_beta", "anthropic", "databricks", "watsonx", "fireworks_ai"]
+ "provider",
+ [
+ "predibase",
+ "vertex_ai_beta",
+ "anthropic",
+ "databricks",
+ "watsonx",
+ "fireworks_ai",
+ ],
)
def test_exception_mapping(provider):
"""
@@ -826,14 +832,14 @@ def test_fireworks_ai_exception_mapping():
2. Text-based rate limit detection (the main issue fixed)
3. Generic 400 errors that should NOT be rate limits
4. ExceptionCheckers utility function
-
+
Related to: https://github.com/BerriAI/litellm/pull/11455
Based on Fireworks AI documentation: https://docs.fireworks.ai/tools-sdks/python-client/api-reference
"""
import litellm
from litellm.llms.fireworks_ai.common_utils import FireworksAIException
from litellm.litellm_core_utils.exception_mapping_utils import ExceptionCheckers
-
+
# Test scenarios covering all important cases
test_scenarios = [
{
@@ -855,57 +861,63 @@ def test_fireworks_ai_exception_mapping():
"expected_exception": litellm.BadRequestError,
},
]
-
+
# Test each scenario
for scenario in test_scenarios:
mock_exception = FireworksAIException(
- status_code=scenario["status_code"],
- message=scenario["message"],
- headers={}
+ status_code=scenario["status_code"], message=scenario["message"], headers={}
)
-
+
try:
response = litellm.completion(
model="fireworks_ai/llama-v3p1-70b-instruct",
messages=[{"role": "user", "content": "Hello"}],
mock_response=mock_exception,
)
- pytest.fail(f"Expected {scenario['expected_exception'].__name__} to be raised")
+ pytest.fail(
+ f"Expected {scenario['expected_exception'].__name__} to be raised"
+ )
except scenario["expected_exception"] as e:
if scenario["expected_exception"] == litellm.RateLimitError:
assert "rate limit" in str(e).lower() or "429" in str(e)
except Exception as e:
- pytest.fail(f"Expected {scenario['expected_exception'].__name__} but got {type(e).__name__}: {e}")
-
+ pytest.fail(
+ f"Expected {scenario['expected_exception'].__name__} but got {type(e).__name__}: {e}"
+ )
+
# Test ExceptionCheckers.is_error_str_rate_limit() method directly
-
+
# Test cases that should return True (rate limit detected)
rate_limit_strings = [
"429 rate limit exceeded",
- "Rate limit exceeded, please try again later",
+ "Rate limit exceeded, please try again later",
"RATE LIMIT ERROR",
"Error 429: rate limit",
'{"error":{"type":"invalid_request_error","message":"rate limit exceeded, please try again later"}}',
"HTTP 429 Too Many Requests",
]
-
+
for error_str in rate_limit_strings:
- assert ExceptionCheckers.is_error_str_rate_limit(error_str), f"Should detect rate limit in: {error_str}"
-
+ assert ExceptionCheckers.is_error_str_rate_limit(
+ error_str
+ ), f"Should detect rate limit in: {error_str}"
+
# Test cases that should return False (not rate limit)
non_rate_limit_strings = [
"400 Bad Request",
- "Authentication failed",
+ "Authentication failed",
"Invalid model specified",
"Context window exceeded",
"Internal server error",
"",
"Some other error message",
]
-
+
for error_str in non_rate_limit_strings:
- assert not ExceptionCheckers.is_error_str_rate_limit(error_str), f"Should NOT detect rate limit in: {error_str}"
-
+ assert not ExceptionCheckers.is_error_str_rate_limit(
+ error_str
+ ), f"Should NOT detect rate limit in: {error_str}"
+
# Test edge cases
assert not ExceptionCheckers.is_error_str_rate_limit(None) # type: ignore
assert not ExceptionCheckers.is_error_str_rate_limit(42) # type: ignore
@@ -1142,6 +1154,7 @@ def test_openai_gateway_timeout_error():
"""
openai_client = OpenAI()
mapped_target = openai_client.chat.completions.with_raw_response # type: ignore
+
def _return_exception(*args, **kwargs):
import datetime
@@ -1175,13 +1188,17 @@ def test_openai_gateway_timeout_error():
setattr(exception, k, v)
raise exception
- try:
+ try:
with patch.object(
mapped_target,
"create",
side_effect=_return_exception,
):
- litellm.completion(model="openai/gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello world"}], client=openai_client)
+ litellm.completion(
+ model="openai/gpt-3.5-turbo",
+ messages=[{"role": "user", "content": "Hello world"}],
+ client=openai_client,
+ )
pytest.fail("Expected to raise Timeout")
except litellm.Timeout as e:
assert e.status_code == 504
@@ -1350,7 +1367,7 @@ def test_context_window_exceeded_error_from_litellm_proxy():
def test_bad_request_error_with_response_without_request():
"""
Test that BadRequestError handles Response objects without a request attribute.
-
+
This simulates a real scenario where a Response is created without a request
(e.g., in tests or when manually creating error responses), and we need to
ensure it doesn't raise RuntimeError when the exception is created.
@@ -1362,8 +1379,7 @@ def test_bad_request_error_with_response_without_request():
# Create a Response without a request (simulates the scenario that was failing)
response_without_request = Response(status_code=400, text="Bad Request")
-
-
+
# Test that extract_and_raise_litellm_exception can handle this
args = {
"response": response_without_request,
@@ -1371,17 +1387,17 @@ def test_bad_request_error_with_response_without_request():
"model": "gpt-3.5-turbo",
"custom_llm_provider": "openai",
}
-
+
# This should raise BadRequestError without RuntimeError
with pytest.raises(litellm.BadRequestError) as exc_info:
extract_and_raise_litellm_exception(**args)
-
+
# Verify the exception was created successfully
error = exc_info.value
assert error is not None
assert error.model == "gpt-3.5-turbo"
assert error.llm_provider == "openai"
-
+
# Verify the exception has a response (should be minimal error response)
assert error.response is not None
# The response should have a request (minimal error response has one)
@@ -1420,6 +1436,3 @@ async def test_exception_bubbling_up(sync_mode, stream_mode, model):
assert exc_info.value.code == "invalid_value"
assert exc_info.value.param is not None
assert exc_info.value.type == "invalid_request_error"
-
-
-
diff --git a/tests/local_testing/test_gcs_bucket.py b/tests/local_testing/test_gcs_bucket.py
index e2f51eb76c7..018a79b53f0 100644
--- a/tests/local_testing/test_gcs_bucket.py
+++ b/tests/local_testing/test_gcs_bucket.py
@@ -22,6 +22,7 @@ from litellm.integrations.gcs_bucket.gcs_bucket import (
)
from litellm.types.utils import StandardCallbackDynamicParams
from unittest.mock import patch
+
verbose_logger.setLevel(logging.DEBUG)
@@ -52,8 +53,8 @@ def load_vertex_ai_credentials():
service_account_key_data = {}
# Update the service_account_key_data with environment variables
- private_key_id = os.environ.get("GCS_PRIVATE_KEY_ID", "")
- private_key = os.environ.get("GCS_PRIVATE_KEY", "")
+ private_key_id = os.environ.get("VERTEX_AI_PRIVATE_KEY_ID", "")
+ private_key = os.environ.get("VERTEX_AI_PRIVATE_KEY", "")
private_key = private_key.replace("\\n", "\n")
service_account_key_data["private_key_id"] = private_key_id
service_account_key_data["private_key"] = private_key
@@ -686,6 +687,7 @@ async def test_basic_gcs_logger_with_folder_in_bucket_name():
if old_bucket_name is not None:
os.environ["GCS_BUCKET_NAME"] = old_bucket_name
+
@pytest.mark.skip(reason="This test is flaky on ci/cd")
def test_create_file_e2e():
"""
@@ -696,6 +698,7 @@ def test_create_file_e2e():
test_file = ("test.wav", test_file_content, "audio/wav")
from litellm import create_file
+
response = create_file(
file=test_file,
purpose="user_data",
@@ -704,6 +707,7 @@ def test_create_file_e2e():
print("response", response)
assert response is not None
+
@pytest.mark.skip(reason="This test is flaky on ci/cd")
def test_create_file_e2e_jsonl():
"""
@@ -714,14 +718,41 @@ def test_create_file_e2e_jsonl():
client = HTTPHandler()
- example_jsonl = [{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gemini-1.5-flash-001", "messages": [{"role": "system", "content": "You are a helpful assistant."},{"role": "user", "content": "Hello world!"}],"max_tokens": 10}},{"custom_id": "request-2", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gemini-1.5-flash-001", "messages": [{"role": "system", "content": "You are an unhelpful assistant."},{"role": "user", "content": "Hello world!"}],"max_tokens": 10}}]
-
+ example_jsonl = [
+ {
+ "custom_id": "request-1",
+ "method": "POST",
+ "url": "/v1/chat/completions",
+ "body": {
+ "model": "gemini-1.5-flash-001",
+ "messages": [
+ {"role": "system", "content": "You are a helpful assistant."},
+ {"role": "user", "content": "Hello world!"},
+ ],
+ "max_tokens": 10,
+ },
+ },
+ {
+ "custom_id": "request-2",
+ "method": "POST",
+ "url": "/v1/chat/completions",
+ "body": {
+ "model": "gemini-1.5-flash-001",
+ "messages": [
+ {"role": "system", "content": "You are an unhelpful assistant."},
+ {"role": "user", "content": "Hello world!"},
+ ],
+ "max_tokens": 10,
+ },
+ },
+ ]
+
# Create and write to the file
file_path = "example.jsonl"
with open(file_path, "w") as f:
for item in example_jsonl:
f.write(json.dumps(item) + "\n")
-
+
# Verify file content
with open(file_path, "r") as f:
content = f.read()
@@ -729,10 +760,11 @@ def test_create_file_e2e_jsonl():
assert len(content) > 0, "File is empty"
from litellm import create_file
+
with patch.object(client, "post") as mock_create_file:
- try:
+ try:
response = create_file(
- file=open(file_path, "rb"),
+ file=open(file_path, "rb"),
purpose="user_data",
custom_llm_provider="vertex_ai",
client=client,
@@ -744,4 +776,7 @@ def test_create_file_e2e_jsonl():
print(f"kwargs: {mock_create_file.call_args.kwargs}")
- assert mock_create_file.call_args.kwargs["data"] is not None and len(mock_create_file.call_args.kwargs["data"]) > 0
\ No newline at end of file
+ assert (
+ mock_create_file.call_args.kwargs["data"] is not None
+ and len(mock_create_file.call_args.kwargs["data"]) > 0
+ )
diff --git a/tests/local_testing/test_loadtest_router.py b/tests/local_testing/test_loadtest_router.py
index 3f6e4af4fb4..3d1062f0d26 100644
--- a/tests/local_testing/test_loadtest_router.py
+++ b/tests/local_testing/test_loadtest_router.py
@@ -39,8 +39,8 @@
# "model_name": "gpt-3.5-turbo",
# "litellm_params": {
# "model": "azure/gpt-4.1-mini",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "api_base": os.getenv("AZURE_API_BASE"),
+# "api_key": os.getenv("AZURE_AI_API_KEY"),
+# "api_base": os.getenv("AZURE_AI_API_BASE"),
# "api_version": os.getenv("AZURE_API_VERSION"),
# },
# },
diff --git a/tests/local_testing/test_prompt_injection_detection.py b/tests/local_testing/test_prompt_injection_detection.py
index c4cc4cde32e..b1a9aff1584 100644
--- a/tests/local_testing/test_prompt_injection_detection.py
+++ b/tests/local_testing/test_prompt_injection_detection.py
@@ -108,9 +108,9 @@ async def test_prompt_injection_llm_eval():
"model_name": "gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
diff --git a/tests/local_testing/test_router.py b/tests/local_testing/test_router.py
index e68d271a113..f7885fb8a03 100644
--- a/tests/local_testing/test_router.py
+++ b/tests/local_testing/test_router.py
@@ -126,7 +126,9 @@ async def test_router_provider_wildcard_routing():
print("response 3 = ", response3)
response4 = await router.acompletion(
- model=os.environ.get("CI_CD_DEFAULT_ANTHROPIC_MODEL", "claude-haiku-4-5-20251001"),
+ model=os.environ.get(
+ "CI_CD_DEFAULT_ANTHROPIC_MODEL", "claude-haiku-4-5-20251001"
+ ),
messages=[{"role": "user", "content": "hello"}],
)
@@ -356,51 +358,6 @@ async def test_router_retries(sync_mode):
print(response.choices[0].message)
-@pytest.mark.parametrize(
- "mistral_api_base",
- [
- "os.environ/AZURE_MISTRAL_API_BASE",
- "https://Mistral-large-nmefg-serverless.eastus2.inference.ai.azure.com/v1/",
- "https://Mistral-large-nmefg-serverless.eastus2.inference.ai.azure.com/v1",
- "https://Mistral-large-nmefg-serverless.eastus2.inference.ai.azure.com/",
- "https://Mistral-large-nmefg-serverless.eastus2.inference.ai.azure.com",
- ],
-)
-@pytest.mark.skip(
- reason="Router no longer creates clients, this is delegated to the provider integration."
-)
-def test_router_azure_ai_studio_init(mistral_api_base):
- router = Router(
- model_list=[
- {
- "model_name": "test-model",
- "litellm_params": {
- "model": "azure/mistral-large-latest",
- "api_key": "os.environ/AZURE_MISTRAL_API_KEY",
- "api_base": mistral_api_base,
- },
- "model_info": {"id": 1234},
- }
- ]
- )
-
- # model_client = router._get_client(
- # deployment={"model_info": {"id": 1234}}, client_type="sync_client", kwargs={}
- # )
- # url = getattr(model_client, "_base_url")
- # uri_reference = str(getattr(url, "_uri_reference"))
-
- # print(f"uri_reference: {uri_reference}")
-
- # assert "/v1/" in uri_reference
- # assert uri_reference.count("v1") == 1
- response = router.completion(
- model="azure/mistral-large-latest",
- messages=[{"role": "user", "content": "Hey, how's it going?"}],
- )
- assert response is not None
-
-
def test_exception_raising():
# this tests if the router raises an exception when invalid params are set
# in this test both deployments have bad keys - Keep this test. It validates if the router raises the most recent exception
@@ -409,8 +366,8 @@ def test_exception_raising():
try:
print("testing if router raises an exception")
- old_api_key = os.environ["AZURE_API_KEY"]
- os.environ["AZURE_API_KEY"] = ""
+ old_api_key = os.environ["AZURE_AI_API_KEY"]
+ os.environ["AZURE_AI_API_KEY"] = ""
model_list = [
{
"model_name": "gpt-3.5-turbo", # openai model name
@@ -418,7 +375,7 @@ def test_exception_raising():
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -446,16 +403,16 @@ def test_exception_raising():
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hello this request will fail"}],
)
- os.environ["AZURE_API_KEY"] = old_api_key
+ os.environ["AZURE_AI_API_KEY"] = old_api_key
pytest.fail(f"Should have raised an Auth Error")
except openai.AuthenticationError:
print(
"Test Passed: Caught an OPENAI AUTH Error, Good job. This is what we needed!"
)
- os.environ["AZURE_API_KEY"] = old_api_key
+ os.environ["AZURE_AI_API_KEY"] = old_api_key
router.reset()
except Exception as e:
- os.environ["AZURE_API_KEY"] = old_api_key
+ os.environ["AZURE_AI_API_KEY"] = old_api_key
print("Got unexpected exception on router!", e)
@@ -530,7 +487,7 @@ def test_call_one_endpoint():
# this test makes a completion calls azure/gpt-4.1-mini, it should work
try:
print("Testing calling a specific deployment")
- old_api_key = os.environ["AZURE_API_KEY"]
+ old_api_key = os.environ["AZURE_AI_API_KEY"]
model_list = [
{
@@ -539,7 +496,7 @@ def test_call_one_endpoint():
"model": "azure/gpt-4.1-mini",
"api_key": old_api_key,
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -548,8 +505,8 @@ def test_call_one_endpoint():
"model_name": "text-embedding-ada-002",
"litellm_params": {
"model": "azure/text-embedding-ada-002",
- "api_key": os.environ["AZURE_API_KEY"],
- "api_base": os.environ["AZURE_API_BASE"],
+ "api_key": os.environ["AZURE_AI_API_KEY"],
+ "api_base": os.environ["AZURE_AI_API_BASE"],
},
"tpm": 100000,
"rpm": 10000,
@@ -562,7 +519,7 @@ def test_call_one_endpoint():
set_verbose=True,
num_retries=1,
) # type: ignore
- old_api_base = os.environ.pop("AZURE_API_BASE", None)
+ old_api_base = os.environ.pop("AZURE_AI_API_BASE", None)
async def call_azure_completion():
response = await router.acompletion(
@@ -584,8 +541,8 @@ def test_call_one_endpoint():
asyncio.run(call_azure_completion())
asyncio.run(call_azure_embedding())
- os.environ["AZURE_API_BASE"] = old_api_base
- os.environ["AZURE_API_KEY"] = old_api_key
+ os.environ["AZURE_AI_API_BASE"] = old_api_base
+ os.environ["AZURE_AI_API_KEY"] = old_api_key
except Exception as e:
print(f"FAILED TEST")
pytest.fail(f"Got unexpected exception on router! - {e}")
@@ -594,7 +551,6 @@ def test_call_one_endpoint():
# test_call_one_endpoint()
-
@pytest.mark.asyncio
@pytest.mark.parametrize("sync_mode", [True, False])
async def test_async_router_context_window_fallback(sync_mode):
@@ -708,9 +664,9 @@ def test_router_context_window_check_pre_call_check_in_group_custom_model_info()
"model_name": "gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
"base_model": "azure/gpt-35-turbo",
"mock_response": "Hello world 1!",
},
@@ -762,9 +718,9 @@ def test_router_context_window_check_pre_call_check():
"model_name": "gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
"base_model": "azure/gpt-35-turbo",
"mock_response": "Hello world 1!",
},
@@ -816,9 +772,9 @@ def test_router_context_window_check_pre_call_check_out_group():
"model_name": "gpt-3.5-turbo-small", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
"base_model": "azure/gpt-35-turbo",
},
},
@@ -896,9 +852,9 @@ def test_router_region_pre_call_check(allowed_model_region):
"model_name": "gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
"base_model": "azure/gpt-35-turbo",
"region_name": allowed_model_region,
},
@@ -1173,8 +1129,8 @@ def test_azure_embedding_on_router():
"model_name": "text-embedding-ada-002",
"litellm_params": {
"model": "azure/text-embedding-ada-002",
- "api_key": os.environ["AZURE_API_KEY"],
- "api_base": os.environ["AZURE_API_BASE"],
+ "api_key": os.environ["AZURE_AI_API_KEY"],
+ "api_base": os.environ["AZURE_AI_API_BASE"],
},
"tpm": 100000,
"rpm": 10000,
@@ -1381,8 +1337,8 @@ def test_reading_keys_os_environ():
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
- "api_key": "os.environ/AZURE_API_KEY",
- "api_base": "os.environ/AZURE_API_BASE",
+ "api_key": "os.environ/AZURE_AI_API_KEY",
+ "api_base": "os.environ/AZURE_AI_API_BASE",
"api_version": "os.environ/AZURE_API_VERSION",
"timeout": "os.environ/AZURE_TIMEOUT",
"stream_timeout": "os.environ/AZURE_STREAM_TIMEOUT",
@@ -1394,11 +1350,11 @@ def test_reading_keys_os_environ():
router = Router(model_list=model_list)
for model in router.model_list:
assert (
- model["litellm_params"]["api_key"] == os.environ["AZURE_API_KEY"]
- ), f"{model['litellm_params']['api_key']} vs {os.environ['AZURE_API_KEY']}"
+ model["litellm_params"]["api_key"] == os.environ["AZURE_AI_API_KEY"]
+ ), f"{model['litellm_params']['api_key']} vs {os.environ['AZURE_AI_API_KEY']}"
assert (
- model["litellm_params"]["api_base"] == os.environ["AZURE_API_BASE"]
- ), f"{model['litellm_params']['api_base']} vs {os.environ['AZURE_API_BASE']}"
+ model["litellm_params"]["api_base"] == os.environ["AZURE_AI_API_BASE"]
+ ), f"{model['litellm_params']['api_base']} vs {os.environ['AZURE_AI_API_BASE']}"
assert (
model["litellm_params"]["api_version"]
== os.environ["AZURE_API_VERSION"]
@@ -1415,8 +1371,8 @@ def test_reading_keys_os_environ():
print("passed testing of reading keys from os.environ")
model_id = model["model_info"]["id"]
async_client: openai.AsyncAzureOpenAI = router.cache.get_cache(f"{model_id}_async_client") # type: ignore
- assert async_client.api_key == os.environ["AZURE_API_KEY"]
- assert async_client.base_url == os.environ["AZURE_API_BASE"]
+ assert async_client.api_key == os.environ["AZURE_AI_API_KEY"]
+ assert async_client.base_url == os.environ["AZURE_AI_API_BASE"]
assert async_client.max_retries == int(
os.environ["AZURE_MAX_RETRIES"]
), f"{async_client.max_retries} vs {os.environ['AZURE_MAX_RETRIES']}"
@@ -1428,8 +1384,8 @@ def test_reading_keys_os_environ():
print("\n Testing async streaming client")
stream_async_client: openai.AsyncAzureOpenAI = router.cache.get_cache(f"{model_id}_stream_async_client") # type: ignore
- assert stream_async_client.api_key == os.environ["AZURE_API_KEY"]
- assert stream_async_client.base_url == os.environ["AZURE_API_BASE"]
+ assert stream_async_client.api_key == os.environ["AZURE_AI_API_KEY"]
+ assert stream_async_client.base_url == os.environ["AZURE_AI_API_BASE"]
assert stream_async_client.max_retries == int(
os.environ["AZURE_MAX_RETRIES"]
), f"{stream_async_client.max_retries} vs {os.environ['AZURE_MAX_RETRIES']}"
@@ -1440,8 +1396,8 @@ def test_reading_keys_os_environ():
print("\n Testing sync client")
client: openai.AzureOpenAI = router.cache.get_cache(f"{model_id}_client") # type: ignore
- assert client.api_key == os.environ["AZURE_API_KEY"]
- assert client.base_url == os.environ["AZURE_API_BASE"]
+ assert client.api_key == os.environ["AZURE_AI_API_KEY"]
+ assert client.base_url == os.environ["AZURE_AI_API_BASE"]
assert client.max_retries == int(
os.environ["AZURE_MAX_RETRIES"]
), f"{client.max_retries} vs {os.environ['AZURE_MAX_RETRIES']}"
@@ -1452,8 +1408,8 @@ def test_reading_keys_os_environ():
print("\n Testing sync stream client")
stream_client: openai.AzureOpenAI = router.cache.get_cache(f"{model_id}_stream_client") # type: ignore
- assert stream_client.api_key == os.environ["AZURE_API_KEY"]
- assert stream_client.base_url == os.environ["AZURE_API_BASE"]
+ assert stream_client.api_key == os.environ["AZURE_AI_API_KEY"]
+ assert stream_client.base_url == os.environ["AZURE_AI_API_BASE"]
assert stream_client.max_retries == int(
os.environ["AZURE_MAX_RETRIES"]
), f"{stream_client.max_retries} vs {os.environ['AZURE_MAX_RETRIES']}"
@@ -1503,7 +1459,7 @@ def test_reading_openai_keys_os_environ():
for model in router.model_list:
assert (
model["litellm_params"]["api_key"] == os.environ["OPENAI_API_KEY"]
- ), f"{model['litellm_params']['api_key']} vs {os.environ['AZURE_API_KEY']}"
+ ), f"{model['litellm_params']['api_key']} vs {os.environ['AZURE_AI_API_KEY']}"
assert float(model["litellm_params"]["timeout"]) == float(
os.environ["AZURE_TIMEOUT"]
), f"{model['litellm_params']['timeout']} vs {os.environ['AZURE_TIMEOUT']}"
@@ -1574,7 +1530,9 @@ def test_router_anthropic_key_dynamic():
{
"model_name": "anthropic-claude",
"litellm_params": {
- "model": os.environ.get("CI_CD_DEFAULT_ANTHROPIC_MODEL", "claude-haiku-4-5-20251001"),
+ "model": os.environ.get(
+ "CI_CD_DEFAULT_ANTHROPIC_MODEL", "claude-haiku-4-5-20251001"
+ ),
"api_key": anthropic_api_key,
},
}
@@ -2273,8 +2231,8 @@ async def test_router_batch_endpoints(provider):
"model_name": "my-custom-name",
"litellm_params": {
"model": "azure/gpt-4o-mini",
- "api_base": os.getenv("AZURE_API_BASE"),
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
},
},
]
@@ -2452,8 +2410,8 @@ def test_is_team_specific_model():
# "model_name": "gpt-3.5-turbo",
# "litellm_params": {
# "model": "azure/gpt-4.1-mini",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "api_base": os.getenv("AZURE_API_BASE"),
+# "api_key": os.getenv("AZURE_AI_API_KEY"),
+# "api_base": os.getenv("AZURE_AI_API_BASE"),
# "tpm": 100000,
# "rpm": 100000,
# },
@@ -2462,8 +2420,8 @@ def test_is_team_specific_model():
# "model_name": "gpt-3.5-turbo",
# "litellm_params": {
# "model": "azure/gpt-4.1-mini",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "api_base": os.getenv("AZURE_API_BASE"),
+# "api_key": os.getenv("AZURE_AI_API_KEY"),
+# "api_base": os.getenv("AZURE_AI_API_BASE"),
# "tpm": 500,
# "rpm": 500,
# },
diff --git a/tests/local_testing/test_router_budget_limiter.py b/tests/local_testing/test_router_budget_limiter.py
index 05ce7c1f53c..1a36e9de8f2 100644
--- a/tests/local_testing/test_router_budget_limiter.py
+++ b/tests/local_testing/test_router_budget_limiter.py
@@ -75,9 +75,9 @@ async def test_provider_budgets_e2e_test():
"model_name": "gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"model_info": {"id": "azure-model-id"},
},
@@ -609,6 +609,7 @@ async def test_deployment_budgets_e2e_test_expect_to_fail():
assert "Exceeded budget for deployment" in str(exc_info.value)
+
@pytest.mark.flaky(retries=6, delay=2)
@pytest.mark.asyncio
async def test_tag_budgets_e2e_test_expect_to_fail():
diff --git a/tests/local_testing/test_router_caching.py b/tests/local_testing/test_router_caching.py
index 6fc220bf728..cb223b661b4 100644
--- a/tests/local_testing/test_router_caching.py
+++ b/tests/local_testing/test_router_caching.py
@@ -268,8 +268,8 @@ async def test_acompletion_caching_on_router_caching_groups():
"model_name": "azure-gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
"api_version": os.getenv("AZURE_API_VERSION"),
},
"tpm": 100000,
diff --git a/tests/local_testing/test_router_client_init.py b/tests/local_testing/test_router_client_init.py
index f2541601c3b..6cf8b4e59b7 100644
--- a/tests/local_testing/test_router_client_init.py
+++ b/tests/local_testing/test_router_client_init.py
@@ -85,7 +85,7 @@ def test_router_init_azure_service_principal_with_secret_with_environment_variab
To allow for local testing without real credentials, first must mock Azure SDK authentication functions
and environment variables.
"""
- monkeypatch.delenv("AZURE_API_KEY", raising=False)
+ monkeypatch.delenv("AZURE_AI_API_KEY", raising=False)
litellm.enable_azure_ad_token_refresh = True
# mock the token provider function
mocked_func_generating_token = MagicMock(return_value="test_token")
@@ -174,9 +174,9 @@ async def test_audio_speech_router():
{
"model_name": "tts",
"litellm_params": {
- "model": "azure/azure-tts",
- "api_base": os.getenv("AZURE_SWEDEN_API_BASE"),
- "api_key": os.getenv("AZURE_SWEDEN_API_KEY"),
+ "model": "azure/tts",
+ "api_base": os.getenv("AZURE_TTS_API_BASE"),
+ "api_key": os.getenv("AZURE_TTS_API_KEY"),
},
},
]
diff --git a/tests/local_testing/test_router_cooldown_handlers.py b/tests/local_testing/test_router_cooldown_handlers.py
index 7be8289abf1..fdc89fc04ed 100644
--- a/tests/local_testing/test_router_cooldown_handlers.py
+++ b/tests/local_testing/test_router_cooldown_handlers.py
@@ -45,9 +45,9 @@ async def test_cooldown_badrequest_error():
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
}
],
diff --git a/tests/local_testing/test_router_debug_logs.py b/tests/local_testing/test_router_debug_logs.py
index 1004e7747ef..0b4771b5267 100644
--- a/tests/local_testing/test_router_debug_logs.py
+++ b/tests/local_testing/test_router_debug_logs.py
@@ -34,9 +34,9 @@ def test_async_fallbacks(caplog):
"model_name": "azure/gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
"mock_response": "Hello world",
},
"tpm": 240000,
diff --git a/tests/local_testing/test_router_fallbacks.py b/tests/local_testing/test_router_fallbacks.py
index c586fa8c93b..383ad104577 100644
--- a/tests/local_testing/test_router_fallbacks.py
+++ b/tests/local_testing/test_router_fallbacks.py
@@ -70,7 +70,7 @@ def test_sync_fallbacks():
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -79,9 +79,9 @@ def test_sync_fallbacks():
"model_name": "azure/gpt-3.5-turbo-context-fallback", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -92,7 +92,7 @@ def test_sync_fallbacks():
"model": "azure/chatgpt-functioncalling",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -132,7 +132,9 @@ def test_sync_fallbacks():
response = router.completion(**kwargs)
print(f"response: {response}")
time.sleep(0.05) # allow a delay as success_callbacks are on a separate thread
- assert customHandler.previous_models == 3 # 1 init call + 2 retries (fallback not counted as previous)
+ assert (
+ customHandler.previous_models == 3
+ ) # 1 init call + 2 retries (fallback not counted as previous)
print("Passed ! Test router_fallbacks: test_sync_fallbacks()")
router.reset()
@@ -153,7 +155,7 @@ async def test_async_fallbacks():
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -162,9 +164,9 @@ async def test_async_fallbacks():
"model_name": "azure/gpt-3.5-turbo-context-fallback", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -175,7 +177,7 @@ async def test_async_fallbacks():
"model": "azure/chatgpt-functioncalling",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -220,7 +222,9 @@ async def test_async_fallbacks():
await asyncio.sleep(
0.05
) # allow a delay as success_callbacks are on a separate thread
- assert customHandler.previous_models == 3 # 1 init call + 2 retries (fallback not counted as previous)
+ assert (
+ customHandler.previous_models == 3
+ ) # 1 init call + 2 retries (fallback not counted as previous)
router.reset()
except litellm.Timeout as e:
pass
@@ -242,7 +246,7 @@ def test_sync_fallbacks_embeddings():
"model": "azure/text-embedding-ada-002",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -292,7 +296,7 @@ async def test_async_fallbacks_embeddings():
"model": "azure/text-embedding-ada-002",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -348,7 +352,7 @@ def test_dynamic_fallbacks_sync():
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -357,9 +361,9 @@ def test_dynamic_fallbacks_sync():
"model_name": "azure/gpt-3.5-turbo-context-fallback", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -370,7 +374,7 @@ def test_dynamic_fallbacks_sync():
"model": "azure/chatgpt-functioncalling",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -403,7 +407,9 @@ def test_dynamic_fallbacks_sync():
response = router.completion(**kwargs)
print(f"response: {response}")
time.sleep(0.05) # allow a delay as success_callbacks are on a separate thread
- assert customHandler.previous_models >= 3 # 1 init call, retries, 1 fallback (count varies with cooldown timing)
+ assert (
+ customHandler.previous_models >= 3
+ ) # 1 init call, retries, 1 fallback (count varies with cooldown timing)
router.reset()
except Exception as e:
pytest.fail(f"An exception occurred - {e}")
@@ -425,7 +431,7 @@ async def test_dynamic_fallbacks_async():
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -434,9 +440,9 @@ async def test_dynamic_fallbacks_async():
"model_name": "azure/gpt-3.5-turbo-context-fallback", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -447,7 +453,7 @@ async def test_dynamic_fallbacks_async():
"model": "azure/chatgpt-functioncalling",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -489,7 +495,9 @@ async def test_dynamic_fallbacks_async():
await asyncio.sleep(
0.05
) # allow a delay as success_callbacks are on a separate thread
- assert customHandler.previous_models >= 3 # 1 init call, retries, 1 fallback (count varies with cooldown timing)
+ assert (
+ customHandler.previous_models >= 3
+ ) # 1 init call, retries, 1 fallback (count varies with cooldown timing)
router.reset()
except Exception as e:
pytest.fail(f"An exception occurred - {e}")
@@ -562,7 +570,7 @@ def test_sync_fallbacks_streaming():
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -571,9 +579,9 @@ def test_sync_fallbacks_streaming():
"model_name": "azure/gpt-3.5-turbo-context-fallback", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -584,7 +592,7 @@ def test_sync_fallbacks_streaming():
"model": "azure/chatgpt-functioncalling",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -643,7 +651,7 @@ async def test_async_fallbacks_max_retries_per_request():
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -652,9 +660,9 @@ async def test_async_fallbacks_max_retries_per_request():
"model_name": "azure/gpt-3.5-turbo-context-fallback", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -665,7 +673,7 @@ async def test_async_fallbacks_max_retries_per_request():
"model": "azure/chatgpt-functioncalling",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -750,9 +758,9 @@ def test_ausage_based_routing_fallbacks():
def get_azure_params(deployment_name: str):
params = {
"model": f"azure/{deployment_name}",
- "api_key": os.environ["AZURE_API_KEY"],
+ "api_key": os.environ["AZURE_AI_API_KEY"],
"api_version": os.environ["AZURE_API_VERSION"],
- "api_base": os.environ["AZURE_API_BASE"],
+ "api_base": os.environ["AZURE_AI_API_BASE"],
}
return params
@@ -855,7 +863,7 @@ def test_custom_cooldown_times():
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 24000000,
},
@@ -863,9 +871,9 @@ def test_custom_cooldown_times():
"model_name": "gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 1,
},
diff --git a/tests/local_testing/test_router_init.py b/tests/local_testing/test_router_init.py
deleted file mode 100644
index e232ff105de..00000000000
--- a/tests/local_testing/test_router_init.py
+++ /dev/null
@@ -1,704 +0,0 @@
-# # this tests if the router is initialized correctly
-# import asyncio
-# import os
-# import sys
-# import time
-# import traceback
-
-# import pytest
-
-# sys.path.insert(
-# 0, os.path.abspath("../..")
-# ) # Adds the parent directory to the system path
-# from collections import defaultdict
-# from concurrent.futures import ThreadPoolExecutor
-
-# from dotenv import load_dotenv
-
-# import litellm
-# from litellm import Router
-
-# load_dotenv()
-
-# # every time we load the router we should have 4 clients:
-# # Async
-# # Sync
-# # Async + Stream
-# # Sync + Stream
-
-
-# def test_init_clients():
-# litellm.set_verbose = True
-# import logging
-
-# from litellm._logging import verbose_router_logger
-
-# verbose_router_logger.setLevel(logging.DEBUG)
-# try:
-# print("testing init 4 clients with diff timeouts")
-# model_list = [
-# {
-# "model_name": "gpt-3.5-turbo",
-# "litellm_params": {
-# "model": "azure/gpt-4.1-mini",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "api_version": os.getenv("AZURE_API_VERSION"),
-# "api_base": os.getenv("AZURE_API_BASE"),
-# "timeout": 0.01,
-# "stream_timeout": 0.000_001,
-# "max_retries": 7,
-# },
-# },
-# ]
-# router = Router(model_list=model_list, set_verbose=True)
-# for elem in router.model_list:
-# model_id = elem["model_info"]["id"]
-# assert router.cache.get_cache(f"{model_id}_client") is not None
-# assert router.cache.get_cache(f"{model_id}_async_client") is not None
-# assert router.cache.get_cache(f"{model_id}_stream_client") is not None
-# assert router.cache.get_cache(f"{model_id}_stream_async_client") is not None
-
-# # check if timeout for stream/non stream clients is set correctly
-# async_client = router.cache.get_cache(f"{model_id}_async_client")
-# stream_async_client = router.cache.get_cache(
-# f"{model_id}_stream_async_client"
-# )
-
-# assert async_client.timeout == 0.01
-# assert stream_async_client.timeout == 0.000_001
-# print(vars(async_client))
-# print()
-# print(async_client._base_url)
-# assert (
-# async_client._base_url
-# == "https://openai-gpt-4-test-v-1.openai.azure.com/openai/"
-# )
-# assert (
-# stream_async_client._base_url
-# == "https://openai-gpt-4-test-v-1.openai.azure.com/openai/"
-# )
-
-# print("PASSED !")
-
-# except Exception as e:
-# traceback.print_exc()
-# pytest.fail(f"Error occurred: {e}")
-
-
-# # test_init_clients()
-
-
-# def test_init_clients_basic():
-# litellm.set_verbose = True
-# try:
-# print("Test basic client init")
-# model_list = [
-# {
-# "model_name": "gpt-3.5-turbo",
-# "litellm_params": {
-# "model": "azure/gpt-4.1-mini",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "api_version": os.getenv("AZURE_API_VERSION"),
-# "api_base": os.getenv("AZURE_API_BASE"),
-# },
-# },
-# ]
-# router = Router(model_list=model_list)
-# for elem in router.model_list:
-# model_id = elem["model_info"]["id"]
-# assert router.cache.get_cache(f"{model_id}_client") is not None
-# assert router.cache.get_cache(f"{model_id}_async_client") is not None
-# assert router.cache.get_cache(f"{model_id}_stream_client") is not None
-# assert router.cache.get_cache(f"{model_id}_stream_async_client") is not None
-# print("PASSED !")
-
-# # see if we can init clients without timeout or max retries set
-# except Exception as e:
-# traceback.print_exc()
-# pytest.fail(f"Error occurred: {e}")
-
-
-# # test_init_clients_basic()
-
-
-# def test_init_clients_basic_azure_cloudflare():
-# # init azure + cloudflare
-# # init OpenAI gpt-3.5
-# # init OpenAI text-embedding
-# # init OpenAI comptaible - Mistral/mistral-medium
-# # init OpenAI compatible - xinference/bge
-# litellm.set_verbose = True
-# try:
-# print("Test basic client init")
-# model_list = [
-# {
-# "model_name": "azure-cloudflare",
-# "litellm_params": {
-# "model": "azure/gpt-4.1-mini",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "api_version": os.getenv("AZURE_API_VERSION"),
-# "api_base": "https://gateway.ai.cloudflare.com/v1/0399b10e77ac6668c80404a5ff49eb37/litellm-test/azure-openai/openai-gpt-4-test-v-1",
-# },
-# },
-# {
-# "model_name": "gpt-openai",
-# "litellm_params": {
-# "model": "gpt-3.5-turbo",
-# "api_key": os.getenv("OPENAI_API_KEY"),
-# },
-# },
-# {
-# "model_name": "text-embedding-ada-002",
-# "litellm_params": {
-# "model": "text-embedding-ada-002",
-# "api_key": os.getenv("OPENAI_API_KEY"),
-# },
-# },
-# {
-# "model_name": "mistral",
-# "litellm_params": {
-# "model": "mistral/mistral-tiny",
-# "api_key": os.getenv("MISTRAL_API_KEY"),
-# },
-# },
-# {
-# "model_name": "bge-base-en",
-# "litellm_params": {
-# "model": "xinference/bge-base-en",
-# "api_base": "http://127.0.0.1:9997/v1",
-# "api_key": os.getenv("OPENAI_API_KEY"),
-# },
-# },
-# ]
-# router = Router(model_list=model_list)
-# for elem in router.model_list:
-# model_id = elem["model_info"]["id"]
-# assert router.cache.get_cache(f"{model_id}_client") is not None
-# assert router.cache.get_cache(f"{model_id}_async_client") is not None
-# assert router.cache.get_cache(f"{model_id}_stream_client") is not None
-# assert router.cache.get_cache(f"{model_id}_stream_async_client") is not None
-# print("PASSED !")
-
-# # see if we can init clients without timeout or max retries set
-# except Exception as e:
-# traceback.print_exc()
-# pytest.fail(f"Error occurred: {e}")
-
-
-# # test_init_clients_basic_azure_cloudflare()
-
-
-# def test_timeouts_router():
-# """
-# Test the timeouts of the router with multiple clients. This HASas to raise a timeout error
-# """
-# import openai
-
-# litellm.set_verbose = True
-# try:
-# print("testing init 4 clients with diff timeouts")
-# model_list = [
-# {
-# "model_name": "gpt-3.5-turbo",
-# "litellm_params": {
-# "model": "azure/gpt-4.1-mini",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "api_version": os.getenv("AZURE_API_VERSION"),
-# "api_base": os.getenv("AZURE_API_BASE"),
-# "timeout": 0.000001,
-# "stream_timeout": 0.000_001,
-# },
-# },
-# ]
-# router = Router(model_list=model_list, num_retries=0)
-
-# print("PASSED !")
-
-# async def test():
-# try:
-# await router.acompletion(
-# model="gpt-3.5-turbo",
-# messages=[
-# {"role": "user", "content": "hello, write a 20 pg essay"}
-# ],
-# )
-# except Exception as e:
-# raise e
-
-# asyncio.run(test())
-# except openai.APITimeoutError as e:
-# print(
-# "Passed: Raised correct exception. Got openai.APITimeoutError\nGood Job", e
-# )
-# print(type(e))
-# pass
-# except Exception as e:
-# pytest.fail(
-# f"Did not raise error `openai.APITimeoutError`. Instead raised error type: {type(e)}, Error: {e}"
-# )
-
-
-# # test_timeouts_router()
-
-
-# def test_stream_timeouts_router():
-# """
-# Test the stream timeouts router. See if it selected the correct client with stream timeout
-# """
-# import openai
-
-# litellm.set_verbose = True
-# try:
-# print("testing init 4 clients with diff timeouts")
-# model_list = [
-# {
-# "model_name": "gpt-3.5-turbo",
-# "litellm_params": {
-# "model": "azure/gpt-4.1-mini",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "api_version": os.getenv("AZURE_API_VERSION"),
-# "api_base": os.getenv("AZURE_API_BASE"),
-# "timeout": 200, # regular calls will not timeout, stream calls will
-# "stream_timeout": 10,
-# },
-# },
-# ]
-# router = Router(model_list=model_list)
-
-# print("PASSED !")
-# data = {
-# "model": "gpt-3.5-turbo",
-# "messages": [{"role": "user", "content": "hello, write a 20 pg essay"}],
-# "stream": True,
-# }
-# selected_client = router._get_client(
-# deployment=router.model_list[0],
-# kwargs=data,
-# client_type=None,
-# )
-# print("Select client timeout", selected_client.timeout)
-# assert selected_client.timeout == 10
-
-# # make actual call
-# response = router.completion(**data)
-
-# for chunk in response:
-# print(f"chunk: {chunk}")
-# except openai.APITimeoutError as e:
-# print(
-# "Passed: Raised correct exception. Got openai.APITimeoutError\nGood Job", e
-# )
-# print(type(e))
-# pass
-# except Exception as e:
-# pytest.fail(
-# f"Did not raise error `openai.APITimeoutError`. Instead raised error type: {type(e)}, Error: {e}"
-# )
-
-
-# # test_stream_timeouts_router()
-
-
-# def test_xinference_embedding():
-# # [Test Init Xinference] this tests if we init xinference on the router correctly
-# # [Test Exception Mapping] tests that xinference is an openai comptiable provider
-# print("Testing init xinference")
-# print(
-# "this tests if we create an OpenAI client for Xinference, with the correct API BASE"
-# )
-
-# model_list = [
-# {
-# "model_name": "xinference",
-# "litellm_params": {
-# "model": "xinference/bge-base-en",
-# "api_base": "os.environ/XINFERENCE_API_BASE",
-# },
-# }
-# ]
-
-# router = Router(model_list=model_list)
-
-# print(router.model_list)
-# print(router.model_list[0])
-
-# assert (
-# router.model_list[0]["litellm_params"]["api_base"] == "http://0.0.0.0:9997"
-# ) # set in env
-
-# openai_client = router._get_client(
-# deployment=router.model_list[0],
-# kwargs={"input": ["hello"], "model": "xinference"},
-# )
-
-# assert openai_client._base_url == "http://0.0.0.0:9997"
-# assert "xinference" in litellm.openai_compatible_providers
-# print("passed")
-
-
-# # test_xinference_embedding()
-
-
-# def test_router_init_gpt_4_vision_enhancements():
-# try:
-# # tests base_url set when any base_url with /openai/deployments passed to router
-# print("Testing Azure GPT_Vision enhancements")
-
-# model_list = [
-# {
-# "model_name": "gpt-4-vision-enhancements",
-# "litellm_params": {
-# "model": "azure/gpt-4-vision",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "base_url": "https://gpt-4-vision-resource.openai.azure.com/openai/deployments/gpt-4-vision/extensions/",
-# "dataSources": [
-# {
-# "type": "AzureComputerVision",
-# "parameters": {
-# "endpoint": "os.environ/AZURE_VISION_ENHANCE_ENDPOINT",
-# "key": "os.environ/AZURE_VISION_ENHANCE_KEY",
-# },
-# }
-# ],
-# },
-# }
-# ]
-
-# router = Router(model_list=model_list)
-
-# print(router.model_list)
-# print(router.model_list[0])
-
-# assert (
-# router.model_list[0]["litellm_params"]["base_url"]
-# == "https://gpt-4-vision-resource.openai.azure.com/openai/deployments/gpt-4-vision/extensions/"
-# ) # set in env
-
-# assert (
-# router.model_list[0]["litellm_params"]["dataSources"][0]["parameters"][
-# "endpoint"
-# ]
-# == os.environ["AZURE_VISION_ENHANCE_ENDPOINT"]
-# )
-
-# assert (
-# router.model_list[0]["litellm_params"]["dataSources"][0]["parameters"][
-# "key"
-# ]
-# == os.environ["AZURE_VISION_ENHANCE_KEY"]
-# )
-
-# azure_client = router._get_client(
-# deployment=router.model_list[0],
-# kwargs={"stream": True, "model": "gpt-4-vision-enhancements"},
-# client_type="async",
-# )
-
-# assert (
-# azure_client._base_url
-# == "https://gpt-4-vision-resource.openai.azure.com/openai/deployments/gpt-4-vision/extensions/"
-# )
-# print("passed")
-# except Exception as e:
-# pytest.fail(f"Error occurred: {e}")
-
-
-# @pytest.mark.parametrize("sync_mode", [True, False])
-# @pytest.mark.asyncio
-# async def test_openai_with_organization(sync_mode):
-# try:
-# print("Testing OpenAI with organization")
-# model_list = [
-# {
-# "model_name": "openai-bad-org",
-# "litellm_params": {
-# "model": "gpt-3.5-turbo",
-# "organization": "org-ikDc4ex8NB",
-# },
-# },
-# {
-# "model_name": "openai-good-org",
-# "litellm_params": {"model": "gpt-3.5-turbo"},
-# },
-# ]
-
-# router = Router(model_list=model_list)
-
-# print(router.model_list)
-# print(router.model_list[0])
-
-# if sync_mode:
-# openai_client = router._get_client(
-# deployment=router.model_list[0],
-# kwargs={"input": ["hello"], "model": "openai-bad-org"},
-# )
-# print(vars(openai_client))
-
-# assert openai_client.organization == "org-ikDc4ex8NB"
-
-# # bad org raises error
-
-# try:
-# response = router.completion(
-# model="openai-bad-org",
-# messages=[{"role": "user", "content": "this is a test"}],
-# )
-# pytest.fail(
-# "Request should have failed - This organization does not exist"
-# )
-# except Exception as e:
-# print("Got exception: " + str(e))
-# assert "header should match organization for API key" in str(
-# e
-# ) or "No such organization" in str(e)
-
-# # good org works
-# response = router.completion(
-# model="openai-good-org",
-# messages=[{"role": "user", "content": "this is a test"}],
-# max_tokens=5,
-# )
-# else:
-# openai_client = router._get_client(
-# deployment=router.model_list[0],
-# kwargs={"input": ["hello"], "model": "openai-bad-org"},
-# client_type="async",
-# )
-# print(vars(openai_client))
-
-# assert openai_client.organization == "org-ikDc4ex8NB"
-
-# # bad org raises error
-
-# try:
-# response = await router.acompletion(
-# model="openai-bad-org",
-# messages=[{"role": "user", "content": "this is a test"}],
-# )
-# pytest.fail(
-# "Request should have failed - This organization does not exist"
-# )
-# except Exception as e:
-# print("Got exception: " + str(e))
-# assert "header should match organization for API key" in str(
-# e
-# ) or "No such organization" in str(e)
-
-# # good org works
-# response = await router.acompletion(
-# model="openai-good-org",
-# messages=[{"role": "user", "content": "this is a test"}],
-# max_tokens=5,
-# )
-
-# except Exception as e:
-# pytest.fail(f"Error occurred: {e}")
-
-
-# def test_init_clients_azure_command_r_plus():
-# # This tests that the router uses the OpenAI client for Azure/Command-R+
-# # For azure/command-r-plus we need to use openai.OpenAI because of how the Azure provider requires requests being sent
-# litellm.set_verbose = True
-# import logging
-
-# from litellm._logging import verbose_router_logger
-
-# verbose_router_logger.setLevel(logging.DEBUG)
-# try:
-# print("testing init 4 clients with diff timeouts")
-# model_list = [
-# {
-# "model_name": "gpt-3.5-turbo",
-# "litellm_params": {
-# "model": "azure/command-r-plus",
-# "api_key": os.getenv("AZURE_COHERE_API_KEY"),
-# "api_base": os.getenv("AZURE_COHERE_API_BASE"),
-# "timeout": 0.01,
-# "stream_timeout": 0.000_001,
-# "max_retries": 7,
-# },
-# },
-# ]
-# router = Router(model_list=model_list, set_verbose=True)
-# for elem in router.model_list:
-# model_id = elem["model_info"]["id"]
-# async_client = router.cache.get_cache(f"{model_id}_async_client")
-# stream_async_client = router.cache.get_cache(
-# f"{model_id}_stream_async_client"
-# )
-# # Assert the Async Clients used are OpenAI clients and not Azure
-# # For using Azure/Command-R-Plus and Azure/Mistral the clients NEED to be OpenAI clients used
-# # this is weirdness introduced on Azure's side
-
-# assert "openai.AsyncOpenAI" in str(async_client)
-# assert "openai.AsyncOpenAI" in str(stream_async_client)
-# print("PASSED !")
-
-# except Exception as e:
-# traceback.print_exc()
-# pytest.fail(f"Error occurred: {e}")
-
-
-# @pytest.mark.asyncio
-# async def test_aaaaatext_completion_with_organization():
-# try:
-# print("Testing Text OpenAI with organization")
-# model_list = [
-# {
-# "model_name": "openai-bad-org",
-# "litellm_params": {
-# "model": "text-completion-openai/gpt-3.5-turbo-instruct",
-# "api_key": os.getenv("OPENAI_API_KEY", None),
-# "organization": "org-ikDc4ex8NB",
-# },
-# },
-# {
-# "model_name": "openai-good-org",
-# "litellm_params": {
-# "model": "text-completion-openai/gpt-3.5-turbo-instruct",
-# "api_key": os.getenv("OPENAI_API_KEY", None),
-# "organization": os.getenv("OPENAI_ORGANIZATION", None),
-# },
-# },
-# ]
-
-# router = Router(model_list=model_list)
-
-# print(router.model_list)
-# print(router.model_list[0])
-
-# openai_client = router._get_client(
-# deployment=router.model_list[0],
-# kwargs={"input": ["hello"], "model": "openai-bad-org"},
-# )
-# print(vars(openai_client))
-
-# assert openai_client.organization == "org-ikDc4ex8NB"
-
-# # bad org raises error
-
-# try:
-# response = await router.atext_completion(
-# model="openai-bad-org",
-# prompt="this is a test",
-# )
-# pytest.fail("Request should have failed - This organization does not exist")
-# except Exception as e:
-# print("Got exception: " + str(e))
-# assert "header should match organization for API key" in str(
-# e
-# ) or "No such organization" in str(e)
-
-# # good org works
-# response = await router.atext_completion(
-# model="openai-good-org",
-# prompt="this is a test",
-# max_tokens=5,
-# )
-# print("working response: ", response)
-
-# except Exception as e:
-# pytest.fail(f"Error occurred: {e}")
-
-
-# def test_init_clients_async_mode():
-# litellm.set_verbose = True
-# import logging
-
-# from litellm._logging import verbose_router_logger
-# from litellm.types.router import RouterGeneralSettings
-
-# verbose_router_logger.setLevel(logging.DEBUG)
-# try:
-# print("testing init 4 clients with diff timeouts")
-# model_list = [
-# {
-# "model_name": "gpt-3.5-turbo",
-# "litellm_params": {
-# "model": "azure/gpt-4.1-mini",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "api_version": os.getenv("AZURE_API_VERSION"),
-# "api_base": os.getenv("AZURE_API_BASE"),
-# "timeout": 0.01,
-# "stream_timeout": 0.000_001,
-# "max_retries": 7,
-# },
-# },
-# ]
-# router = Router(
-# model_list=model_list,
-# set_verbose=True,
-# router_general_settings=RouterGeneralSettings(async_only_mode=True),
-# )
-# for elem in router.model_list:
-# model_id = elem["model_info"]["id"]
-
-# # sync clients not initialized in async_only_mode=True
-# assert router.cache.get_cache(f"{model_id}_client") is None
-# assert router.cache.get_cache(f"{model_id}_stream_client") is None
-
-# # only async clients initialized in async_only_mode=True
-# assert router.cache.get_cache(f"{model_id}_async_client") is not None
-# assert router.cache.get_cache(f"{model_id}_stream_async_client") is not None
-# except Exception as e:
-# pytest.fail(f"Error occurred: {e}")
-
-
-# @pytest.mark.parametrize(
-# "environment,expected_models",
-# [
-# ("development", ["gpt-3.5-turbo"]),
-# ("production", ["gpt-4", "gpt-3.5-turbo", "gpt-4o"]),
-# ],
-# )
-# def test_init_router_with_supported_environments(environment, expected_models):
-# """
-# Tests that the correct models are setup on router when LITELLM_ENVIRONMENT is set
-# """
-# os.environ["LITELLM_ENVIRONMENT"] = environment
-# model_list = [
-# {
-# "model_name": "gpt-3.5-turbo",
-# "litellm_params": {
-# "model": "azure/gpt-4.1-mini",
-# "api_key": os.getenv("AZURE_API_KEY"),
-# "api_version": os.getenv("AZURE_API_VERSION"),
-# "api_base": os.getenv("AZURE_API_BASE"),
-# "timeout": 0.01,
-# "stream_timeout": 0.000_001,
-# "max_retries": 7,
-# },
-# "model_info": {"supported_environments": ["development", "production"]},
-# },
-# {
-# "model_name": "gpt-4",
-# "litellm_params": {
-# "model": "openai/gpt-4",
-# "api_key": os.getenv("OPENAI_API_KEY"),
-# "timeout": 0.01,
-# "stream_timeout": 0.000_001,
-# "max_retries": 7,
-# },
-# "model_info": {"supported_environments": ["production"]},
-# },
-# {
-# "model_name": "gpt-4o",
-# "litellm_params": {
-# "model": "openai/gpt-4o",
-# "api_key": os.getenv("OPENAI_API_KEY"),
-# "timeout": 0.01,
-# "stream_timeout": 0.000_001,
-# "max_retries": 7,
-# },
-# "model_info": {"supported_environments": ["production"]},
-# },
-# ]
-# router = Router(model_list=model_list, set_verbose=True)
-# _model_list = router.get_model_names()
-
-# print("model_list: ", _model_list)
-# print("expected_models: ", expected_models)
-
-# assert set(_model_list) == set(expected_models)
-
-# os.environ.pop("LITELLM_ENVIRONMENT")
diff --git a/tests/local_testing/test_router_timeout.py b/tests/local_testing/test_router_timeout.py
index 1d09f1f1e0f..943a5413d60 100644
--- a/tests/local_testing/test_router_timeout.py
+++ b/tests/local_testing/test_router_timeout.py
@@ -31,8 +31,8 @@ def test_router_timeouts():
"model_name": "openai-gpt-4",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
- "api_key": "os.environ/AZURE_API_KEY",
- "api_base": "os.environ/AZURE_API_BASE",
+ "api_key": "os.environ/AZURE_AI_API_KEY",
+ "api_base": "os.environ/AZURE_AI_API_BASE",
"api_version": "os.environ/AZURE_API_VERSION",
},
"tpm": 80000,
diff --git a/tests/local_testing/test_router_utils.py b/tests/local_testing/test_router_utils.py
index 4f7f53cef02..f2fd2fdf559 100644
--- a/tests/local_testing/test_router_utils.py
+++ b/tests/local_testing/test_router_utils.py
@@ -35,7 +35,7 @@ def test_returned_settings():
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -99,7 +99,7 @@ def test_update_kwargs_before_fallbacks_unit_test():
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
}
],
@@ -136,7 +136,7 @@ async def test_update_kwargs_before_fallbacks(call_type):
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
}
],
@@ -266,6 +266,7 @@ async def test_call_router_callbacks_on_success():
)
assert increment["increment_value"] == 1
+
@pytest.mark.serial
@pytest.mark.asyncio
async def test_call_router_callbacks_on_failure():
@@ -486,7 +487,9 @@ def test_router_get_deployment_credentials_with_provider():
)
# Test getting credentials by model_id
- credentials = router.get_deployment_credentials_with_provider(model_id="openai-deployment-1")
+ credentials = router.get_deployment_credentials_with_provider(
+ model_id="openai-deployment-1"
+ )
assert credentials is not None
assert credentials["api_key"] == "sk-test-123"
assert credentials["custom_llm_provider"] == "openai"
@@ -499,14 +502,16 @@ def test_router_get_deployment_credentials_with_provider():
assert credentials2["custom_llm_provider"] == "anthropic"
# Test with non-existent model
- credentials3 = router.get_deployment_credentials_with_provider(model_id="non-existent")
+ credentials3 = router.get_deployment_credentials_with_provider(
+ model_id="non-existent"
+ )
assert credentials3 is None
def test_router_get_deployment_credentials_with_provider_wildcard():
"""
Test that get_deployment_credentials_with_provider handles wildcard patterns.
-
+
When a model like openai/gpt-4o is requested and the config has openai/*,
the method should resolve the wildcard pattern and return credentials.
"""
@@ -533,20 +538,26 @@ def test_router_get_deployment_credentials_with_provider_wildcard():
)
# Test wildcard pattern matching for OpenAI
- credentials = router.get_deployment_credentials_with_provider(model_id="openai/gpt-4o")
+ credentials = router.get_deployment_credentials_with_provider(
+ model_id="openai/gpt-4o"
+ )
assert credentials is not None
assert credentials["api_key"] == "sk-wildcard-123"
assert credentials["custom_llm_provider"] == "openai"
assert credentials["api_base"] == "https://api.openai.com/v1"
# Test wildcard pattern matching for Anthropic
- credentials2 = router.get_deployment_credentials_with_provider(model_id="anthropic/claude-3-opus")
+ credentials2 = router.get_deployment_credentials_with_provider(
+ model_id="anthropic/claude-3-opus"
+ )
assert credentials2 is not None
assert credentials2["api_key"] == "sk-ant-wildcard-456"
assert credentials2["custom_llm_provider"] == "anthropic"
# Test with non-matching model
- credentials3 = router.get_deployment_credentials_with_provider(model_id="vertex_ai/gemini-pro")
+ credentials3 = router.get_deployment_credentials_with_provider(
+ model_id="vertex_ai/gemini-pro"
+ )
assert credentials3 is None
diff --git a/tests/local_testing/test_streaming.py b/tests/local_testing/test_streaming.py
index 56f0e5fe826..97d6875f161 100644
--- a/tests/local_testing/test_streaming.py
+++ b/tests/local_testing/test_streaming.py
@@ -552,7 +552,6 @@ async def test_completion_predibase_streaming(sync_mode):
pytest.fail(f"Error occurred: {e}")
-
def test_completion_azure_function_calling_stream():
try:
litellm.set_verbose = False
@@ -1655,80 +1654,9 @@ def test_sagemaker_weird_response():
# test_sagemaker_weird_response()
-@pytest.mark.skip(reason="Move to being a mock endpoint")
-@pytest.mark.asyncio
-async def test_sagemaker_streaming_async():
- try:
- messages = [{"role": "user", "content": "Hey, how's it going?"}]
- litellm.set_verbose = True
- response = await litellm.acompletion(
- model="sagemaker/jumpstart-dft-hf-llm-mistral-7b-ins-20240329-150233",
- model_id="huggingface-llm-mistral-7b-instruct-20240329-150233",
- messages=messages,
- temperature=0.2,
- max_tokens=80,
- aws_region_name=os.getenv("AWS_REGION_NAME_2"),
- aws_access_key_id=os.getenv("AWS_ACCESS_KEY_ID_2"),
- aws_secret_access_key=os.getenv("AWS_SECRET_ACCESS_KEY_2"),
- stream=True,
- )
- # Add any assertions here to check the response
- print(response)
- complete_response = ""
- has_finish_reason = False
- # Add any assertions here to check the response
- idx = 0
- async for chunk in response:
- # print
- chunk, finished = streaming_format_tests(idx, chunk)
- has_finish_reason = finished
- complete_response += chunk
- if finished:
- break
- idx += 1
- if has_finish_reason is False:
- raise Exception("finish reason not set for last chunk")
- if complete_response.strip() == "":
- raise Exception("Empty response received")
- print(f"completion_response: {complete_response}")
- except Exception as e:
- pytest.fail(f"An exception occurred - {str(e)}")
-
-
# asyncio.run(test_sagemaker_streaming_async())
-@pytest.mark.skip(reason="costly sagemaker deployment. Move to mock implementation")
-def test_completion_sagemaker_stream():
- try:
- response = completion(
- model="sagemaker/jumpstart-dft-hf-llm-mistral-7b-ins-20240329-150233",
- model_id="huggingface-llm-mistral-7b-instruct-20240329-150233",
- messages=messages,
- temperature=0.2,
- max_tokens=80,
- aws_region_name=os.getenv("AWS_REGION_NAME_2"),
- aws_access_key_id=os.getenv("AWS_ACCESS_KEY_ID_2"),
- aws_secret_access_key=os.getenv("AWS_SECRET_ACCESS_KEY_2"),
- stream=True,
- )
- complete_response = ""
- has_finish_reason = False
- # Add any assertions here to check the response
- for idx, chunk in enumerate(response):
- chunk, finished = streaming_format_tests(idx, chunk)
- has_finish_reason = finished
- if finished:
- break
- complete_response += chunk
- if has_finish_reason is False:
- raise Exception("finish reason not set for last chunk")
- if complete_response.strip() == "":
- raise Exception("Empty response received")
- except Exception as e:
- pytest.fail(f"Error occurred: {e}")
-
-
@pytest.mark.skip(reason="Account deleted by IBM.")
@pytest.mark.asyncio
async def test_completion_watsonx_stream():
@@ -2725,8 +2653,8 @@ def test_azure_streaming_and_function_calling():
tool_choice="auto",
messages=messages,
stream=True,
- api_base=os.getenv("AZURE_API_BASE"),
- api_key=os.getenv("AZURE_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
api_version="2024-02-15-preview",
)
# Add any assertions here to check the response
@@ -2796,8 +2724,8 @@ async def test_azure_astreaming_and_function_calling():
tool_choice="auto",
messages=messages,
stream=True,
- api_base=os.getenv("AZURE_API_BASE"),
- api_key=os.getenv("AZURE_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
api_version="2024-02-15-preview",
caching=True,
)
@@ -2827,8 +2755,8 @@ async def test_azure_astreaming_and_function_calling():
tool_choice="auto",
messages=messages,
stream=True,
- api_base=os.getenv("AZURE_API_BASE"),
- api_key=os.getenv("AZURE_API_KEY"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
api_version="2024-02-15-preview",
caching=True,
)
@@ -3109,7 +3037,9 @@ def test_unit_test_custom_stream_wrapper_repeating_chunk(
print(f"expected_chunk_fail: {expected_chunk_fail}")
if (loop_amount > litellm.REPEATED_STREAMING_CHUNK_LIMIT) and expected_chunk_fail:
- with pytest.raises((litellm.InternalServerError, litellm.exceptions.MidStreamFallbackError)):
+ with pytest.raises(
+ (litellm.InternalServerError, litellm.exceptions.MidStreamFallbackError)
+ ):
for chunk in response:
continue
else:
diff --git a/tests/local_testing/test_timeout.py b/tests/local_testing/test_timeout.py
index 4128a595d76..e2ec9ba9a09 100644
--- a/tests/local_testing/test_timeout.py
+++ b/tests/local_testing/test_timeout.py
@@ -111,8 +111,8 @@ def test_hanging_request_azure():
"model_name": "azure-gpt",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
- "api_base": os.environ["AZURE_API_BASE"],
- "api_key": os.environ["AZURE_API_KEY"],
+ "api_base": os.environ["AZURE_AI_API_BASE"],
+ "api_key": os.environ["AZURE_AI_API_KEY"],
},
},
{
@@ -175,8 +175,8 @@ def test_hanging_request_openai():
"model_name": "azure-gpt",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
- "api_base": os.environ["AZURE_API_BASE"],
- "api_key": os.environ["AZURE_API_KEY"],
+ "api_base": os.environ["AZURE_AI_API_BASE"],
+ "api_key": os.environ["AZURE_AI_API_KEY"],
},
},
{
diff --git a/tests/local_testing/test_tpm_rpm_routing_v2.py b/tests/local_testing/test_tpm_rpm_routing_v2.py
index a3218bc987a..c7449ef6e2a 100644
--- a/tests/local_testing/test_tpm_rpm_routing_v2.py
+++ b/tests/local_testing/test_tpm_rpm_routing_v2.py
@@ -39,9 +39,7 @@ from create_mock_standard_logging_payload import create_standard_logging_payload
def test_tpm_rpm_updated():
test_cache = DualCache()
- lowest_tpm_logger = LowestTPMLoggingHandler(
- router_cache=test_cache
- )
+ lowest_tpm_logger = LowestTPMLoggingHandler(router_cache=test_cache)
model_group = "gpt-3.5-turbo"
deployment_id = "1234"
deployment = "azure/gpt-4.1-mini"
@@ -108,9 +106,7 @@ def test_get_available_deployments():
"model_info": {"id": "5678"},
},
]
- lowest_tpm_logger = LowestTPMLoggingHandler(
- router_cache=test_cache
- )
+ lowest_tpm_logger = LowestTPMLoggingHandler(router_cache=test_cache)
model_group = "gpt-3.5-turbo"
## DEPLOYMENT 1 ##
total_tokens = 50
@@ -669,9 +665,7 @@ def test_return_potential_deployments():
"""
test_cache = DualCache()
- lowest_tpm_logger = LowestTPMLoggingHandler(
- router_cache=test_cache
- )
+ lowest_tpm_logger = LowestTPMLoggingHandler(router_cache=test_cache)
args: Dict = {
"healthy_deployments": [
@@ -731,8 +725,8 @@ async def test_tpm_rpm_routing_model_name_checks():
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
"mock_response": "Hey, how's it going?",
},
}
diff --git a/tests/local_testing/vertex_key.json b/tests/local_testing/vertex_key.json
index 45ca6acc010..800969fb305 100644
--- a/tests/local_testing/vertex_key.json
+++ b/tests/local_testing/vertex_key.json
@@ -1,13 +1,13 @@
{
"type": "service_account",
- "project_id": "pathrise-convert-1606954137718",
+ "project_id": "litellm-ci-cd",
"private_key_id": "",
"private_key": "",
- "client_email": "ci-cd-723@pathrise-convert-1606954137718.iam.gserviceaccount.com",
- "client_id": "109577393201924326488",
+ "client_email": "test-litellm-ci-cd@litellm-ci-cd.iam.gserviceaccount.com",
+ "client_id": "116563532503305622785",
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
"token_uri": "https://oauth2.googleapis.com/token",
"auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs",
- "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/ci-cd-723%40pathrise-convert-1606954137718.iam.gserviceaccount.com",
+ "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/test-litellm-ci-cd%40litellm-ci-cd.iam.gserviceaccount.com",
"universe_domain": "googleapis.com"
}
diff --git a/tests/logging_callback_tests/test_alerting.py b/tests/logging_callback_tests/test_alerting.py
index f77cbcb4bb1..86588bbd14b 100644
--- a/tests/logging_callback_tests/test_alerting.py
+++ b/tests/logging_callback_tests/test_alerting.py
@@ -641,7 +641,7 @@ async def test_outage_alerting_called(
"model_name": model,
"litellm_params": {
"model": model,
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_base": api_base,
"vertex_location": vertex_location,
"vertex_project": vertex_project,
@@ -749,7 +749,7 @@ async def test_region_outage_alerting_called(
"model_name": model,
"litellm_params": {
"model": model,
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_base": api_base,
"vertex_location": vertex_location,
"vertex_project": vertex_project,
@@ -760,7 +760,7 @@ async def test_region_outage_alerting_called(
"model_name": model,
"litellm_params": {
"model": model,
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_base": api_base,
"vertex_location": vertex_location,
"vertex_project": "vertex_project-2",
@@ -788,40 +788,6 @@ async def test_region_outage_alerting_called(
mock_send_alert.assert_not_called()
-@pytest.mark.asyncio
-@pytest.mark.skip(reason="test only needs to run locally ")
-async def test_alerting():
- router = litellm.Router(
- model_list=[
- {
- "model_name": "gpt-3.5-turbo",
- "litellm_params": {
- "model": "gpt-3.5-turbo",
- "api_key": "bad_key",
- },
- }
- ],
- debug_level="DEBUG",
- set_verbose=True,
- alerting_config=AlertingConfig(
- alerting_threshold=10, # threshold for slow / hanging llm responses (in seconds). Defaults to 300 seconds
- webhook_url=os.getenv(
- "SLACK_WEBHOOK_URL"
- ), # webhook you want to send alerts to
- ),
- )
- try:
- await router.acompletion(
- model="gpt-3.5-turbo",
- messages=[{"role": "user", "content": "Hey, how's it going?"}],
- )
-
- except Exception:
- pass
- finally:
- await asyncio.sleep(3)
-
-
@pytest.mark.asyncio
async def test_langfuse_trace_id():
"""
@@ -868,7 +834,9 @@ async def test_langfuse_trace_id():
returned_trace_id = trace_url.split("/")[-1]
- assert returned_trace_id == litellm_logging_obj._get_trace_id(service_name="langfuse")
+ assert returned_trace_id == litellm_logging_obj._get_trace_id(
+ service_name="langfuse"
+ )
@pytest.mark.asyncio
@@ -1007,7 +975,7 @@ async def test_soft_budget_alerts():
# Verify alert message contains correct percentage
alert_message = mock_send_alert.call_args[1]["message"]
-
+
print("GOT MESSAGE\n\n", alert_message)
expected_message = (
@@ -1077,10 +1045,10 @@ key_no_max_budget_info = CallInfo(
async def test_soft_budget_alerts_webhook(entity_info):
"""
Tests that soft budget alerts are triggered for different entity types.
-
+
Tests:
- Key with max budget
- - Team
+ - Team
- User
- Key without max budget
"""
@@ -1097,7 +1065,7 @@ async def test_soft_budget_alerts_webhook(entity_info):
# Verify the webhook event
call_args = mock_send_alert.call_args[1]
logged_webhook_event: WebhookEvent = call_args["user_info"]
-
+
# Validate the webhook event has all expected fields
assert logged_webhook_event.spend == entity_info.spend
assert logged_webhook_event.soft_budget == entity_info.soft_budget
@@ -1106,10 +1074,3 @@ async def test_soft_budget_alerts_webhook(entity_info):
assert logged_webhook_event.user_email == entity_info.user_email
assert logged_webhook_event.key_alias == entity_info.key_alias
assert logged_webhook_event.event_group == entity_info.event_group
-
-
-
-
-
-
-
\ No newline at end of file
diff --git a/tests/logging_callback_tests/test_amazing_s3_logs.py b/tests/logging_callback_tests/test_amazing_s3_logs.py
index f074026926a..c1e6884278c 100644
--- a/tests/logging_callback_tests/test_amazing_s3_logs.py
+++ b/tests/logging_callback_tests/test_amazing_s3_logs.py
@@ -74,15 +74,13 @@ async def test_basic_s3_logging(sync_mode, streaming):
s3.delete_object(Bucket="load-testing-oct", Key=key)
-
@pytest.mark.asyncio
-@pytest.mark.parametrize(
- "streaming", [(True)]
-)
+@pytest.mark.parametrize("streaming", [True])
@pytest.mark.flaky(retries=3, delay=1)
async def test_basic_s3_v2_logging(streaming):
from blockbuster import BlockBuster
from litellm.integrations.s3_v2 import S3Logger
+
s3_v2_logger = S3Logger(s3_flush_interval=1)
litellm.callbacks = [s3_v2_logger]
blockbuster = BlockBuster()
@@ -120,7 +118,7 @@ async def test_basic_s3_v2_logging(streaming):
print(f"all_s3_keys: {all_s3_keys}")
- #assert that atlest one key has response.id in it
+ # assert that atlest one key has response.id in it
assert any(response_id in key for key in all_s3_keys)
s3 = boto3.client("s3")
# delete all objects
@@ -134,22 +132,22 @@ async def test_basic_s3_v2_logging_failure():
"""Test that S3 v2 logger makes httpx PUT request when logging failures"""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.integrations.s3_v2 import S3Logger
-
+
# Create S3 logger with short flush interval
s3_v2_logger = S3Logger(s3_flush_interval=1)
-
+
# Mock the httpx client to capture the PUT request
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.raise_for_status = MagicMock()
-
+
s3_v2_logger.async_httpx_client = AsyncMock()
s3_v2_logger.async_httpx_client.put.return_value = mock_response
-
+
# Track the upload method calls
original_upload = s3_v2_logger.async_upload_data_to_s3
upload_called = False
-
+
async def mock_upload(batch_logging_element):
nonlocal upload_called
upload_called = True
@@ -157,12 +155,12 @@ async def test_basic_s3_v2_logging_failure():
url = f"https://test-bucket.s3.us-west-2.amazonaws.com/{batch_logging_element.s3_object_key}"
headers = {"Content-Type": "application/json"}
data = '{"model": "gpt-4o-mini"}'
-
+
# Make the actual httpx call we want to test
await s3_v2_logger.async_httpx_client.put(url=url, headers=headers, data=data)
-
+
s3_v2_logger.async_upload_data_to_s3 = mock_upload
-
+
# Configure S3 callback params
litellm.callbacks = [s3_v2_logger]
litellm.s3_callback_params = {
@@ -172,7 +170,7 @@ async def test_basic_s3_v2_logging_failure():
"s3_region_name": "us-west-2",
}
litellm.set_verbose = True
-
+
# Trigger a failure by using invalid API key
try:
response = await litellm.acompletion(
@@ -182,33 +180,33 @@ async def test_basic_s3_v2_logging_failure():
)
except Exception as e:
print(f"Expected error: {e}")
-
+
# Wait for logger to process the failure
await asyncio.sleep(5)
-
+
# Verify that our mock upload was called
assert upload_called, "S3 upload method was not called"
print("✓ S3 upload method was called")
-
+
# Verify that httpx PUT was called
s3_v2_logger.async_httpx_client.put.assert_called()
-
+
# Get the call arguments to verify the S3 URL
call_args = s3_v2_logger.async_httpx_client.put.call_args
assert call_args is not None
- url = call_args[1]['url'] if 'url' in call_args[1] else call_args[0][0]
-
+ url = call_args[1]["url"] if "url" in call_args[1] else call_args[0][0]
+
# Verify the URL contains expected S3 endpoint
assert "test-bucket.s3.us-west-2.amazonaws.com" in url
print(f"✓ S3 PUT request made to: {url}")
-
+
# Verify headers include expected content type
- headers = call_args[1]['headers']
- assert headers['Content-Type'] == 'application/json'
+ headers = call_args[1]["headers"]
+ assert headers["Content-Type"] == "application/json"
print("✓ S3 request headers are correct")
-
+
# Verify JSON data was included
- data = call_args[1]['data']
+ data = call_args[1]["data"]
assert data is not None
assert '"model": "gpt-4o-mini"' in data
print("✓ S3 request data contains expected log payload")
@@ -411,83 +409,19 @@ async def make_async_calls():
return total_time
-@pytest.mark.skip(reason="flaky test on ci/cd")
-def test_s3_logging_r2():
- # all s3 requests need to be in one test function
- # since we are modifying stdout, and pytests runs tests in parallel
- # on circle ci - we only test litellm.acompletion()
- try:
- # redirect stdout to log_file
- # litellm.cache = litellm.Cache(
- # type="s3", s3_bucket_name="litellm-r2-bucket", s3_region_name="us-west-2"
- # )
- litellm.set_verbose = True
- from litellm._logging import verbose_logger
- import logging
-
- verbose_logger.setLevel(level=logging.DEBUG)
-
- litellm.success_callback = ["s3"]
- litellm.s3_callback_params = {
- "s3_bucket_name": "litellm-r2-bucket",
- "s3_aws_secret_access_key": "os.environ/R2_S3_ACCESS_KEY",
- "s3_aws_access_key_id": "os.environ/R2_S3_ACCESS_ID",
- "s3_endpoint_url": "os.environ/R2_S3_URL",
- "s3_region_name": "os.environ/R2_S3_REGION_NAME",
- }
- print("Testing async s3 logging")
-
- expected_keys = []
-
- import time
-
- curr_time = str(time.time())
-
- async def _test():
- return await litellm.acompletion(
- model="gpt-3.5-turbo",
- messages=[{"role": "user", "content": f"This is a test {curr_time}"}],
- max_tokens=10,
- temperature=0.7,
- user="ishaan-2",
- )
-
- response = asyncio.run(_test())
- print(f"response: {response}")
- expected_keys.append(response.id)
-
- import boto3
-
- s3 = boto3.client(
- "s3",
- endpoint_url=os.getenv("R2_S3_URL"),
- region_name=os.getenv("R2_S3_REGION_NAME"),
- aws_access_key_id=os.getenv("R2_S3_ACCESS_ID"),
- aws_secret_access_key=os.getenv("R2_S3_ACCESS_KEY"),
- )
-
- bucket_name = "litellm-r2-bucket"
- # List objects in the bucket
- response = s3.list_objects(Bucket=bucket_name)
-
- except Exception as e:
- pytest.fail(f"An exception occurred - {e}")
- finally:
- # post, close log file and verify
- # Reset stdout to the original value
- print("Passed! Testing async s3 logging")
-
from litellm.integrations.s3_v2 import S3Logger
+
class TestS3Logger(S3Logger):
def __init__(self, *args, **kwargs):
self.recorded_requests = {}
self.logged_standard_logging_payload: Optional[StandardLoggingPayload] = None
super().__init__(*args, **kwargs)
-
+
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
self.recorded_requests[response_obj["id"]] = start_time
print("recorded request", self.recorded_requests)
self.logged_standard_logging_payload = kwargs["standard_logging_object"]
- return await super().async_log_success_event(kwargs, response_obj, start_time, end_time)
-
+ return await super().async_log_success_event(
+ kwargs, response_obj, start_time, end_time
+ )
diff --git a/tests/logging_callback_tests/test_custom_callback_router.py b/tests/logging_callback_tests/test_custom_callback_router.py
index f6c7f2fa023..63d8b14f488 100644
--- a/tests/logging_callback_tests/test_custom_callback_router.py
+++ b/tests/logging_callback_tests/test_custom_callback_router.py
@@ -267,7 +267,10 @@ class CompletionCustomHandler(
try:
print("CompletionCustomHandler.async_log_success_event, kwargs: ", kwargs)
self.states.append("async_success")
- print("############### CompletionCustomHandler async success, kwargs: ", kwargs)
+ print(
+ "############### CompletionCustomHandler async success, kwargs: ",
+ kwargs,
+ )
## START TIME
assert isinstance(start_time, datetime)
## END TIME
@@ -396,9 +399,9 @@ async def test_async_chat_azure():
"model_name": "gpt-4.1-nano", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"model_info": {"base_model": "azure/gpt-4.1-mini"},
"tpm": 240000,
@@ -443,7 +446,7 @@ async def test_async_chat_azure():
"model": "azure/gpt-4o-new-test",
"api_key": "my-bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -483,9 +486,9 @@ async def test_async_embedding_azure():
"model_name": "azure-embedding-model", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/text-embedding-ada-002",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -506,7 +509,7 @@ async def test_async_embedding_azure():
"model": "azure/text-embedding-ada-002",
"api_key": "my-bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -549,7 +552,7 @@ async def test_async_chat_azure_with_fallbacks():
"model": "azure/gpt-4.1-mini",
"api_key": "my-bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -608,9 +611,9 @@ async def test_async_completion_azure_caching():
"model_name": "gpt-4.1-nano", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
@@ -664,23 +667,23 @@ async def test_async_completion_azure_caching_streaming():
)
litellm.callbacks = [customHandler_caching]
unique_time = uuid.uuid4()
-
+
# Use Router instead of direct litellm.acompletion to get router-specific metadata
model_list = [
{
"model_name": "gpt-4.1-nano",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
},
]
router = Router(model_list=model_list)
-
+
response1 = await router.acompletion(
model="gpt-4.1-nano",
messages=[
@@ -725,12 +728,16 @@ async def test_async_embedding_azure_caching():
port=os.environ["REDIS_PORT"],
password=os.environ["REDIS_PASSWORD"],
)
- router = Router(model_list=[{
- "model_name": "text-embedding-ada-002",
- "litellm_params": {
- "model": "openai/text-embedding-ada-002",
- },
- }])
+ router = Router(
+ model_list=[
+ {
+ "model_name": "text-embedding-ada-002",
+ "litellm_params": {
+ "model": "openai/text-embedding-ada-002",
+ },
+ }
+ ]
+ )
litellm.callbacks = [customHandler_caching]
unique_time = time.time()
response1 = await router.aembedding(
@@ -818,4 +825,3 @@ async def test_rate_limit_error_callback():
assert "original_model_group" in mock_client.call_args.kwargs
assert mock_client.call_args.kwargs["original_model_group"] == "my-test-gpt"
-
diff --git a/tests/ocr_tests/vertex_key.json b/tests/ocr_tests/vertex_key.json
index 45ca6acc010..800969fb305 100644
--- a/tests/ocr_tests/vertex_key.json
+++ b/tests/ocr_tests/vertex_key.json
@@ -1,13 +1,13 @@
{
"type": "service_account",
- "project_id": "pathrise-convert-1606954137718",
+ "project_id": "litellm-ci-cd",
"private_key_id": "",
"private_key": "",
- "client_email": "ci-cd-723@pathrise-convert-1606954137718.iam.gserviceaccount.com",
- "client_id": "109577393201924326488",
+ "client_email": "test-litellm-ci-cd@litellm-ci-cd.iam.gserviceaccount.com",
+ "client_id": "116563532503305622785",
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
"token_uri": "https://oauth2.googleapis.com/token",
"auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs",
- "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/ci-cd-723%40pathrise-convert-1606954137718.iam.gserviceaccount.com",
+ "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/test-litellm-ci-cd%40litellm-ci-cd.iam.gserviceaccount.com",
"universe_domain": "googleapis.com"
}
diff --git a/tests/old_proxy_tests/tests/load_test_q.py b/tests/old_proxy_tests/tests/load_test_q.py
index a8f2c0a322d..89137c306a7 100644
--- a/tests/old_proxy_tests/tests/load_test_q.py
+++ b/tests/old_proxy_tests/tests/load_test_q.py
@@ -26,7 +26,7 @@ config = {
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
- "api_key": os.environ["AZURE_API_KEY"],
+ "api_key": os.environ["AZURE_AI_API_KEY"],
"api_base": "https://openai-gpt-4-test-v-1.openai.azure.com/",
"api_version": "2023-07-01-preview",
},
@@ -34,7 +34,7 @@ config = {
]
}
print("STARTING LOAD TEST Q")
-print(os.environ["AZURE_API_KEY"])
+print(os.environ["AZURE_AI_API_KEY"])
response = requests.post(
url=f"{base_url}/key/generate",
diff --git a/tests/pass_through_tests/vertex_key.json b/tests/pass_through_tests/vertex_key.json
index 45ca6acc010..800969fb305 100644
--- a/tests/pass_through_tests/vertex_key.json
+++ b/tests/pass_through_tests/vertex_key.json
@@ -1,13 +1,13 @@
{
"type": "service_account",
- "project_id": "pathrise-convert-1606954137718",
+ "project_id": "litellm-ci-cd",
"private_key_id": "",
"private_key": "",
- "client_email": "ci-cd-723@pathrise-convert-1606954137718.iam.gserviceaccount.com",
- "client_id": "109577393201924326488",
+ "client_email": "test-litellm-ci-cd@litellm-ci-cd.iam.gserviceaccount.com",
+ "client_id": "116563532503305622785",
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
"token_uri": "https://oauth2.googleapis.com/token",
"auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs",
- "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/ci-cd-723%40pathrise-convert-1606954137718.iam.gserviceaccount.com",
+ "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/test-litellm-ci-cd%40litellm-ci-cd.iam.gserviceaccount.com",
"universe_domain": "googleapis.com"
}
diff --git a/tests/pass_through_unit_tests/messages_api_structured_output/test_azure_anthropic_structured_output.py b/tests/pass_through_unit_tests/messages_api_structured_output/test_azure_anthropic_structured_output.py
index da46016b358..b2470bf6b67 100644
--- a/tests/pass_through_unit_tests/messages_api_structured_output/test_azure_anthropic_structured_output.py
+++ b/tests/pass_through_unit_tests/messages_api_structured_output/test_azure_anthropic_structured_output.py
@@ -30,7 +30,7 @@ class TestAzureAnthropicStructuredOutput(BaseAnthropicMessagesStructuredOutputTe
return "azure_ai/claude-opus-4-5"
def get_api_base(self) -> Optional[str]:
- return "https://krish-mh44t553-eastus2.services.ai.azure.com/"
+ return "https://krris-mnb3t0vd-swedencentral.services.ai.azure.com"
def get_api_key(self) -> Optional[str]:
- return os.environ.get("AZURE_ANTHROPIC_API_KEY")
\ No newline at end of file
+ return os.environ.get("AZURE_ANTHROPIC_API_KEY")
diff --git a/tests/pass_through_unit_tests/test_unit_test_passthrough_router.py b/tests/pass_through_unit_tests/test_unit_test_passthrough_router.py
index 8e016b68d05..3f62f557915 100644
--- a/tests/pass_through_unit_tests/test_unit_test_passthrough_router.py
+++ b/tests/pass_through_unit_tests/test_unit_test_passthrough_router.py
@@ -102,7 +102,7 @@ class TestPassthroughEndpointRouter(unittest.TestCase):
mock_get_secret.return_value = "env_azure_key"
result = self.router.get_credentials("azure", None)
self.assertEqual(result, "env_azure_key")
- mock_get_secret.assert_called_once_with("AZURE_API_KEY")
+ mock_get_secret.assert_called_once_with("AZURE_AI_API_KEY")
def test_default_env_variable_method(self):
"""
diff --git a/tests/proxy_unit_tests/example_config_yaml/azure_config.yaml b/tests/proxy_unit_tests/example_config_yaml/azure_config.yaml
index 0a015aefde8..05ba0c9bf54 100644
--- a/tests/proxy_unit_tests/example_config_yaml/azure_config.yaml
+++ b/tests/proxy_unit_tests/example_config_yaml/azure_config.yaml
@@ -4,12 +4,12 @@ model_list:
model: azure/gpt-4.1-mini
api_base: https://openai-gpt-4-test-v-1.openai.azure.com/
api_version: "2023-05-15"
- api_key: os.environ/AZURE_API_KEY
+ api_key: os.environ/AZURE_AI_API_KEY
tpm: 20_000
- model_name: gpt-4-team2
litellm_params:
model: azure/gpt-4
- api_key: os.environ/AZURE_API_KEY
+ api_key: os.environ/AZURE_AI_API_KEY
api_base: https://openai-gpt-4-test-v-2.openai.azure.com/
tpm: 100_000
diff --git a/tests/proxy_unit_tests/test_configs/test_bad_config.yaml b/tests/proxy_unit_tests/test_configs/test_bad_config.yaml
index 4a70886a93b..4bc4c7cc541 100644
--- a/tests/proxy_unit_tests/test_configs/test_bad_config.yaml
+++ b/tests/proxy_unit_tests/test_configs/test_bad_config.yaml
@@ -6,16 +6,16 @@ model_list:
- model_name: working-azure-gpt-3.5-turbo
litellm_params:
model: azure/gpt-4.1-mini
- api_base: os.environ/AZURE_API_BASE
- api_key: os.environ/AZURE_API_KEY
+ api_base: os.environ/AZURE_AI_API_BASE
+ api_key: os.environ/AZURE_AI_API_KEY
- model_name: azure-gpt-3.5-turbo
litellm_params:
model: azure/gpt-4.1-mini
- api_base: os.environ/AZURE_API_BASE
+ api_base: os.environ/AZURE_AI_API_BASE
api_key: bad-key
- model_name: azure-embedding
litellm_params:
model: azure/text-embedding-ada-002
- api_base: os.environ/AZURE_API_BASE
+ api_base: os.environ/AZURE_AI_API_BASE
api_key: bad-key
\ No newline at end of file
diff --git a/tests/proxy_unit_tests/test_configs/test_cloudflare_azure_with_cache_config.yaml b/tests/proxy_unit_tests/test_configs/test_cloudflare_azure_with_cache_config.yaml
index 99028356183..24240008fe2 100644
--- a/tests/proxy_unit_tests/test_configs/test_cloudflare_azure_with_cache_config.yaml
+++ b/tests/proxy_unit_tests/test_configs/test_cloudflare_azure_with_cache_config.yaml
@@ -3,7 +3,7 @@ model_list:
litellm_params:
model: azure/gpt-4.1-mini
api_base: https://gateway.ai.cloudflare.com/v1/0399b10e77ac6668c80404a5ff49eb37/litellm-test/azure-openai/openai-gpt-4-test-v-1
- api_key: os.environ/AZURE_API_KEY
+ api_key: os.environ/AZURE_AI_API_KEY
api_version: 2023-07-01-preview
litellm_settings:
diff --git a/tests/proxy_unit_tests/test_configs/test_config_no_auth.yaml b/tests/proxy_unit_tests/test_configs/test_config_no_auth.yaml
index cdc447a5ee6..f4896217049 100644
--- a/tests/proxy_unit_tests/test_configs/test_config_no_auth.yaml
+++ b/tests/proxy_unit_tests/test_configs/test_config_no_auth.yaml
@@ -11,7 +11,7 @@ model_list:
model_name: azure-model
- litellm_params:
api_base: https://gateway.ai.cloudflare.com/v1/0399b10e77ac6668c80404a5ff49eb37/litellm-test/azure-openai/openai-gpt-4-test-v-1
- api_key: os.environ/AZURE_API_KEY
+ api_key: os.environ/AZURE_AI_API_KEY
model: azure/gpt-4.1-mini
model_name: azure-cloudflare-model
- litellm_params:
@@ -49,8 +49,8 @@ model_list:
id: 79fc75bf-8e1b-47d5-8d24-9365a854af03
model_name: test_openai_models
- litellm_params:
- api_base: os.environ/AZURE_API_BASE
- api_key: os.environ/AZURE_API_KEY
+ api_base: os.environ/AZURE_AI_API_BASE
+ api_key: os.environ/AZURE_AI_API_KEY
api_version: 2023-07-01-preview
model: azure/text-embedding-ada-002
model_info:
@@ -94,16 +94,16 @@ model_list:
mode: image_generation
model_name: dall-e-3
- litellm_params:
- api_base: os.environ/AZURE_API_BASE
- api_key: os.environ/AZURE_API_KEY
+ api_base: os.environ/AZURE_AI_API_BASE
+ api_key: os.environ/AZURE_AI_API_KEY
api_version: 2023-06-01-preview
model: azure/
model_info:
mode: image_generation
model_name: dall-e-2
- litellm_params:
- api_base: os.environ/AZURE_API_BASE
- api_key: os.environ/AZURE_API_KEY
+ api_base: os.environ/AZURE_AI_API_BASE
+ api_key: os.environ/AZURE_AI_API_KEY
api_version: 2023-07-01-preview
model: azure/text-embedding-ada-002
model_info:
diff --git a/tests/proxy_unit_tests/test_jwt.py b/tests/proxy_unit_tests/test_jwt.py
index 24cf15a3214..a5be1a3a42d 100644
--- a/tests/proxy_unit_tests/test_jwt.py
+++ b/tests/proxy_unit_tests/test_jwt.py
@@ -1331,6 +1331,9 @@ def test_jwt_handler_is_jwt_static_method():
# Test with empty string
assert JWTHandler.is_jwt("") == False
+ # Test with None (missing Authorization header)
+ assert JWTHandler.is_jwt(None) == False
+
@pytest.mark.parametrize(
"requested_model, should_work",
diff --git a/tests/proxy_unit_tests/test_proxy_pass_user_config.py b/tests/proxy_unit_tests/test_proxy_pass_user_config.py
index 4828014e335..6beb86eca72 100644
--- a/tests/proxy_unit_tests/test_proxy_pass_user_config.py
+++ b/tests/proxy_unit_tests/test_proxy_pass_user_config.py
@@ -54,8 +54,9 @@ def client_no_auth():
@pytest.mark.skipif(
- os.environ.get("AZURE_API_KEY") is None or os.environ.get("OPENAI_API_KEY") is None,
- reason="AZURE_API_KEY or OPENAI_API_KEY not set - skipping integration test"
+ os.environ.get("AZURE_AI_API_KEY") is None
+ or os.environ.get("OPENAI_API_KEY") is None,
+ reason="AZURE_AI_API_KEY or OPENAI_API_KEY not set - skipping integration test",
)
def test_chat_completion(client_no_auth):
global headers
@@ -69,9 +70,9 @@ def test_chat_completion(client_no_auth):
model_name="user-azure-instance",
litellm_params=CompletionRequest(
model="azure/gpt-4.1-mini",
- api_key=os.getenv("AZURE_API_KEY"),
+ api_key=os.getenv("AZURE_AI_API_KEY"),
api_version=os.getenv("AZURE_API_VERSION"),
- api_base=os.getenv("AZURE_API_BASE"),
+ api_base=os.getenv("AZURE_AI_API_BASE"),
timeout=10,
),
tpm=240000,
diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py
index 61a2f3055af..59b9297d1b2 100644
--- a/tests/proxy_unit_tests/test_proxy_server.py
+++ b/tests/proxy_unit_tests/test_proxy_server.py
@@ -119,7 +119,7 @@ def fake_env_vars(monkeypatch):
# Set some fake environment variables
monkeypatch.setenv("OPENAI_API_KEY", "fake_openai_api_key")
monkeypatch.setenv("OPENAI_API_BASE", "http://fake-openai-api-base")
- monkeypatch.setenv("AZURE_API_BASE", "http://fake-azure-api-base")
+ monkeypatch.setenv("AZURE_AI_API_BASE", "http://fake-azure-api-base")
monkeypatch.setenv("AZURE_OPENAI_API_KEY", "fake_azure_openai_api_key")
monkeypatch.setenv("AZURE_SWEDEN_API_BASE", "http://fake-azure-sweden-api-base")
monkeypatch.setenv("REDIS_HOST", "localhost")
@@ -178,7 +178,7 @@ def test_chat_completion(mock_acompletion, client_no_auth):
def test_chat_completion_malformed_messages_returns_400(client_no_auth):
"""
Test that malformed messages (strings instead of dicts) return 400 instead of 500.
-
+
This test verifies that when a client sends messages as raw strings instead of
{role, content} objects, LiteLLM returns a 400 invalid_request_error instead
of a 500 Internal Server Error.
@@ -188,33 +188,41 @@ def test_chat_completion_malformed_messages_returns_400(client_no_auth):
# Test data with malformed messages (string instead of dict)
test_data = {
"model": "gpt-3.5-turbo",
- "messages": ["hi how are you"], # Invalid: should be [{"role": "user", "content": "hi how are you"}]
+ "messages": [
+ "hi how are you"
+ ], # Invalid: should be [{"role": "user", "content": "hi how are you"}]
}
print("testing proxy server with malformed messages")
- response = client_no_auth.post("/v1/chat/completions", json=test_data, headers=headers)
-
+ response = client_no_auth.post(
+ "/v1/chat/completions", json=test_data, headers=headers
+ )
+
print(f"response status: {response.status_code}")
print(f"response text: {response.text}")
-
+
# Should return 400, not 500
- assert response.status_code == 400, f"Expected 400, got {response.status_code}. Response: {response.text}"
-
+ assert (
+ response.status_code == 400
+ ), f"Expected 400, got {response.status_code}. Response: {response.text}"
+
# Verify error format
result = response.json()
assert "error" in result, "Response should contain 'error' key"
error = result["error"]
-
+
# Verify error type and message
- assert error.get("type") == "invalid_request_error" or error.get("type") is None, \
- f"Expected invalid_request_error or None, got {error.get('type')}"
- assert error.get("code") == "400" or error.get("code") == 400, \
- f"Expected code 400, got {error.get('code')}"
-
+ assert (
+ error.get("type") == "invalid_request_error" or error.get("type") is None
+ ), f"Expected invalid_request_error or None, got {error.get('type')}"
+ assert (
+ error.get("code") == "400" or error.get("code") == 400
+ ), f"Expected code 400, got {error.get('code')}"
+
# Error message should indicate invalid request format
error_message = error.get("message", "")
assert len(error_message) > 0, "Error message should not be empty"
-
+
except Exception as e:
pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}")
@@ -342,7 +350,7 @@ def test_chat_completion_forward_llm_provider_auth_headers(
"""
Test that LLM provider auth headers (x-api-key, x-goog-api-key) are forwarded
when forward_llm_provider_auth_headers=True.
-
+
This allows clients to send their own LLM provider API keys through the proxy.
"""
try:
@@ -351,7 +359,7 @@ def test_chat_completion_forward_llm_provider_auth_headers(
gs["forward_client_headers_to_llm_api"] = True
gs["forward_llm_provider_auth_headers"] = forward_llm_auth_headers
setattr(litellm.proxy.proxy_server, "general_settings", gs)
-
+
# Test data
test_data = {
"model": "gpt-3.5-turbo",
@@ -360,7 +368,7 @@ def test_chat_completion_forward_llm_provider_auth_headers(
],
"max_tokens": 10,
}
-
+
# Headers including LLM provider auth
request_headers = {
"Authorization": "Bearer sk-proxy-auth-123", # Proxy auth (should be stripped)
@@ -368,17 +376,17 @@ def test_chat_completion_forward_llm_provider_auth_headers(
"x-goog-api-key": "google-api-key-123", # Google API key
"X-Custom-Header": "custom-value", # Custom header (should be forwarded)
}
-
+
# Make request
response = client_no_auth.post(
"/v1/chat/completions", json=test_data, headers=request_headers
)
-
+
assert response.status_code == 200
-
+
# Check forwarded headers
forwarded_headers = mock_acompletion.call_args.kwargs.get("headers", {})
-
+
if forward_llm_auth_headers:
# LLM provider auth headers should be forwarded
assert "x-api-key" in forwarded_headers
@@ -389,19 +397,23 @@ def test_chat_completion_forward_llm_provider_auth_headers(
# LLM provider auth headers should be stripped
assert "x-api-key" not in forwarded_headers
assert "x-goog-api-key" not in forwarded_headers
-
+
# Custom headers should always be forwarded (when forward_client_headers_to_llm_api=True)
assert "x-custom-header" in forwarded_headers
assert forwarded_headers["x-custom-header"] == "custom-value"
-
+
# Proxy Authorization should never be forwarded
assert "authorization" not in forwarded_headers
-
- print(f"✓ Test passed with forward_llm_provider_auth_headers={forward_llm_auth_headers}")
+
+ print(
+ f"✓ Test passed with forward_llm_provider_auth_headers={forward_llm_auth_headers}"
+ )
print(f" Forwarded headers: {list(forwarded_headers.keys())}")
-
+
except Exception as e:
- pytest.fail(f"Test failed with forward_llm_auth_headers={forward_llm_auth_headers}: {str(e)}")
+ pytest.fail(
+ f"Test failed with forward_llm_auth_headers={forward_llm_auth_headers}: {str(e)}"
+ )
finally:
# Clean up
gs = getattr(litellm.proxy.proxy_server, "general_settings")
@@ -2406,9 +2418,7 @@ async def test_run_background_health_check_reflects_llm_model_list(monkeypatch):
test_model_list_2 = [{"model_name": "model-b"}]
called_model_lists = []
- async def fake_perform_health_check(
- model_list, details, max_concurrency=None
- ):
+ async def fake_perform_health_check(model_list, details, max_concurrency=None):
called_model_lists.append(copy.deepcopy(model_list))
return (["healthy"], ["unhealthy"])
@@ -2452,13 +2462,14 @@ async def test_background_health_check_skip_disabled_models(monkeypatch):
test_model_list = [
{"model_name": "model-a"},
- {"model_name": "model-b", "model_info": {"disable_background_health_check": True}},
+ {
+ "model_name": "model-b",
+ "model_info": {"disable_background_health_check": True},
+ },
]
called_model_lists = []
- async def fake_perform_health_check(
- model_list, details, max_concurrency=None
- ):
+ async def fake_perform_health_check(model_list, details, max_concurrency=None):
called_model_lists.append(copy.deepcopy(model_list))
return (["healthy"], [])
@@ -2500,15 +2511,15 @@ def test_get_timeout_from_request():
@pytest.mark.parametrize(
"ui_exists, ui_has_content",
[
- (True, True), # UI path exists and has content
+ (True, True), # UI path exists and has content
(True, False), # UI path exists but is empty
- (False, False), # UI path doesn't exist
+ (False, False), # UI path doesn't exist
],
)
def test_non_root_ui_path_logic(monkeypatch, tmp_path, ui_exists, ui_has_content):
"""
Test the non-root Docker UI path detection logic.
-
+
Tests that when LITELLM_NON_ROOT is set to "true":
- If UI path exists and has content, it should be used
- If UI path doesn't exist or is empty, proper error logging occurs
@@ -2516,44 +2527,54 @@ def test_non_root_ui_path_logic(monkeypatch, tmp_path, ui_exists, ui_has_content
import tempfile
import shutil
from unittest.mock import MagicMock
-
+
# Create a temporary directory to act as /tmp/litellm_ui
test_ui_path = tmp_path / "litellm_ui"
-
+
if ui_exists:
test_ui_path.mkdir(parents=True, exist_ok=True)
if ui_has_content:
# Create some dummy files to simulate built UI
(test_ui_path / "index.html").write_text("")
(test_ui_path / "app.js").write_text("console.log('test');")
-
+
# Mock the environment variable and os.path operations
monkeypatch.setenv("LITELLM_NON_ROOT", "true")
-
+
# Create a mock logger to capture log messages
mock_logger = MagicMock()
-
+
# We need to reimport or reload the relevant code section
# Since this is module-level code, we'll test the logic directly
ui_path = None
non_root_ui_path = str(test_ui_path)
-
+
# Simulate the logic from proxy_server.py lines 909-920
if os.getenv("LITELLM_NON_ROOT", "").lower() == "true":
if os.path.exists(non_root_ui_path) and os.listdir(non_root_ui_path):
- mock_logger.info(f"Using pre-built UI for non-root Docker: {non_root_ui_path}")
- mock_logger.info(f"UI files found: {len(os.listdir(non_root_ui_path))} items")
+ mock_logger.info(
+ f"Using pre-built UI for non-root Docker: {non_root_ui_path}"
+ )
+ mock_logger.info(
+ f"UI files found: {len(os.listdir(non_root_ui_path))} items"
+ )
ui_path = non_root_ui_path
else:
- mock_logger.error(f"UI not found at {non_root_ui_path}. UI will not be available.")
- mock_logger.error(f"Path exists: {os.path.exists(non_root_ui_path)}, Has content: {os.path.exists(non_root_ui_path) and bool(os.listdir(non_root_ui_path))}")
-
+ mock_logger.error(
+ f"UI not found at {non_root_ui_path}. UI will not be available."
+ )
+ mock_logger.error(
+ f"Path exists: {os.path.exists(non_root_ui_path)}, Has content: {os.path.exists(non_root_ui_path) and bool(os.listdir(non_root_ui_path))}"
+ )
+
# Verify behavior based on test parameters
if ui_exists and ui_has_content:
# UI should be found and used
assert ui_path == non_root_ui_path
assert mock_logger.info.call_count == 2
- mock_logger.info.assert_any_call(f"Using pre-built UI for non-root Docker: {non_root_ui_path}")
+ mock_logger.info.assert_any_call(
+ f"Using pre-built UI for non-root Docker: {non_root_ui_path}"
+ )
# Verify the second info call mentions the number of items
info_calls = [call[0][0] for call in mock_logger.info.call_args_list]
assert any("UI files found:" in call and "items" in call for call in info_calls)
@@ -2562,7 +2583,9 @@ def test_non_root_ui_path_logic(monkeypatch, tmp_path, ui_exists, ui_has_content
# UI should not be found, error should be logged
assert ui_path is None
assert mock_logger.error.call_count == 2
- mock_logger.error.assert_any_call(f"UI not found at {non_root_ui_path}. UI will not be available.")
+ mock_logger.error.assert_any_call(
+ f"UI not found at {non_root_ui_path}. UI will not be available."
+ )
# Verify the second error call has path existence info
error_calls = [call[0][0] for call in mock_logger.error.call_args_list]
assert any("Path exists:" in call for call in error_calls)
@@ -2574,17 +2597,17 @@ async def test_get_config_callbacks_with_all_types(client_no_auth):
"""
Test that /get/config/callbacks returns all three callback types:
- success_callback with type="success"
- - failure_callback with type="failure"
+ - failure_callback with type="failure"
- callbacks (success_and_failure) with type="success_and_failure"
"""
from litellm.proxy.proxy_server import ProxyConfig
-
+
# Create a mock config with all three callback types
mock_config_data = {
"litellm_settings": {
"success_callback": ["langfuse", "braintrust"],
"failure_callback": ["sentry"],
- "callbacks": ["otel", "langsmith"]
+ "callbacks": ["otel", "langsmith"],
},
"environment_variables": {
"LANGFUSE_PUBLIC_KEY": "test-public-key",
@@ -2595,51 +2618,53 @@ async def test_get_config_callbacks_with_all_types(client_no_auth):
"OTEL_ENDPOINT": "http://localhost:4317",
"LANGSMITH_API_KEY": "test-langsmith-key",
},
- "general_settings": {}
+ "general_settings": {},
}
-
+
proxy_config = getattr(litellm.proxy.proxy_server, "proxy_config")
-
+
with patch.object(
proxy_config, "get_config", new=AsyncMock(return_value=mock_config_data)
):
response = client_no_auth.get("/get/config/callbacks")
-
+
assert response.status_code == 200
result = response.json()
-
+
# Verify response structure
assert "status" in result
assert result["status"] == "success"
assert "callbacks" in result
-
+
callbacks = result["callbacks"]
-
+
# Verify we have all 5 callbacks (2 success + 1 failure + 2 success_and_failure)
assert len(callbacks) == 5
-
+
# Group callbacks by type
success_callbacks = [cb for cb in callbacks if cb.get("type") == "success"]
failure_callbacks = [cb for cb in callbacks if cb.get("type") == "failure"]
- success_and_failure_callbacks = [cb for cb in callbacks if cb.get("type") == "success_and_failure"]
-
+ success_and_failure_callbacks = [
+ cb for cb in callbacks if cb.get("type") == "success_and_failure"
+ ]
+
# Verify all callbacks have required fields
for callback in callbacks:
assert "name" in callback
assert "variables" in callback
assert "type" in callback
assert callback["type"] in ["success", "failure", "success_and_failure"]
-
+
# Verify success callbacks
assert len(success_callbacks) == 2
success_names = [cb["name"] for cb in success_callbacks]
assert "langfuse" in success_names
assert "braintrust" in success_names
-
+
# Verify failure callbacks
assert len(failure_callbacks) == 1
assert failure_callbacks[0]["name"] == "sentry"
-
+
# Verify success_and_failure callbacks
assert len(success_and_failure_callbacks) == 2
success_and_failure_names = [cb["name"] for cb in success_and_failure_callbacks]
@@ -2654,13 +2679,13 @@ async def test_get_config_callbacks_environment_variables(client_no_auth):
for each callback type. Values are returned as-is from the config (no decryption).
"""
from litellm.proxy.proxy_server import ProxyConfig
-
+
# Create a mock config with callbacks and their env vars
mock_config_data = {
"litellm_settings": {
"success_callback": ["langfuse"],
"failure_callback": [],
- "callbacks": ["otel"]
+ "callbacks": ["otel"],
},
"environment_variables": {
"LANGFUSE_PUBLIC_KEY": "test-public-key",
@@ -2670,21 +2695,21 @@ async def test_get_config_callbacks_environment_variables(client_no_auth):
"OTEL_ENDPOINT": "http://localhost:4317",
"OTEL_HEADERS": "key=value",
},
- "general_settings": {}
+ "general_settings": {},
}
-
+
proxy_config = getattr(litellm.proxy.proxy_server, "proxy_config")
-
+
with patch.object(
proxy_config, "get_config", new=AsyncMock(return_value=mock_config_data)
):
response = client_no_auth.get("/get/config/callbacks")
-
+
assert response.status_code == 200
result = response.json()
-
+
callbacks = result["callbacks"]
-
+
# Find langfuse callback (success type)
langfuse_callback = next(
(cb for cb in callbacks if cb["name"] == "langfuse"), None
@@ -2692,7 +2717,7 @@ async def test_get_config_callbacks_environment_variables(client_no_auth):
assert langfuse_callback is not None
assert langfuse_callback["type"] == "success"
assert "variables" in langfuse_callback
-
+
# Verify langfuse env vars are present (values returned as-is, no decryption)
langfuse_vars = langfuse_callback["variables"]
assert "LANGFUSE_PUBLIC_KEY" in langfuse_vars
@@ -2701,15 +2726,13 @@ async def test_get_config_callbacks_environment_variables(client_no_auth):
assert langfuse_vars["LANGFUSE_SECRET_KEY"] == "test-secret-key"
assert "LANGFUSE_HOST" in langfuse_vars
assert langfuse_vars["LANGFUSE_HOST"] == "https://cloud.langfuse.com"
-
+
# Find otel callback (success_and_failure type)
- otel_callback = next(
- (cb for cb in callbacks if cb["name"] == "otel"), None
- )
+ otel_callback = next((cb for cb in callbacks if cb["name"] == "otel"), None)
assert otel_callback is not None
assert otel_callback["type"] == "success_and_failure"
assert "variables" in otel_callback
-
+
# Verify otel env vars are present
otel_vars = otel_callback["variables"]
assert "OTEL_EXPORTER" in otel_vars
diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py
index 00d4cd24e4b..6c6ec7bcd60 100644
--- a/tests/proxy_unit_tests/test_proxy_utils.py
+++ b/tests/proxy_unit_tests/test_proxy_utils.py
@@ -2044,6 +2044,58 @@ def test_update_model_if_team_alias_exists(data, user_api_key_dict, expected_mod
assert test_data.get("model") == expected_model
+def test_team_alias_stale_bypass_disabled_by_default(monkeypatch):
+ monkeypatch.delenv("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", raising=False)
+ import litellm.proxy.litellm_pre_call_utils as pre_call_utils
+ from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists
+
+ # Reset module-level cache to ensure test isolation
+ pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None
+
+ class _MockRouter:
+ team_model_to_deployment_indices = {("team-1", "gpt-4o"): [0]}
+
+ test_data = {"model": "gpt-4o"}
+ user_api_key_dict = UserAPIKeyAuth(
+ api_key="test_key",
+ team_id="team-1",
+ team_model_aliases={"gpt-4o": "model_name_team-1_legacy-uuid"},
+ )
+
+ with patch("litellm.proxy.proxy_server.llm_router", _MockRouter()):
+ _update_model_if_team_alias_exists(
+ data=test_data, user_api_key_dict=user_api_key_dict
+ )
+
+ assert test_data.get("model") == "model_name_team-1_legacy-uuid"
+
+
+def test_team_alias_stale_bypass_enabled_by_flag(monkeypatch):
+ import litellm.proxy.litellm_pre_call_utils as pre_call_utils
+ from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists
+
+ # Reset module-level cache to ensure test isolation
+ pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None
+
+ class _MockRouter:
+ team_model_to_deployment_indices = {("team-1", "gpt-4o"): [0]}
+
+ test_data = {"model": "gpt-4o"}
+ user_api_key_dict = UserAPIKeyAuth(
+ api_key="test_key",
+ team_id="team-1",
+ team_model_aliases={"gpt-4o": "model_name_team-1_legacy-uuid"},
+ )
+ monkeypatch.setenv("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", "true")
+
+ with patch("litellm.proxy.proxy_server.llm_router", _MockRouter()):
+ _update_model_if_team_alias_exists(
+ data=test_data, user_api_key_dict=user_api_key_dict
+ )
+
+ assert test_data.get("model") == "gpt-4o"
+
+
@pytest.fixture
def mock_prisma_client():
client = MagicMock()
diff --git a/tests/proxy_unit_tests/vertex_key.json b/tests/proxy_unit_tests/vertex_key.json
index 45ca6acc010..800969fb305 100644
--- a/tests/proxy_unit_tests/vertex_key.json
+++ b/tests/proxy_unit_tests/vertex_key.json
@@ -1,13 +1,13 @@
{
"type": "service_account",
- "project_id": "pathrise-convert-1606954137718",
+ "project_id": "litellm-ci-cd",
"private_key_id": "",
"private_key": "",
- "client_email": "ci-cd-723@pathrise-convert-1606954137718.iam.gserviceaccount.com",
- "client_id": "109577393201924326488",
+ "client_email": "test-litellm-ci-cd@litellm-ci-cd.iam.gserviceaccount.com",
+ "client_id": "116563532503305622785",
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
"token_uri": "https://oauth2.googleapis.com/token",
"auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs",
- "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/ci-cd-723%40pathrise-convert-1606954137718.iam.gserviceaccount.com",
+ "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/test-litellm-ci-cd%40litellm-ci-cd.iam.gserviceaccount.com",
"universe_domain": "googleapis.com"
}
diff --git a/tests/router_unit_tests/test_get_model_list_alias_optimization.py b/tests/router_unit_tests/test_get_model_list_alias_optimization.py
index 31d992b6646..2c2df3be945 100644
--- a/tests/router_unit_tests/test_get_model_list_alias_optimization.py
+++ b/tests/router_unit_tests/test_get_model_list_alias_optimization.py
@@ -44,7 +44,7 @@ def test_map_team_model_should_not_iterate_aliases_for_non_alias_team_model_name
{f"alias-{idx}": "gpt-4" for idx in range(200)}
)
- assert (
- router.map_team_model(team_model_name="team-model", team_id="team-1")
- == "gpt-3.5-turbo"
- )
+ # map_team_model should return the public name unchanged (not the internal UUID name)
+ # so the router can find all sibling deployments via team_id filtering
+ result = router.map_team_model(team_model_name="team-model", team_id="team-1")
+ assert result == "team-model", f"Expected public name 'team-model', got {result}"
diff --git a/tests/router_unit_tests/test_router_index_management.py b/tests/router_unit_tests/test_router_index_management.py
index 90d98b8ab0a..2694c62827c 100644
--- a/tests/router_unit_tests/test_router_index_management.py
+++ b/tests/router_unit_tests/test_router_index_management.py
@@ -118,6 +118,28 @@ class TestRouterIndexManagement:
assert router.model_id_to_deployment_index_map["id-2"] == 1
assert router.model_id_to_deployment_index_map["id-3"] == 2
+ def test_update_team_model_index(self, router):
+ """Test _update_team_model_index updates team_model_to_deployment_indices."""
+ model = {
+ "model_name": "team-alias",
+ "model_info": {
+ "id": "dep-1",
+ "team_id": "team-abc",
+ "team_public_model_name": "gpt-4o",
+ },
+ }
+ router._update_team_model_index(model, 0)
+ assert router.team_model_to_deployment_indices[("team-abc", "gpt-4o")] == [0]
+ router._update_team_model_index(model, 2)
+ assert router.team_model_to_deployment_indices[("team-abc", "gpt-4o")] == [0, 2]
+
+ router._update_team_model_index(
+ {"model_name": "x", "model_info": {"id": "dep-2"}}, 5
+ )
+ assert router.team_model_to_deployment_indices == {
+ ("team-abc", "gpt-4o"): [0, 2],
+ }
+
def test_has_model_id(self, router):
"""Test has_model_id function for O(1) membership check"""
# Setup: Add models to router
diff --git a/tests/store_model_in_db_tests/test_adding_passthrough_model.py b/tests/store_model_in_db_tests/test_adding_passthrough_model.py
index e901be5bd74..80cd7e1aabf 100644
--- a/tests/store_model_in_db_tests/test_adding_passthrough_model.py
+++ b/tests/store_model_in_db_tests/test_adding_passthrough_model.py
@@ -1,12 +1,12 @@
"""
Test adding a pass through assemblyai model + api key + api base to the db
-wait 20 seconds
-make request
+wait 20 seconds
+make request
-Cases to cover
-1. user points api base to /assemblyai
+Cases to cover
+1. user points api base to /assemblyai
2. user points api base to /asssemblyai/us
-3. user points api base to /assemblyai/eu
+3. user points api base to /assemblyai/eu
4. Bad API Key / credential - 401
"""
@@ -21,7 +21,7 @@ TEST_MASTER_KEY = "sk-1234"
PROXY_BASE_URL = "http://0.0.0.0:4000"
US_BASE_URL = f"{PROXY_BASE_URL}/assemblyai"
EU_BASE_URL = f"{PROXY_BASE_URL}/eu.assemblyai"
-ASSEMBLYAI_API_KEY_ENV_VAR = "TEST_SPECIAL_ASSEMBLYAI_API_KEY"
+ASSEMBLYAI_API_KEY_ENV_VAR = "ASSEMBLYAI_API_KEY"
def _delete_all_assemblyai_models_from_db():
diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py
index 8bc6ffc0505..e40543e01a0 100644
--- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py
+++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py
@@ -2127,9 +2127,10 @@ def test_convert_chat_completion_file_type_to_input_file():
}
]
- input_items, instructions = (
- handler.convert_chat_completion_messages_to_responses_api(messages)
- )
+ (
+ input_items,
+ instructions,
+ ) = handler.convert_chat_completion_messages_to_responses_api(messages)
assert len(input_items) == 1
msg = input_items[0]
@@ -2176,11 +2177,257 @@ def test_convert_chat_completion_file_type_with_file_id():
}
]
- input_items, instructions = (
- handler.convert_chat_completion_messages_to_responses_api(messages)
- )
+ (
+ input_items,
+ instructions,
+ ) = handler.convert_chat_completion_messages_to_responses_api(messages)
content = input_items[0]["content"]
assert content[1]["type"] == "input_file"
assert content[1]["file_id"] == "file-abc123"
assert "file_data" not in content[1]
+
+
+# =============================================================================
+# Tests for reasoning_items round-trip (encrypted_content preservation)
+# =============================================================================
+
+
+def test_reasoning_items_non_streaming_round_trip():
+ """
+ Non-streaming: verify that reasoning_items (with encrypted_content) are:
+ 1. Extracted from ResponseReasoningItem and attached to the Message.
+ 2. Emitted as a 'reasoning' input item when the assistant message is
+ passed back to convert_chat_completion_messages_to_responses_api.
+ """
+ from unittest.mock import Mock
+
+ from openai.types.responses import ResponseOutputMessage, ResponseOutputText
+ from openai.types.responses.response_reasoning_item import (
+ ResponseReasoningItem,
+ Summary,
+ )
+
+ from litellm.completion_extras.litellm_responses_transformation.transformation import (
+ LiteLLMResponsesTransformationHandler,
+ )
+ from litellm.types.llms.openai import (
+ InputTokensDetails,
+ OutputTokensDetails,
+ ResponseAPIUsage,
+ ResponsesAPIResponse,
+ )
+ from litellm.types.utils import ModelResponse, Usage
+
+ handler = LiteLLMResponsesTransformationHandler()
+
+ encrypted = "gAAAAABpw5abc123FAKE=="
+ summary_text = "**Thinking about it**\n\nSome reasoning here."
+
+ reasoning_item = ResponseReasoningItem(
+ id="rs_test001",
+ summary=[Summary(text=summary_text, type="summary_text")],
+ type="reasoning",
+ content=None,
+ encrypted_content=encrypted,
+ status=None,
+ )
+ output_message = ResponseOutputMessage(
+ id="msg_test001",
+ content=[
+ ResponseOutputText(
+ annotations=[],
+ text="The answer is 42.",
+ type="output_text",
+ logprobs=[],
+ )
+ ],
+ role="assistant",
+ status="completed",
+ type="message",
+ )
+ usage = ResponseAPIUsage(
+ input_tokens=10,
+ input_tokens_details=InputTokensDetails(
+ audio_tokens=None, cached_tokens=0, text_tokens=None
+ ),
+ output_tokens=20,
+ output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None),
+ total_tokens=30,
+ cost=None,
+ )
+ raw_response = ResponsesAPIResponse(
+ id="resp_test001",
+ created_at=1234567890,
+ error=None,
+ incomplete_details=None,
+ instructions=None,
+ metadata={},
+ model="gpt-5-mini",
+ object="response",
+ output=[reasoning_item, output_message],
+ parallel_tool_calls=True,
+ temperature=1.0,
+ tool_choice="auto",
+ tools=[],
+ top_p=1.0,
+ max_output_tokens=None,
+ previous_response_id=None,
+ reasoning={"effort": "low", "summary": "detailed"},
+ status="completed",
+ text={"format": {"type": "text"}, "verbosity": "medium"},
+ truncation="disabled",
+ usage=usage,
+ user=None,
+ store=True,
+ background=False,
+ billing={"payer": "developer"},
+ max_tool_calls=None,
+ prompt_cache_key=None,
+ safety_identifier=None,
+ service_tier="default",
+ top_logprobs=0,
+ )
+ model_response = ModelResponse(
+ id="chatcmpl-test001",
+ created=1234567890,
+ model=None,
+ object="chat.completion",
+ system_fingerprint=None,
+ choices=[],
+ usage=Usage(completion_tokens=0, prompt_tokens=0, total_tokens=0),
+ )
+
+ result = handler.transform_response(
+ model="gpt-5-mini",
+ raw_response=raw_response,
+ model_response=model_response,
+ logging_obj=Mock(),
+ request_data={"model": "gpt-5-mini"},
+ messages=[{"role": "user", "content": "What is the answer?"}],
+ optional_params={},
+ litellm_params={},
+ encoding=Mock(),
+ )
+
+ # ── Part 1: reasoning_items on the response message ──────────────────────
+ assert len(result.choices) == 1
+ msg = result.choices[0].message
+
+ assert (
+ msg.reasoning_content == summary_text
+ ), "reasoning_content should equal summary text"
+
+ assert msg.reasoning_items is not None, "reasoning_items should be set"
+ assert len(msg.reasoning_items) == 1
+ ri = msg.reasoning_items[0]
+ assert ri["type"] == "reasoning"
+ assert ri["id"] == "rs_test001"
+ assert ri["encrypted_content"] == encrypted, "encrypted_content must be preserved"
+ assert len(ri["summary"]) == 1
+ assert ri["summary"][0]["text"] == summary_text
+
+ # ── Part 2: reasoning item round-trips through message history ────────────
+ history = [
+ {"role": "user", "content": "What is the answer?"},
+ {
+ "role": "assistant",
+ "content": msg.content,
+ "reasoning_items": msg.reasoning_items,
+ },
+ {"role": "user", "content": "Can you elaborate?"},
+ ]
+ input_items, _ = handler.convert_chat_completion_messages_to_responses_api(history)
+
+ # The reasoning input item must appear before the assistant message item
+ types = [item.get("type") for item in input_items]
+ assert (
+ "reasoning" in types
+ ), "reasoning input item must be emitted for the assistant turn"
+
+ reasoning_input = next(
+ item for item in input_items if item.get("type") == "reasoning"
+ )
+ assert reasoning_input["id"] == "rs_test001"
+ assert reasoning_input["encrypted_content"] == encrypted
+ assert reasoning_input["summary"][0]["text"] == summary_text
+
+ # reasoning item must come before the assistant message item
+ reasoning_idx = types.index("reasoning")
+ assistant_msg_idx = next(
+ i
+ for i, item in enumerate(input_items)
+ if item.get("type") == "message" and item.get("role") == "assistant"
+ )
+ assert (
+ reasoning_idx < assistant_msg_idx
+ ), "reasoning input item must precede the assistant message item"
+
+
+def test_reasoning_items_streaming_emitted_on_response_completed():
+ """
+ Streaming: verify that reasoning_items (with encrypted_content) are emitted
+ on the delta of the response.completed chunk, enabling the caller to
+ round-trip them in subsequent requests.
+ """
+ from litellm.completion_extras.litellm_responses_transformation.transformation import (
+ OpenAiResponsesToChatCompletionStreamIterator,
+ )
+
+ iterator = OpenAiResponsesToChatCompletionStreamIterator(
+ streaming_response=None, sync_stream=True
+ )
+
+ encrypted = "gAAAAABpw5xyz987FAKE=="
+ summary_text = "**Reasoning summary**\n\nModel thought about this carefully."
+
+ chunk = {
+ "type": "response.completed",
+ "response": {
+ "id": "resp_stream001",
+ "status": "completed",
+ "output": [
+ {
+ "type": "reasoning",
+ "id": "rs_stream001",
+ "encrypted_content": encrypted,
+ "summary": [{"type": "summary_text", "text": summary_text}],
+ },
+ {
+ "type": "message",
+ "id": "msg_stream001",
+ "role": "assistant",
+ "content": [{"type": "output_text", "text": "The answer."}],
+ "status": "completed",
+ },
+ ],
+ "usage": {
+ "input_tokens": 10,
+ "output_tokens": 5,
+ "total_tokens": 15,
+ "input_tokens_details": {"cached_tokens": 0},
+ "output_tokens_details": {"reasoning_tokens": 0},
+ },
+ },
+ }
+
+ result = iterator.chunk_parser(chunk)
+
+ assert len(result.choices) == 1
+ delta = result.choices[0].delta
+
+ # finish_reason must be set (response is complete)
+ assert result.choices[0].finish_reason == "stop"
+
+ # reasoning_items must be on the delta
+ assert (
+ getattr(delta, "reasoning_items", None) is not None
+ ), "reasoning_items must be present on the response.completed delta"
+ assert len(delta.reasoning_items) == 1
+ ri = delta.reasoning_items[0]
+ assert ri["type"] == "reasoning"
+ assert ri["id"] == "rs_stream001"
+ assert (
+ ri["encrypted_content"] == encrypted
+ ), "encrypted_content must be preserved in streaming"
+ assert ri["summary"][0]["text"] == summary_text
diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_cancel_expected_output.json b/tests/test_litellm/expected_fine_tuning_api/azure_cancel_expected_output.json
new file mode 100644
index 00000000000..d0a519fcd11
--- /dev/null
+++ b/tests/test_litellm/expected_fine_tuning_api/azure_cancel_expected_output.json
@@ -0,0 +1,20 @@
+{
+ "id": "ftjob-azure-create-123",
+ "object": "fine_tuning.job",
+ "created_at": 1735689600,
+ "model": "davinci-002",
+ "status": "cancelled",
+ "fine_tuned_model": null,
+ "training_file": "file-5e4b20ecbd724182b9964f3cd2ab7212",
+ "hyperparameters": {
+ "n_epochs": 3,
+ "batch_size": null,
+ "learning_rate_multiplier": null
+ },
+ "organization_id": "",
+ "result_files": [],
+ "validation_file": null,
+ "trained_tokens": null,
+ "estimated_finish": null,
+ "error": null
+}
diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_cancel_raw_response.json b/tests/test_litellm/expected_fine_tuning_api/azure_cancel_raw_response.json
new file mode 100644
index 00000000000..0093a04f708
--- /dev/null
+++ b/tests/test_litellm/expected_fine_tuning_api/azure_cancel_raw_response.json
@@ -0,0 +1,18 @@
+{
+ "id": "ftjob-azure-create-123",
+ "object": "fine_tuning.job",
+ "created_at": 1735689600,
+ "model": "davinci-002",
+ "status": "canceled",
+ "fine_tuned_model": null,
+ "training_file": "file-5e4b20ecbd724182b9964f3cd2ab7212",
+ "hyperparameters": {
+ "n_epochs": 3
+ },
+ "organization_id": null,
+ "result_files": null,
+ "validation_file": null,
+ "trained_tokens": null,
+ "estimated_finish": null,
+ "error": null
+}
diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_cancel_request.json b/tests/test_litellm/expected_fine_tuning_api/azure_cancel_request.json
new file mode 100644
index 00000000000..bdbbdbad074
--- /dev/null
+++ b/tests/test_litellm/expected_fine_tuning_api/azure_cancel_request.json
@@ -0,0 +1,3 @@
+{
+ "fine_tuning_job_id": "ftjob-azure-create-123"
+}
diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_create_expected_output.json b/tests/test_litellm/expected_fine_tuning_api/azure_create_expected_output.json
new file mode 100644
index 00000000000..201ae3047d3
--- /dev/null
+++ b/tests/test_litellm/expected_fine_tuning_api/azure_create_expected_output.json
@@ -0,0 +1,20 @@
+{
+ "id": "ftjob-azure-create-123",
+ "object": "fine_tuning.job",
+ "created_at": 1735689600,
+ "model": "davinci-002",
+ "status": "running",
+ "fine_tuned_model": null,
+ "training_file": "file-5e4b20ecbd724182b9964f3cd2ab7212",
+ "hyperparameters": {
+ "n_epochs": 3,
+ "batch_size": null,
+ "learning_rate_multiplier": null
+ },
+ "organization_id": "",
+ "result_files": [],
+ "validation_file": null,
+ "trained_tokens": null,
+ "estimated_finish": null,
+ "error": null
+}
diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_create_raw_response.json b/tests/test_litellm/expected_fine_tuning_api/azure_create_raw_response.json
new file mode 100644
index 00000000000..857b2499389
--- /dev/null
+++ b/tests/test_litellm/expected_fine_tuning_api/azure_create_raw_response.json
@@ -0,0 +1,18 @@
+{
+ "id": "ftjob-azure-create-123",
+ "object": "fine_tuning.job",
+ "created_at": 1735689600,
+ "model": "davinci-002",
+ "status": "running",
+ "fine_tuned_model": null,
+ "training_file": "file-5e4b20ecbd724182b9964f3cd2ab7212",
+ "hyperparameters": {
+ "n_epochs": 3
+ },
+ "organization_id": null,
+ "result_files": null,
+ "validation_file": null,
+ "trained_tokens": null,
+ "estimated_finish": null,
+ "error": null
+}
diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_create_request.json b/tests/test_litellm/expected_fine_tuning_api/azure_create_request.json
new file mode 100644
index 00000000000..57319b2698c
--- /dev/null
+++ b/tests/test_litellm/expected_fine_tuning_api/azure_create_request.json
@@ -0,0 +1,8 @@
+{
+ "model": "gpt-35-turbo-1106",
+ "training_file": "file-5e4b20ecbd724182b9964f3cd2ab7212",
+ "hyperparameters": {},
+ "extra_body": {
+ "trainingType": 1
+ }
+}
diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_list_raw_response.json b/tests/test_litellm/expected_fine_tuning_api/azure_list_raw_response.json
new file mode 100644
index 00000000000..9986e7d4c54
--- /dev/null
+++ b/tests/test_litellm/expected_fine_tuning_api/azure_list_raw_response.json
@@ -0,0 +1,20 @@
+{
+ "object": "list",
+ "data": [
+ {
+ "id": "ftjob-azure-create-123",
+ "object": "fine_tuning.job",
+ "created_at": 1735689600,
+ "model": "davinci-002",
+ "status": "running"
+ },
+ {
+ "id": "ftjob-azure-prev-000",
+ "object": "fine_tuning.job",
+ "created_at": 1735603200,
+ "model": "davinci-002",
+ "status": "succeeded"
+ }
+ ],
+ "has_more": false
+}
diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_list_request.json b/tests/test_litellm/expected_fine_tuning_api/azure_list_request.json
new file mode 100644
index 00000000000..6bfbb2ffec7
--- /dev/null
+++ b/tests/test_litellm/expected_fine_tuning_api/azure_list_request.json
@@ -0,0 +1,4 @@
+{
+ "after": "ftjob-azure-prev-000",
+ "limit": 2
+}
diff --git a/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py b/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py
index cd76ba1e863..9bcf08fdd71 100644
--- a/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py
+++ b/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py
@@ -2,7 +2,7 @@
Unit tests for Prometheus user and team count metrics
"""
from datetime import datetime, timezone
-from unittest.mock import MagicMock, patch
+from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from prometheus_client import REGISTRY
@@ -523,3 +523,193 @@ async def test_set_user_budget_metrics_after_api_request_inf_when_genuinely_no_b
assert actual_value == float("inf"), (
"remaining_user_budget_metric should be +Inf when user truly has no budget"
)
+
+
+# ---------------------------------------------------------------------------
+# Org budget metric tests
+# ---------------------------------------------------------------------------
+
+
+def test_org_budget_metrics_initialized(prometheus_logger):
+ """Test that the 3 org budget gauge metrics are initialized."""
+ assert hasattr(prometheus_logger, "litellm_remaining_org_budget_metric")
+ assert hasattr(prometheus_logger, "litellm_org_max_budget_metric")
+ assert hasattr(prometheus_logger, "litellm_org_budget_remaining_hours_metric")
+ assert prometheus_logger.litellm_remaining_org_budget_metric is not None
+ assert prometheus_logger.litellm_org_max_budget_metric is not None
+ assert prometheus_logger.litellm_org_budget_remaining_hours_metric is not None
+
+
+def test_set_org_budget_metrics_remaining_budget(prometheus_logger):
+ """_set_org_budget_metrics sets remaining budget gauge correctly."""
+ prometheus_logger.litellm_remaining_org_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_max_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_budget_remaining_hours_metric = MagicMock()
+
+ prometheus_logger._set_org_budget_metrics(
+ org_id="org-abc",
+ org_alias="my-org",
+ spend=200.0,
+ max_budget=500.0,
+ budget_reset_at=None,
+ )
+
+ set_call = prometheus_logger.litellm_remaining_org_budget_metric.labels().set
+ set_call.assert_called_once()
+ actual = set_call.call_args[0][0]
+ assert abs(actual - 300.0) < 0.01, f"Expected 300.0, got {actual}"
+
+
+def test_set_org_budget_metrics_max_budget(prometheus_logger):
+ """_set_org_budget_metrics sets max budget gauge when max_budget is not None."""
+ prometheus_logger.litellm_remaining_org_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_max_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_budget_remaining_hours_metric = MagicMock()
+
+ prometheus_logger._set_org_budget_metrics(
+ org_id="org-abc",
+ org_alias="my-org",
+ spend=100.0,
+ max_budget=1000.0,
+ budget_reset_at=None,
+ )
+
+ prometheus_logger.litellm_org_max_budget_metric.labels().set.assert_called_once_with(
+ 1000.0
+ )
+
+
+def test_set_org_budget_metrics_no_max_budget(prometheus_logger):
+ """_set_org_budget_metrics does not set max budget gauge when max_budget is None."""
+ prometheus_logger.litellm_remaining_org_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_max_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_budget_remaining_hours_metric = MagicMock()
+
+ prometheus_logger._set_org_budget_metrics(
+ org_id="org-abc",
+ org_alias="my-org",
+ spend=50.0,
+ max_budget=None,
+ budget_reset_at=None,
+ )
+
+ prometheus_logger.litellm_org_max_budget_metric.labels().set.assert_not_called()
+
+
+def test_set_org_budget_metrics_remaining_hours(prometheus_logger):
+ """_set_org_budget_metrics sets remaining hours gauge when budget_reset_at is set."""
+ prometheus_logger.litellm_remaining_org_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_max_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_budget_remaining_hours_metric = MagicMock()
+
+ future_reset = datetime(2099, 1, 1, tzinfo=timezone.utc)
+ prometheus_logger._set_org_budget_metrics(
+ org_id="org-abc",
+ org_alias="my-org",
+ spend=10.0,
+ max_budget=500.0,
+ budget_reset_at=future_reset,
+ )
+
+ prometheus_logger.litellm_org_budget_remaining_hours_metric.labels().set.assert_called_once()
+
+
+@pytest.mark.asyncio
+async def test_set_org_budget_metrics_after_api_request(prometheus_logger):
+ """_set_org_budget_metrics_after_api_request uses cache helper and accounts for response_cost."""
+ import sys
+
+ prometheus_logger.litellm_remaining_org_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_max_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_budget_remaining_hours_metric = MagicMock()
+
+ budget_mock = MagicMock()
+ budget_mock.max_budget = 1000.0
+ budget_mock.budget_reset_at = datetime(2099, 1, 1, tzinfo=timezone.utc)
+
+ org_mock = MagicMock()
+ org_mock.organization_id = "org-xyz"
+ org_mock.organization_alias = "test-org"
+ org_mock.spend = 300.0
+ org_mock.litellm_budget_table = budget_mock
+
+ mock_prisma = MagicMock()
+ mock_proxy_server = MagicMock()
+ mock_proxy_server.prisma_client = mock_prisma
+ mock_proxy_server.user_api_key_cache = MagicMock()
+
+ with (
+ patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}),
+ patch(
+ "litellm.proxy.auth.auth_checks.get_org_object",
+ AsyncMock(return_value=org_mock),
+ ),
+ ):
+ await prometheus_logger._set_org_budget_metrics_after_api_request(
+ org_id="org-xyz",
+ response_cost=50.0,
+ )
+
+ # remaining budget should reflect spend + response_cost (300 + 50 = 350, remaining = 1000 - 350 = 650)
+ remaining_call = prometheus_logger.litellm_remaining_org_budget_metric.labels().set.call_args
+ assert remaining_call is not None
+ assert remaining_call[0][0] == pytest.approx(650.0)
+
+ prometheus_logger.litellm_org_max_budget_metric.labels().set.assert_called_once_with(
+ 1000.0
+ )
+ prometheus_logger.litellm_org_budget_remaining_hours_metric.labels().set.assert_called_once()
+
+
+@pytest.mark.asyncio
+async def test_set_org_budget_metrics_after_api_request_no_org_id(prometheus_logger):
+ """_set_org_budget_metrics_after_api_request is a no-op when org_id is None."""
+ prometheus_logger.litellm_remaining_org_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_max_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_budget_remaining_hours_metric = MagicMock()
+
+ await prometheus_logger._set_org_budget_metrics_after_api_request(
+ org_id=None,
+ response_cost=1.0,
+ )
+
+ prometheus_logger.litellm_remaining_org_budget_metric.labels().set.assert_not_called()
+ prometheus_logger.litellm_org_max_budget_metric.labels().set.assert_not_called()
+ prometheus_logger.litellm_org_budget_remaining_hours_metric.labels().set.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_initialize_org_budget_metrics(prometheus_logger):
+ """_initialize_org_budget_metrics fetches all orgs and sets gauges for each."""
+ import sys
+
+ prometheus_logger.litellm_remaining_org_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_max_budget_metric = MagicMock()
+ prometheus_logger.litellm_org_budget_remaining_hours_metric = MagicMock()
+
+ budget_mock = MagicMock()
+ budget_mock.max_budget = 500.0
+ budget_mock.budget_reset_at = None
+
+ org_mock = MagicMock()
+ org_mock.organization_id = "org-init"
+ org_mock.organization_alias = "init-org"
+ org_mock.spend = 100.0
+ org_mock.litellm_budget_table = budget_mock
+
+ mock_prisma = MagicMock()
+ mock_prisma.db.litellm_organizationtable.find_many = AsyncMock(
+ return_value=[org_mock]
+ )
+ mock_prisma.db.litellm_organizationtable.count = AsyncMock(return_value=1)
+
+ mock_proxy_server = MagicMock()
+ mock_proxy_server.prisma_client = mock_prisma
+
+ with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}):
+ await prometheus_logger._initialize_org_budget_metrics()
+
+ prometheus_logger.litellm_remaining_org_budget_metric.labels().set.assert_called_once()
+ prometheus_logger.litellm_org_max_budget_metric.labels().set.assert_called_once_with(
+ 500.0
+ )
diff --git a/tests/test_litellm/llms/azure/test_azure_common_utils.py b/tests/test_litellm/llms/azure/test_azure_common_utils.py
index d689c676580..9ed801b360e 100644
--- a/tests/test_litellm/llms/azure/test_azure_common_utils.py
+++ b/tests/test_litellm/llms/azure/test_azure_common_utils.py
@@ -460,7 +460,7 @@ async def test_ensure_initialize_azure_sdk_client_always_used(call_type):
"api_key": "test-api-key",
"api_version": os.getenv("AZURE_API_VERSION", "2023-05-15"),
"api_base": os.getenv(
- "AZURE_API_BASE", "https://test.openai.azure.com"
+ "AZURE_AI_API_BASE", "https://test.openai.azure.com"
),
},
}
@@ -539,7 +539,11 @@ async def test_ensure_initialize_azure_sdk_client_always_used(call_type):
patch_target = (
"litellm.rerank_api.main.azure_rerank.initialize_azure_sdk_client"
)
- elif call_type == CallTypes.acreate_batch or call_type == CallTypes.aretrieve_batch or call_type == CallTypes.acancel_batch:
+ elif (
+ call_type == CallTypes.acreate_batch
+ or call_type == CallTypes.aretrieve_batch
+ or call_type == CallTypes.acancel_batch
+ ):
patch_target = (
"litellm.batches.main.azure_batches_instance.initialize_azure_sdk_client"
)
@@ -570,7 +574,9 @@ async def test_ensure_initialize_azure_sdk_client_always_used(call_type):
or call_type == CallTypes.avideo_extension
):
# Skip video call types as they don't use Azure SDK client initialization
- pytest.skip(f"Skipping {call_type.value} because Azure video calls don't use initialize_azure_sdk_client")
+ pytest.skip(
+ f"Skipping {call_type.value} because Azure video calls don't use initialize_azure_sdk_client"
+ )
elif (
call_type == CallTypes.alist_containers
or call_type == CallTypes.aretrieve_container
@@ -580,13 +586,26 @@ async def test_ensure_initialize_azure_sdk_client_always_used(call_type):
or call_type == CallTypes.aupload_container_file
):
# Skip container call types as they're not supported for Azure (only OpenAI)
- pytest.skip(f"Skipping {call_type.value} because Azure doesn't support container operations")
- elif call_type == CallTypes.avector_store_file_create or call_type == CallTypes.avector_store_file_list or call_type == CallTypes.avector_store_file_retrieve or call_type == CallTypes.avector_store_file_content or call_type == CallTypes.avector_store_file_update or call_type == CallTypes.avector_store_file_delete:
+ pytest.skip(
+ f"Skipping {call_type.value} because Azure doesn't support container operations"
+ )
+ elif (
+ call_type == CallTypes.avector_store_file_create
+ or call_type == CallTypes.avector_store_file_list
+ or call_type == CallTypes.avector_store_file_retrieve
+ or call_type == CallTypes.avector_store_file_content
+ or call_type == CallTypes.avector_store_file_update
+ or call_type == CallTypes.avector_store_file_delete
+ ):
# Skip vector store file call types as they're not supported for Azure (only OpenAI)
- pytest.skip(f"Skipping {call_type.value} because Azure doesn't support vector store file operations")
+ pytest.skip(
+ f"Skipping {call_type.value} because Azure doesn't support vector store file operations"
+ )
elif call_type == CallTypes.aocr or call_type == CallTypes.ocr:
# Skip OCR call types as they don't use Azure SDK client initialization
- pytest.skip(f"Skipping {call_type.value} because OCR calls don't use initialize_azure_sdk_client")
+ pytest.skip(
+ f"Skipping {call_type.value} because OCR calls don't use initialize_azure_sdk_client"
+ )
# Mock the initialize_azure_sdk_client function
with patch(patch_target) as mock_init_azure:
# Also mock async_function_with_fallbacks to prevent actual API calls
@@ -651,7 +670,7 @@ async def test_ensure_initialize_azure_sdk_client_always_used_azure_text(call_ty
"api_key": "test-api-key",
"api_version": os.getenv("AZURE_API_VERSION", "2023-05-15"),
"api_base": os.getenv(
- "AZURE_API_BASE", "https://test.openai.azure.com"
+ "AZURE_AI_API_BASE", "https://test.openai.azure.com"
),
},
}
@@ -767,7 +786,7 @@ AZURE_API_FUNCTION_PARAMS = [
"speech",
False,
{
- "model": "azure/tts-1",
+ "model": "azure/tts",
"input": "Hello, this is a test of text to speech",
"voice": "alloy",
"api_key": "test-api-key",
@@ -1434,43 +1453,44 @@ def test_token_provider_raises_exception(setup_mocks):
def test_get_azure_ad_token_provider_with_default_azure_credential():
"""
- Test that get_azure_ad_token_provider correctly uses DefaultAzureCredential
+ Test that get_azure_ad_token_provider correctly uses DefaultAzureCredential
when explicitly specified as the credential type. This verifies that the function
can dynamically instantiate DefaultAzureCredential and return a working token provider.
"""
# Mock Azure identity classes
- with patch('azure.identity.DefaultAzureCredential') as mock_default_cred, \
- patch('azure.identity.get_bearer_token_provider') as mock_token_provider:
-
+ with patch("azure.identity.DefaultAzureCredential") as mock_default_cred, patch(
+ "azure.identity.get_bearer_token_provider"
+ ) as mock_token_provider:
# Configure mocks
mock_credential_instance = MagicMock()
mock_default_cred.return_value = mock_credential_instance
mock_token_provider.return_value = lambda: "test-default-azure-token"
-
+
# Test with DefaultAzureCredential specified explicitly
token_provider = get_azure_ad_token_provider(
azure_scope="https://cognitiveservices.azure.com/.default",
- azure_credential=AzureCredentialType.DefaultAzureCredential
+ azure_credential=AzureCredentialType.DefaultAzureCredential,
)
-
+
# Verify DefaultAzureCredential was instantiated
mock_default_cred.assert_called_once_with()
-
+
# Verify get_bearer_token_provider was called with the right parameters
mock_token_provider.assert_called_once_with(
- mock_credential_instance,
- "https://cognitiveservices.azure.com/.default"
+ mock_credential_instance, "https://cognitiveservices.azure.com/.default"
)
-
+
# Verify the returned token provider works
token = token_provider()
assert token == "test-default-azure-token"
-def test_get_azure_ad_token_fallback_to_default_azure_credential(setup_mocks, monkeypatch):
+def test_get_azure_ad_token_fallback_to_default_azure_credential(
+ setup_mocks, monkeypatch
+):
"""
- Test that get_azure_ad_token falls back to DefaultAzureCredential when the
- service principal method fails but token refresh is enabled. This tests the
+ Test that get_azure_ad_token falls back to DefaultAzureCredential when the
+ service principal method fails but token refresh is enabled. This tests the
complete fallback flow from service principal to DefaultAzureCredential.
"""
# Clear environment variables that might interfere
@@ -1486,7 +1506,7 @@ def test_get_azure_ad_token_fallback_to_default_azure_credential(setup_mocks, mo
# Enable token refresh
setup_mocks["litellm"].enable_azure_ad_token_refresh = True
- # Configure get_azure_ad_token_provider to fail first (service principal)
+ # Configure get_azure_ad_token_provider to fail first (service principal)
# but succeed on second call (DefaultAzureCredential)
def mock_token_provider_side_effect(*args, **kwargs):
# If called with azure_credential=DefaultAzureCredential, return a working provider
@@ -1512,19 +1532,22 @@ def test_get_azure_ad_token_fallback_to_default_azure_credential(setup_mocks, mo
# 1. First with just azure_scope (service principal attempt)
# 2. Second with azure_credential=DefaultAzureCredential (fallback)
assert setup_mocks["token_provider"].call_count == 2
-
+
# Verify the calls were made with expected parameters
calls = setup_mocks["token_provider"].call_args_list
-
+
# First call should be service principal attempt (no azure_credential)
first_call_kwargs = calls[0][1]
assert "azure_scope" in first_call_kwargs
assert first_call_kwargs.get("azure_credential") is None
-
+
# Second call should be DefaultAzureCredential attempt
second_call_kwargs = calls[1][1]
assert "azure_scope" in second_call_kwargs
- assert second_call_kwargs.get("azure_credential") == AzureCredentialType.DefaultAzureCredential
+ assert (
+ second_call_kwargs.get("azure_credential")
+ == AzureCredentialType.DefaultAzureCredential
+ )
# Verify the token is what we expect from our DefaultAzureCredential mock
assert token == "mock-default-azure-credential-token"
@@ -1584,9 +1607,13 @@ def test_azure_v1_api_uses_openai_client(api_version):
)
# Should be OpenAI client, not AzureOpenAI
- assert isinstance(client, OpenAI), f"Expected OpenAI client for api_version={api_version}"
+ assert isinstance(
+ client, OpenAI
+ ), f"Expected OpenAI client for api_version={api_version}"
# base_url should be /openai/v1/ (not /deployments/)
- assert "/openai/v1/" in str(client.base_url), f"base_url should contain /openai/v1/, got {client.base_url}"
+ assert "/openai/v1/" in str(
+ client.base_url
+ ), f"base_url should contain /openai/v1/, got {client.base_url}"
# Test async client
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
@@ -1606,9 +1633,13 @@ def test_azure_v1_api_uses_openai_client(api_version):
)
# Should be AsyncOpenAI client, not AsyncAzureOpenAI
- assert isinstance(async_client, AsyncOpenAI), f"Expected AsyncOpenAI client for api_version={api_version}"
+ assert isinstance(
+ async_client, AsyncOpenAI
+ ), f"Expected AsyncOpenAI client for api_version={api_version}"
# base_url should be /openai/v1/
- assert "/openai/v1/" in str(async_client.base_url), f"base_url should contain /openai/v1/, got {async_client.base_url}"
+ assert "/openai/v1/" in str(
+ async_client.base_url
+ ), f"base_url should contain /openai/v1/, got {async_client.base_url}"
def test_azure_traditional_api_uses_azure_openai_client():
@@ -1643,7 +1674,9 @@ def test_azure_traditional_api_uses_azure_openai_client():
)
# Should be AzureOpenAI client
- assert isinstance(client, AzureOpenAI), f"Expected AzureOpenAI client for api_version={api_version}"
+ assert isinstance(
+ client, AzureOpenAI
+ ), f"Expected AzureOpenAI client for api_version={api_version}"
# Test async client
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
@@ -1663,4 +1696,6 @@ def test_azure_traditional_api_uses_azure_openai_client():
)
# Should be AsyncAzureOpenAI client
- assert isinstance(async_client, AsyncAzureOpenAI), f"Expected AsyncAzureOpenAI client for api_version={api_version}"
+ assert isinstance(
+ async_client, AsyncAzureOpenAI
+ ), f"Expected AsyncAzureOpenAI client for api_version={api_version}"
diff --git a/tests/test_litellm/llms/azure/test_azure_fine_tuning_api.py b/tests/test_litellm/llms/azure/test_azure_fine_tuning_api.py
new file mode 100644
index 00000000000..8d008d6c071
--- /dev/null
+++ b/tests/test_litellm/llms/azure/test_azure_fine_tuning_api.py
@@ -0,0 +1,150 @@
+import json
+from pathlib import Path
+from unittest.mock import AsyncMock, patch
+
+import pytest
+from openai import AsyncAzureOpenAI
+
+import litellm
+from litellm.llms.azure.fine_tuning.handler import AzureOpenAIFineTuningAPI
+
+
+def _expected_dir() -> Path:
+ return Path(__file__).resolve().parent.parent.parent / "expected_fine_tuning_api"
+
+
+def _load_json(file_name: str) -> dict:
+ path = _expected_dir() / file_name
+ assert path.exists(), f"Expected fixture file not found: {path}"
+ with open(path) as f:
+ return json.load(f)
+
+
+class _MockSDKResponse:
+ def __init__(self, payload: dict):
+ self._payload = payload
+
+ def model_dump(self) -> dict:
+ return self._payload
+
+
+def _mock_azure_client(
+ create_payload: dict | None = None,
+ list_payload: dict | None = None,
+ cancel_payload: dict | None = None,
+):
+ client = AsyncAzureOpenAI(
+ api_key="test-key",
+ api_version="2024-10-21",
+ azure_endpoint="https://exampleopenaiendpoint-production.up.railway.app",
+ )
+ client.fine_tuning.jobs.create = AsyncMock(
+ return_value=(
+ _MockSDKResponse(create_payload) if create_payload is not None else None
+ )
+ ) # type: ignore[method-assign]
+ client.fine_tuning.jobs.list = AsyncMock(
+ return_value=list_payload
+ ) # type: ignore[method-assign]
+ client.fine_tuning.jobs.cancel = AsyncMock(
+ return_value=(
+ _MockSDKResponse(cancel_payload) if cancel_payload is not None else None
+ )
+ ) # type: ignore[method-assign]
+ return client
+
+
+@pytest.mark.asyncio
+async def test_azure_acreate_fine_tuning_job_request_and_output_match_expected_json():
+ expected_request = _load_json("azure_create_request.json")
+ raw_response = _load_json("azure_create_raw_response.json")
+ expected_output = _load_json("azure_create_expected_output.json")
+
+ mock_client = _mock_azure_client(create_payload=raw_response)
+
+ with patch.object(
+ AzureOpenAIFineTuningAPI, "get_openai_client", return_value=mock_client
+ ):
+ response = await litellm.acreate_fine_tuning_job(
+ model="gpt-35-turbo-1106",
+ training_file="file-5e4b20ecbd724182b9964f3cd2ab7212",
+ custom_llm_provider="azure",
+ api_base="https://exampleopenaiendpoint-production.up.railway.app",
+ api_key="test-key",
+ api_version="2024-10-21",
+ )
+
+ request_kwargs = mock_client.fine_tuning.jobs.create.call_args.kwargs
+ assert request_kwargs == expected_request
+
+ response_dict = response.model_dump(exclude={"_hidden_params"})
+ for key, expected_value in expected_output.items():
+ assert key in response_dict, f"Missing key in response: {key}"
+ assert response_dict[key] == expected_value
+
+ assert response.id is not None
+ assert response.model == "davinci-002"
+
+
+@pytest.mark.asyncio
+async def test_azure_alist_fine_tuning_jobs_request_matches_expected_json():
+ expected_request = _load_json("azure_list_request.json")
+ raw_list_response = _load_json("azure_list_raw_response.json")
+
+ mock_client = _mock_azure_client(list_payload=raw_list_response)
+
+ with patch.object(
+ AzureOpenAIFineTuningAPI, "get_openai_client", return_value=mock_client
+ ):
+ response = await litellm.alist_fine_tuning_jobs(
+ after=expected_request["after"],
+ limit=expected_request["limit"],
+ custom_llm_provider="azure",
+ api_base="https://exampleopenaiendpoint-production.up.railway.app",
+ api_key="test-key",
+ api_version="2024-10-21",
+ )
+
+ request_kwargs = mock_client.fine_tuning.jobs.list.call_args.kwargs
+ assert request_kwargs == expected_request
+ assert response == raw_list_response
+
+
+@pytest.mark.asyncio
+async def test_azure_acancel_fine_tuning_job_request_and_output_match_expected_json():
+ expected_request = _load_json("azure_cancel_request.json")
+ raw_response = _load_json("azure_cancel_raw_response.json")
+ expected_output = _load_json("azure_cancel_expected_output.json")
+
+ mock_client = _mock_azure_client(cancel_payload=raw_response)
+
+ with patch.object(
+ AzureOpenAIFineTuningAPI, "get_openai_client", return_value=mock_client
+ ):
+ response = await litellm.acancel_fine_tuning_job(
+ fine_tuning_job_id=expected_request["fine_tuning_job_id"],
+ custom_llm_provider="azure",
+ api_base="https://exampleopenaiendpoint-production.up.railway.app",
+ api_key="test-key",
+ api_version="2024-10-21",
+ )
+
+ request_kwargs = mock_client.fine_tuning.jobs.cancel.call_args.kwargs
+ assert request_kwargs == expected_request
+
+ response_dict = response.model_dump(exclude={"_hidden_params"})
+ for key, expected_value in expected_output.items():
+ assert key in response_dict, f"Missing key in response: {key}"
+ assert response_dict[key] == expected_value
+
+ assert response.status == "cancelled"
+
+
+def test_azure_trainingtype_defaults_to_one():
+ handler = AzureOpenAIFineTuningAPI()
+ create_data = {"model": "gpt-4o-mini", "training_file": "file-test"}
+
+ handler._ensure_training_type(create_data)
+
+ assert "extra_body" in create_data
+ assert create_data["extra_body"]["trainingType"] == 1
diff --git a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_count_tokens_transformation.py b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_count_tokens_transformation.py
index 78806831685..d66798a5725 100644
--- a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_count_tokens_transformation.py
+++ b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_count_tokens_transformation.py
@@ -3,6 +3,7 @@ Tests for Azure AI Anthropic CountTokens transformation.
Verifies that the CountTokens API uses the correct authentication headers.
"""
+
import os
import sys
@@ -40,7 +41,7 @@ class TestAzureAIAnthropicCountTokensConfig:
assert headers["anthropic-version"] == "2023-06-01"
assert "anthropic-beta" in headers
- def test_get_required_headers_includes_azure_api_key(self):
+ def test_get_required_headers_includes_AZURE_AI_API_KEY(self):
"""
Test that get_required_headers includes Azure api-key header.
diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py
index e9aaa97a421..867a3e61bbd 100644
--- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py
+++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py
@@ -6,9 +6,7 @@ import sys
import pytest
from fastapi.testclient import TestClient
-sys.path.insert(
- 0, os.path.abspath("../../../../..")
-) # Adds the parent directory to the system path
+sys.path.insert(0, os.path.abspath("../../../../..")) # Adds the parent directory to the system path
from unittest.mock import MagicMock, patch
import litellm
@@ -37,10 +35,7 @@ def test_transform_usage():
)
assert openai_usage.completion_tokens == usage["outputTokens"]
assert openai_usage.total_tokens == usage["totalTokens"]
- assert (
- openai_usage.prompt_tokens_details.cached_tokens
- == usage["cacheReadInputTokens"]
- )
+ assert openai_usage.prompt_tokens_details.cached_tokens == usage["cacheReadInputTokens"]
assert openai_usage._cache_creation_input_tokens == usage["cacheWriteInputTokens"]
assert openai_usage._cache_read_input_tokens == usage["cacheReadInputTokens"]
# completion_tokens_details should always be populated
@@ -194,14 +189,10 @@ def test_apply_tool_call_transformation_if_needed():
role="user",
content=json.dumps(tool_response),
)
- transformed_message, _ = config.apply_tool_call_transformation_if_needed(
- message, tool_calls
- )
+ transformed_message, _ = config.apply_tool_call_transformation_if_needed(message, tool_calls)
assert len(transformed_message.tool_calls) == 1
assert transformed_message.tool_calls[0].function.name == "test_function"
- assert transformed_message.tool_calls[0].function.arguments == json.dumps(
- tool_response["parameters"]
- )
+ assert transformed_message.tool_calls[0].function.arguments == json.dumps(tool_response["parameters"])
def test_transform_tool_call_with_cache_control():
@@ -250,12 +241,7 @@ def test_transform_tool_call_with_cache_control():
print(function_out_msg)
assert function_out_msg["toolSpec"]["name"] == "get_location"
assert function_out_msg["toolSpec"]["description"] == "Get the user's location"
- assert (
- function_out_msg["toolSpec"]["inputSchema"]["json"]["properties"]["location"][
- "type"
- ]
- == "string"
- )
+ assert function_out_msg["toolSpec"]["inputSchema"]["json"]["properties"]["location"]["type"] == "string"
transformed_cache_msg = result["toolConfig"]["tools"][1]
assert "cachePoint" in transformed_cache_msg
@@ -285,6 +271,7 @@ def test_reasoning_with_forced_tool_choice_switches_to_auto():
assert optional_params["tool_choice"] == {"auto": {}}
+
def test_get_supported_openai_params():
config = AmazonConverseConfig()
supported_params = config.get_supported_openai_params(
@@ -307,15 +294,13 @@ def test_get_supported_openai_params_bedrock_converse():
for model in litellm.BEDROCK_CONVERSE_MODELS:
print(f"Testing model: {model}")
config = AmazonConverseConfig()
- supported_params_without_prefix = config.get_supported_openai_params(
- model=model
- )
+ supported_params_without_prefix = config.get_supported_openai_params(model=model)
- supported_params_with_prefix = config.get_supported_openai_params(
- model=f"bedrock/converse/{model}"
- )
+ supported_params_with_prefix = config.get_supported_openai_params(model=f"bedrock/converse/{model}")
- assert set(supported_params_without_prefix) == set(supported_params_with_prefix), f"Supported params mismatch for model: {model}. Without prefix: {supported_params_without_prefix}, With prefix: {supported_params_with_prefix}"
+ assert set(supported_params_without_prefix) == set(supported_params_with_prefix), (
+ f"Supported params mismatch for model: {model}. Without prefix: {supported_params_without_prefix}, With prefix: {supported_params_with_prefix}"
+ )
print(f"✅ Passed for model: {model}")
@@ -382,7 +367,7 @@ def test_transform_response_with_computer_use_tool():
},
}
}
- ]
+ ],
}
},
"stopReason": "tool_use",
@@ -396,10 +381,12 @@ def test_transform_response_with_computer_use_tool():
"cacheWriteInputTokens": 0,
},
}
+
# Mock httpx.Response
class MockResponse:
def json(self):
return response_json
+
@property
def text(self):
return json.dumps(response_json)
@@ -468,12 +455,10 @@ def test_transform_response_with_bash_tool():
"toolUse": {
"toolUseId": "tooluse_456",
"name": "bash",
- "input": {
- "command": "ls -la *.py"
- },
+ "input": {"command": "ls -la *.py"},
}
}
- ]
+ ],
}
},
"stopReason": "tool_use",
@@ -487,10 +472,12 @@ def test_transform_response_with_bash_tool():
"cacheWriteInputTokens": 0,
},
}
+
# Mock httpx.Response
class MockResponse:
def json(self):
return response_json
+
@property
def text(self):
return json.dumps(response_json)
@@ -549,10 +536,11 @@ def test_transform_response_with_structured_response_being_called():
"name": "json_tool_call",
"input": {
"Current_Temperature": 62,
- "Weather_Explanation": "San Francisco typically has mild, cool weather year-round due to its coastal location and marine influence. The city is known for its fog, moderate temperatures, and relatively stable climate with little seasonal variation."},
+ "Weather_Explanation": "San Francisco typically has mild, cool weather year-round due to its coastal location and marine influence. The city is known for its fog, moderate temperatures, and relatively stable climate with little seasonal variation.",
+ },
}
}
- ]
+ ],
}
},
"stopReason": "tool_use",
@@ -566,10 +554,12 @@ def test_transform_response_with_structured_response_being_called():
"cacheWriteInputTokens": 0,
},
}
+
# Mock httpx.Response
class MockResponse:
def json(self):
return response_json
+
@property
def text(self):
return json.dumps(response_json)
@@ -580,49 +570,43 @@ def test_transform_response_with_structured_response_being_called():
"json_mode": True,
"tools": [
{
- 'type': 'function',
- 'function': {
- 'name': 'get_weather',
- 'description': 'Get the current weather in a given location',
- 'parameters': {
- 'type': 'object',
- 'properties': {
- 'location': {
- 'type': 'string',
- 'description': 'The city and state, e.g. San Francisco, CA'
- },
- 'unit': {
- 'type': 'string',
- 'enum': ['celsius', 'fahrenheit']
- }
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "description": "Get the current weather in a given location",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "location": {"type": "string", "description": "The city and state, e.g. San Francisco, CA"},
+ "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]},
},
- 'required': ['location']
- }
- }
+ "required": ["location"],
+ },
+ },
},
{
- 'type': 'function',
- 'function': {
- 'name': 'json_tool_call',
- 'parameters': {
- '$schema': 'http://json-schema.org/draft-07/schema#',
- 'type': 'object',
- 'required': ['Weather_Explanation', 'Current_Temperature'],
- 'properties': {
- 'Weather_Explanation': {
- 'type': ['string', 'null'],
- 'description': '1-2 sentences explaining the weather in the location'
+ "type": "function",
+ "function": {
+ "name": "json_tool_call",
+ "parameters": {
+ "$schema": "http://json-schema.org/draft-07/schema#",
+ "type": "object",
+ "required": ["Weather_Explanation", "Current_Temperature"],
+ "properties": {
+ "Weather_Explanation": {
+ "type": ["string", "null"],
+ "description": "1-2 sentences explaining the weather in the location",
+ },
+ "Current_Temperature": {
+ "type": ["number", "null"],
+ "description": "Current temperature in the location",
},
- 'Current_Temperature': {
- 'type': ['number', 'null'],
- 'description': 'Current temperature in the location'
- }
},
- 'additionalProperties': False
- }
- }
- }
- ]
+ "additionalProperties": False,
+ },
+ },
+ },
+ ],
}
# Call the transformation logic
result = config._transform_response(
@@ -641,7 +625,11 @@ def test_transform_response_with_structured_response_being_called():
assert result.choices[0].message.tool_calls is None
assert result.choices[0].message.content is not None
- assert result.choices[0].message.content == '{"Current_Temperature": 62, "Weather_Explanation": "San Francisco typically has mild, cool weather year-round due to its coastal location and marine influence. The city is known for its fog, moderate temperatures, and relatively stable climate with little seasonal variation."}'
+ assert (
+ result.choices[0].message.content
+ == '{"Current_Temperature": 62, "Weather_Explanation": "San Francisco typically has mild, cool weather year-round due to its coastal location and marine influence. The city is known for its fog, moderate temperatures, and relatively stable climate with little seasonal variation."}'
+ )
+
def test_transform_response_with_structured_response_calling_tool():
"""Test response transformation with structured response."""
@@ -650,28 +638,20 @@ def test_transform_response_with_structured_response_calling_tool():
# Simulate a Bedrock Converse response with a bash tool call
response_json = {
- "metrics": {
- "latencyMs": 1148
- },
+ "metrics": {"latencyMs": 1148},
"output": {
- "message":
- {
+ "message": {
"content": [
- {
- "text": "I\'ll check the current weather in San Francisco for you."
- },
+ {"text": "I'll check the current weather in San Francisco for you."},
{
"toolUse": {
- "input": {
- "location": "San Francisco, CA",
- "unit": "celsius"
- },
+ "input": {"location": "San Francisco, CA", "unit": "celsius"},
"name": "get_weather",
- "toolUseId": "tooluse_oKk__QrqSUmufMw3Q7vGaQ"
+ "toolUseId": "tooluse_oKk__QrqSUmufMw3Q7vGaQ",
}
- }
+ },
],
- "role": "assistant"
+ "role": "assistant",
}
},
"stopReason": "tool_use",
@@ -682,13 +662,15 @@ def test_transform_response_with_structured_response_calling_tool():
"cacheWriteInputTokens": 0,
"inputTokens": 534,
"outputTokens": 69,
- "totalTokens": 603
- }
+ "totalTokens": 603,
+ },
}
+
# Mock httpx.Response
class MockResponse:
def json(self):
return response_json
+
@property
def text(self):
return json.dumps(response_json)
@@ -699,49 +681,43 @@ def test_transform_response_with_structured_response_calling_tool():
"json_mode": True,
"tools": [
{
- 'type': 'function',
- 'function': {
- 'name': 'get_weather',
- 'description': 'Get the current weather in a given location',
- 'parameters': {
- 'type': 'object',
- 'properties': {
- 'location': {
- 'type': 'string',
- 'description': 'The city and state, e.g. San Francisco, CA'
- },
- 'unit': {
- 'type': 'string',
- 'enum': ['celsius', 'fahrenheit']
- }
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "description": "Get the current weather in a given location",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "location": {"type": "string", "description": "The city and state, e.g. San Francisco, CA"},
+ "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]},
},
- 'required': ['location']
- }
- }
+ "required": ["location"],
+ },
+ },
},
{
- 'type': 'function',
- 'function': {
- 'name': 'json_tool_call',
- 'parameters': {
- '$schema': 'http://json-schema.org/draft-07/schema#',
- 'type': 'object',
- 'required': ['Weather_Explanation', 'Current_Temperature'],
- 'properties': {
- 'Weather_Explanation': {
- 'type': ['string', 'null'],
- 'description': '1-2 sentences explaining the weather in the location'
+ "type": "function",
+ "function": {
+ "name": "json_tool_call",
+ "parameters": {
+ "$schema": "http://json-schema.org/draft-07/schema#",
+ "type": "object",
+ "required": ["Weather_Explanation", "Current_Temperature"],
+ "properties": {
+ "Weather_Explanation": {
+ "type": ["string", "null"],
+ "description": "1-2 sentences explaining the weather in the location",
+ },
+ "Current_Temperature": {
+ "type": ["number", "null"],
+ "description": "Current temperature in the location",
},
- 'Current_Temperature': {
- 'type': ['number', 'null'],
- 'description': 'Current temperature in the location'
- }
},
- 'additionalProperties': False
- }
- }
- }
- ]
+ "additionalProperties": False,
+ },
+ },
+ },
+ ],
}
# Call the transformation logic
result = config._transform_response(
@@ -760,7 +736,10 @@ def test_transform_response_with_structured_response_calling_tool():
assert result.choices[0].message.tool_calls is not None
assert len(result.choices[0].message.tool_calls) == 1
assert result.choices[0].message.tool_calls[0].function.name == "get_weather"
- assert result.choices[0].message.tool_calls[0].function.arguments == '{"location": "San Francisco, CA", "unit": "celsius"}'
+ assert (
+ result.choices[0].message.tool_calls[0].function.arguments
+ == '{"location": "San Francisco, CA", "unit": "celsius"}'
+ )
@pytest.mark.asyncio
@@ -775,12 +754,7 @@ async def test_bedrock_bash_tool_acompletion():
}
]
- messages = [
- {
- "role": "user",
- "content": "run ls command and find all python files"
- }
- ]
+ messages = [{"role": "user", "content": "run ls command and find all python files"}]
try:
response = await litellm.acompletion(
@@ -788,7 +762,7 @@ async def test_bedrock_bash_tool_acompletion():
messages=messages,
tools=tools,
# Using dummy API key - test should fail with auth error, proving request formatting works
- api_key="dummy-key-for-testing"
+ api_key="dummy-key-for-testing",
)
# If we get here, something's wrong - we expect an auth error
assert False, "Expected authentication error but got successful response"
@@ -797,8 +771,16 @@ async def test_bedrock_bash_tool_acompletion():
# Check if it's an expected authentication/credentials error
auth_error_indicators = [
- "credentials", "authentication", "unauthorized", "access denied",
- "aws", "region", "profile", "token", "invalid", "signature"
+ "credentials",
+ "authentication",
+ "unauthorized",
+ "access denied",
+ "aws",
+ "region",
+ "profile",
+ "token",
+ "invalid",
+ "signature",
]
if any(auth_error in error_str for auth_error in auth_error_indicators):
@@ -828,17 +810,14 @@ async def test_bedrock_computer_use_acompletion():
{
"role": "user",
"content": [
- {
- "type": "text",
- "text": "Go to the bedrock console"
- },
+ {"type": "text", "text": "Go to the bedrock console"},
{
"type": "image_url",
"image_url": {
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg=="
- }
- }
- ]
+ },
+ },
+ ],
}
]
@@ -848,7 +827,7 @@ async def test_bedrock_computer_use_acompletion():
messages=messages,
tools=tools,
# Using dummy API key - test should fail with auth error, proving request formatting works
- api_key="dummy-key-for-testing"
+ api_key="dummy-key-for-testing",
)
# If we get here, something's wrong - we expect an auth error
assert False, "Expected authentication error but got successful response"
@@ -857,8 +836,16 @@ async def test_bedrock_computer_use_acompletion():
# Check if it's an expected authentication/credentials error
auth_error_indicators = [
- "credentials", "authentication", "unauthorized", "access denied",
- "aws", "region", "profile", "token", "invalid", "signature"
+ "credentials",
+ "authentication",
+ "unauthorized",
+ "access denied",
+ "aws",
+ "region",
+ "profile",
+ "token",
+ "invalid",
+ "signature",
]
if any(auth_error in error_str for auth_error in auth_error_indicators):
@@ -886,15 +873,10 @@ async def test_transformation_directly():
{
"type": "bash_20241022",
"name": "bash",
- }
+ },
]
- messages = [
- {
- "role": "user",
- "content": "run ls command and find all python files"
- }
- ]
+ messages = [{"role": "user", "content": "run ls command and find all python files"}]
# Transform request
request_data = config.transform_request(
@@ -902,7 +884,7 @@ async def test_transformation_directly():
messages=messages,
optional_params={"tools": tools},
litellm_params={},
- headers={}
+ headers={},
)
# Verify the structure
@@ -994,16 +976,11 @@ def test_transform_request_with_multiple_tools():
},
"required": ["location"],
},
- }
- }
+ },
+ },
]
- messages = [
- {
- "role": "user",
- "content": "run ls command and find all python files"
- }
- ]
+ messages = [{"role": "user", "content": "run ls command and find all python files"}]
# Transform request
request_data = config.transform_request(
@@ -1011,7 +988,7 @@ def test_transform_request_with_multiple_tools():
messages=messages,
optional_params={"tools": tools},
litellm_params={},
- headers={}
+ headers={},
)
# Verify the structure
@@ -1054,17 +1031,14 @@ def test_transform_request_with_computer_tool_only():
{
"role": "user",
"content": [
- {
- "type": "text",
- "text": "Go to the bedrock console"
- },
+ {"type": "text", "text": "Go to the bedrock console"},
{
"type": "image_url",
"image_url": {
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg=="
- }
- }
- ]
+ },
+ },
+ ],
}
]
@@ -1074,7 +1048,7 @@ def test_transform_request_with_computer_tool_only():
messages=messages,
optional_params={"tools": tools},
litellm_params={},
- headers={}
+ headers={},
)
# Verify the structure
@@ -1102,12 +1076,7 @@ def test_transform_request_with_bash_tool_only():
}
]
- messages = [
- {
- "role": "user",
- "content": "run ls command and find all python files"
- }
- ]
+ messages = [{"role": "user", "content": "run ls command and find all python files"}]
# Transform request
request_data = config.transform_request(
@@ -1115,7 +1084,7 @@ def test_transform_request_with_bash_tool_only():
messages=messages,
optional_params={"tools": tools},
litellm_params={},
- headers={}
+ headers={},
)
# Verify the structure
@@ -1143,12 +1112,7 @@ def test_transform_request_with_text_editor_tool():
}
]
- messages = [
- {
- "role": "user",
- "content": "Edit this text file"
- }
- ]
+ messages = [{"role": "user", "content": "Edit this text file"}]
# Transform request
request_data = config.transform_request(
@@ -1156,7 +1120,7 @@ def test_transform_request_with_text_editor_tool():
messages=messages,
optional_params={"tools": tools},
litellm_params={},
- headers={}
+ headers={},
)
# Verify the structure
@@ -1194,16 +1158,11 @@ def test_transform_request_with_function_tool():
},
"required": ["location"],
},
- }
+ },
}
]
- messages = [
- {
- "role": "user",
- "content": "What's the weather like in San Francisco?"
- }
- ]
+ messages = [{"role": "user", "content": "What's the weather like in San Francisco?"}]
# Transform request
request_data = config.transform_request(
@@ -1211,7 +1170,7 @@ def test_transform_request_with_function_tool():
messages=messages,
optional_params={"tools": tools},
litellm_params={},
- headers={}
+ headers={},
)
# Verify the structure
@@ -1247,7 +1206,7 @@ def test_map_openai_params_with_response_format():
},
"required": ["location"],
},
- }
+ },
}
]
@@ -1279,7 +1238,7 @@ def test_map_openai_params_with_response_format():
non_default_params={"response_format": json_schema},
optional_params={"tools": tools},
model="eu.anthropic.claude-sonnet-4-20250514-v1:0",
- drop_params=False
+ drop_params=False,
)
assert "tools" in optional_params
@@ -1299,31 +1258,21 @@ async def test_assistant_message_cache_control():
# Test assistant message with string content and cache_control
messages = [
{"role": "user", "content": "Hello"},
- {
- "role": "assistant",
- "content": "Hi there!",
- "cache_control": {"type": "ephemeral"}
- }
+ {"role": "assistant", "content": "Hi there!", "cache_control": {"type": "ephemeral"}},
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
assert result == async_result
async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
assert result == async_result
@@ -1353,26 +1302,16 @@ async def test_assistant_message_list_content_cache_control():
{"role": "user", "content": "Hello"},
{
"role": "assistant",
- "content": [
- {
- "type": "text",
- "text": "This should be cached",
- "cache_control": {"type": "ephemeral"}
- }
- ]
- }
+ "content": [{"type": "text", "text": "This should be cached", "cache_control": {"type": "ephemeral"}}],
+ },
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
assert result == async_result
@@ -1399,36 +1338,22 @@ async def test_tool_message_cache_control():
"role": "assistant",
"content": None,
"tool_calls": [
- {
- "id": "call_123",
- "type": "function",
- "function": {"name": "get_weather", "arguments": "{}"}
- }
- ]
+ {"id": "call_123", "type": "function", "function": {"name": "get_weather", "arguments": "{}"}}
+ ],
},
{
"role": "tool",
"tool_call_id": "call_123",
- "content": [
- {
- "type": "text",
- "text": "Weather data: sunny, 25°C",
- "cache_control": {"type": "ephemeral"}
- }
- ]
- }
+ "content": [{"type": "text", "text": "Weather data: sunny, 25°C", "cache_control": {"type": "ephemeral"}}],
+ },
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
assert result == async_result
@@ -1463,31 +1388,23 @@ async def test_tool_message_string_content_cache_control():
"role": "assistant",
"content": None,
"tool_calls": [
- {
- "id": "call_123",
- "type": "function",
- "function": {"name": "get_weather", "arguments": "{}"}
- }
- ]
+ {"id": "call_123", "type": "function", "function": {"name": "get_weather", "arguments": "{}"}}
+ ],
},
{
"role": "tool",
"tool_call_id": "call_123",
"content": "Weather: sunny, 25°C",
- "cache_control": {"type": "ephemeral"}
- }
+ "cache_control": {"type": "ephemeral"},
+ },
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
assert result == async_result
@@ -1523,22 +1440,18 @@ async def test_assistant_tool_calls_cache_control():
"id": "call_proxy_123",
"type": "function",
"function": {"name": "calc", "arguments": "{}"},
- "cache_control": {"type": "ephemeral"}
+ "cache_control": {"type": "ephemeral"},
}
- ]
- }
+ ],
+ },
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
assert result == async_result
@@ -1575,28 +1488,24 @@ async def test_multiple_tool_calls_with_mixed_cache_control():
"id": "call_1",
"type": "function",
"function": {"name": "calc", "arguments": '{"expr": "2+2"}'},
- "cache_control": {"type": "ephemeral"}
+ "cache_control": {"type": "ephemeral"},
},
{
"id": "call_2",
"type": "function",
- "function": {"name": "calc", "arguments": '{"expr": "3+3"}'}
+ "function": {"name": "calc", "arguments": '{"expr": "3+3"}'},
# No cache_control
- }
- ]
- }
+ },
+ ],
+ },
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
assert result == async_result
@@ -1632,20 +1541,16 @@ async def test_no_cache_control_no_cache_point():
{
"role": "tool",
"tool_call_id": "call_123",
- "content": "Tool result" # No cache_control
- }
+ "content": "Tool result", # No cache_control
+ },
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
- messages=messages,
- model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
assert result == async_result
@@ -1665,6 +1570,7 @@ async def test_no_cache_control_no_cache_point():
# Guarded Text Feature Tests
# ============================================================================
+
def test_guarded_text_wraps_in_guardrail_converse_content():
"""Test that guarded_text content type gets wrapped in guardContent blocks."""
from litellm.litellm_core_utils.prompt_templates.factory import (
@@ -1677,15 +1583,13 @@ def test_guarded_text_wraps_in_guardrail_converse_content():
"content": [
{"type": "text", "text": "Regular text content"},
{"type": "guarded_text", "text": "This should be guarded"},
- {"type": "text", "text": "More regular text"}
- ]
+ {"type": "text", "text": "More regular text"},
+ ],
}
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="us.amazon.nova-pro-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="us.amazon.nova-pro-v1:0", llm_provider="bedrock_converse"
)
# Should have 1 message
@@ -1705,6 +1609,7 @@ def test_guarded_text_wraps_in_guardrail_converse_content():
assert "guardContent" in content[1]
assert content[1]["guardContent"]["text"]["text"] == "This should be guarded"
+
def test_guarded_text_with_system_messages():
"""Test guarded_text with system messages using the full transformation."""
config = AmazonConverseConfig()
@@ -1715,24 +1620,22 @@ def test_guarded_text_with_system_messages():
"role": "user",
"content": [
{"type": "text", "text": "What is the main topic of this legal document?"},
- {"type": "guarded_text", "text": "This is a set of very long instructions that you will follow. Here is a legal document that you will use to answer the user's question."}
- ]
- }
+ {
+ "type": "guarded_text",
+ "text": "This is a set of very long instructions that you will follow. Here is a legal document that you will use to answer the user's question.",
+ },
+ ],
+ },
]
- optional_params = {
- "guardrailConfig": {
- "guardrailIdentifier": "gr-abc123",
- "guardrailVersion": "DRAFT"
- }
- }
+ optional_params = {"guardrailConfig": {"guardrailIdentifier": "gr-abc123", "guardrailVersion": "DRAFT"}}
result = config._transform_request(
model="us.amazon.nova-pro-v1:0",
messages=messages,
optional_params=optional_params,
litellm_params={},
- headers={}
+ headers={},
)
# Should have system content blocks
@@ -1755,7 +1658,10 @@ def test_guarded_text_with_system_messages():
assert content[0]["text"] == "What is the main topic of this legal document?"
# Second should be guardContent
assert "guardContent" in content[1]
- assert content[1]["guardContent"]["text"]["text"] == "This is a set of very long instructions that you will follow. Here is a legal document that you will use to answer the user's question."
+ assert (
+ content[1]["guardContent"]["text"]["text"]
+ == "This is a set of very long instructions that you will follow. Here is a legal document that you will use to answer the user's question."
+ )
def test_guarded_text_with_mixed_content_types():
@@ -1770,15 +1676,13 @@ def test_guarded_text_with_mixed_content_types():
"content": [
{"type": "text", "text": "Look at this image"},
{"type": "image_url", "image_url": {"url": "data:image/png;base64,test"}},
- {"type": "guarded_text", "text": "This sensitive content should be guarded"}
- ]
+ {"type": "guarded_text", "text": "This sensitive content should be guarded"},
+ ],
}
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="us.amazon.nova-pro-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="us.amazon.nova-pro-v1:0", llm_provider="bedrock_converse"
)
# Should have 1 message
@@ -1800,6 +1704,7 @@ def test_guarded_text_with_mixed_content_types():
assert "guardContent" in content[2]
assert content[2]["guardContent"]["text"]["text"] == "This sensitive content should be guarded"
+
@pytest.mark.asyncio
async def test_async_guarded_text():
"""Test async version of guarded_text processing."""
@@ -1810,17 +1715,12 @@ async def test_async_guarded_text():
messages = [
{
"role": "user",
- "content": [
- {"type": "text", "text": "Hello"},
- {"type": "guarded_text", "text": "This should be guarded"}
- ]
+ "content": [{"type": "text", "text": "Hello"}, {"type": "guarded_text", "text": "This should be guarded"}],
}
]
result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
- messages=messages,
- model="us.amazon.nova-pro-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="us.amazon.nova-pro-v1:0", llm_provider="bedrock_converse"
)
# Should have 1 message
@@ -1851,31 +1751,21 @@ def test_guarded_text_with_tool_calls():
"role": "user",
"content": [
{"type": "text", "text": "What's the weather?"},
- {"type": "guarded_text", "text": "Please be careful with sensitive information"}
- ]
+ {"type": "guarded_text", "text": "Please be careful with sensitive information"},
+ ],
},
{
"role": "assistant",
"content": None,
"tool_calls": [
- {
- "id": "call_123",
- "type": "function",
- "function": {"name": "get_weather", "arguments": "{}"}
- }
- ]
+ {"id": "call_123", "type": "function", "function": {"name": "get_weather", "arguments": "{}"}}
+ ],
},
- {
- "role": "tool",
- "tool_call_id": "call_123",
- "content": "It's sunny and 25°C"
- }
+ {"role": "tool", "tool_call_id": "call_123", "content": "It's sunny and 25°C"},
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="us.amazon.nova-pro-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="us.amazon.nova-pro-v1:0", llm_provider="bedrock_converse"
)
# Should have 3 messages
@@ -1909,26 +1799,18 @@ def test_guarded_text_guardrail_config_preserved():
messages = [
{
"role": "user",
- "content": [
- {"type": "text", "text": "Hello"},
- {"type": "guarded_text", "text": "This should be guarded"}
- ]
+ "content": [{"type": "text", "text": "Hello"}, {"type": "guarded_text", "text": "This should be guarded"}],
}
]
- optional_params = {
- "guardrailConfig": {
- "guardrailIdentifier": "gr-abc123",
- "guardrailVersion": "DRAFT"
- }
- }
+ optional_params = {"guardrailConfig": {"guardrailIdentifier": "gr-abc123", "guardrailVersion": "DRAFT"}}
result = config._transform_request(
model="us.amazon.nova-pro-v1:0",
messages=messages,
optional_params=optional_params,
litellm_params={},
- headers={}
+ headers={},
)
# GuardrailConfig should be present at top level
@@ -1946,23 +1828,10 @@ def test_auto_convert_last_user_message_to_guarded_text():
config = AmazonConverseConfig()
messages = [
- {
- "role": "user",
- "content": [
- {
- "type": "text",
- "text": "What is the main topic of this legal document?"
- }
- ]
- }
+ {"role": "user", "content": [{"type": "text", "text": "What is the main topic of this legal document?"}]}
]
- optional_params = {
- "guardrailConfig": {
- "guardrailIdentifier": "gr-abc123",
- "guardrailVersion": "1"
- }
- }
+ optional_params = {"guardrailConfig": {"guardrailIdentifier": "gr-abc123", "guardrailVersion": "1"}}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
@@ -1979,19 +1848,9 @@ def test_auto_convert_last_user_message_string_content():
"""Test that last user message with string content is automatically converted to guarded_text when guardrailConfig is present."""
config = AmazonConverseConfig()
- messages = [
- {
- "role": "user",
- "content": "What is the main topic of this legal document?"
- }
- ]
+ messages = [{"role": "user", "content": "What is the main topic of this legal document?"}]
- optional_params = {
- "guardrailConfig": {
- "guardrailIdentifier": "gr-abc123",
- "guardrailVersion": "1"
- }
- }
+ optional_params = {"guardrailConfig": {"guardrailIdentifier": "gr-abc123", "guardrailVersion": "1"}}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
@@ -2009,15 +1868,7 @@ def test_no_conversion_when_no_guardrail_config():
config = AmazonConverseConfig()
messages = [
- {
- "role": "user",
- "content": [
- {
- "type": "text",
- "text": "What is the main topic of this legal document?"
- }
- ]
- }
+ {"role": "user", "content": [{"type": "text", "text": "What is the main topic of this legal document?"}]}
]
optional_params = {}
@@ -2033,24 +1884,9 @@ def test_no_conversion_when_guarded_text_already_present():
"""Test that no conversion happens when guarded_text is already present in the last user message."""
config = AmazonConverseConfig()
- messages = [
- {
- "role": "user",
- "content": [
- {
- "type": "guarded_text",
- "text": "This is already guarded"
- }
- ]
- }
- ]
+ messages = [{"role": "user", "content": [{"type": "guarded_text", "text": "This is already guarded"}]}]
- optional_params = {
- "guardrailConfig": {
- "guardrailIdentifier": "gr-abc123",
- "guardrailVersion": "1"
- }
- }
+ optional_params = {"guardrailConfig": {"guardrailIdentifier": "gr-abc123", "guardrailVersion": "1"}}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
@@ -2067,24 +1903,13 @@ def test_auto_convert_with_mixed_content():
{
"role": "user",
"content": [
- {
- "type": "text",
- "text": "What is the main topic of this legal document?"
- },
- {
- "type": "image_url",
- "image_url": {"url": "https://example.com/image.jpg"}
- }
- ]
+ {"type": "text", "text": "What is the main topic of this legal document?"},
+ {"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}},
+ ],
}
]
- optional_params = {
- "guardrailConfig": {
- "guardrailIdentifier": "gr-abc123",
- "guardrailVersion": "1"
- }
- }
+ optional_params = {"guardrailConfig": {"guardrailIdentifier": "gr-abc123", "guardrailVersion": "1"}}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
@@ -2108,23 +1933,10 @@ def test_auto_convert_in_full_transformation():
config = AmazonConverseConfig()
messages = [
- {
- "role": "user",
- "content": [
- {
- "type": "text",
- "text": "What is the main topic of this legal document?"
- }
- ]
- }
+ {"role": "user", "content": [{"type": "text", "text": "What is the main topic of this legal document?"}]}
]
- optional_params = {
- "guardrailConfig": {
- "guardrailIdentifier": "gr-abc123",
- "guardrailVersion": "1"
- }
- }
+ optional_params = {"guardrailConfig": {"guardrailIdentifier": "gr-abc123", "guardrailVersion": "1"}}
# Test the full transformation
result = config._transform_request(
@@ -2132,7 +1944,7 @@ def test_auto_convert_in_full_transformation():
messages=messages,
optional_params=optional_params,
litellm_params={},
- headers={}
+ headers={},
)
# Verify the transformation worked
@@ -2152,45 +1964,13 @@ def test_convert_consecutive_user_messages_to_guarded_text():
config = AmazonConverseConfig()
messages = [
- {
- "role": "user",
- "content": [
- {
- "type": "text",
- "text": "First user message"
- }
- ]
- },
- {
- "role": "assistant",
- "content": "Assistant response"
- },
- {
- "role": "user",
- "content": [
- {
- "type": "text",
- "text": "Second user message"
- }
- ]
- },
- {
- "role": "user",
- "content": [
- {
- "type": "text",
- "text": "Third user message"
- }
- ]
- }
+ {"role": "user", "content": [{"type": "text", "text": "First user message"}]},
+ {"role": "assistant", "content": "Assistant response"},
+ {"role": "user", "content": [{"type": "text", "text": "Second user message"}]},
+ {"role": "user", "content": [{"type": "text", "text": "Third user message"}]},
]
- optional_params = {
- "guardrailConfig": {
- "guardrailIdentifier": "gr-abc123",
- "guardrailVersion": "1"
- }
- }
+ optional_params = {"guardrailConfig": {"guardrailIdentifier": "gr-abc123", "guardrailVersion": "1"}}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
@@ -2223,41 +2003,12 @@ def test_convert_all_user_messages_when_all_consecutive():
config = AmazonConverseConfig()
messages = [
- {
- "role": "user",
- "content": [
- {
- "type": "text",
- "text": "First user message"
- }
- ]
- },
- {
- "role": "user",
- "content": [
- {
- "type": "text",
- "text": "Second user message"
- }
- ]
- },
- {
- "role": "user",
- "content": [
- {
- "type": "text",
- "text": "Third user message"
- }
- ]
- }
+ {"role": "user", "content": [{"type": "text", "text": "First user message"}]},
+ {"role": "user", "content": [{"type": "text", "text": "Second user message"}]},
+ {"role": "user", "content": [{"type": "text", "text": "Third user message"}]},
]
- optional_params = {
- "guardrailConfig": {
- "guardrailIdentifier": "gr-abc123",
- "guardrailVersion": "1"
- }
- }
+ optional_params = {"guardrailConfig": {"guardrailIdentifier": "gr-abc123", "guardrailVersion": "1"}}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
@@ -2279,26 +2030,12 @@ def test_convert_consecutive_user_messages_with_string_content():
config = AmazonConverseConfig()
messages = [
- {
- "role": "assistant",
- "content": "Assistant response"
- },
- {
- "role": "user",
- "content": "First user message"
- },
- {
- "role": "user",
- "content": "Second user message"
- }
+ {"role": "assistant", "content": "Assistant response"},
+ {"role": "user", "content": "First user message"},
+ {"role": "user", "content": "Second user message"},
]
- optional_params = {
- "guardrailConfig": {
- "guardrailIdentifier": "gr-abc123",
- "guardrailVersion": "1"
- }
- }
+ optional_params = {"guardrailConfig": {"guardrailIdentifier": "gr-abc123", "guardrailVersion": "1"}}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
@@ -2327,32 +2064,11 @@ def test_skip_consecutive_user_messages_with_existing_guarded_text():
config = AmazonConverseConfig()
messages = [
- {
- "role": "user",
- "content": [
- {
- "type": "guarded_text",
- "text": "Already guarded"
- }
- ]
- },
- {
- "role": "user",
- "content": [
- {
- "type": "text",
- "text": "Should be converted"
- }
- ]
- }
+ {"role": "user", "content": [{"type": "guarded_text", "text": "Already guarded"}]},
+ {"role": "user", "content": [{"type": "text", "text": "Should be converted"}]},
]
- optional_params = {
- "guardrailConfig": {
- "guardrailIdentifier": "gr-abc123",
- "guardrailVersion": "1"
- }
- }
+ optional_params = {"guardrailConfig": {"guardrailIdentifier": "gr-abc123", "guardrailVersion": "1"}}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
@@ -2384,11 +2100,7 @@ def test_request_metadata_transformation():
"""Test that requestMetadata is properly transformed to top-level field."""
config = AmazonConverseConfig()
- request_metadata = {
- "cost_center": "engineering",
- "user_id": "user123",
- "session_id": "sess_abc123"
- }
+ request_metadata = {"cost_center": "engineering", "user_id": "user123", "session_id": "sess_abc123"}
messages = [
{"role": "user", "content": "Hello!"},
@@ -2400,7 +2112,7 @@ def test_request_metadata_transformation():
messages=messages,
optional_params={"requestMetadata": request_metadata},
litellm_params={},
- headers={}
+ headers={},
)
# Verify that requestMetadata appears as top-level field
@@ -2426,7 +2138,7 @@ def test_request_metadata_validation():
messages=messages,
optional_params={"requestMetadata": valid_metadata},
litellm_params={},
- headers={}
+ headers={},
)
# Test too many items (max 16)
@@ -2438,7 +2150,7 @@ def test_request_metadata_validation():
messages=messages,
optional_params={"requestMetadata": too_many_items},
litellm_params={},
- headers={}
+ headers={},
)
assert False, "Should have raised validation error for too many items"
except Exception as e:
@@ -2461,7 +2173,7 @@ def test_request_metadata_key_constraints():
messages=messages,
optional_params={"requestMetadata": invalid_metadata},
litellm_params={},
- headers={}
+ headers={},
)
assert False, "Should have raised validation error for key too long"
except Exception as e:
@@ -2476,7 +2188,7 @@ def test_request_metadata_key_constraints():
messages=messages,
optional_params={"requestMetadata": invalid_metadata},
litellm_params={},
- headers={}
+ headers={},
)
assert False, "Should have raised validation error for empty key"
except Exception as e:
@@ -2499,7 +2211,7 @@ def test_request_metadata_value_constraints():
messages=messages,
optional_params={"requestMetadata": invalid_metadata},
litellm_params={},
- headers={}
+ headers={},
)
assert False, "Should have raised validation error for value too long"
except Exception as e:
@@ -2514,7 +2226,7 @@ def test_request_metadata_value_constraints():
messages=messages,
optional_params={"requestMetadata": valid_metadata},
litellm_params={},
- headers={}
+ headers={},
)
@@ -2537,7 +2249,7 @@ def test_request_metadata_character_pattern():
messages=messages,
optional_params={"requestMetadata": valid_metadata},
litellm_params={},
- headers={}
+ headers={},
)
@@ -2545,10 +2257,7 @@ def test_request_metadata_with_other_params():
"""Test that requestMetadata works alongside other parameters."""
config = AmazonConverseConfig()
- request_metadata = {
- "experiment": "test_A",
- "user_type": "premium"
- }
+ request_metadata = {"experiment": "test_A", "user_type": "premium"}
messages = [
{"role": "user", "content": "What's the weather?"},
@@ -2562,12 +2271,10 @@ def test_request_metadata_with_other_params():
"description": "Get the current weather",
"parameters": {
"type": "object",
- "properties": {
- "location": {"type": "string"}
- },
- "required": ["location"]
- }
- }
+ "properties": {"location": {"type": "string"}},
+ "required": ["location"],
+ },
+ },
}
]
@@ -2575,14 +2282,9 @@ def test_request_metadata_with_other_params():
request_data = config.transform_request(
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
messages=messages,
- optional_params={
- "requestMetadata": request_metadata,
- "tools": tools,
- "max_tokens": 100,
- "temperature": 0.7
- },
+ optional_params={"requestMetadata": request_metadata, "tools": tools, "max_tokens": 100, "temperature": 0.7},
litellm_params={},
- headers={}
+ headers={},
)
# Verify requestMetadata is at top level
@@ -2607,7 +2309,7 @@ def test_request_metadata_empty():
messages=messages,
optional_params={"requestMetadata": {}},
litellm_params={},
- headers={}
+ headers={},
)
assert "requestMetadata" in request_data
@@ -2626,7 +2328,7 @@ def test_request_metadata_not_provided():
messages=messages,
optional_params={},
litellm_params={},
- headers={}
+ headers={},
)
# requestMetadata should not be in the request
@@ -2649,16 +2351,14 @@ def test_empty_assistant_message_handling():
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": ""}, # Empty content
- {"role": "user", "content": "How are you?"}
+ {"role": "user", "content": "How are you?"},
]
# Use patch to ensure we modify the litellm reference that factory.py actually uses
# This avoids issues with module reloading during parallel test execution
with patch.object(factory_module.litellm, "modify_params", True):
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
# Should have 3 messages: user, assistant (with placeholder), user
@@ -2676,13 +2376,11 @@ def test_empty_assistant_message_handling():
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": " "}, # Whitespace-only content
- {"role": "user", "content": "How are you?"}
+ {"role": "user", "content": "How are you?"},
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
# Assistant message should have placeholder text instead of whitespace
@@ -2693,13 +2391,11 @@ def test_empty_assistant_message_handling():
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": [{"type": "text", "text": ""}]}, # Empty text in list
- {"role": "user", "content": "How are you?"}
+ {"role": "user", "content": "How are you?"},
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
# Assistant message should have placeholder text instead of empty text
@@ -2710,13 +2406,11 @@ def test_empty_assistant_message_handling():
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "I'm doing well, thank you!"}, # Normal content
- {"role": "user", "content": "How are you?"}
+ {"role": "user", "content": "How are you?"},
]
result = _bedrock_converse_messages_pt(
- messages=messages,
- model="anthropic.claude-3-5-sonnet-20240620-v1:0",
- llm_provider="bedrock_converse"
+ messages=messages, model="anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse"
)
# Assistant message should keep original content
@@ -2828,6 +2522,7 @@ def test_thinking_with_max_completion_tokens():
assert result["thinking"]["type"] == "enabled"
assert result["thinking"]["budget_tokens"] == 5000
+
def test_drop_thinking_param_when_thinking_blocks_missing():
"""
Test that thinking param is dropped when modify_params=True and
@@ -2868,24 +2563,21 @@ def test_drop_thinking_param_when_thinking_blocks_missing():
optional_params = {"thinking": {"type": "enabled", "budget_tokens": 1000}}
# Verify the condition is detected
- assert last_assistant_with_tool_calls_has_no_thinking_blocks(
- messages_without_thinking_blocks
- ), "Should detect missing thinking_blocks"
+ assert last_assistant_with_tool_calls_has_no_thinking_blocks(messages_without_thinking_blocks), (
+ "Should detect missing thinking_blocks"
+ )
# Simulate what _transform_request_helper does
if (
optional_params.get("thinking") is not None
and messages_without_thinking_blocks is not None
- and last_assistant_with_tool_calls_has_no_thinking_blocks(
- messages_without_thinking_blocks
- )
+ and last_assistant_with_tool_calls_has_no_thinking_blocks(messages_without_thinking_blocks)
):
if litellm.modify_params:
optional_params.pop("thinking", None)
assert "thinking" not in optional_params, (
- "thinking param should be dropped when modify_params=True "
- "and thinking_blocks are missing"
+ "thinking param should be dropped when modify_params=True and thinking_blocks are missing"
)
# Test case 2: thinking should NOT be dropped when thinking_blocks are present
@@ -2901,29 +2593,23 @@ def test_drop_thinking_param_when_thinking_blocks_missing():
"function": {"name": "search", "arguments": "{}"},
}
],
- "thinking_blocks": [
- {"type": "thinking", "thinking": "Let me search for weather..."}
- ],
+ "thinking_blocks": [{"type": "thinking", "thinking": "Let me search for weather..."}],
},
{"role": "tool", "content": "Weather is sunny", "tool_call_id": "call_123"},
]
- optional_params_with_thinking = {
- "thinking": {"type": "enabled", "budget_tokens": 1000}
- }
+ optional_params_with_thinking = {"thinking": {"type": "enabled", "budget_tokens": 1000}}
# Verify the condition is NOT detected when thinking_blocks are present
- assert not last_assistant_with_tool_calls_has_no_thinking_blocks(
- messages_with_thinking_blocks
- ), "Should NOT detect missing thinking_blocks when they are present"
+ assert not last_assistant_with_tool_calls_has_no_thinking_blocks(messages_with_thinking_blocks), (
+ "Should NOT detect missing thinking_blocks when they are present"
+ )
# Simulate what _transform_request_helper does
if (
optional_params_with_thinking.get("thinking") is not None
and messages_with_thinking_blocks is not None
- and last_assistant_with_tool_calls_has_no_thinking_blocks(
- messages_with_thinking_blocks
- )
+ and last_assistant_with_tool_calls_has_no_thinking_blocks(messages_with_thinking_blocks)
):
if litellm.modify_params:
optional_params_with_thinking.pop("thinking", None)
@@ -2935,24 +2621,18 @@ def test_drop_thinking_param_when_thinking_blocks_missing():
# Test case 3: thinking should NOT be dropped when modify_params=False
litellm.modify_params = False
- optional_params_no_modify = {
- "thinking": {"type": "enabled", "budget_tokens": 1000}
- }
+ optional_params_no_modify = {"thinking": {"type": "enabled", "budget_tokens": 1000}}
# Simulate what _transform_request_helper does
if (
optional_params_no_modify.get("thinking") is not None
and messages_without_thinking_blocks is not None
- and last_assistant_with_tool_calls_has_no_thinking_blocks(
- messages_without_thinking_blocks
- )
+ and last_assistant_with_tool_calls_has_no_thinking_blocks(messages_without_thinking_blocks)
):
if litellm.modify_params:
optional_params_no_modify.pop("thinking", None)
- assert "thinking" in optional_params_no_modify, (
- "thinking param should NOT be dropped when modify_params=False"
- )
+ assert "thinking" in optional_params_no_modify, "thinking param should NOT be dropped when modify_params=False"
finally:
# Restore original modify_params setting
@@ -2960,46 +2640,53 @@ def test_drop_thinking_param_when_thinking_blocks_missing():
def test_supports_native_structured_outputs():
- """Test model detection for native structured outputs support."""
- config = AmazonConverseConfig()
+ """Test model detection for native structured outputs support.
- # Supported models
- assert config._supports_native_structured_outputs(
- "anthropic.claude-sonnet-4-5-20250929-v1:0"
- )
- assert config._supports_native_structured_outputs(
- "anthropic.claude-haiku-4-5-20251001-v1:0"
- )
- assert config._supports_native_structured_outputs(
- "anthropic.claude-opus-4-6-v1:0"
- )
- assert config._supports_native_structured_outputs(
- "eu.anthropic.claude-opus-4-5-20260101-v1:0"
- )
- assert config._supports_native_structured_outputs("qwen.qwen3-235b-instruct-v1:0")
- assert config._supports_native_structured_outputs("mistral.mistral-large-3-v1:0")
- assert config._supports_native_structured_outputs("deepseek.deepseek-v3.1-v1:0")
+ Support is driven by the ``supports_native_structured_output`` flag in the
+ cost JSON (litellm.model_cost), not a hardcoded model set.
+ """
+ old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
+ old_cost = litellm.model_cost
+ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
+ litellm.model_cost = litellm.get_model_cost_map(url="")
+ try:
+ config = AmazonConverseConfig()
- # Unsupported models — should fall back to tool-call approach
- assert not config._supports_native_structured_outputs(
- "anthropic.claude-3-5-sonnet-20241022-v2:0"
- )
- assert not config._supports_native_structured_outputs(
- "anthropic.claude-sonnet-4-20250514-v1:0"
- )
- assert not config._supports_native_structured_outputs(
- "meta.llama3-3-70b-instruct-v1:0"
- )
- assert not config._supports_native_structured_outputs(
- "amazon.nova-pro-v1:0"
- )
- # Excluded despite AWS listing them: broken constrained decoding on Bedrock
- assert not config._supports_native_structured_outputs(
- "openai.gpt-oss-120b-1:0"
- )
- assert not config._supports_native_structured_outputs(
- "mistral.magistral-small-2509"
- )
+ # Supported models (have supports_native_structured_output=true in cost JSON)
+ assert config._supports_native_structured_outputs("anthropic.claude-sonnet-4-5-20250929-v1:0")
+ assert config._supports_native_structured_outputs("anthropic.claude-haiku-4-5-20251001-v1:0")
+ assert config._supports_native_structured_outputs("anthropic.claude-opus-4-6-v1")
+ # Regional prefix is stripped by get_bedrock_base_model
+ assert config._supports_native_structured_outputs("eu.anthropic.claude-opus-4-5-20251101-v1:0")
+ # Claude 4.6 Sonnet
+ assert config._supports_native_structured_outputs("anthropic.claude-sonnet-4-6")
+ assert config._supports_native_structured_outputs("us.anthropic.claude-sonnet-4-6")
+ # Non-Anthropic models
+ assert config._supports_native_structured_outputs("qwen.qwen3-235b-a22b-2507-v1:0")
+ assert config._supports_native_structured_outputs("mistral.mistral-large-3-675b-instruct")
+ assert config._supports_native_structured_outputs("minimax.minimax-m2")
+ assert config._supports_native_structured_outputs("moonshot.kimi-k2-thinking")
+ assert config._supports_native_structured_outputs("nvidia.nemotron-nano-3-30b")
+ # DeepSeek: old substring "deepseek-v3.1" didn't match real ID
+ assert config._supports_native_structured_outputs("deepseek.v3-v1:0")
+
+ # Unsupported models -- should fall back to tool-call approach
+ assert not config._supports_native_structured_outputs("anthropic.claude-3-5-sonnet-20241022-v2:0")
+ assert not config._supports_native_structured_outputs("anthropic.claude-sonnet-4-20250514-v1:0")
+ assert not config._supports_native_structured_outputs("meta.llama3-3-70b-instruct-v1:0")
+ assert not config._supports_native_structured_outputs("amazon.nova-pro-v1:0")
+ # Excluded: broken constrained decoding on Bedrock
+ assert not config._supports_native_structured_outputs("openai.gpt-oss-120b-1:0")
+ assert not config._supports_native_structured_outputs("mistral.magistral-small-2509")
+ # Excluded: ignores schema or broken on Bedrock
+ assert not config._supports_native_structured_outputs("google.gemma-3-27b-it")
+ assert not config._supports_native_structured_outputs("nvidia.nemotron-nano-12b-v2")
+ finally:
+ litellm.model_cost = old_cost
+ if old_env is None:
+ os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
+ else:
+ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_create_output_config_for_response_format():
@@ -3039,49 +2726,57 @@ def test_create_output_config_for_response_format():
def test_translate_response_format_native_output_config():
"""For supported models, _translate_response_format_param should produce outputConfig."""
- config = AmazonConverseConfig()
+ old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
+ old_cost = litellm.model_cost
+ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
+ litellm.model_cost = litellm.get_model_cost_map(url="")
+ try:
+ config = AmazonConverseConfig()
- response_format = {
- "type": "json_schema",
- "json_schema": {
- "name": "WeatherResult",
- "description": "Weather info",
- "schema": {
- "type": "object",
- "properties": {
- "temp": {"type": "number"},
+ response_format = {
+ "type": "json_schema",
+ "json_schema": {
+ "name": "WeatherResult",
+ "description": "Weather info",
+ "schema": {
+ "type": "object",
+ "properties": {
+ "temp": {"type": "number"},
+ },
+ "required": ["temp"],
},
- "required": ["temp"],
},
- },
- }
+ }
- optional_params: dict = {}
- result = config._translate_response_format_param(
- value=response_format,
- model="anthropic.claude-sonnet-4-5-20250929-v1:0",
- optional_params=optional_params,
- non_default_params={"response_format": response_format},
- is_thinking_enabled=False,
- )
+ optional_params: dict = {}
+ result = config._translate_response_format_param(
+ value=response_format,
+ model="anthropic.claude-sonnet-4-5-20250929-v1:0",
+ optional_params=optional_params,
+ non_default_params={"response_format": response_format},
+ is_thinking_enabled=False,
+ )
- # Should have outputConfig, NOT tools
- assert "outputConfig" in result
- assert "tools" not in result
- assert "tool_choice" not in result
- assert result["json_mode"] is True
- # No fake_stream for native approach
- assert "fake_stream" not in result
+ # Should have outputConfig, NOT tools
+ assert "outputConfig" in result
+ assert "tools" not in result
+ assert "tool_choice" not in result
+ assert result["json_mode"] is True
+ # No fake_stream for native approach
+ assert "fake_stream" not in result
- # Verify the schema content (additionalProperties: false is added by normalization)
- schema_str = result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["schema"]
- parsed_schema = json.loads(schema_str)
- expected_schema = {**response_format["json_schema"]["schema"], "additionalProperties": False}
- assert parsed_schema == expected_schema
- assert (
- result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["name"]
- == "WeatherResult"
- )
+ # Verify the schema content (additionalProperties: false is added by normalization)
+ schema_str = result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["schema"]
+ parsed_schema = json.loads(schema_str)
+ expected_schema = {**response_format["json_schema"]["schema"], "additionalProperties": False}
+ assert parsed_schema == expected_schema
+ assert result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["name"] == "WeatherResult"
+ finally:
+ litellm.model_cost = old_cost
+ if old_env is None:
+ os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
+ else:
+ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_translate_response_format_fallback_tool_call():
@@ -3118,42 +2813,53 @@ def test_translate_response_format_fallback_tool_call():
def test_native_structured_output_no_fake_stream():
"""When using native structured outputs with streaming, fake_stream should NOT be set."""
- config = AmazonConverseConfig()
+ old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
+ old_cost = litellm.model_cost
+ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
+ litellm.model_cost = litellm.get_model_cost_map(url="")
+ try:
+ config = AmazonConverseConfig()
- response_format = {
- "type": "json_schema",
- "json_schema": {
- "name": "Result",
- "schema": {
- "type": "object",
- "properties": {
- "answer": {"type": "string"},
+ response_format = {
+ "type": "json_schema",
+ "json_schema": {
+ "name": "Result",
+ "schema": {
+ "type": "object",
+ "properties": {
+ "answer": {"type": "string"},
+ },
},
},
- },
- }
+ }
- optional_params: dict = {}
- result = config._translate_response_format_param(
- value=response_format,
- model="anthropic.claude-sonnet-4-5-20250929-v1:0",
- optional_params=optional_params,
- non_default_params={"response_format": response_format, "stream": True},
- is_thinking_enabled=False,
- )
+ optional_params: dict = {}
+ result = config._translate_response_format_param(
+ value=response_format,
+ model="anthropic.claude-sonnet-4-5-20250929-v1:0",
+ optional_params=optional_params,
+ non_default_params={"response_format": response_format, "stream": True},
+ is_thinking_enabled=False,
+ )
- assert "outputConfig" in result
- assert result["json_mode"] is True
- # No fake_stream for native approach
- assert "fake_stream" not in result
+ assert "outputConfig" in result
+ assert result["json_mode"] is True
+ # No fake_stream for native approach
+ assert "fake_stream" not in result
- # Verify the schema content
- schema_str = result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["schema"]
- assert json.loads(schema_str) == {
- "type": "object",
- "properties": {"answer": {"type": "string"}},
- "additionalProperties": False,
- }
+ # Verify the schema content
+ schema_str = result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["schema"]
+ assert json.loads(schema_str) == {
+ "type": "object",
+ "properties": {"answer": {"type": "string"}},
+ "additionalProperties": False,
+ }
+ finally:
+ litellm.model_cost = old_cost
+ if old_env is None:
+ os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
+ else:
+ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_transform_request_with_output_config():
@@ -3227,11 +2933,7 @@ def test_transform_response_native_structured_output():
"output": {
"message": {
"role": "assistant",
- "content": [
- {
- "text": '{"temp": 62, "description": "Mild and foggy"}'
- }
- ],
+ "content": [{"text": '{"temp": 62, "description": "Mild and foggy"}'}],
}
},
"stopReason": "end_turn",
@@ -3388,23 +3090,34 @@ def test_add_additional_properties_definitions():
def test_json_object_no_schema_falls_back_to_tool_call():
"""response_format: {type: json_object} with no schema should use tool-call fallback,
even for models that support native structured outputs."""
- config = AmazonConverseConfig()
- optional_params: dict = {}
- non_default_params = {"response_format": {"type": "json_object"}}
+ old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
+ old_cost = litellm.model_cost
+ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
+ litellm.model_cost = litellm.get_model_cost_map(url="")
+ try:
+ config = AmazonConverseConfig()
+ optional_params: dict = {}
+ non_default_params = {"response_format": {"type": "json_object"}}
- result = config._translate_response_format_param(
- value=non_default_params["response_format"],
- model="anthropic.claude-sonnet-4-5-20250929-v1:0",
- optional_params=optional_params,
- non_default_params=non_default_params,
- is_thinking_enabled=False,
- )
+ result = config._translate_response_format_param(
+ value=non_default_params["response_format"],
+ model="anthropic.claude-sonnet-4-5-20250929-v1:0",
+ optional_params=optional_params,
+ non_default_params=non_default_params,
+ is_thinking_enabled=False,
+ )
- # Should NOT use native outputConfig (no schema provided)
- assert "outputConfig" not in result
- # Should use tool-call fallback
- assert "tools" in result
- assert result["json_mode"] is True
+ # Should NOT use native outputConfig (no schema provided)
+ assert "outputConfig" not in result
+ # Should use tool-call fallback
+ assert "tools" in result
+ assert result["json_mode"] is True
+ finally:
+ litellm.model_cost = old_cost
+ if old_env is None:
+ os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
+ else:
+ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_output_config_applies_additional_properties():
@@ -3427,7 +3140,6 @@ def test_output_config_applies_additional_properties():
assert parsed["properties"]["nested"]["additionalProperties"] is False
-
_TOOL_PARAM = [
{
"type": "function",
@@ -3505,9 +3217,7 @@ def test_parallel_tool_calls_older_model_drops_disable_flag():
class TestBedrockMinThinkingBudgetTokens:
"""Test that thinking.budget_tokens is clamped to the Bedrock minimum (1024)."""
- def _map_params(
- self, thinking_value, model="anthropic.claude-3-7-sonnet-20250219-v1:0"
- ):
+ def _map_params(self, thinking_value, model="anthropic.claude-3-7-sonnet-20250219-v1:0"):
"""Helper to call map_openai_params with the given thinking value."""
config = AmazonConverseConfig()
non_default_params = {"thinking": thinking_value}
@@ -3545,6 +3255,7 @@ class TestBedrockMinThinkingBudgetTokens:
)
assert "thinking" not in result or result.get("thinking") is None
+
def test_transform_response_with_both_json_tool_call_and_real_tool():
"""
When Bedrock returns BOTH json_tool_call AND a real tool (get_weather),
@@ -3733,9 +3444,7 @@ def test_streaming_filters_json_tool_call_with_real_tools():
# Chunk 2: json_tool_call delta — should become text, not tool_use
json_delta = ContentBlockDeltaEvent(toolUse={"input": '{"temp": 62}'})
- text_2, tool_use_2, _, _, _ = decoder._handle_converse_delta_event(
- json_delta, index=0
- )
+ text_2, tool_use_2, _, _, _ = decoder._handle_converse_delta_event(json_delta, index=0)
assert text_2 == '{"temp": 62}'
assert tool_use_2 is None
@@ -3758,12 +3467,8 @@ def test_streaming_filters_json_tool_call_with_real_tools():
assert decoder.tool_calls_index == 0
# Chunk 5: real tool delta
- real_delta = ContentBlockDeltaEvent(
- toolUse={"input": '{"location": "SF"}'}
- )
- text_5, tool_use_5, _, _, _ = decoder._handle_converse_delta_event(
- real_delta, index=1
- )
+ real_delta = ContentBlockDeltaEvent(toolUse={"input": '{"location": "SF"}'})
+ text_5, tool_use_5, _, _, _ = decoder._handle_converse_delta_event(real_delta, index=1)
assert text_5 == ""
assert tool_use_5 is not None
assert tool_use_5["function"]["arguments"] == '{"location": "SF"}'
@@ -3796,9 +3501,7 @@ def test_streaming_without_json_mode_passes_all_tools():
# json_tool_call delta — should be a tool_use, not text
json_delta = ContentBlockDeltaEvent(toolUse={"input": '{"data": 1}'})
- text, tool_use_delta, _, _, _ = decoder._handle_converse_delta_event(
- json_delta, index=0
- )
+ text, tool_use_delta, _, _, _ = decoder._handle_converse_delta_event(json_delta, index=0)
assert text == ""
assert tool_use_delta is not None
assert tool_use_delta["function"]["arguments"] == '{"data": 1}'
@@ -3899,4 +3602,3 @@ def test_cache_control_injection_tool_config_not_added_without_injection_point()
tools = result["toolConfig"]["tools"]
# No cachePoint should be appended
assert all("cachePoint" not in tool for tool in tools)
-
diff --git a/tests/test_litellm/llms/gemini/files/test_gemini_files_transformation.py b/tests/test_litellm/llms/gemini/files/test_gemini_files_transformation.py
index a5f72fc08c3..6cc97cd95e6 100644
--- a/tests/test_litellm/llms/gemini/files/test_gemini_files_transformation.py
+++ b/tests/test_litellm/llms/gemini/files/test_gemini_files_transformation.py
@@ -3,10 +3,10 @@ Test Google AI Studio (Gemini) files transformation functionality
"""
import os
-import pytest
from unittest.mock import Mock, patch
import httpx
+import pytest
from litellm.llms.gemini.files.transformation import GoogleAIStudioFilesHandler
from litellm.types.llms.openai import OpenAIFileObject
@@ -23,7 +23,7 @@ class TestGoogleAIStudioFilesTransformation:
"""
Test that transform_retrieve_file_request returns empty params dict
to avoid 'Content-Type' query parameter error
-
+
Regression test for: https://github.com/BerriAI/litellm/issues/XXX
When retrieving a file, the API was incorrectly trying to pass Content-Type
as a query parameter, which Gemini API rejected.
@@ -37,14 +37,19 @@ class TestGoogleAIStudioFilesTransformation:
litellm_params=litellm_params,
)
- # Verify URL is constructed correctly with API key
- assert "key=test-api-key" in url
- assert file_id in url
+ # Verify URL is constructed exactly as required:
+ # https://generativelanguage.googleapis.com/v1beta/files/{file_id}?key=API_KEY
+ assert (
+ url
+ == "https://generativelanguage.googleapis.com/v1beta/files/test123?key=test-api-key"
+ )
# CRITICAL: params should be empty dict, not contain Content-Type or any other params
# These would be incorrectly interpreted as query parameters
assert params == {}, f"Expected empty params dict, got: {params}"
- assert "Content-Type" not in params, "Content-Type should not be in query params"
+ assert (
+ "Content-Type" not in params
+ ), "Content-Type should not be in query params"
def test_transform_retrieve_file_request_with_file_name_only(self):
"""
@@ -59,17 +64,44 @@ class TestGoogleAIStudioFilesTransformation:
litellm_params=litellm_params,
)
- # Verify URL is constructed correctly
- assert "generativelanguage.googleapis.com" in url
- assert file_id in url
- assert "key=test-api-key" in url
+ # Verify URL is constructed exactly as required:
+ # https://generativelanguage.googleapis.com/v1beta/files/{file_id}?key=API_KEY
+ assert (
+ url
+ == "https://generativelanguage.googleapis.com/v1beta/files/test123?key=test-api-key"
+ )
# CRITICAL: params should be empty dict
assert params == {}, f"Expected empty params dict, got: {params}"
- assert "Content-Type" not in params, "Content-Type should not be in query params"
+ assert (
+ "Content-Type" not in params
+ ), "Content-Type should not be in query params"
- @patch.dict('os.environ', {}, clear=True)
- @patch('litellm.llms.gemini.common_utils.get_secret_str', return_value=None)
+ def test_transform_retrieve_file_request_with_raw_id_only(self):
+ """
+ Regression guard for the exact retrieval URL format.
+
+ If someone changes the method and stops producing:
+ https://generativelanguage.googleapis.com/v1beta/files/{file_id}?key=API_KEY
+ this test should fail.
+ """
+ file_id = "cctqueckiggb"
+ litellm_params = {"api_key": "test-api-key"}
+
+ url, params = self.handler.transform_retrieve_file_request(
+ file_id=file_id,
+ optional_params={},
+ litellm_params=litellm_params,
+ )
+
+ assert (
+ url
+ == "https://generativelanguage.googleapis.com/v1beta/files/cctqueckiggb?key=test-api-key"
+ )
+ assert params == {}
+
+ @patch.dict("os.environ", {}, clear=True)
+ @patch("litellm.llms.gemini.common_utils.get_secret_str", return_value=None)
def test_transform_retrieve_file_request_missing_api_key(self, mock_get_secret):
"""Test that transform_retrieve_file_request raises error when API key is missing"""
file_id = "files/test123"
@@ -178,7 +210,7 @@ class TestGoogleAIStudioFilesTransformation:
def test_transform_retrieve_file_response_missing_createTime(self):
"""
Test that transform_retrieve_file_response raises proper error when createTime is missing
-
+
This tests the error scenario that occurs when API returns an error response
without the expected file metadata fields.
"""
@@ -221,14 +253,15 @@ class TestGoogleAIStudioFilesTransformation:
assert "x-goog-api-key" in result_headers
assert result_headers["x-goog-api-key"] == api_key
- @patch.dict('os.environ', {}, clear=True)
- @patch('litellm.llms.gemini.common_utils.get_secret_str', return_value=None)
+ @patch.dict("os.environ", {}, clear=True)
+ @patch("litellm.llms.gemini.common_utils.get_secret_str", return_value=None)
def test_validate_environment_missing_api_key(self, mock_get_secret):
"""Test that validate_environment raises error when API key is missing"""
headers = {}
with pytest.raises(
- ValueError, match="GEMINI_API_KEY is required for Google AI Studio file operations"
+ ValueError,
+ match="GEMINI_API_KEY is required for Google AI Studio file operations",
):
self.handler.validate_environment(
headers=headers,
@@ -243,7 +276,7 @@ class TestGoogleAIStudioFilesTransformation:
"""Test that get_complete_url constructs proper upload URL"""
api_base = "https://generativelanguage.googleapis.com"
api_key = "test-api-key"
-
+
url = self.handler.get_complete_url(
api_base=api_base,
api_key=api_key,
@@ -274,7 +307,7 @@ class TestGoogleAIStudioFilesTransformation:
# Verify URL extraction
assert "files/test123" in url
assert "generativelanguage.googleapis.com" in url
-
+
# Params should be empty (API key goes in header via validate_environment)
assert params == {}
diff --git a/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py b/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py
index 69741cdec6f..cc0a32d2ce6 100644
--- a/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py
+++ b/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py
@@ -10,6 +10,7 @@ sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
+import litellm
from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig
from litellm.types.llms.openai import OpenAIRealtimeStreamSessionEvents
@@ -227,3 +228,17 @@ def test_gemini_realtime_transformation_generation_complete():
contains_audio_delta = True
break
assert contains_audio_delta, "Expected audio delta event"
+
+
+def test_gemini_3_1_flash_live_preview_model_cost_map_entry():
+ for key in (
+ "gemini-3.1-flash-live-preview",
+ "gemini/gemini-3.1-flash-live-preview",
+ ):
+ assert key in litellm.model_cost
+ info = litellm.model_cost[key]
+ assert "/v1/realtime" in info.get("supported_endpoints", [])
+ assert info.get("max_input_tokens") == 131072
+ assert info.get("max_output_tokens") == 65536
+ assert "video" in info.get("supported_modalities", [])
+ assert info.get("supports_function_calling") is True
diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py
index d3214a88018..802868aa7a9 100644
--- a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py
+++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py
@@ -336,7 +336,9 @@ class TestOpenAIResponsesAPIConfig:
)
assert isinstance(result, ImageGenerationPartialImageEvent)
- assert result.type == ResponsesAPIStreamEvents.IMAGE_GENERATION_PARTIAL_IMAGE
+ assert (
+ result.type == ResponsesAPIStreamEvents.IMAGE_GENERATION_PARTIAL_IMAGE
+ )
assert result.partial_image_index == idx
assert result.b64_json == chunk["b64_json"]
@@ -689,9 +691,7 @@ class TestTransformListInputItemsRequest:
def test_openai_transform_compact_response_api_request_query_params_preserved(self):
"""Test compact URL construction preserves query params and appends path."""
# Setup
- azure_style_api_base = (
- "https://test.openai.azure.com/openai/responses?api-version=2024-05-01-preview"
- )
+ azure_style_api_base = "https://test.openai.azure.com/openai/responses?api-version=2024-05-01-preview"
# Execute
url, data = self.openai_config.transform_compact_response_api_request(
@@ -731,12 +731,12 @@ class TestTransformListInputItemsRequest:
def test_azure_transform_list_input_items_request_minimal(self):
"""Test Azure implementation with minimal parameters"""
# Setup
- azure_api_base = "https://test.openai.azure.com/openai/responses?api-version=2024-05-01-preview"
+ AZURE_AI_API_BASE = "https://test.openai.azure.com/openai/responses?api-version=2024-05-01-preview"
# Execute
url, params = self.azure_config.transform_list_input_items_request(
response_id=self.response_id,
- api_base=azure_api_base,
+ api_base=AZURE_AI_API_BASE,
litellm_params=self.litellm_params,
headers=self.headers,
)
@@ -749,12 +749,12 @@ class TestTransformListInputItemsRequest:
def test_azure_transform_list_input_items_request_url_construction(self):
"""Test Azure implementation URL construction with response_id in path"""
# Setup
- azure_api_base = "https://test.openai.azure.com/openai/responses?api-version=2024-05-01-preview"
+ AZURE_AI_API_BASE = "https://test.openai.azure.com/openai/responses?api-version=2024-05-01-preview"
# Execute
url, params = self.azure_config.transform_list_input_items_request(
response_id=self.response_id,
- api_base=azure_api_base,
+ api_base=AZURE_AI_API_BASE,
litellm_params=self.litellm_params,
headers=self.headers,
)
@@ -768,12 +768,12 @@ class TestTransformListInputItemsRequest:
def test_azure_transform_list_input_items_request_with_all_params(self):
"""Test Azure implementation with all optional parameters"""
# Setup
- azure_api_base = "https://test.openai.azure.com/openai/responses?api-version=2024-05-01-preview"
+ AZURE_AI_API_BASE = "https://test.openai.azure.com/openai/responses?api-version=2024-05-01-preview"
# Execute
url, params = self.azure_config.transform_list_input_items_request(
response_id=self.response_id,
- api_base=azure_api_base,
+ api_base=AZURE_AI_API_BASE,
litellm_params=self.litellm_params,
headers=self.headers,
after="cursor_after_123",
@@ -1128,9 +1128,9 @@ class TestPhaseParameter:
phase = getattr(output_item, "phase", None)
expected = "commentary" if idx == 0 else "final_answer"
- assert phase == expected, (
- f"output[{idx}] phase={phase!r}, expected {expected!r}"
- )
+ assert (
+ phase == expected
+ ), f"output[{idx}] phase={phase!r}, expected {expected!r}"
def test_streaming_output_item_done_preserves_phase(self):
"""OutputItemDoneEvent must preserve phase on its item."""
diff --git a/tests/test_litellm/llms/openrouter/test_openrouter_provider_routing.py b/tests/test_litellm/llms/openrouter/test_openrouter_provider_routing.py
index 72cf2eec371..0815b15c873 100644
--- a/tests/test_litellm/llms/openrouter/test_openrouter_provider_routing.py
+++ b/tests/test_litellm/llms/openrouter/test_openrouter_provider_routing.py
@@ -80,7 +80,10 @@ class TestOpenRouterNativeModelRouting:
"input_model,expected_model",
[
("openrouter/anthropic/claude-3-haiku", "anthropic/claude-3-haiku"),
- ("openrouter/meta-llama/llama-3-70b-instruct", "meta-llama/llama-3-70b-instruct"),
+ (
+ "openrouter/meta-llama/llama-3-70b-instruct",
+ "meta-llama/llama-3-70b-instruct",
+ ),
],
)
def test_regular_models_still_strip_normally(self, input_model, expected_model):
@@ -88,3 +91,12 @@ class TestOpenRouterNativeModelRouting:
result_model, provider, _, _ = litellm.get_llm_provider(model=input_model)
assert provider == "openrouter"
assert result_model == expected_model
+
+ def test_wildcard_deployment_strips_routing_prefix(self):
+ """openrouter/* proxy deployments pass custom_llm_provider; strip LiteLLM prefix."""
+ result_model, provider, _, _ = litellm.get_llm_provider(
+ model="openrouter/anthropic/claude-3.5-sonnet",
+ custom_llm_provider="openrouter",
+ )
+ assert provider == "openrouter"
+ assert result_model == "anthropic/claude-3.5-sonnet"
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py
index 3c8e1c75559..74117c01463 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py
@@ -454,6 +454,10 @@ class TestAgentHealthCheck:
self.admin_client = _make_app_with_role(LitellmUserRoles.PROXY_ADMIN)
self.mock_registry = MagicMock()
monkeypatch.setattr(ar_mod, "global_agent_registry", self.mock_registry)
+ # Ensure prisma_client is None so the endpoint skips DB queries.
+ # In CI with parallel workers, a MagicMock can leak from other test
+ # scopes, causing "object MagicMock can't be used in 'await'" errors.
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
def _make_agent(self, agent_id: str, url: str | None = None) -> AgentResponse:
card = _sample_agent_card_params()
diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py
index 69188fd200e..bd659ed518f 100644
--- a/tests/test_litellm/proxy/auth/test_auth_checks.py
+++ b/tests/test_litellm/proxy/auth/test_auth_checks.py
@@ -29,10 +29,13 @@ from litellm.proxy._types import (
from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
_can_object_call_vector_stores,
+ _check_team_member_budget,
_get_fuzzy_user_object,
_get_team_db_check,
_log_budget_lookup_failure,
+ _team_max_budget_check,
_virtual_key_max_budget_alert_check,
+ _virtual_key_max_budget_check,
_virtual_key_soft_budget_check,
get_key_object,
get_user_object,
@@ -1629,3 +1632,151 @@ async def test_custom_auth_common_checks_opt_in():
parent_otel_span=None,
)
mock_common.assert_called_once()
+
+
+# =====================================================================
+# Spend counter budget check tests (v2 — Redis-backed spend counters)
+# =====================================================================
+
+
+@pytest.mark.asyncio
+async def test_virtual_key_budget_check_reads_from_spend_counter():
+ """Budget check should use get_current_spend when counter exists,
+ even if cached object shows lower spend."""
+ from litellm.proxy.utils import ProxyLogging
+
+ valid_token = UserAPIKeyAuth(
+ token="test-hashed-token",
+ spend=0.0, # stale — counter has 1.5
+ max_budget=1.0,
+ user_id="test-user",
+ )
+
+ proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
+ proxy_logging_obj.budget_alerts = AsyncMock()
+
+ async def mock_get_current_spend(counter_key, fallback_spend):
+ if counter_key == "spend:key:test-hashed-token":
+ return 1.5
+ return fallback_spend
+
+ with patch(
+ "litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
+ ):
+ with pytest.raises(litellm.BudgetExceededError) as exc_info:
+ await _virtual_key_max_budget_check(
+ valid_token=valid_token,
+ proxy_logging_obj=proxy_logging_obj,
+ )
+ assert exc_info.value.current_cost == 1.5
+ assert exc_info.value.max_budget == 1.0
+
+
+@pytest.mark.asyncio
+async def test_virtual_key_budget_check_fallback_no_counter():
+ """When counter doesn't exist, budget check should fall back
+ to cached object's spend via fallback_spend."""
+ from litellm.proxy.utils import ProxyLogging
+
+ valid_token = UserAPIKeyAuth(
+ token="test-hashed-token",
+ spend=15.0,
+ max_budget=10.0,
+ user_id="test-user",
+ )
+
+ proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
+ proxy_logging_obj.budget_alerts = AsyncMock()
+
+ # get_current_spend returns fallback_spend when no counter exists
+ async def mock_get_current_spend(counter_key, fallback_spend):
+ return fallback_spend
+
+ with patch(
+ "litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
+ ):
+ with pytest.raises(litellm.BudgetExceededError) as exc_info:
+ await _virtual_key_max_budget_check(
+ valid_token=valid_token,
+ proxy_logging_obj=proxy_logging_obj,
+ )
+ assert exc_info.value.current_cost == 15.0
+
+
+@pytest.mark.asyncio
+async def test_team_budget_check_reads_from_spend_counter():
+ """Team budget check should use get_current_spend when counter exists."""
+ from litellm.proxy.utils import ProxyLogging
+
+ team_object = LiteLLM_TeamTable(
+ team_id="test-team",
+ spend=0.0, # stale
+ max_budget=1.0,
+ )
+ valid_token = UserAPIKeyAuth(token="test-token", team_id="test-team")
+
+ proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
+ proxy_logging_obj.budget_alerts = AsyncMock()
+
+ async def mock_get_current_spend(counter_key, fallback_spend):
+ if counter_key == "spend:team:test-team":
+ return 1.5
+ return fallback_spend
+
+ with patch(
+ "litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
+ ):
+ with pytest.raises(litellm.BudgetExceededError) as exc_info:
+ await _team_max_budget_check(
+ team_object=team_object,
+ valid_token=valid_token,
+ proxy_logging_obj=proxy_logging_obj,
+ )
+ assert exc_info.value.current_cost == 1.5
+
+
+@pytest.mark.asyncio
+async def test_team_member_budget_check_reads_from_spend_counter():
+ """Team member budget check should use get_current_spend when counter exists."""
+ from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership
+ from litellm.proxy.utils import ProxyLogging
+
+ team_object = LiteLLM_TeamTable(team_id="test-team")
+ user_object = LiteLLM_UserTable(user_id="test-user")
+ valid_token = UserAPIKeyAuth(
+ token="test-token",
+ user_id="test-user",
+ team_id="test-team",
+ )
+
+ team_membership = LiteLLM_TeamMembership(
+ user_id="test-user",
+ team_id="test-team",
+ spend=0.0, # stale
+ litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
+ )
+
+ proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
+
+ async def mock_get_current_spend(counter_key, fallback_spend):
+ if counter_key == "spend:team_member:test-user:test-team":
+ return 1.5
+ return fallback_spend
+
+ with patch(
+ "litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
+ ), patch(
+ "litellm.proxy.auth.auth_checks.get_team_membership",
+ new_callable=AsyncMock,
+ return_value=team_membership,
+ ):
+ with pytest.raises(litellm.BudgetExceededError) as exc_info:
+ await _check_team_member_budget(
+ team_object=team_object,
+ user_object=user_object,
+ valid_token=valid_token,
+ prisma_client=MagicMock(),
+ user_api_key_cache=MagicMock(),
+ proxy_logging_obj=proxy_logging_obj,
+ )
+ assert exc_info.value.current_cost == 1.5
diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py
index 11939f0fddd..ada67fbba88 100644
--- a/tests/test_litellm/proxy/auth/test_handle_jwt.py
+++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py
@@ -339,6 +339,123 @@ async def test_sync_user_role_and_teams():
assert set(user.teams) == {"team1", "team2"}
+@pytest.mark.asyncio
+async def test_sync_user_role_and_teams_cache_invalidation_on_role_change():
+ """Test that user cache is updated when role changes."""
+ mock_cache = AsyncMock()
+
+ jwt_handler = JWTHandler()
+ jwt_handler.update_environment(
+ prisma_client=None,
+ user_api_key_cache=AsyncMock(),
+ litellm_jwtauth=LiteLLM_JWTAuth(
+ jwt_litellm_role_map=[
+ JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)
+ ],
+ roles_jwt_field="roles",
+ team_ids_jwt_field="my_id_teams",
+ sync_user_role_and_teams=True,
+ ),
+ )
+
+ token = {"roles": ["ADMIN"], "my_id_teams": ["team1"]}
+ user = LiteLLM_UserTable(
+ user_id="u1",
+ user_role=LitellmUserRoles.INTERNAL_USER.value,
+ teams=["team1"], # teams already match — only role differs
+ )
+
+ prisma = AsyncMock()
+ prisma.db.litellm_usertable.update = AsyncMock()
+
+ await JWTAuthManager.sync_user_role_and_teams(
+ jwt_handler, token, user, prisma, user_api_key_cache=mock_cache
+ )
+
+ mock_cache.async_set_cache.assert_called_once()
+ call_kwargs = mock_cache.async_set_cache.call_args
+ assert call_kwargs.kwargs["key"] == "u1"
+ assert call_kwargs.kwargs["value"]["user_role"] == LitellmUserRoles.PROXY_ADMIN.value
+
+
+@pytest.mark.asyncio
+async def test_sync_user_role_and_teams_cache_invalidation_on_team_change():
+ """Test that user cache is updated when team memberships change."""
+ mock_cache = AsyncMock()
+
+ jwt_handler = JWTHandler()
+ jwt_handler.update_environment(
+ prisma_client=None,
+ user_api_key_cache=AsyncMock(),
+ litellm_jwtauth=LiteLLM_JWTAuth(
+ jwt_litellm_role_map=[
+ JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)
+ ],
+ roles_jwt_field="roles",
+ team_ids_jwt_field="my_id_teams",
+ sync_user_role_and_teams=True,
+ ),
+ )
+
+ token = {"roles": ["ADMIN"], "my_id_teams": ["team1", "team2"]}
+ user = LiteLLM_UserTable(
+ user_id="u1",
+ user_role=LitellmUserRoles.PROXY_ADMIN.value, # role already matches
+ teams=["team2"], # teams differ
+ )
+
+ prisma = AsyncMock()
+ prisma.db.litellm_usertable.update = AsyncMock()
+
+ with patch(
+ "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership",
+ new_callable=AsyncMock,
+ ):
+ await JWTAuthManager.sync_user_role_and_teams(
+ jwt_handler, token, user, prisma, user_api_key_cache=mock_cache
+ )
+
+ mock_cache.async_set_cache.assert_called_once()
+ call_kwargs = mock_cache.async_set_cache.call_args
+ assert call_kwargs.kwargs["key"] == "u1"
+ assert set(call_kwargs.kwargs["value"]["teams"]) == {"team1", "team2"}
+
+
+@pytest.mark.asyncio
+async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes():
+ """Test that cache is NOT written when role and teams already match."""
+ mock_cache = AsyncMock()
+
+ jwt_handler = JWTHandler()
+ jwt_handler.update_environment(
+ prisma_client=None,
+ user_api_key_cache=AsyncMock(),
+ litellm_jwtauth=LiteLLM_JWTAuth(
+ jwt_litellm_role_map=[
+ JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)
+ ],
+ roles_jwt_field="roles",
+ team_ids_jwt_field="my_id_teams",
+ sync_user_role_and_teams=True,
+ ),
+ )
+
+ token = {"roles": ["ADMIN"], "my_id_teams": ["team1"]}
+ user = LiteLLM_UserTable(
+ user_id="u1",
+ user_role=LitellmUserRoles.PROXY_ADMIN.value,
+ teams=["team1"],
+ )
+
+ prisma = AsyncMock()
+
+ await JWTAuthManager.sync_user_role_and_teams(
+ jwt_handler, token, user, prisma, user_api_key_cache=mock_cache
+ )
+
+ mock_cache.async_set_cache.assert_not_called()
+
+
@pytest.mark.asyncio
async def test_map_jwt_role_to_litellm_role():
"""Test JWT role mapping to LiteLLM roles with various patterns"""
diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
index 81ca758983b..ca70b4bfa8e 100644
--- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
+++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
@@ -28,7 +28,7 @@ def test_get_api_key():
assert get_api_key(
custom_litellm_key_header=None,
api_key=bearer_token,
- azure_api_key_header=None,
+ AZURE_AI_API_KEY_header=None,
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
@@ -59,7 +59,7 @@ def test_get_api_key_with_custom_litellm_key_header(
assert get_api_key(
custom_litellm_key_header=custom_litellm_key_header,
api_key=None,
- azure_api_key_header=None,
+ AZURE_AI_API_KEY_header=None,
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
@@ -334,7 +334,6 @@ async def test_proxy_admin_expired_key_from_cache():
"litellm.proxy.auth.user_api_key_auth._delete_cache_key_object",
new_callable=AsyncMock,
) as mock_delete_cache:
-
mock_get_key_object.return_value = expired_token
# Set attributes on proxy_server module (these are imported inside _user_api_key_auth_builder)
@@ -372,7 +371,7 @@ async def test_proxy_admin_expired_key_from_cache():
await _user_api_key_auth_builder(
request=request,
api_key=f"Bearer {api_key}", # Add Bearer prefix
- azure_api_key_header="",
+ AZURE_AI_API_KEY_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
@@ -567,6 +566,10 @@ class TestJWTOAuth2Coexistence:
assert JWTHandler.is_jwt("Bearer token") is False
assert JWTHandler.is_jwt("two.parts") is False
+ def test_is_jwt_returns_false_for_none(self):
+ """None token (missing Authorization header) should not be treated as JWT."""
+ assert JWTHandler.is_jwt(None) is False
+
@pytest.mark.asyncio
async def test_both_enabled_opaque_token_uses_oauth2(self):
"""
@@ -605,7 +608,6 @@ class TestJWTOAuth2Coexistence:
"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(),
@@ -670,7 +672,6 @@ class TestJWTOAuth2Coexistence:
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(),
@@ -722,7 +723,6 @@ class TestJWTOAuth2Coexistence:
new_callable=AsyncMock,
return_value=mock_oauth2_response,
) as mock_oauth2:
-
result = await user_api_key_auth(
request=mock_request,
api_key=f"Bearer {jwt_like_token}",
@@ -842,7 +842,7 @@ async def test_user_api_key_auth_builder_no_blocking_calls():
await _user_api_key_auth_builder(
request=request,
api_key=f"Bearer {api_key}",
- azure_api_key_header="",
+ AZURE_AI_API_KEY_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py
index e358cbe3be4..0f90d236aed 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py
@@ -85,6 +85,7 @@ async def test_ui_view_users_proxy_admin_no_org_filter(mocker):
Proxy admin: find_many is called without organization_memberships in where.
"""
mock_prisma_client = mocker.MagicMock()
+
async def mock_find_many(*args, **kwargs):
assert "organization_memberships" not in (kwargs.get("where") or {})
return []
@@ -327,6 +328,7 @@ async def test_ui_view_users_flag_on_team_admin_non_org_team_403(mocker):
Flag ON, team admin for non-org team: returns 403.
"""
from fastapi import HTTPException
+
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
mock_prisma_client = mocker.MagicMock()
@@ -372,9 +374,7 @@ async def test_ui_view_users_flag_on_team_admin_non_org_team_403(mocker):
with pytest.raises(HTTPException) as exc_info:
await ui_view_users(
- user_api_key_dict=UserAPIKeyAuth(
- user_id="team-admin-user", user_role=None
- ),
+ user_api_key_dict=UserAPIKeyAuth(user_id="team-admin-user", user_role=None),
user_id=None,
user_email="u",
team_id=tid,
@@ -633,7 +633,9 @@ async def test_get_users_includes_timestamps(mocker):
# Call get_users function directly with proxy admin auth
admin_key = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
- response = await get_users(page=1, page_size=1, user_api_key_dict=admin_key, organization_ids=None)
+ response = await get_users(
+ page=1, page_size=1, user_api_key_dict=admin_key, organization_ids=None
+ )
print("user /list response: ", response)
@@ -855,7 +857,9 @@ async def test_new_user_non_admin_cannot_create_admin(mocker):
# Verify the exception details
assert exc_info.value.code == 403 or exc_info.value.code == "403"
- assert "Only proxy admins can create administrative users" in str(exc_info.value.message)
+ assert "Only proxy admins can create administrative users" in str(
+ exc_info.value.message
+ )
assert "proxy_admin" in str(exc_info.value.message)
assert "proxy_admin_viewer" in str(exc_info.value.message)
assert str(LitellmUserRoles.PROXY_ADMIN) in str(exc_info.value.message)
@@ -896,14 +900,14 @@ async def test_user_info_url_encoding_plus_character(mocker):
# Mock the prisma client
mock_prisma_client = mocker.MagicMock()
-
+
# Create a real LiteLLM_UserTable instance (BaseModel) so isinstance check passes
mock_user = LiteLLM_UserTable(
user_id="machine-user+alp-air-admin-b58-b@tempus.com",
user_email="machine-user+alp-air-admin-b58-b@tempus.com",
teams=[],
)
-
+
# Mock get_data to return user when called with user_id, empty list for keys
async def mock_get_data(*args, **kwargs):
if kwargs.get("table_name") == "key":
@@ -913,7 +917,7 @@ async def test_user_info_url_encoding_plus_character(mocker):
elif kwargs.get("user_id") is not None:
return mock_user
return None
-
+
mock_prisma_client.get_data = mocker.AsyncMock(side_effect=mock_get_data)
# Mock list_team to return None (patch it from where it's imported)
@@ -941,7 +945,7 @@ async def test_user_info_url_encoding_plus_character(mocker):
"machine-user alp-air-admin-b58-b@tempus.com" # What FastAPI gives us
)
expected_user_id = "machine-user+alp-air-admin-b58-b@tempus.com"
-
+
response = await user_info(
user_id=decoded_user_id,
user_api_key_dict=mock_user_api_key_dict,
@@ -955,7 +959,7 @@ async def test_user_info_url_encoding_plus_character(mocker):
if call.kwargs.get("user_id") and not call.kwargs.get("table_name"):
user_call = call
break
-
+
assert user_call is not None, "get_data should be called with user_id"
assert user_call.kwargs["user_id"] == expected_user_id
@@ -972,7 +976,7 @@ async def test_user_info_nonexistent_user(mocker):
# Mock the prisma client
mock_prisma_client = mocker.MagicMock()
-
+
# Mock get_data to return None (user doesn't exist)
async def mock_get_data(*args, **kwargs):
if kwargs.get("table_name") == "key":
@@ -980,7 +984,7 @@ async def test_user_info_nonexistent_user(mocker):
elif kwargs.get("user_id") is not None:
return None # User not found
return None
-
+
mock_prisma_client.get_data = mocker.AsyncMock(side_effect=mock_get_data)
# Patch the prisma client import in the endpoint
@@ -996,7 +1000,7 @@ async def test_user_info_nonexistent_user(mocker):
# Call user_info function with a non-existent user_id
nonexistent_user_id = "nonexistent-user@example.com"
-
+
# Should raise ProxyException with 404 status code (HTTPException is converted by decorator)
with pytest.raises(ProxyException) as exc_info:
await user_info(
@@ -1370,9 +1374,7 @@ async def test_check_duplicate_user_id(mocker):
await _check_duplicate_user_id("existing-user-id", mock_prisma_client)
assert exc_info.value.status_code == 409
- assert "User with id existing-user-id already exists" in str(
- exc_info.value.detail
- )
+ assert "User with id existing-user-id already exists" in str(exc_info.value.detail)
# No duplicate should pass
async def mock_find_first_no_duplicate(*args, **kwargs):
@@ -1393,7 +1395,7 @@ async def test_check_duplicate_user_id(mocker):
def test_process_keys_for_user_info_filters_dashboard_keys(monkeypatch):
"""
Test that _process_keys_for_user_info filters out keys with team_id='litellm-dashboard'
-
+
UI session tokens (team_id='litellm-dashboard') should be excluded from user info responses
to prevent confusion, as these are automatically created during dashboard login.
"""
@@ -1412,7 +1414,7 @@ def test_process_keys_for_user_info_filters_dashboard_keys(monkeypatch):
"user_id": "test-user",
"key_alias": "dashboard-session-key",
}
-
+
mock_key_regular = MagicMock()
mock_key_regular.model_dump.return_value = {
"token": "sk-regular-token",
@@ -1420,7 +1422,7 @@ def test_process_keys_for_user_info_filters_dashboard_keys(monkeypatch):
"user_id": "test-user",
"key_alias": "regular-key",
}
-
+
mock_key_no_team = MagicMock()
mock_key_no_team.model_dump.return_value = {
"token": "sk-no-team-token",
@@ -1446,20 +1448,24 @@ def test_process_keys_for_user_info_filters_dashboard_keys(monkeypatch):
# Verify that dashboard key is filtered out
assert len(result) == 2, "Should return 2 keys (dashboard key filtered out)"
-
+
# Verify dashboard key is not in results
result_team_ids = [key.get("team_id") for key in result]
- assert UI_SESSION_TOKEN_TEAM_ID not in result_team_ids, "Dashboard key should be filtered out"
-
+ assert (
+ UI_SESSION_TOKEN_TEAM_ID not in result_team_ids
+ ), "Dashboard key should be filtered out"
+
# Verify regular keys are included
assert "regular-team" in result_team_ids, "Regular team key should be included"
assert None in result_team_ids, "No-team key should be included"
-
+
# Verify the correct keys are returned
result_tokens = [key.get("token") for key in result]
assert "sk-regular-token" in result_tokens, "Regular key should be included"
assert "sk-no-team-token" in result_tokens, "No-team key should be included"
- assert "sk-dashboard-token" not in result_tokens, "Dashboard key should not be included"
+ assert (
+ "sk-dashboard-token" not in result_tokens
+ ), "Dashboard key should not be included"
def test_process_keys_for_user_info_handles_none_keys(monkeypatch):
@@ -1558,7 +1564,13 @@ async def test_get_users_user_id_partial_match(mocker):
admin_key = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
captured_where_conditions.clear()
- await get_users(user_ids="test-user", page=1, page_size=1, user_api_key_dict=admin_key, organization_ids=None)
+ await get_users(
+ user_ids="test-user",
+ page=1,
+ page_size=1,
+ user_api_key_dict=admin_key,
+ organization_ids=None,
+ )
assert "user_id" in captured_where_conditions
assert "contains" in captured_where_conditions["user_id"]
@@ -1566,7 +1578,13 @@ async def test_get_users_user_id_partial_match(mocker):
assert captured_where_conditions["user_id"]["mode"] == "insensitive"
captured_where_conditions.clear()
- await get_users(user_ids="user1,user2,user3", page=1, page_size=1, user_api_key_dict=admin_key, organization_ids=None)
+ await get_users(
+ user_ids="user1,user2,user3",
+ page=1,
+ page_size=1,
+ user_api_key_dict=admin_key,
+ organization_ids=None,
+ )
assert "user_id" in captured_where_conditions
assert "in" in captured_where_conditions["user_id"]
@@ -1578,7 +1596,7 @@ def test_update_internal_user_params_reset_max_budget_with_none():
Test that _update_internal_user_params allows setting max_budget to None.
This verifies the fix for unsetting/resetting the budget to unlimited.
"""
-
+
# Case 1: max_budget is explicitly None in the input dictionary
data_json = {"max_budget": None, "user_id": "test_user"}
data = UpdateUserRequest(max_budget=None, user_id="test_user")
@@ -1610,7 +1628,7 @@ def test_update_internal_user_params_ignores_other_nones():
def test_update_internal_user_params_keeps_original_max_budget_when_not_provided():
"""
- Test that _update_internal_user_params does not include max_budget
+ Test that _update_internal_user_params does not include max_budget
when it's not provided in the request (should keep original value).
"""
# Create test data without max_budget
@@ -1631,7 +1649,7 @@ def test_generate_request_base_validator():
Test that GenerateRequestBase validator converts empty string to None for max_budget
"""
from litellm.proxy._types import GenerateRequestBase
-
+
# Test with empty string
req = GenerateRequestBase(max_budget="")
assert req.max_budget is None
@@ -1662,9 +1680,7 @@ async def test_get_user_daily_activity_non_admin_cannot_view_other_users(monkeyp
# Mock the prisma client so the DB-not-connected check passes
mock_prisma_client = MagicMock()
- monkeypatch.setattr(
- "litellm.proxy.proxy_server.prisma_client", mock_prisma_client
- )
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
# Non-admin caller
non_admin_key_dict = UserAPIKeyAuth(
@@ -1731,9 +1747,7 @@ async def test_get_user_daily_activity_aggregated_admin_global_view(monkeypatch)
# Mock the prisma client
mock_prisma_client = MagicMock()
- monkeypatch.setattr(
- "litellm.proxy.proxy_server.prisma_client", mock_prisma_client
- )
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
# Mock the downstream helper so we don't need a real DB
mock_response = MagicMock()
@@ -1846,7 +1860,9 @@ async def test_delete_user_cleans_up_created_by_invitation_links(mocker):
call_kwargs = mock_prisma_client.db.litellm_invitationlink.delete_many.call_args
where_clause = call_kwargs.kwargs.get("where") or call_kwargs[1].get("where")
- assert "OR" in where_clause, "Should use OR to match user_id, created_by, and updated_by"
+ assert (
+ "OR" in where_clause
+ ), "Should use OR to match user_id, created_by, and updated_by"
or_conditions = where_clause["OR"]
assert len(or_conditions) == 3, "Should have 3 OR conditions"
@@ -2188,9 +2204,20 @@ async def test_user_info_v2_response_shape(mocker):
# Verify all expected fields are present
response_dict = response.model_dump()
expected_fields = {
- "user_id", "user_email", "user_alias", "user_role", "spend",
- "max_budget", "models", "budget_duration", "budget_reset_at",
- "metadata", "created_at", "updated_at", "sso_user_id", "teams",
+ "user_id",
+ "user_email",
+ "user_alias",
+ "user_role",
+ "spend",
+ "max_budget",
+ "models",
+ "budget_duration",
+ "budget_reset_at",
+ "metadata",
+ "created_at",
+ "updated_at",
+ "sso_user_id",
+ "teams",
}
assert set(response_dict.keys()) == expected_fields
@@ -2418,4 +2445,74 @@ async def test_user_info_v2_url_encoding_plus_character(mocker):
)
assert isinstance(response, UserInfoV2Response)
- assert response.user_id == expected_user_id
\ No newline at end of file
+ assert response.user_id == expected_user_id
+
+
+class TestGetUserIdFromRequestValidation:
+ """Tests for user_id input validation in get_user_id_from_request."""
+
+ def _make_request(self, query_string: str):
+ from unittest.mock import MagicMock
+
+ from starlette.requests import Request
+
+ request = MagicMock(spec=Request)
+ request.url.query = query_string
+ return request
+
+ def test_valid_uuid(self):
+ from litellm.proxy.management_endpoints.internal_user_endpoints import (
+ get_user_id_from_request,
+ )
+
+ request = self._make_request("user_id=550e8400-e29b-41d4-a716-446655440000")
+ result = get_user_id_from_request(request)
+ assert result == "550e8400-e29b-41d4-a716-446655440000"
+
+ def test_valid_email(self):
+ from litellm.proxy.management_endpoints.internal_user_endpoints import (
+ get_user_id_from_request,
+ )
+
+ request = self._make_request("user_id=user%40example.com")
+ result = get_user_id_from_request(request)
+ assert result == "user@example.com"
+
+ def test_rejects_overlong_user_id(self):
+ from litellm.proxy.management_endpoints.internal_user_endpoints import (
+ get_user_id_from_request,
+ )
+
+ long_id = "a" * 513
+ request = self._make_request(f"user_id={long_id}")
+ result = get_user_id_from_request(request)
+ assert result is None
+
+ def test_rejects_null_byte(self):
+ from litellm.proxy.management_endpoints.internal_user_endpoints import (
+ get_user_id_from_request,
+ )
+
+ request = self._make_request("user_id=admin%00evil")
+ result = get_user_id_from_request(request)
+ assert result is None
+
+ def test_rejects_control_characters(self):
+ from litellm.proxy.management_endpoints.internal_user_endpoints import (
+ get_user_id_from_request,
+ )
+
+ # Tab character (0x09)
+ request = self._make_request("user_id=admin%09evil")
+ result = get_user_id_from_request(request)
+ assert result is None
+
+ def test_allows_512_char_user_id(self):
+ from litellm.proxy.management_endpoints.internal_user_endpoints import (
+ get_user_id_from_request,
+ )
+
+ exact_id = "a" * 512
+ request = self._make_request(f"user_id={exact_id}")
+ result = get_user_id_from_request(request)
+ assert result == exact_id
diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py
index f3c89003105..2e566ab6222 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py
@@ -1,13 +1,14 @@
import json
import os
import sys
-from litellm._uuid import uuid
from typing import Dict, Optional
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi.testclient import TestClient
+from litellm._uuid import uuid
+
sys.path.insert(
0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path
@@ -27,9 +28,15 @@ from litellm.types.router import Deployment, LiteLLM_Params, updateDeployment
class MockPrismaClient:
- def __init__(self, team_exists: bool = True, user_admin: bool = True):
+ def __init__(
+ self,
+ team_exists: bool = True,
+ user_admin: bool = True,
+ sibling_deployments: list = None,
+ ):
self.team_exists = team_exists
self.user_admin = user_admin
+ self.sibling_deployments = sibling_deployments or []
self.db = self
async def find_unique(self, where):
@@ -45,10 +52,53 @@ class MockPrismaClient:
)
return None
+ async def find_many(self, where):
+ # Filter sibling deployments by team_id if where clause specifies it
+ if not self.sibling_deployments:
+ return []
+
+ # Extract team_id from where clause if present
+ team_id_filter = None
+ if where and "model_info" in where:
+ model_info_filter = where["model_info"]
+ if isinstance(model_info_filter, dict) and "path" in model_info_filter:
+ if (
+ model_info_filter["path"] == ["team_id"]
+ and "equals" in model_info_filter
+ ):
+ team_id_filter = model_info_filter["equals"]
+
+ # Filter deployments by team_id if specified
+ if team_id_filter:
+
+ def _get_team_id(model_info):
+ if isinstance(model_info, dict):
+ return model_info.get("team_id")
+ if isinstance(model_info, str):
+ try:
+ parsed = json.loads(model_info)
+ except (TypeError, ValueError):
+ return None
+ if isinstance(parsed, dict):
+ return parsed.get("team_id")
+ return None
+
+ return [
+ d
+ for d in self.sibling_deployments
+ if _get_team_id(d.model_info) == team_id_filter
+ ]
+
+ return self.sibling_deployments
+
@property
def litellm_teamtable(self):
return self
+ @property
+ def litellm_proxymodeltable(self):
+ return self
+
class MockLLMRouter:
def __init__(self):
@@ -399,7 +449,9 @@ class TestClearCache:
"""
Test that clear_cache clears DB models and preserves config models.
"""
- from litellm.proxy.management_endpoints.model_management_endpoints import clear_cache
+ from litellm.proxy.management_endpoints.model_management_endpoints import (
+ clear_cache,
+ )
# Create mock router with mixed DB and config models
mock_router = MagicMock()
@@ -407,18 +459,18 @@ class TestClearCache:
{
"model_name": "gpt-4",
"model_info": {"id": "db-model-1", "db_model": True},
- "litellm_params": {"model": "gpt-4"}
+ "litellm_params": {"model": "gpt-4"},
},
{
- "model_name": "gpt-3.5-turbo",
+ "model_name": "gpt-3.5-turbo",
"model_info": {"id": "config-model-1", "db_model": False},
- "litellm_params": {"model": "gpt-3.5-turbo"}
+ "litellm_params": {"model": "gpt-3.5-turbo"},
},
{
"model_name": "claude-3",
"model_info": {"id": "db-model-2", "db_model": True},
- "litellm_params": {"model": "claude-3"}
- }
+ "litellm_params": {"model": "claude-3"},
+ },
]
mock_router.delete_deployment = MagicMock(return_value=True)
mock_router.auto_routers = MagicMock()
@@ -466,8 +518,8 @@ class TestUpdatePublicModelGroups:
"""
import litellm
from litellm.proxy.management_endpoints.model_management_endpoints import (
- update_public_model_groups,
UpdatePublicModelGroupsRequest,
+ update_public_model_groups,
)
old_db_models = ["db-model-1", "db-model-2"]
@@ -525,7 +577,10 @@ class TestUpdatePublicModelGroups:
)
old_links = {"Old Doc": "https://old.example.com"}
- new_links = {"New Doc": "https://new.example.com", "API Ref": "https://api.example.com"}
+ new_links = {
+ "New Doc": "https://new.example.com",
+ "API Ref": "https://api.example.com",
+ }
async def mock_get_config(*args, **kwargs):
litellm.public_model_groups_links = old_links
@@ -558,6 +613,161 @@ class TestUpdatePublicModelGroups:
litellm.public_model_groups_links = original_value
+class TestTeamModelSiblingRouting:
+ """
+ Verify that sibling team deployments (same public model name, different
+ api_base) are all reachable through routing — no alias overwrite, no
+ collapse to a single deployment.
+ """
+
+ @pytest.mark.asyncio
+ async def test_no_model_aliases_written_for_team_models(self):
+ """
+ _add_team_model_to_db must NOT write model_aliases (which caused
+ the second sibling to overwrite the first). It should only call
+ team_model_add to register the public name on the team's models list.
+ """
+ from litellm.proxy.management_endpoints.model_management_endpoints import (
+ _add_team_model_to_db,
+ )
+ from litellm.types.router import ModelInfo
+
+ team_id = "team_no_alias"
+ public_name = "gpt-4.1-mini"
+
+ async def mock_add_model_to_db(model_params, user_api_key_dict, prisma_client):
+ return MagicMock(model_id=str(uuid.uuid4()))
+
+ mock_team_model_add = AsyncMock()
+
+ user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
+ prisma_client = MockPrismaClient(team_exists=True)
+
+ for api_base in ["https://eastus.example.com", "https://westus.example.com"]:
+ dep = Deployment(
+ model_name=public_name,
+ litellm_params=LiteLLM_Params(
+ model="azure/gpt-4o-mini",
+ api_key="key",
+ api_base=api_base,
+ ),
+ model_info=ModelInfo(team_id=team_id),
+ )
+ with patch(
+ "litellm.proxy.management_endpoints.model_management_endpoints._add_model_to_db",
+ side_effect=mock_add_model_to_db,
+ ), patch(
+ "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add",
+ mock_team_model_add,
+ ):
+ await _add_team_model_to_db(
+ model_params=dep,
+ user_api_key_dict=user,
+ prisma_client=prisma_client,
+ )
+
+ assert mock_team_model_add.call_count == 2
+
+ @pytest.mark.asyncio
+ async def test_router_finds_all_sibling_team_deployments(self):
+ """
+ When two team deployments share team_public_model_name="gpt-4.1-mini",
+ the router's _common_checks_available_deployment must return BOTH as
+ healthy_deployments (not collapse to one).
+ """
+ import litellm
+
+ team_id = "teamA"
+ public_name = "gpt-4.1-mini"
+
+ router = litellm.Router(
+ model_list=[
+ {
+ "model_name": f"model_name_{team_id}_uuid1",
+ "litellm_params": {
+ "model": "azure/gpt-4o-mini",
+ "api_key": "key-1",
+ "api_base": "https://eastus.openai.azure.com",
+ },
+ "model_info": {
+ "team_id": team_id,
+ "team_public_model_name": public_name,
+ },
+ },
+ {
+ "model_name": f"model_name_{team_id}_uuid2",
+ "litellm_params": {
+ "model": "azure/gpt-4o-mini",
+ "api_key": "key-2",
+ "api_base": "https://westus.openai.azure.com",
+ },
+ "model_info": {
+ "team_id": team_id,
+ "team_public_model_name": public_name,
+ },
+ },
+ {
+ "model_name": "global-gpt-4o",
+ "litellm_params": {
+ "model": "azure/gpt-4o",
+ "api_key": "global-key",
+ "api_base": "https://global.openai.azure.com",
+ },
+ "model_info": {}, # No team_id - global deployment
+ },
+ ],
+ )
+
+ # map_team_model should return the public name (not an internal UUID)
+ result = router.map_team_model(public_name, team_id)
+ assert result == public_name
+
+ # _common_checks_available_deployment should return both deployments
+ model, healthy = router._common_checks_available_deployment(
+ model=public_name,
+ request_kwargs={"metadata": {"user_api_key_team_id": team_id}},
+ )
+ assert isinstance(healthy, list)
+ assert len(healthy) == 2
+ api_bases = {d["litellm_params"]["api_base"] for d in healthy}
+ assert api_bases == {
+ "https://eastus.openai.azure.com",
+ "https://westus.openai.azure.com",
+ }
+
+ def test_global_deployments_accessible_to_teams(self):
+ """Test that global deployments (no team_id) are accessible to all teams"""
+ import litellm
+
+ router = litellm.Router(
+ model_list=[
+ {
+ "model_name": "global-gpt-4o",
+ "litellm_params": {
+ "model": "azure/gpt-4o",
+ "api_key": "global-key",
+ "api_base": "https://global.openai.azure.com",
+ },
+ "model_info": {}, # No team_id - global deployment
+ },
+ ],
+ )
+
+ # Global deployment should be accessible when team_id is provided
+ deployments = router._get_all_deployments(
+ model_name="global-gpt-4o", team_id="teamA"
+ )
+ assert len(deployments) == 1
+ assert deployments[0]["model_name"] == "global-gpt-4o"
+
+ # should_include_deployment should return True for global deployments
+ assert router.should_include_deployment(
+ model_name="global-gpt-4o",
+ model={"model_name": "global-gpt-4o", "model_info": {}},
+ team_id="teamA",
+ )
+
+
class TestTeamModelUpdate:
"""Test team model update handles team_id consistently with model creation"""
@@ -591,10 +801,10 @@ class TestTeamModelUpdate:
"litellm.proxy.proxy_server.premium_user",
True,
), patch(
- "litellm.proxy.management_endpoints.model_management_endpoints.update_team"
- ) as mock_update_team, patch(
"litellm.proxy.management_endpoints.model_management_endpoints.team_model_add"
- ) as mock_team_model_add:
+ ) as mock_team_model_add, patch(
+ "litellm.proxy.management_endpoints.model_management_endpoints.update_team"
+ ) as mock_update_team:
result = await _update_team_model_in_db(
db_model=db_model,
patch_data=patch_data,
@@ -604,8 +814,201 @@ class TestTeamModelUpdate:
assert result.get("model_name", "").startswith("model_name_test_team_123_")
assert "team_public_model_name" in str(result.get("model_info", ""))
- mock_update_team.assert_called_once()
+ # team_model_add must be called to add public name to team's models list
mock_team_model_add.assert_called_once()
+ # update_team (model_aliases write) must NOT be called in the new implementation
+ mock_update_team.assert_not_called()
+
+ @pytest.mark.asyncio
+ async def test_rename_preserves_old_name_when_siblings_exist(self):
+ """Test that renaming a deployment preserves old public name when sibling deployments still use it"""
+ from unittest.mock import MagicMock
+
+ from litellm.proxy.management_endpoints.model_management_endpoints import (
+ _update_existing_team_model_assignment,
+ )
+ from litellm.types.router import ModelInfo
+
+ # Create a deployment being renamed
+ db_model = Deployment(
+ model_name="model_name_team_123_uuid1",
+ litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"),
+ model_info=ModelInfo(
+ team_id="team_123", team_public_model_name="old-public-name"
+ ),
+ )
+
+ # Create a sibling deployment that still uses the old public name
+ sibling_deployment = MagicMock()
+ sibling_deployment.model_name = "model_name_team_123_uuid2"
+ sibling_deployment.model_info = {
+ "team_id": "team_123",
+ "team_public_model_name": "old-public-name",
+ }
+
+ prisma_client = MockPrismaClient(
+ team_exists=True, sibling_deployments=[sibling_deployment]
+ )
+
+ patch_data = updateDeployment(
+ model_name="new-public-name",
+ model_info=ModelInfo(team_id="team_123"),
+ )
+
+ user_api_key_dict = UserAPIKeyAuth(
+ user_id="test_user",
+ user_role=LitellmUserRoles.PROXY_ADMIN,
+ )
+
+ with patch(
+ "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete"
+ ) as mock_delete, patch(
+ "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add"
+ ) as mock_add:
+ await _update_existing_team_model_assignment(
+ team_id="team_123",
+ public_model_name="new-public-name",
+ db_model=db_model,
+ patch_data=patch_data,
+ user_api_key_dict=user_api_key_dict,
+ prisma_client=prisma_client, # type: ignore
+ )
+
+ # team_model_delete should NOT be called because sibling exists
+ mock_delete.assert_not_called()
+ # team_model_add should be called to add new public name
+ mock_add.assert_called_once()
+
+ @pytest.mark.asyncio
+ async def test_first_time_public_name_assignment_adds_team_model(self):
+ """If existing team deployment had no public name, first assignment must call team_model_add."""
+ from litellm.proxy.management_endpoints.model_management_endpoints import (
+ _update_existing_team_model_assignment,
+ )
+ from litellm.types.router import ModelInfo
+
+ db_model = Deployment(
+ model_name="model_name_team_123_uuid1",
+ litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"),
+ model_info=ModelInfo(team_id="team_123"),
+ )
+
+ patch_data = updateDeployment(
+ model_name="new-public-name",
+ model_info=ModelInfo(team_id="team_123"),
+ )
+
+ user_api_key_dict = UserAPIKeyAuth(
+ user_id="test_user",
+ user_role=LitellmUserRoles.PROXY_ADMIN,
+ )
+
+ with patch(
+ "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete"
+ ) as mock_delete, patch(
+ "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add"
+ ) as mock_add:
+ await _update_existing_team_model_assignment(
+ team_id="team_123",
+ public_model_name="new-public-name",
+ db_model=db_model,
+ patch_data=patch_data,
+ user_api_key_dict=user_api_key_dict,
+ prisma_client=None,
+ )
+
+ mock_add.assert_called_once()
+ mock_delete.assert_not_called()
+
+ @pytest.mark.asyncio
+ async def test_rename_with_prisma_none_clears_patch_model_name(self):
+ """Rename path must clear patch_data.model_name even when prisma is unavailable (P1)."""
+ from litellm.proxy.management_endpoints.model_management_endpoints import (
+ _update_existing_team_model_assignment,
+ )
+ from litellm.types.router import ModelInfo
+
+ db_model = Deployment(
+ model_name="model_name_team_123_uuid1",
+ litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"),
+ model_info=ModelInfo(
+ team_id="team_123", team_public_model_name="old-public-name"
+ ),
+ )
+ patch_data = updateDeployment(
+ model_name="new-public-name",
+ model_info=ModelInfo(team_id="team_123"),
+ )
+ user_api_key_dict = UserAPIKeyAuth(
+ user_id="test_user",
+ user_role=LitellmUserRoles.PROXY_ADMIN,
+ )
+
+ await _update_existing_team_model_assignment(
+ team_id="team_123",
+ public_model_name="new-public-name",
+ db_model=db_model,
+ patch_data=patch_data,
+ user_api_key_dict=user_api_key_dict,
+ prisma_client=None,
+ )
+
+ assert patch_data.model_name is None
+
+ @pytest.mark.asyncio
+ async def test_rename_handles_legacy_string_model_info(self):
+ """Test rename path handles legacy string-encoded model_info rows without crashing."""
+ from unittest.mock import MagicMock
+
+ from litellm.proxy.management_endpoints.model_management_endpoints import (
+ _update_existing_team_model_assignment,
+ )
+ from litellm.types.router import ModelInfo
+
+ db_model = Deployment(
+ model_name="model_name_team_123_uuid1",
+ litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"),
+ model_info=ModelInfo(
+ team_id="team_123", team_public_model_name="old-public-name"
+ ),
+ )
+
+ sibling_deployment = MagicMock()
+ sibling_deployment.model_name = "model_name_team_123_uuid2"
+ sibling_deployment.model_info = (
+ '{"team_id":"team_123","team_public_model_name":"old-public-name"}'
+ )
+
+ prisma_client = MockPrismaClient(
+ team_exists=True, sibling_deployments=[sibling_deployment]
+ )
+
+ patch_data = updateDeployment(
+ model_name="new-public-name",
+ model_info=ModelInfo(team_id="team_123"),
+ )
+
+ user_api_key_dict = UserAPIKeyAuth(
+ user_id="test_user",
+ user_role=LitellmUserRoles.PROXY_ADMIN,
+ )
+
+ with patch(
+ "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete"
+ ) as mock_delete, patch(
+ "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add"
+ ) as mock_add:
+ await _update_existing_team_model_assignment(
+ team_id="team_123",
+ public_model_name="new-public-name",
+ db_model=db_model,
+ patch_data=patch_data,
+ user_api_key_dict=user_api_key_dict,
+ prisma_client=prisma_client, # type: ignore
+ )
+
+ mock_delete.assert_not_called()
+ mock_add.assert_called_once()
@pytest.mark.asyncio
async def test_patch_model_with_team_id_validates_permissions(self):
@@ -657,27 +1060,37 @@ class TestModelInfoEndpoint:
user_id="test_user",
api_key="test_key",
models=["gpt-4", "claude-3"],
- team_models=["gpt-3.5-turbo"]
+ team_models=["gpt-3.5-turbo"],
)
- with patch("litellm.proxy.proxy_server.llm_router") as mock_router, \
- patch("litellm.proxy.proxy_server.get_key_models") as mock_get_key_models, \
- patch("litellm.proxy.proxy_server.get_team_models") as mock_get_team_models, \
- patch("litellm.proxy.proxy_server.get_complete_model_list") as mock_get_complete_models, \
- patch("litellm.get_llm_provider") as mock_get_provider:
-
+ with patch("litellm.proxy.proxy_server.llm_router") as mock_router, patch(
+ "litellm.proxy.proxy_server.get_key_models"
+ ) as mock_get_key_models, patch(
+ "litellm.proxy.proxy_server.get_team_models"
+ ) as mock_get_team_models, patch(
+ "litellm.proxy.proxy_server.get_complete_model_list"
+ ) as mock_get_complete_models, patch(
+ "litellm.get_llm_provider"
+ ) as mock_get_provider:
# Setup mocks
- mock_router.get_model_names.return_value = ["gpt-4", "claude-3", "gpt-3.5-turbo"]
+ mock_router.get_model_names.return_value = [
+ "gpt-4",
+ "claude-3",
+ "gpt-3.5-turbo",
+ ]
mock_router.get_model_access_groups.return_value = {}
mock_get_key_models.return_value = ["gpt-4", "claude-3"]
mock_get_team_models.return_value = ["gpt-3.5-turbo"]
- mock_get_complete_models.return_value = ["gpt-4", "claude-3", "gpt-3.5-turbo"]
+ mock_get_complete_models.return_value = [
+ "gpt-4",
+ "claude-3",
+ "gpt-3.5-turbo",
+ ]
mock_get_provider.return_value = (None, "openai", None, None)
# Test accessible model
result = await model_info(
- model_id="gpt-4",
- user_api_key_dict=user_api_key_dict
+ model_id="gpt-4", user_api_key_dict=user_api_key_dict
)
assert result["id"] == "gpt-4"
@@ -688,22 +1101,25 @@ class TestModelInfoEndpoint:
@pytest.mark.asyncio
async def test_model_info_inaccessible_model_returns_404(self):
"""Test model_info returns 404 for inaccessible models"""
- from litellm.proxy.proxy_server import model_info
from fastapi import HTTPException
+ from litellm.proxy.proxy_server import model_info
+
# Mock user with limited access
user_api_key_dict = UserAPIKeyAuth(
user_id="test_user",
api_key="test_key",
models=["gpt-4"], # Only has access to gpt-4
- team_models=[]
+ team_models=[],
)
- with patch("litellm.proxy.proxy_server.llm_router") as mock_router, \
- patch("litellm.proxy.proxy_server.get_key_models") as mock_get_key_models, \
- patch("litellm.proxy.proxy_server.get_team_models") as mock_get_team_models, \
- patch("litellm.proxy.proxy_server.get_complete_model_list") as mock_get_complete_models:
-
+ with patch("litellm.proxy.proxy_server.llm_router") as mock_router, patch(
+ "litellm.proxy.proxy_server.get_key_models"
+ ) as mock_get_key_models, patch(
+ "litellm.proxy.proxy_server.get_team_models"
+ ) as mock_get_team_models, patch(
+ "litellm.proxy.proxy_server.get_complete_model_list"
+ ) as mock_get_complete_models:
# Setup mocks - user only has access to gpt-4
mock_router.get_model_names.return_value = ["gpt-4", "claude-3"]
mock_router.get_model_access_groups.return_value = {}
@@ -715,32 +1131,35 @@ class TestModelInfoEndpoint:
with pytest.raises(HTTPException) as exc_info:
await model_info(
model_id="claude-3", # Not in user's accessible models
- user_api_key_dict=user_api_key_dict
+ user_api_key_dict=user_api_key_dict,
)
-
+
assert exc_info.value.status_code == 404
assert "does not exist or is not accessible" in exc_info.value.detail
- @pytest.mark.asyncio
+ @pytest.mark.asyncio
async def test_model_info_team_model_access(self):
"""Test model_info works with team model access"""
from litellm.proxy.proxy_server import model_info
-
+
# Mock user with team access
user_api_key_dict = UserAPIKeyAuth(
user_id="test_user",
- api_key="test_key",
+ api_key="test_key",
team_id="test_team",
models=[], # No direct key models
- team_models=["team-model-1"]
+ team_models=["team-model-1"],
)
- with patch("litellm.proxy.proxy_server.llm_router") as mock_router, \
- patch("litellm.proxy.proxy_server.get_key_models") as mock_get_key_models, \
- patch("litellm.proxy.proxy_server.get_team_models") as mock_get_team_models, \
- patch("litellm.proxy.proxy_server.get_complete_model_list") as mock_get_complete_models, \
- patch("litellm.get_llm_provider") as mock_get_provider:
-
+ with patch("litellm.proxy.proxy_server.llm_router") as mock_router, patch(
+ "litellm.proxy.proxy_server.get_key_models"
+ ) as mock_get_key_models, patch(
+ "litellm.proxy.proxy_server.get_team_models"
+ ) as mock_get_team_models, patch(
+ "litellm.proxy.proxy_server.get_complete_model_list"
+ ) as mock_get_complete_models, patch(
+ "litellm.get_llm_provider"
+ ) as mock_get_provider:
# Setup mocks
mock_router.get_model_names.return_value = ["team-model-1"]
mock_router.get_model_access_groups.return_value = {}
@@ -751,10 +1170,9 @@ class TestModelInfoEndpoint:
# Test team model access
result = await model_info(
- model_id="team-model-1",
- user_api_key_dict=user_api_key_dict
+ model_id="team-model-1", user_api_key_dict=user_api_key_dict
)
assert result["id"] == "team-model-1"
- assert result["object"] == "model"
+ assert result["object"] == "model"
assert result["owned_by"] == "custom"
diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
index ff636ca04ae..f9c7cefcc4c 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
@@ -24,6 +24,7 @@ from litellm.proxy.management_endpoints.ui_sso import (
MicrosoftSSOHandler,
SSOAuthenticationHandler,
_setup_team_mappings,
+ _sync_user_role_from_jwt_role_map,
determine_role_from_groups,
normalize_email,
process_sso_jwt_access_token,
@@ -1321,7 +1322,7 @@ async def test_get_generic_sso_response_with_additional_headers():
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
):
# Act
- result, received_response = await get_generic_sso_response(
+ result, received_response, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id=generic_client_id,
@@ -1383,7 +1384,7 @@ async def test_get_generic_sso_response_with_empty_headers():
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
):
# Act
- result, received_response = await get_generic_sso_response(
+ result, received_response, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id=generic_client_id,
@@ -5254,3 +5255,159 @@ class TestValidateReturnTo:
)
SSOAuthenticationHandler._validate_return_to("https://cp.example.com:3000/ui")
+
+class TestSyncUserRoleFromJwtRoleMap:
+ """Tests for _sync_user_role_from_jwt_role_map."""
+
+ @staticmethod
+ def _make_jwt_handler():
+ from litellm.caching.caching import DualCache
+ from litellm.proxy._types import (
+ JWTLiteLLMRoleMap,
+ LiteLLM_JWTAuth,
+ LitellmUserRoles,
+ )
+
+ handler = JWTHandler()
+ handler.update_environment(
+ prisma_client=None,
+ user_api_key_cache=DualCache(),
+ litellm_jwtauth=LiteLLM_JWTAuth(
+ roles_jwt_field="custom_roles",
+ user_id_upsert=True,
+ sync_user_role_and_teams=True,
+ jwt_litellm_role_map=[
+ JWTLiteLLMRoleMap(
+ jwt_role="my-admin",
+ litellm_role=LitellmUserRoles.PROXY_ADMIN,
+ ),
+ JWTLiteLLMRoleMap(
+ jwt_role="my-viewer",
+ litellm_role=LitellmUserRoles.INTERNAL_USER,
+ ),
+ ],
+ ),
+ )
+ return handler
+
+ @staticmethod
+ def _make_sso_values(user_role=None):
+ from litellm.proxy._types import SSOUserDefinedValues
+
+ user_id = "testuser@example.com"
+ return SSOUserDefinedValues(
+ models=[],
+ user_id=user_id,
+ user_email=user_id,
+ user_role=user_role,
+ max_budget=None,
+ budget_duration=None,
+ )
+
+ @pytest.mark.asyncio
+ async def test_stripped_response_has_no_roles(self):
+ """Bug repro: stripped received_response lacks role claims."""
+ from litellm.caching.caching import DualCache
+
+ handler = self._make_jwt_handler()
+ sso_values = self._make_sso_values()
+
+ await _sync_user_role_from_jwt_role_map(
+ jwt_handler=handler,
+ received_response={"token_type": "Bearer", "expires_in": 3600},
+ user_info=None,
+ prisma_client=AsyncMock(),
+ user_api_key_cache=DualCache(),
+ user_defined_values=sso_values,
+ )
+
+ assert sso_values["user_role"] is None
+
+ @pytest.mark.asyncio
+ async def test_decoded_access_token_maps_role(self):
+ """Decoded JWT payload with role claims maps correctly."""
+ from litellm.caching.caching import DualCache
+ from litellm.proxy._types import LitellmUserRoles
+
+ handler = self._make_jwt_handler()
+ sso_values = self._make_sso_values()
+
+ await _sync_user_role_from_jwt_role_map(
+ jwt_handler=handler,
+ received_response={"sub": "testuser@example.com", "custom_roles": ["my-admin"]},
+ user_info=None,
+ prisma_client=AsyncMock(),
+ user_api_key_cache=DualCache(),
+ user_defined_values=sso_values,
+ )
+
+ assert sso_values["user_role"] == LitellmUserRoles.PROXY_ADMIN.value
+
+ @pytest.mark.asyncio
+ async def test_existing_user_role_updated_in_db_and_cache(self):
+ """Existing user with stale role gets updated in DB and cache."""
+ from litellm.caching.caching import DualCache
+ from litellm.proxy._types import LitellmUserRoles
+
+ handler = self._make_jwt_handler()
+ cache = DualCache()
+ prisma = AsyncMock()
+ prisma.db.litellm_usertable.update = AsyncMock()
+ user_id = "testuser@example.com"
+
+ existing_user = LiteLLM_UserTable(
+ user_id=user_id,
+ user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
+ )
+ await cache.async_set_cache(key=user_id, value=existing_user.model_dump(), ttl=60)
+
+ sso_values = self._make_sso_values(
+ user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
+ )
+
+ await _sync_user_role_from_jwt_role_map(
+ jwt_handler=handler,
+ received_response={"sub": user_id, "custom_roles": ["my-admin"]},
+ user_info=existing_user,
+ prisma_client=prisma,
+ user_api_key_cache=cache,
+ user_defined_values=sso_values,
+ )
+
+ prisma.db.litellm_usertable.update.assert_called_once_with(
+ where={"user_id": user_id},
+ data={"user_role": LitellmUserRoles.PROXY_ADMIN.value},
+ )
+ assert existing_user.user_role == LitellmUserRoles.PROXY_ADMIN.value
+ assert sso_values["user_role"] == LitellmUserRoles.PROXY_ADMIN.value
+
+ @pytest.mark.asyncio
+ async def test_same_role_no_db_write(self):
+ """No DB update when the mapped role matches the existing role."""
+ from litellm.caching.caching import DualCache
+ from litellm.proxy._types import LitellmUserRoles
+
+ handler = self._make_jwt_handler()
+ prisma = AsyncMock()
+ prisma.db.litellm_usertable.update = AsyncMock()
+
+ existing_user = LiteLLM_UserTable(
+ user_id="testuser@example.com",
+ user_role=LitellmUserRoles.PROXY_ADMIN.value,
+ )
+
+ sso_values = self._make_sso_values(
+ user_role=LitellmUserRoles.PROXY_ADMIN.value,
+ )
+
+ await _sync_user_role_from_jwt_role_map(
+ jwt_handler=handler,
+ received_response={"sub": "testuser@example.com", "custom_roles": ["my-admin"]},
+ user_info=existing_user,
+ prisma_client=prisma,
+ user_api_key_cache=DualCache(),
+ user_defined_values=sso_values,
+ )
+
+ prisma.db.litellm_usertable.update.assert_not_called()
+
diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py
index ed5b8cbd81c..c6a03cf4ecd 100644
--- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py
+++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py
@@ -34,8 +34,8 @@ def llm_router() -> Router:
"model_name": "azure-gpt-3-5-turbo",
"litellm_params": {
"model": "azure/chatgpt-v-2",
- "api_key": "azure_api_key",
- "api_base": "azure_api_base",
+ "api_key": "AZURE_AI_API_KEY",
+ "api_base": "AZURE_AI_API_BASE",
"api_version": "azure_api_version",
},
"model_info": {
@@ -106,16 +106,18 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router:
Asserts 'create_file' is called with the correct arguments
"""
import litellm
+ import litellm.proxy.proxy_server as ps
from litellm import Router
from litellm.proxy._types import LitellmUserRoles
- import litellm.proxy.proxy_server as ps
from litellm.proxy.utils import ProxyLogging
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
# Mock create_file as an async function
- mock_create_file = mocker.patch("litellm.files.main.create_file", new=mocker.AsyncMock())
+ mock_create_file = mocker.patch(
+ "litellm.files.main.create_file", new=mocker.AsyncMock()
+ )
proxy_logging_obj = ProxyLogging(
user_api_key_cache=DualCache(default_in_memory_ttl=1)
@@ -127,7 +129,14 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router:
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
class DummyManagedFiles(BaseFileEndpoints):
- async def acreate_file(self, llm_router, create_file_request, target_model_names_list, litellm_parent_otel_span, user_api_key_dict):
+ async def acreate_file(
+ self,
+ llm_router,
+ create_file_request,
+ target_model_names_list,
+ litellm_parent_otel_span,
+ user_api_key_dict,
+ ):
# Handle both dict and object forms of create_file_request
if isinstance(create_file_request, dict):
file_data = create_file_request.get("file")
@@ -135,12 +144,12 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router:
else:
file_data = create_file_request.file
purpose_data = create_file_request.purpose
-
+
# Call the mocked litellm.files.main.create_file to ensure asserts work
await litellm.files.main.create_file(
custom_llm_provider="azure",
model="azure/chatgpt-v-2",
- api_key="azure_api_key",
+ api_key="AZURE_AI_API_KEY",
file=file_data[1],
purpose=purpose_data,
)
@@ -153,6 +162,7 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router:
)
# Return a dummy response object as needed by the test
from litellm.types.llms.openai import OpenAIFileObject
+
return OpenAIFileObject(
id="dummy-id",
object="file",
@@ -162,17 +172,21 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router:
purpose=purpose_data,
status="uploaded",
)
-
+
async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router):
raise NotImplementedError("Not implemented for test")
-
+
async def afile_list(self, purpose, litellm_parent_otel_span):
raise NotImplementedError("Not implemented for test")
-
- async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_delete(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
-
- async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_content(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
# Manually add the hook to the proxy_hook_mapping
@@ -214,7 +228,7 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router:
if (
kwargs.get("custom_llm_provider") == "azure"
and kwargs.get("model") == "azure/chatgpt-v-2"
- and kwargs.get("api_key") == "azure_api_key"
+ and kwargs.get("api_key") == "AZURE_AI_API_KEY"
):
azure_call_found = True
break
@@ -245,8 +259,8 @@ def test_target_storage_invokes_storage_backend(
"""
Ensure target_storage is parsed and invokes the storage backend service.
"""
- from litellm.proxy._types import LitellmUserRoles
import litellm.proxy.proxy_server as ps
+ from litellm.proxy._types import LitellmUserRoles
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
@@ -304,8 +318,8 @@ def test_target_storage_with_target_models(
"""
Ensure target_storage and target_model_names are parsed and passed through.
"""
- from litellm.proxy._types import LitellmUserRoles
import litellm.proxy.proxy_server as ps
+ from litellm.proxy._types import LitellmUserRoles
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
@@ -611,7 +625,9 @@ def test_create_file_for_each_model(
assert openai_call_found, "OpenAI call not found with expected parameters"
-def test_create_file_with_expires_after(mocker: MockerFixture, monkeypatch, llm_router: Router):
+def test_create_file_with_expires_after(
+ mocker: MockerFixture, monkeypatch, llm_router: Router
+):
"""
Test that expires_after is properly parsed and passed through when creating a file
"""
@@ -624,18 +640,25 @@ def test_create_file_with_expires_after(mocker: MockerFixture, monkeypatch, llm_
proxy_logging_obj._add_proxy_hooks(llm_router)
class DummyManagedFiles(BaseFileEndpoints):
- async def acreate_file(self, llm_router, create_file_request, target_model_names_list, litellm_parent_otel_span, user_api_key_dict):
+ async def acreate_file(
+ self,
+ llm_router,
+ create_file_request,
+ target_model_names_list,
+ litellm_parent_otel_span,
+ user_api_key_dict,
+ ):
# Verify expires_after is in the request
if isinstance(create_file_request, dict):
expires_after = create_file_request.get("expires_after")
else:
expires_after = getattr(create_file_request, "expires_after", None)
-
+
# Verify expires_after was passed correctly
assert expires_after is not None, "expires_after should be in the request"
assert expires_after["anchor"] == "created_at"
assert expires_after["seconds"] == 2592000
-
+
# Return a dummy response
return OpenAIFileObject(
id="file-abc123",
@@ -646,17 +669,21 @@ def test_create_file_with_expires_after(mocker: MockerFixture, monkeypatch, llm_
purpose="fine-tune",
status="uploaded",
)
-
+
async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router):
raise NotImplementedError("Not implemented for test")
-
+
async def afile_list(self, purpose, litellm_parent_otel_span):
raise NotImplementedError("Not implemented for test")
-
- async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_delete(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
-
- async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_content(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
proxy_logging_obj.proxy_hook_mapping["managed_files"] = DummyManagedFiles()
@@ -688,7 +715,9 @@ def test_create_file_with_expires_after(mocker: MockerFixture, monkeypatch, llm_
assert result["purpose"] == "fine-tune"
-def test_create_file_with_expires_after_missing_anchor(mocker: MockerFixture, monkeypatch, llm_router: Router):
+def test_create_file_with_expires_after_missing_anchor(
+ mocker: MockerFixture, monkeypatch, llm_router: Router
+):
"""
Test that an error is returned when expires_after[anchor] is missing
"""
@@ -717,10 +746,15 @@ def test_create_file_with_expires_after_missing_anchor(mocker: MockerFixture, mo
assert response.status_code == 400
error_detail = response.json()
- assert "expires_after" in error_detail["error"]["message"].lower() or "both" in error_detail["error"]["message"].lower()
+ assert (
+ "expires_after" in error_detail["error"]["message"].lower()
+ or "both" in error_detail["error"]["message"].lower()
+ )
-def test_create_file_with_expires_after_missing_seconds(mocker: MockerFixture, monkeypatch, llm_router: Router):
+def test_create_file_with_expires_after_missing_seconds(
+ mocker: MockerFixture, monkeypatch, llm_router: Router
+):
"""
Test that an error is returned when expires_after[seconds] is missing
"""
@@ -749,10 +783,15 @@ def test_create_file_with_expires_after_missing_seconds(mocker: MockerFixture, m
assert response.status_code == 400
error_detail = response.json()
- assert "expires_after" in error_detail["error"]["message"].lower() or "both" in error_detail["error"]["message"].lower()
+ assert (
+ "expires_after" in error_detail["error"]["message"].lower()
+ or "both" in error_detail["error"]["message"].lower()
+ )
-def test_create_file_with_expires_after_valid_values(mocker: MockerFixture, monkeypatch, llm_router: Router):
+def test_create_file_with_expires_after_valid_values(
+ mocker: MockerFixture, monkeypatch, llm_router: Router
+):
"""
Test that expires_after works with valid anchor and seconds values
"""
@@ -765,18 +804,25 @@ def test_create_file_with_expires_after_valid_values(mocker: MockerFixture, monk
proxy_logging_obj._add_proxy_hooks(llm_router)
class DummyManagedFiles(BaseFileEndpoints):
- async def acreate_file(self, llm_router, create_file_request, target_model_names_list, litellm_parent_otel_span, user_api_key_dict):
+ async def acreate_file(
+ self,
+ llm_router,
+ create_file_request,
+ target_model_names_list,
+ litellm_parent_otel_span,
+ user_api_key_dict,
+ ):
# Verify expires_after is in the request
if isinstance(create_file_request, dict):
expires_after = create_file_request.get("expires_after")
else:
expires_after = getattr(create_file_request, "expires_after", None)
-
+
# Verify expires_after was passed correctly
assert expires_after is not None, "expires_after should be in the request"
assert expires_after["anchor"] == "created_at"
assert expires_after["seconds"] == 3600
-
+
return OpenAIFileObject(
id="file-abc123",
object="file",
@@ -786,17 +832,21 @@ def test_create_file_with_expires_after_valid_values(mocker: MockerFixture, monk
purpose="fine-tune",
status="uploaded",
)
-
+
async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router):
raise NotImplementedError("Not implemented for test")
-
+
async def afile_list(self, purpose, litellm_parent_otel_span):
raise NotImplementedError("Not implemented for test")
-
- async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_delete(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
-
- async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_content(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
proxy_logging_obj.proxy_hook_mapping["managed_files"] = DummyManagedFiles()
@@ -827,7 +877,9 @@ def test_create_file_with_expires_after_valid_values(mocker: MockerFixture, monk
assert result["purpose"] == "fine-tune"
-def test_create_file_without_expires_after(mocker: MockerFixture, monkeypatch, llm_router: Router):
+def test_create_file_without_expires_after(
+ mocker: MockerFixture, monkeypatch, llm_router: Router
+):
"""
Test that file creation works normally without expires_after
"""
@@ -840,16 +892,25 @@ def test_create_file_without_expires_after(mocker: MockerFixture, monkeypatch, l
proxy_logging_obj._add_proxy_hooks(llm_router)
class DummyManagedFiles(BaseFileEndpoints):
- async def acreate_file(self, llm_router, create_file_request, target_model_names_list, litellm_parent_otel_span, user_api_key_dict):
+ async def acreate_file(
+ self,
+ llm_router,
+ create_file_request,
+ target_model_names_list,
+ litellm_parent_otel_span,
+ user_api_key_dict,
+ ):
# Verify expires_after is None when not provided
if isinstance(create_file_request, dict):
expires_after = create_file_request.get("expires_after")
else:
expires_after = getattr(create_file_request, "expires_after", None)
-
+
# expires_after should be None when not provided
- assert expires_after is None, "expires_after should be None when not provided"
-
+ assert (
+ expires_after is None
+ ), "expires_after should be None when not provided"
+
return OpenAIFileObject(
id="file-abc123",
object="file",
@@ -859,17 +920,21 @@ def test_create_file_without_expires_after(mocker: MockerFixture, monkeypatch, l
purpose="fine-tune",
status="uploaded",
)
-
+
async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router):
raise NotImplementedError("Not implemented for test")
-
+
async def afile_list(self, purpose, litellm_parent_otel_span):
raise NotImplementedError("Not implemented for test")
-
- async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_delete(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
-
- async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_content(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
proxy_logging_obj.proxy_hook_mapping["managed_files"] = DummyManagedFiles()
@@ -898,11 +963,13 @@ def test_create_file_without_expires_after(mocker: MockerFixture, monkeypatch, l
assert result["purpose"] == "fine-tune"
-def test_managed_files_with_loadbalancing(mocker: MockerFixture, monkeypatch, llm_router: Router):
+def test_managed_files_with_loadbalancing(
+ mocker: MockerFixture, monkeypatch, llm_router: Router
+):
"""
Test that managed files work with loadbalancing when both target_model_names
and enable_loadbalancing_on_batch_endpoints are enabled.
-
+
This ensures that the priority order is correct:
- managed files should take precedence over deprecated loadbalancing
- managed files internally use llm_router.acreate_file() which provides loadbalancing
@@ -912,28 +979,34 @@ def test_managed_files_with_loadbalancing(mocker: MockerFixture, monkeypatch, ll
# Enable loadbalancing on batch endpoints
monkeypatch.setattr("litellm.enable_loadbalancing_on_batch_endpoints", True)
-
+
proxy_logging_obj = ProxyLogging(
user_api_key_cache=DualCache(default_in_memory_ttl=1)
)
proxy_logging_obj._add_proxy_hooks(llm_router)
-
+
# Track calls to verify loadbalancing through router
router_acreate_file_calls = []
-
+
class ManagedFilesWithLoadbalancing(BaseFileEndpoints):
- async def acreate_file(self, llm_router, create_file_request, target_model_names_list, litellm_parent_otel_span, user_api_key_dict):
+ async def acreate_file(
+ self,
+ llm_router,
+ create_file_request,
+ target_model_names_list,
+ litellm_parent_otel_span,
+ user_api_key_dict,
+ ):
# Verify we receive the target model names
- assert len(target_model_names_list) > 0, "Should have target_model_names_list"
-
+ assert (
+ len(target_model_names_list) > 0
+ ), "Should have target_model_names_list"
+
# Simulate what managed files does - call llm_router.acreate_file for each model
# This is where loadbalancing happens internally
for model in target_model_names_list:
- router_acreate_file_calls.append({
- "model": model,
- "via_router": True
- })
-
+ router_acreate_file_calls.append({"model": model, "via_router": True})
+
# Return a managed file ID (base64 encoded)
return OpenAIFileObject(
id="litellm_managed_file_abc123",
@@ -944,23 +1017,29 @@ def test_managed_files_with_loadbalancing(mocker: MockerFixture, monkeypatch, ll
purpose="batch",
status="uploaded",
)
-
+
async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router):
raise NotImplementedError("Not implemented for test")
-
+
async def afile_list(self, purpose, litellm_parent_otel_span):
raise NotImplementedError("Not implemented for test")
-
- async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_delete(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
-
- async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_content(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
-
+
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
- proxy_logging_obj.proxy_hook_mapping["managed_files"] = ManagedFilesWithLoadbalancing()
+ proxy_logging_obj.proxy_hook_mapping[
+ "managed_files"
+ ] = ManagedFilesWithLoadbalancing()
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj
@@ -971,12 +1050,12 @@ def test_managed_files_with_loadbalancing(mocker: MockerFixture, monkeypatch, ll
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="test-key", user_role=LitellmUserRoles.PROXY_ADMIN
)
-
+
try:
# Create batch file content
test_file_content = b'{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]}}'
test_file = ("batch_data.jsonl", test_file_content, "application/jsonl")
-
+
# Make request with both target_model_names AND enable_loadbalancing_on_batch_endpoints
response = client.post(
"/v1/files",
@@ -987,7 +1066,7 @@ def test_managed_files_with_loadbalancing(mocker: MockerFixture, monkeypatch, ll
},
headers={"Authorization": "Bearer test-key"},
)
-
+
# Verify success
assert response.status_code == 200, response.text
finally:
@@ -995,13 +1074,17 @@ def test_managed_files_with_loadbalancing(mocker: MockerFixture, monkeypatch, ll
result = response.json()
assert result["id"] == "litellm_managed_file_abc123"
assert result["purpose"] == "batch"
-
+
# Verify that managed files was called (via router for loadbalancing)
# This proves that managed files took precedence over deprecated loadbalancing
- assert len(router_acreate_file_calls) == 2, "Should have called router for both models"
+ assert (
+ len(router_acreate_file_calls) == 2
+ ), "Should have called router for both models"
assert router_acreate_file_calls[0]["model"] == "azure-gpt-3-5-turbo"
assert router_acreate_file_calls[1]["model"] == "gpt-3.5-turbo"
- assert all(call["via_router"] for call in router_acreate_file_calls), "All calls should go through router"
+ assert all(
+ call["via_router"] for call in router_acreate_file_calls
+ ), "All calls should go through router"
def test_create_file_with_nested_litellm_metadata(
@@ -1009,22 +1092,29 @@ def test_create_file_with_nested_litellm_metadata(
):
"""
Test that nested litellm_metadata is correctly parsed from form data in bracket notation.
-
+
Regression test for: litellm_metadata[spend_logs_metadata][owner] format should be
correctly parsed into nested dictionary structure.
"""
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.types.llms.openai import OpenAIFileObject
-
+
proxy_logging_obj = ProxyLogging(
user_api_key_cache=DualCache(default_in_memory_ttl=1)
)
proxy_logging_obj._add_proxy_hooks(llm_router)
-
+
captured_litellm_metadata = {}
-
+
class DummyManagedFiles(BaseFileEndpoints):
- async def acreate_file(self, llm_router, create_file_request, target_model_names_list, litellm_parent_otel_span, user_api_key_dict):
+ async def acreate_file(
+ self,
+ llm_router,
+ create_file_request,
+ target_model_names_list,
+ litellm_parent_otel_span,
+ user_api_key_dict,
+ ):
# Capture litellm_metadata for verification
if isinstance(create_file_request, dict):
captured_litellm_metadata.update(
@@ -1034,7 +1124,7 @@ def test_create_file_with_nested_litellm_metadata(
captured_litellm_metadata.update(
getattr(create_file_request, "litellm_metadata", {})
)
-
+
return OpenAIFileObject(
id="file-test-123",
object="file",
@@ -1044,28 +1134,32 @@ def test_create_file_with_nested_litellm_metadata(
purpose="fine-tune",
status="uploaded",
)
-
+
async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router):
raise NotImplementedError("Not implemented for test")
-
+
async def afile_list(self, purpose, litellm_parent_otel_span):
raise NotImplementedError("Not implemented for test")
-
- async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_delete(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
-
- async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_content(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
-
+
proxy_logging_obj.proxy_hook_mapping["managed_files"] = DummyManagedFiles()
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj
)
-
+
test_file_content = b'{"prompt": "Hello", "completion": "Hi"}'
test_file = ("test.jsonl", test_file_content, "application/jsonl")
-
+
# Test with nested litellm_metadata in bracket notation
response = client.post(
"/v1/files",
@@ -1080,12 +1174,12 @@ def test_create_file_with_nested_litellm_metadata(
},
headers={"Authorization": "Bearer test-key"},
)
-
+
# Verify success
assert response.status_code == 200
result = response.json()
assert result["id"] == "file-test-123"
-
+
# Verify nested metadata was correctly parsed
assert "spend_logs_metadata" in captured_litellm_metadata
assert captured_litellm_metadata["spend_logs_metadata"]["owner"] == "john_doe"
@@ -1099,26 +1193,33 @@ def test_create_file_with_deep_nested_litellm_metadata(
):
"""
Test that deeply nested litellm_metadata is correctly parsed from form data.
-
+
Regression test for: litellm_metadata[a][b][c] format should be correctly parsed.
"""
+ import litellm.proxy.proxy_server as ps
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.proxy._types import LitellmUserRoles
- import litellm.proxy.proxy_server as ps
from litellm.types.llms.openai import OpenAIFileObject
-
+
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
-
+
proxy_logging_obj = ProxyLogging(
user_api_key_cache=DualCache(default_in_memory_ttl=1)
)
proxy_logging_obj._add_proxy_hooks(llm_router)
-
+
captured_litellm_metadata = {}
-
+
class DummyManagedFiles(BaseFileEndpoints):
- async def acreate_file(self, llm_router, create_file_request, target_model_names_list, litellm_parent_otel_span, user_api_key_dict):
+ async def acreate_file(
+ self,
+ llm_router,
+ create_file_request,
+ target_model_names_list,
+ litellm_parent_otel_span,
+ user_api_key_dict,
+ ):
if isinstance(create_file_request, dict):
captured_litellm_metadata.update(
create_file_request.get("litellm_metadata", {})
@@ -1127,7 +1228,7 @@ def test_create_file_with_deep_nested_litellm_metadata(
captured_litellm_metadata.update(
getattr(create_file_request, "litellm_metadata", {})
)
-
+
return OpenAIFileObject(
id="file-test-456",
object="file",
@@ -1137,33 +1238,37 @@ def test_create_file_with_deep_nested_litellm_metadata(
purpose="batch",
status="uploaded",
)
-
+
async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router):
raise NotImplementedError("Not implemented for test")
-
+
async def afile_list(self, purpose, litellm_parent_otel_span):
raise NotImplementedError("Not implemented for test")
-
- async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_delete(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
-
- async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data):
+
+ async def afile_content(
+ self, file_id, litellm_parent_otel_span, llm_router, **data
+ ):
raise NotImplementedError("Not implemented for test")
-
+
proxy_logging_obj.proxy_hook_mapping["managed_files"] = DummyManagedFiles()
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj
)
-
+
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"
)
-
+
try:
test_file_content = b'{"custom_id": "req-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo"}}'
test_file = ("nested.jsonl", test_file_content, "application/jsonl")
-
+
# Test with deeply nested metadata
response = client.post(
"/v1/files",
@@ -1177,12 +1282,12 @@ def test_create_file_with_deep_nested_litellm_metadata(
},
headers={"Authorization": "Bearer test-key"},
)
-
+
# Verify success
assert response.status_code == 200, response.text
result = response.json()
assert result["id"] == "file-test-456"
-
+
# Verify deeply nested metadata was correctly parsed
assert "config" in captured_litellm_metadata
assert "database" in captured_litellm_metadata["config"]
@@ -1356,7 +1461,9 @@ def test_file_team_injects_when_caller_sends_nothing(
# ---------------------------------------------------------------------------
-def _post_file_raw(monkeypatch, llm_router: Router, team_metadata: dict, form_data: dict):
+def _post_file_raw(
+ monkeypatch, llm_router: Router, team_metadata: dict, form_data: dict
+):
"""POST /v1/files and return the raw response (no status assertion)."""
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py
index 340d8df4dd1..11bae9009a2 100644
--- a/tests/test_litellm/proxy/test_proxy_server.py
+++ b/tests/test_litellm/proxy/test_proxy_server.py
@@ -104,10 +104,7 @@ def test_login_v2_returns_redirect_url_and_sets_cookie(monkeypatch):
)
assert response.status_code == 200
- assert response.json() == {
- "redirect_url": "http://testserver/ui/?login=success",
- "token": "signed-token",
- }
+ assert response.json() == {"redirect_url": "http://testserver/ui/?login=success"}
assert response.cookies.get("token") == "signed-token"
mock_authenticate_user.assert_awaited_once_with(
@@ -516,15 +513,11 @@ def test_restructure_ui_html_files_handles_nested_routes(tmp_path):
assert (ui_root / "home" / "index.html").read_text() == "home"
assert not (ui_root / "mcp" / "oauth" / "callback.html").exists()
assert (
- (ui_root / "mcp" / "oauth" / "callback" / "index.html").read_text()
- == "callback"
- )
+ ui_root / "mcp" / "oauth" / "callback" / "index.html"
+ ).read_text() == "callback"
assert (ui_root / "existing" / "index.html").read_text() == "keep"
assert (ui_root / "_next" / "ignore.html").read_text() == "asset"
- assert (
- (ui_root / "litellm-asset-prefix" / "ignore.html").read_text()
- == "asset"
- )
+ assert (ui_root / "litellm-asset-prefix" / "ignore.html").read_text() == "asset"
def test_ui_extensionless_route_requires_restructure(tmp_path):
@@ -541,9 +534,7 @@ def test_ui_extensionless_route_requires_restructure(tmp_path):
(ui_root / "login.html").write_text("login")
fastapi_app = FastAPI()
- fastapi_app.mount(
- "/ui", StaticFiles(directory=str(ui_root), html=True), name="ui"
- )
+ fastapi_app.mount("/ui", StaticFiles(directory=str(ui_root), html=True), name="ui")
client = TestClient(fastapi_app)
assert client.get("/ui/login.html").status_code == 200
@@ -564,37 +555,37 @@ def test_restructure_always_happens(monkeypatch):
"""
# Test Case 1: is_non_root is True - restructuring happens in /var/lib/litellm/ui
monkeypatch.setenv("LITELLM_NON_ROOT", "true")
-
+
runtime_ui_path = "/var/lib/litellm/ui"
packaged_ui_path = "/some/packaged/ui/path"
-
+
# Simulate the logic from proxy_server.py
is_non_root = os.getenv("LITELLM_NON_ROOT", "").lower() == "true"
if is_non_root:
ui_path = runtime_ui_path
else:
ui_path = packaged_ui_path
-
+
# Restructuring always happens now, regardless of ui_path vs packaged_ui_path
should_restructure = True
-
+
assert is_non_root is True
assert should_restructure is True
assert ui_path == runtime_ui_path
-
+
# Test Case 2: is_non_root is False - restructuring happens directly in packaged_ui_path
monkeypatch.delenv("LITELLM_NON_ROOT", raising=False)
-
+
# Simulate the logic from proxy_server.py
is_non_root = os.getenv("LITELLM_NON_ROOT", "").lower() == "true"
if is_non_root:
ui_path = runtime_ui_path
else:
ui_path = packaged_ui_path
-
+
# Restructuring always happens now, even when ui_path == packaged_ui_path
should_restructure = True
-
+
assert is_non_root is False
assert should_restructure is True
assert ui_path == packaged_ui_path
@@ -691,9 +682,7 @@ def test_update_config_fields_deep_merge_db_wins():
"hidden": True,
},
# Demonstrate that None values from DB are skipped (preserve existing)
- "legacy-sonnet": {
- "hidden": None # should not clobber current True
- },
+ "legacy-sonnet": {"hidden": None}, # should not clobber current True
}
}
@@ -743,9 +732,7 @@ def test_get_config_custom_callback_api_env_vars(monkeypatch):
mock_router = MagicMock()
mock_router.get_settings.return_value = {}
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
- monkeypatch.setattr(
- proxy_config, "get_config", AsyncMock(return_value=config_data)
- )
+ monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data))
# Bypass auth dependency
original_overrides = app.dependency_overrides.copy()
@@ -923,7 +910,9 @@ def test_embedding_input_array_of_tokens(client_no_auth):
assert response.status_code == 200
result = response.json()
print(len(result["data"][0]["embedding"]))
- assert len(result["data"][0]["embedding"]) > 10 # this usually has len==1536 so
+ assert (
+ len(result["data"][0]["embedding"]) > 10
+ ) # this usually has len==1536 so
except Exception as e:
pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}")
@@ -1180,7 +1169,6 @@ async def test_delete_deployment_type_mismatch():
with patch("litellm.proxy.proxy_server.llm_router", mock_llm_router), patch(
"litellm.proxy.proxy_server.user_config_file_path", "test_config.yaml"
):
-
# Call the function under test
deleted_count = await pc._delete_deployment(db_models=[])
@@ -1322,6 +1310,7 @@ def test_normalize_datetime_for_sorting():
# Test Case 6: Timezone-aware datetime object (non-UTC)
from datetime import timedelta
+
aware_dt = datetime(2024, 1, 15, 10, 30, 0, tzinfo=timezone(timedelta(hours=5)))
result = _normalize_datetime_for_sorting(aware_dt)
assert result is not None
@@ -1574,6 +1563,91 @@ async def test_load_environment_variables_litellm_license_and_edge_cases():
assert "FAILED_SECRET" not in os.environ
+@pytest.mark.asyncio
+async def test_load_environment_variables_blocks_dangerous_keys():
+ """
+ Test that _load_environment_variables rejects dangerous env var keys
+ like PATH, LD_PRELOAD, PYTHONPATH, etc.
+ """
+ import logging
+
+ from litellm.proxy.proxy_server import ProxyConfig
+
+ proxy_config = ProxyConfig()
+
+ original_path = os.environ.get("PATH", "")
+
+ test_config = {
+ "environment_variables": {
+ "PATH": "/tmp/evil",
+ "LD_PRELOAD": "/tmp/evil.so",
+ "PYTHONPATH": "/tmp/evil",
+ "SAFE_CUSTOM_VAR": "safe_value",
+ }
+ }
+
+ with patch.dict(os.environ, {}, clear=False):
+ proxy_config._load_environment_variables(test_config)
+
+ # Blocked keys should not be set to the attacker value
+ assert os.environ.get("PATH") != "/tmp/evil"
+ assert (
+ "LD_PRELOAD" not in os.environ or os.environ["LD_PRELOAD"] != "/tmp/evil.so"
+ )
+ assert os.environ.get("PYTHONPATH") != "/tmp/evil"
+
+ # Safe keys should still be set
+ assert os.environ["SAFE_CUSTOM_VAR"] == "safe_value"
+
+
+@pytest.mark.asyncio
+async def test_load_environment_variables_allows_proxy_keys():
+ """
+ Test that HTTP_PROXY/HTTPS_PROXY are allowed since they are commonly used
+ in corporate environments to route outbound API calls.
+ """
+ from litellm.proxy.proxy_server import ProxyConfig
+
+ proxy_config = ProxyConfig()
+
+ test_config = {
+ "environment_variables": {
+ "HTTP_PROXY": "http://corp-proxy:8080",
+ "HTTPS_PROXY": "http://corp-proxy:8080",
+ }
+ }
+
+ with patch.dict(os.environ, {}, clear=False):
+ proxy_config._load_environment_variables(test_config)
+
+ assert os.environ["HTTP_PROXY"] == "http://corp-proxy:8080"
+ assert os.environ["HTTPS_PROXY"] == "http://corp-proxy:8080"
+
+
+@pytest.mark.asyncio
+async def test_load_environment_variables_blocks_no_proxy():
+ """
+ Test that NO_PROXY/no_proxy are blocked to prevent bypassing proxy-based
+ network monitoring.
+ """
+ from litellm.proxy.proxy_server import ProxyConfig
+
+ proxy_config = ProxyConfig()
+
+ test_config = {
+ "environment_variables": {
+ "NO_PROXY": "internal-service",
+ "no_proxy": "internal-service",
+ }
+ }
+
+ with patch.dict(os.environ, {}, clear=False):
+ proxy_config._load_environment_variables(test_config)
+
+ assert os.environ.get("NO_PROXY") != "internal-service"
+ assert os.environ.get("no_proxy") != "internal-service"
+
+
@pytest.mark.asyncio
async def test_write_config_to_file(monkeypatch):
"""
@@ -1882,7 +1956,6 @@ async def test_chat_completion_result_no_nested_none_values():
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing",
return_value=mock_base_processor,
):
-
# Call the chat_completion function
result = await chat_completion(
request=mock_request,
@@ -2027,9 +2100,7 @@ class TestPriceDataReloadAPI:
# Mock the database connection
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
- mock_prisma.db.litellm_config.find_unique = AsyncMock(
- return_value=None
- )
+ mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None)
response = client_with_auth.post("/reload/model_cost_map")
@@ -2372,8 +2443,13 @@ class TestPriceDataReloadIntegration:
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
# Simulate existing config with a schedule
mock_existing = MagicMock()
- mock_existing.param_value = {"interval_hours": 12, "force_reload": False}
- mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_existing)
+ mock_existing.param_value = {
+ "interval_hours": 12,
+ "force_reload": False,
+ }
+ mock_prisma.db.litellm_config.find_unique = AsyncMock(
+ return_value=mock_existing
+ )
mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None)
response = client.post("/reload/model_cost_map")
@@ -2415,7 +2491,9 @@ class TestPriceDataReloadIntegration:
) as mock_reload:
mock_reload.return_value = {"anthropic": {"beta_header": "test-value"}}
- asyncio.run(proxy_config._check_and_reload_anthropic_beta_headers(mock_prisma))
+ asyncio.run(
+ proxy_config._check_and_reload_anthropic_beta_headers(mock_prisma)
+ )
# Verify the upsert update branch preserves interval_hours
mock_prisma.db.litellm_config.upsert.assert_called()
@@ -2456,7 +2534,9 @@ class TestPriceDataReloadIntegration:
# Simulate existing config with a schedule
mock_existing = MagicMock()
mock_existing.param_value = {"interval_hours": 8, "force_reload": False}
- mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_existing)
+ mock_prisma.db.litellm_config.find_unique = AsyncMock(
+ return_value=mock_existing
+ )
mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None)
response = client.post("/reload/anthropic_beta_headers")
@@ -2821,7 +2901,7 @@ async def test_model_info_v1_oci_secrets_not_leaked():
mock_user_api_key_dict.api_key = "test-key"
mock_user_api_key_dict.team_models = []
mock_user_api_key_dict.models = ["oci-grok-test"]
-
+
# Mock model data with OCI sensitive information
mock_model_data = {
"model_name": "oci-grok-test",
@@ -2834,59 +2914,73 @@ async def test_model_info_v1_oci_secrets_not_leaked():
"oci_tenancy": "ocid1.tenancy.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk",
"oci_key_file": "/path/to/oci_api_key.pem",
"oci_compartment_id": "ocid1.compartment.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk",
- "drop_params": True
+ "drop_params": True,
},
- "model_info": {
- "mode": "completion",
- "id": "test-model-id"
- }
+ "model_info": {"mode": "completion", "id": "test-model-id"},
}
-
+
# Mock the llm_router to return our test data
mock_router = MagicMock()
mock_router.get_model_names.return_value = ["oci-grok-test"]
mock_router.get_model_access_groups.return_value = {}
mock_router.get_model_list.return_value = [mock_model_data]
-
+
# Mock global variables
- with patch("litellm.proxy.proxy_server.llm_router", mock_router), \
- patch("litellm.proxy.proxy_server.llm_model_list", [mock_model_data]), \
- patch("litellm.proxy.proxy_server.general_settings", {"infer_model_from_keys": False}), \
- patch("litellm.proxy.proxy_server.user_model", None):
-
+ with patch("litellm.proxy.proxy_server.llm_router", mock_router), patch(
+ "litellm.proxy.proxy_server.llm_model_list", [mock_model_data]
+ ), patch(
+ "litellm.proxy.proxy_server.general_settings", {"infer_model_from_keys": False}
+ ), patch(
+ "litellm.proxy.proxy_server.user_model", None
+ ):
# Call the model_info_v1 endpoint
result = await model_info_v1(
- user_api_key_dict=mock_user_api_key_dict,
- litellm_model_id=None
+ user_api_key_dict=mock_user_api_key_dict, litellm_model_id=None
)
-
+
# Verify the result structure
assert "data" in result
assert len(result["data"]) == 1
-
+
model_info = result["data"][0]
litellm_params = model_info["litellm_params"]
-
+
# Verify that sensitive OCI fields are masked
assert "****" in litellm_params["oci_key"], "oci_key should be masked"
- assert "****" in litellm_params["oci_fingerprint"], "oci_fingerprint should be masked"
+ assert (
+ "****" in litellm_params["oci_fingerprint"]
+ ), "oci_fingerprint should be masked"
assert "****" in litellm_params["oci_tenancy"], "oci_tenancy should be masked"
assert "****" in litellm_params["oci_key_file"], "oci_key_file should be masked"
-
+
# Verify that non-sensitive fields are NOT masked
- assert litellm_params["model"] == "oci/xai.grok-4", "model field should not be masked"
- assert litellm_params["oci_region"] == "us-phoenix-1", "oci_region should not be masked"
+ assert (
+ litellm_params["model"] == "oci/xai.grok-4"
+ ), "model field should not be masked"
+ assert (
+ litellm_params["oci_region"] == "us-phoenix-1"
+ ), "oci_region should not be masked"
assert litellm_params["drop_params"] is True, "drop_params should not be masked"
-
+
# Verify the model field specifically is not masked (this was the original issue)
- assert "****" not in litellm_params["model"], "model field should never be masked"
- assert litellm_params["model"].startswith("oci/"), "model should retain its full value"
-
+ assert (
+ "****" not in litellm_params["model"]
+ ), "model field should never be masked"
+ assert litellm_params["model"].startswith(
+ "oci/"
+ ), "model should retain its full value"
+
# Verify that actual secret values are not present in the response
result_str = str(result)
- assert "ocid1.api_key.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk" not in result_str
+ assert (
+ "ocid1.api_key.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk"
+ not in result_str
+ )
assert "aa:bb:cc:dd:ee:ff:11:22:33:44:55:66:77:88:99:00" not in result_str
- assert "ocid1.tenancy.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk" not in result_str
+ assert (
+ "ocid1.tenancy.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk"
+ not in result_str
+ )
assert "/path/to/oci_api_key.pem" not in result_str
@@ -2898,17 +2992,17 @@ def test_add_callback_from_db_to_in_memory_litellm_callbacks():
from unittest.mock import MagicMock, patch
from litellm.proxy.proxy_server import ProxyConfig
-
+
proxy_config = ProxyConfig()
-
+
# Mock the callback manager
mock_callback_manager = MagicMock()
-
+
with patch("litellm.proxy.proxy_server.litellm") as mock_litellm:
# Set up mock litellm attributes
mock_litellm._known_custom_logger_compatible_callbacks = []
mock_litellm.logging_callback_manager = mock_callback_manager
-
+
# Test Case 1: Add success callback
mock_success_callbacks = []
proxy_config._add_callback_from_db_to_in_memory_litellm_callbacks(
@@ -2916,9 +3010,11 @@ def test_add_callback_from_db_to_in_memory_litellm_callbacks():
event_types=["success"],
existing_callbacks=mock_success_callbacks,
)
- mock_callback_manager.add_litellm_success_callback.assert_called_once_with("prometheus")
+ mock_callback_manager.add_litellm_success_callback.assert_called_once_with(
+ "prometheus"
+ )
mock_callback_manager.reset_mock()
-
+
# Test Case 2: Add failure callback
mock_failure_callbacks = []
proxy_config._add_callback_from_db_to_in_memory_litellm_callbacks(
@@ -2926,9 +3022,11 @@ def test_add_callback_from_db_to_in_memory_litellm_callbacks():
event_types=["failure"],
existing_callbacks=mock_failure_callbacks,
)
- mock_callback_manager.add_litellm_failure_callback.assert_called_once_with("langfuse")
+ mock_callback_manager.add_litellm_failure_callback.assert_called_once_with(
+ "langfuse"
+ )
mock_callback_manager.reset_mock()
-
+
# Test Case 3: Add callback for both success and failure
mock_callbacks = []
proxy_config._add_callback_from_db_to_in_memory_litellm_callbacks(
@@ -2938,7 +3036,7 @@ def test_add_callback_from_db_to_in_memory_litellm_callbacks():
)
mock_callback_manager.add_litellm_callback.assert_called_once_with("s3")
mock_callback_manager.reset_mock()
-
+
# Test Case 4: Don't add callback if it already exists
existing_callbacks_with_item = ["prometheus"]
proxy_config._add_callback_from_db_to_in_memory_litellm_callbacks(
@@ -2952,7 +3050,7 @@ def test_add_callback_from_db_to_in_memory_litellm_callbacks():
def test_should_load_db_object_with_supported_db_objects():
"""
Test _should_load_db_object method with supported_db_objects configuration.
-
+
Verifies that when supported_db_objects is set, only specified object types
are loaded from the database.
"""
@@ -3056,8 +3154,12 @@ async def test_tag_cache_update_called():
"spend": 10.0,
}
- with patch.object(cache, "async_get_cache", new=AsyncMock(return_value=mock_tag_obj)) as mock_get_cache:
- with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache:
+ with patch.object(
+ cache, "async_get_cache", new=AsyncMock(return_value=mock_tag_obj)
+ ) as mock_get_cache:
+ with patch.object(
+ cache, "async_set_cache_pipeline", new=AsyncMock()
+ ) as mock_set_cache:
await litellm.proxy.proxy_server.update_cache(
token=None,
user_id=None,
@@ -3108,8 +3210,12 @@ async def test_tag_cache_update_multiple_tags():
return mock_tag2_obj
return None
- with patch.object(cache, "async_get_cache", new=AsyncMock(side_effect=mock_get_cache_side_effect)) as mock_get_cache:
- with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache:
+ with patch.object(
+ cache, "async_get_cache", new=AsyncMock(side_effect=mock_get_cache_side_effect)
+ ) as mock_get_cache:
+ with patch.object(
+ cache, "async_set_cache_pipeline", new=AsyncMock()
+ ) as mock_set_cache:
await litellm.proxy.proxy_server.update_cache(
token=None,
user_id=None,
@@ -3130,7 +3236,9 @@ async def test_tag_cache_update_multiple_tags():
assert len(cache_list) == 2
- tag_updates = {cache_key: cache_value for cache_key, cache_value in cache_list}
+ tag_updates = {
+ cache_key: cache_value for cache_key, cache_value in cache_list
+ }
assert "tag:tag1" in tag_updates
assert "tag:tag2" in tag_updates
assert tag_updates["tag:tag1"]["spend"] == 15.0
@@ -3250,7 +3358,9 @@ async def test_init_sso_settings_in_db_error_handling():
assert True
except Exception as e:
# The exception should be caught and logged, not propagated
- pytest.fail(f"Exception should have been caught and logged, but was raised: {e}")
+ pytest.fail(
+ f"Exception should have been caught and logged, but was raised: {e}"
+ )
@pytest.mark.asyncio
@@ -3353,11 +3463,15 @@ def test_get_prompt_spec_for_db_prompt_with_versions():
}
# Test version 1
- prompt_spec_v1 = proxy_config._get_prompt_spec_for_db_prompt(db_prompt=mock_prompt_v1)
+ prompt_spec_v1 = proxy_config._get_prompt_spec_for_db_prompt(
+ db_prompt=mock_prompt_v1
+ )
assert prompt_spec_v1.prompt_id == "chat_prompt.v1"
# Test version 2
- prompt_spec_v2 = proxy_config._get_prompt_spec_for_db_prompt(db_prompt=mock_prompt_v2)
+ prompt_spec_v2 = proxy_config._get_prompt_spec_for_db_prompt(
+ db_prompt=mock_prompt_v2
+ )
assert prompt_spec_v2.prompt_id == "chat_prompt.v2"
@@ -3372,15 +3486,15 @@ def test_root_redirect_when_docs_url_not_root_and_redirect_url_set(monkeypatch):
config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml"
# Ensure docs are mounted on a non-root path to trigger redirect logic
monkeypatch.setenv("DOCS_URL", "/docs")
-
+
test_redirect_url = "/ui"
monkeypatch.setenv("ROOT_REDIRECT_URL", test_redirect_url)
-
+
asyncio.run(initialize(config=config_fp, debug=True))
-
+
docs_url = _get_docs_url()
root_redirect_url = os.getenv("ROOT_REDIRECT_URL")
-
+
# Remove any existing "/" route that might interfere
routes_to_remove = []
for route in app.routes:
@@ -3389,16 +3503,17 @@ def test_root_redirect_when_docs_url_not_root_and_redirect_url_set(monkeypatch):
routes_to_remove.append(route)
elif not hasattr(route, "methods"): # Catch-all routes
routes_to_remove.append(route)
-
+
for route in routes_to_remove:
app.routes.remove(route)
-
+
# Add the redirect route if conditions are met (matching the actual implementation)
if docs_url != "/" and root_redirect_url:
+
@app.get("/", include_in_schema=False)
async def root_redirect():
return RedirectResponse(url=root_redirect_url)
-
+
client = TestClient(app)
response = client.get("/", follow_redirects=False)
assert response.status_code == 307
@@ -3422,12 +3537,13 @@ async def test_get_image_non_root_uses_var_lib_assets_dir(monkeypatch):
def exists_side_effect(path):
return False if path == "/var/lib/litellm/assets" else True
- with patch("litellm.proxy.proxy_server.os.makedirs") as mock_makedirs, \
- patch("litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect), \
- patch("litellm.proxy.proxy_server.os.access", return_value=True), \
- patch("litellm.proxy.proxy_server.os.getenv") as mock_getenv, \
- patch("litellm.proxy.proxy_server.FileResponse") as mock_file_response:
-
+ with patch("litellm.proxy.proxy_server.os.makedirs") as mock_makedirs, patch(
+ "litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect
+ ), patch("litellm.proxy.proxy_server.os.access", return_value=True), patch(
+ "litellm.proxy.proxy_server.os.getenv"
+ ) as mock_getenv, patch(
+ "litellm.proxy.proxy_server.FileResponse"
+ ) as mock_file_response:
# Setup mock_getenv to return empty string for UI_LOGO_PATH
def getenv_side_effect(key, default=""):
if key == "UI_LOGO_PATH":
@@ -3471,12 +3587,13 @@ async def test_get_image_non_root_fallback_to_default_logo(monkeypatch):
return True
# Mock os.path operations
- with patch("litellm.proxy.proxy_server.os.makedirs") as mock_makedirs, \
- patch("litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect), \
- patch("litellm.proxy.proxy_server.os.access", return_value=True), \
- patch("litellm.proxy.proxy_server.os.getenv") as mock_getenv, \
- patch("litellm.proxy.proxy_server.FileResponse") as mock_file_response:
-
+ with patch("litellm.proxy.proxy_server.os.makedirs") as mock_makedirs, patch(
+ "litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect
+ ), patch("litellm.proxy.proxy_server.os.access", return_value=True), patch(
+ "litellm.proxy.proxy_server.os.getenv"
+ ) as mock_getenv, patch(
+ "litellm.proxy.proxy_server.FileResponse"
+ ) as mock_file_response:
# Setup mock_getenv
def getenv_side_effect(key, default=""):
if key == "UI_LOGO_PATH":
@@ -3495,8 +3612,9 @@ async def test_get_image_non_root_fallback_to_default_logo(monkeypatch):
# Verify that exists was called to check /var/lib/litellm/assets/logo.jpg
assets_logo_path = "/var/lib/litellm/assets/logo.jpg"
- assert any(assets_logo_path in str(call) for call in exists_calls), \
- f"Should check if {assets_logo_path} exists"
+ assert any(
+ assets_logo_path in str(call) for call in exists_calls
+ ), f"Should check if {assets_logo_path} exists"
# Verify FileResponse was called (with fallback logo)
assert mock_file_response.called, "FileResponse should be called"
@@ -3516,11 +3634,11 @@ async def test_get_image_root_case_uses_current_dir(monkeypatch):
monkeypatch.delenv("UI_LOGO_PATH", raising=False)
# Mock os.path operations
- with patch("litellm.proxy.proxy_server.os.makedirs") as mock_makedirs, \
- patch("litellm.proxy.proxy_server.os.path.exists", return_value=True), \
- patch("litellm.proxy.proxy_server.os.getenv") as mock_getenv, \
- patch("litellm.proxy.proxy_server.FileResponse") as mock_file_response:
-
+ with patch("litellm.proxy.proxy_server.os.makedirs") as mock_makedirs, patch(
+ "litellm.proxy.proxy_server.os.path.exists", return_value=True
+ ), patch("litellm.proxy.proxy_server.os.getenv") as mock_getenv, patch(
+ "litellm.proxy.proxy_server.FileResponse"
+ ) as mock_file_response:
# Setup mock_getenv
def getenv_side_effect(key, default=""):
if key == "UI_LOGO_PATH":
@@ -3536,10 +3654,13 @@ async def test_get_image_root_case_uses_current_dir(monkeypatch):
# Verify makedirs was NOT called with /var/lib/litellm/assets (should not create it for root case)
var_lib_assets_calls = [
- call for call in mock_makedirs.call_args_list
+ call
+ for call in mock_makedirs.call_args_list
if "/var/lib/litellm/assets" in str(call)
]
- assert len(var_lib_assets_calls) == 0, "Should not create /var/lib/litellm/assets for root case"
+ assert (
+ len(var_lib_assets_calls) == 0
+ ), "Should not create /var/lib/litellm/assets for root case"
# Verify FileResponse was called
assert mock_file_response.called, "FileResponse should be called"
@@ -3569,13 +3690,14 @@ async def test_get_image_custom_local_logo_bypasses_cache(monkeypatch):
calls_to_file_response.append(path)
return MagicMock()
- with patch("litellm.proxy.proxy_server.os.path.exists", return_value=True), \
- patch("litellm.proxy.proxy_server.os.access", return_value=True), \
- patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response):
-
+ with patch("litellm.proxy.proxy_server.os.path.exists", return_value=True), patch(
+ "litellm.proxy.proxy_server.os.access", return_value=True
+ ), patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response):
await get_image()
- assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once"
+ assert (
+ len(calls_to_file_response) == 1
+ ), "FileResponse should be called exactly once"
assert calls_to_file_response[0] == "/app/custom_logo.jpg", (
f"Expected custom logo path, got {calls_to_file_response[0]}. "
"A stale cached_logo.jpg may have been returned instead."
@@ -3602,17 +3724,18 @@ async def test_get_image_default_logo_still_uses_cache(monkeypatch):
calls_to_file_response.append(path)
return MagicMock()
- with patch("litellm.proxy.proxy_server.os.path.exists", return_value=True), \
- patch("litellm.proxy.proxy_server.os.access", return_value=True), \
- patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response):
-
+ with patch("litellm.proxy.proxy_server.os.path.exists", return_value=True), patch(
+ "litellm.proxy.proxy_server.os.access", return_value=True
+ ), patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response):
await get_image()
- assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once"
+ assert (
+ len(calls_to_file_response) == 1
+ ), "FileResponse should be called exactly once"
served_path = calls_to_file_response[0]
- assert served_path.endswith("cached_logo.jpg"), (
- f"Expected cached_logo.jpg for default logo, got {served_path}"
- )
+ assert served_path.endswith(
+ "cached_logo.jpg"
+ ), f"Expected cached_logo.jpg for default logo, got {served_path}"
@pytest.mark.asyncio
@@ -3641,20 +3764,23 @@ async def test_get_image_custom_logo_missing_falls_through_to_default(monkeypatc
return False
return True
- with patch("litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect), \
- patch("litellm.proxy.proxy_server.os.access", return_value=True), \
- patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response):
-
+ with patch(
+ "litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect
+ ), patch("litellm.proxy.proxy_server.os.access", return_value=True), patch(
+ "litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response
+ ):
await get_image()
- assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once"
+ assert (
+ len(calls_to_file_response) == 1
+ ), "FileResponse should be called exactly once"
served_path = calls_to_file_response[0]
- assert served_path != "/app/nonexistent_logo.jpg", (
- "Should not attempt to serve a non-existent custom logo"
- )
- assert served_path.endswith("cached_logo.jpg"), (
- f"Expected fallback to cached_logo.jpg, got {served_path}"
- )
+ assert (
+ served_path != "/app/nonexistent_logo.jpg"
+ ), "Should not attempt to serve a non-existent custom logo"
+ assert served_path.endswith(
+ "cached_logo.jpg"
+ ), f"Expected fallback to cached_logo.jpg, got {served_path}"
@pytest.mark.asyncio
@@ -3686,20 +3812,23 @@ async def test_get_image_custom_logo_missing_no_cache_serves_default(monkeypatch
return False
return True
- with patch("litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect), \
- patch("litellm.proxy.proxy_server.os.access", return_value=True), \
- patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response):
-
+ with patch(
+ "litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect
+ ), patch("litellm.proxy.proxy_server.os.access", return_value=True), patch(
+ "litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response
+ ):
await get_image()
- assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once"
+ assert (
+ len(calls_to_file_response) == 1
+ ), "FileResponse should be called exactly once"
served_path = calls_to_file_response[0]
- assert served_path != "/app/nonexistent_logo.jpg", (
- "Should not attempt to serve a non-existent custom logo"
- )
- assert served_path.endswith("logo.jpg"), (
- f"Expected fallback to default logo.jpg, got {served_path}"
- )
+ assert (
+ served_path != "/app/nonexistent_logo.jpg"
+ ), "Should not attempt to serve a non-existent custom logo"
+ assert served_path.endswith(
+ "logo.jpg"
+ ), f"Expected fallback to default logo.jpg, got {served_path}"
def test_get_config_normalizes_string_callbacks(monkeypatch):
@@ -3721,9 +3850,7 @@ def test_get_config_normalizes_string_callbacks(monkeypatch):
mock_router = MagicMock()
mock_router.get_settings.return_value = {}
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
- monkeypatch.setattr(
- proxy_config, "get_config", AsyncMock(return_value=config_data)
- )
+ monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data))
original_overrides = app.dependency_overrides.copy()
app.dependency_overrides[user_api_key_auth] = lambda: MagicMock()
@@ -4127,9 +4254,9 @@ async def test_update_general_settings_store_model_in_db_false():
proxy_config = ProxyConfig()
- with patch(
- "litellm.proxy.proxy_server.store_model_in_db", True
- ), patch("litellm.proxy.proxy_server.general_settings", {}):
+ with patch("litellm.proxy.proxy_server.store_model_in_db", True), patch(
+ "litellm.proxy.proxy_server.general_settings", {}
+ ):
await proxy_config._update_general_settings(
db_general_settings={"store_model_in_db": False}
)
@@ -4150,9 +4277,9 @@ async def test_update_general_settings_store_model_in_db_string_normalization():
proxy_config = ProxyConfig()
# Test "true" string
- with patch(
- "litellm.proxy.proxy_server.store_model_in_db", False
- ), patch("litellm.proxy.proxy_server.general_settings", {}):
+ with patch("litellm.proxy.proxy_server.store_model_in_db", False), patch(
+ "litellm.proxy.proxy_server.general_settings", {}
+ ):
await proxy_config._update_general_settings(
db_general_settings={"store_model_in_db": "true"}
)
@@ -4161,9 +4288,9 @@ async def test_update_general_settings_store_model_in_db_string_normalization():
assert ps.store_model_in_db is True
# Test "True" string
- with patch(
- "litellm.proxy.proxy_server.store_model_in_db", False
- ), patch("litellm.proxy.proxy_server.general_settings", {}):
+ with patch("litellm.proxy.proxy_server.store_model_in_db", False), patch(
+ "litellm.proxy.proxy_server.general_settings", {}
+ ):
await proxy_config._update_general_settings(
db_general_settings={"store_model_in_db": "True"}
)
@@ -4172,9 +4299,9 @@ async def test_update_general_settings_store_model_in_db_string_normalization():
assert ps.store_model_in_db is True
# Test "false" string
- with patch(
- "litellm.proxy.proxy_server.store_model_in_db", True
- ), patch("litellm.proxy.proxy_server.general_settings", {}):
+ with patch("litellm.proxy.proxy_server.store_model_in_db", True), patch(
+ "litellm.proxy.proxy_server.general_settings", {}
+ ):
await proxy_config._update_general_settings(
db_general_settings={"store_model_in_db": "false"}
)
@@ -4194,9 +4321,9 @@ async def test_update_general_settings_store_model_in_db_none_keeps_current():
proxy_config = ProxyConfig()
# When current is True and DB sends None, should stay True
- with patch(
- "litellm.proxy.proxy_server.store_model_in_db", True
- ), patch("litellm.proxy.proxy_server.general_settings", {}):
+ with patch("litellm.proxy.proxy_server.store_model_in_db", True), patch(
+ "litellm.proxy.proxy_server.general_settings", {}
+ ):
await proxy_config._update_general_settings(
db_general_settings={"store_model_in_db": None}
)
@@ -4205,9 +4332,9 @@ async def test_update_general_settings_store_model_in_db_none_keeps_current():
assert ps.store_model_in_db is True
# When current is False and DB sends None, should stay False
- with patch(
- "litellm.proxy.proxy_server.store_model_in_db", False
- ), patch("litellm.proxy.proxy_server.general_settings", {}):
+ with patch("litellm.proxy.proxy_server.store_model_in_db", False), patch(
+ "litellm.proxy.proxy_server.general_settings", {}
+ ):
await proxy_config._update_general_settings(
db_general_settings={"store_model_in_db": None}
)
@@ -4238,13 +4365,9 @@ async def test_store_model_in_db_db_override_when_config_false():
mock_proxy_logging.slack_alerting_instance = MagicMock()
mock_proxy_config = AsyncMock()
- with patch(
- "litellm.proxy.proxy_server.proxy_config", mock_proxy_config
- ), patch(
+ with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config), patch(
"litellm.proxy.proxy_server.store_model_in_db", False
- ), patch(
- "litellm.proxy.proxy_server.get_secret_bool", return_value=False
- ):
+ ), patch("litellm.proxy.proxy_server.get_secret_bool", return_value=False):
await ProxyStartupEvent.initialize_scheduled_background_jobs(
general_settings={},
prisma_client=mock_prisma_client,
@@ -4282,13 +4405,9 @@ async def test_store_model_in_db_db_check_skipped_when_already_true(monkeypatch)
mock_proxy_logging.slack_alerting_instance = MagicMock()
mock_proxy_config = AsyncMock()
- with patch(
- "litellm.proxy.proxy_server.proxy_config", mock_proxy_config
- ), patch(
+ with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config), patch(
"litellm.proxy.proxy_server.store_model_in_db", True
- ), patch(
- "litellm.proxy.proxy_server.get_secret_bool", return_value=True
- ):
+ ), patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True):
await ProxyStartupEvent.initialize_scheduled_background_jobs(
general_settings={},
prisma_client=mock_prisma_client,
@@ -4328,13 +4447,9 @@ async def test_store_model_in_db_db_failure_graceful(monkeypatch):
mock_proxy_logging.slack_alerting_instance = MagicMock()
mock_proxy_config = AsyncMock()
- with patch(
- "litellm.proxy.proxy_server.proxy_config", mock_proxy_config
- ), patch(
+ with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config), patch(
"litellm.proxy.proxy_server.store_model_in_db", False
- ), patch(
- "litellm.proxy.proxy_server.get_secret_bool", return_value=False
- ):
+ ), patch("litellm.proxy.proxy_server.get_secret_bool", return_value=False):
# Should not raise an exception
await ProxyStartupEvent.initialize_scheduled_background_jobs(
general_settings={},
@@ -4352,3 +4467,183 @@ async def test_store_model_in_db_db_failure_graceful(monkeypatch):
# add_deployment should NOT have been called since store_model_in_db is False
mock_proxy_config.add_deployment.assert_not_called()
+
+
+# =====================================================================
+# Spend counter tests (v2 — Redis-backed spend counters)
+# =====================================================================
+
+
+@pytest.mark.asyncio
+async def test_get_current_spend_reads_redis_first():
+ """get_current_spend should prefer Redis over in-memory."""
+ from litellm.caching.dual_cache import DualCache
+
+ counter_cache = DualCache()
+
+ # In-memory has stale value
+ counter_cache.in_memory_cache.set_cache(key="spend:key:test", value=0.30)
+
+ # Mock Redis with cross-pod authoritative value
+ mock_redis = AsyncMock()
+ mock_redis.async_get_cache = AsyncMock(return_value=0.90)
+ counter_cache.redis_cache = mock_redis
+
+ import litellm.proxy.proxy_server as ps
+
+ original = ps.spend_counter_cache
+ ps.spend_counter_cache = counter_cache
+
+ try:
+ from litellm.proxy.proxy_server import get_current_spend
+
+ result = await get_current_spend(
+ counter_key="spend:key:test",
+ fallback_spend=0.0,
+ )
+ # Should return Redis value (0.90), not in-memory (0.30)
+ assert result == 0.90
+ mock_redis.async_get_cache.assert_called_once_with(key="spend:key:test")
+ finally:
+ ps.spend_counter_cache = original
+
+
+@pytest.mark.asyncio
+async def test_get_current_spend_fallback_to_in_memory():
+ """When Redis is not configured, get_current_spend uses in-memory."""
+ from litellm.caching.dual_cache import DualCache
+
+ counter_cache = DualCache() # no redis_cache
+ counter_cache.in_memory_cache.set_cache(key="spend:key:test", value=0.50)
+
+ import litellm.proxy.proxy_server as ps
+
+ original = ps.spend_counter_cache
+ ps.spend_counter_cache = counter_cache
+
+ try:
+ from litellm.proxy.proxy_server import get_current_spend
+
+ result = await get_current_spend(
+ counter_key="spend:key:test",
+ fallback_spend=0.0,
+ )
+ assert result == 0.50
+ finally:
+ ps.spend_counter_cache = original
+
+
+@pytest.mark.asyncio
+async def test_increment_spend_counters_initializes_and_increments():
+ """Counter should initialize from cached object spend, then increment.
+
+ Uses a pre-hashed token to match production: metadata["user_api_key"]
+ is always hashed by the auth flow before reaching the cost callback.
+ """
+ from litellm.caching.dual_cache import DualCache
+ from litellm.proxy._types import LiteLLM_VerificationTokenView, hash_token
+
+ key_cache = DualCache()
+ counter_cache = DualCache()
+
+ # In production, the auth flow hashes the raw key before it reaches
+ # the cost callback. Simulate that by passing the hashed token.
+ hashed_token = hash_token("sk-test-token-for-counter")
+
+ # Simulate a cached key object with existing spend from DB
+ cached_key = LiteLLM_VerificationTokenView(
+ token=hashed_token,
+ spend=5.0,
+ max_budget=10.0,
+ )
+ key_cache.in_memory_cache.set_cache(key=hashed_token, value=cached_key)
+
+ import litellm.proxy.proxy_server as ps
+
+ original_key_cache = ps.user_api_key_cache
+ original_counter_cache = ps.spend_counter_cache
+ ps.user_api_key_cache = key_cache
+ ps.spend_counter_cache = counter_cache
+
+ try:
+ from litellm.proxy.proxy_server import increment_spend_counters
+
+ # Pass pre-hashed token (as the cost callback would in production)
+ await increment_spend_counters(
+ token=hashed_token,
+ team_id=None,
+ user_id=None,
+ response_cost=0.50,
+ )
+
+ # Counter should be: base(5.0) + increment(0.50) = 5.50
+ counter = counter_cache.in_memory_cache.get_cache(
+ key=f"spend:key:{hashed_token}"
+ )
+ assert counter == 5.50
+
+ # Second increment — counter already exists, just increment
+ await increment_spend_counters(
+ token=hashed_token,
+ team_id=None,
+ user_id=None,
+ response_cost=0.25,
+ )
+
+ counter = counter_cache.in_memory_cache.get_cache(
+ key=f"spend:key:{hashed_token}"
+ )
+ assert counter == 5.75
+ finally:
+ ps.user_api_key_cache = original_key_cache
+ ps.spend_counter_cache = original_counter_cache
+
+
+@pytest.mark.asyncio
+async def test_increment_spend_counters_team_and_member():
+ """Counter should track team and team member spend separately."""
+ from litellm.caching.dual_cache import DualCache
+ from litellm.proxy._types import LiteLLM_TeamTable
+
+ key_cache = DualCache()
+ counter_cache = DualCache()
+
+ # Cached team object
+ team_obj = LiteLLM_TeamTable(team_id="team-1", spend=2.0)
+ key_cache.in_memory_cache.set_cache(key="team_id:team-1", value=team_obj)
+
+ # Cached team membership
+ key_cache.in_memory_cache.set_cache(
+ key="team_membership:user-1:team-1",
+ value={"user_id": "user-1", "team_id": "team-1", "spend": 1.0},
+ )
+
+ import litellm.proxy.proxy_server as ps
+
+ original_key_cache = ps.user_api_key_cache
+ original_counter_cache = ps.spend_counter_cache
+ ps.user_api_key_cache = key_cache
+ ps.spend_counter_cache = counter_cache
+
+ try:
+ from litellm.proxy.proxy_server import increment_spend_counters
+
+ await increment_spend_counters(
+ token=None,
+ team_id="team-1",
+ user_id="user-1",
+ response_cost=0.30,
+ )
+
+ team_counter = counter_cache.in_memory_cache.get_cache(
+ key="spend:team:team-1"
+ )
+ assert team_counter == 2.30
+
+ member_counter = counter_cache.in_memory_cache.get_cache(
+ key="spend:team_member:user-1:team-1"
+ )
+ assert member_counter == 1.30
+ finally:
+ ps.user_api_key_cache = original_key_cache
+ ps.spend_counter_cache = original_counter_cache
diff --git a/tests/test_litellm/router_utils/test_health_check_routing.py b/tests/test_litellm/router_utils/test_health_check_routing.py
new file mode 100644
index 00000000000..f40144b44c9
--- /dev/null
+++ b/tests/test_litellm/router_utils/test_health_check_routing.py
@@ -0,0 +1,197 @@
+"""
+Tests for health-check-driven routing filter in the Router.
+"""
+
+import time
+
+import pytest
+
+from litellm.caching.caching import DualCache
+from litellm.router_utils.health_state_cache import DeploymentHealthCache
+
+
+def _make_deployment(model_id: str, model_name: str = "gpt-4") -> dict:
+ """Helper to create a deployment dict for testing."""
+ return {
+ "model_name": model_name,
+ "litellm_params": {"model": model_name, "api_key": "fake"},
+ "model_info": {"id": model_id},
+ }
+
+
+def _make_health_cache(
+ unhealthy_ids: set = None, staleness_threshold: float = 60.0
+) -> DeploymentHealthCache:
+ """Create a health cache pre-populated with unhealthy deployment IDs."""
+ cache = DualCache()
+ health_cache = DeploymentHealthCache(
+ cache=cache, staleness_threshold=staleness_threshold
+ )
+ if unhealthy_ids:
+ now = time.time()
+ states = {}
+ for uid in unhealthy_ids:
+ states[uid] = {
+ "is_healthy": False,
+ "timestamp": now,
+ "reason": "test_unhealthy",
+ }
+ health_cache.set_deployment_health_states(states)
+ return health_cache
+
+
+class TestFilterHealthCheckUnhealthyDeployments:
+ """Test the sync filter method."""
+
+ def _make_router_like(self, enable: bool, health_cache: DeploymentHealthCache):
+ """Create a minimal object that behaves like Router for filter testing."""
+
+ class FakeRouter:
+ def __init__(self):
+ self.enable_health_check_routing = enable
+ self.health_state_cache = health_cache
+
+ # Import the actual method and bind it
+ from litellm.router import Router
+
+ fake = FakeRouter()
+ # Use the unbound method
+ fake._filter_health_check_unhealthy_deployments = (
+ Router._filter_health_check_unhealthy_deployments.__get__(fake, FakeRouter)
+ )
+ return fake
+
+ def test_filter_removes_unhealthy_deployments(self):
+ """Unhealthy deployments should be removed from candidates."""
+ health_cache = _make_health_cache(unhealthy_ids={"deploy-2"})
+ router = self._make_router_like(enable=True, health_cache=health_cache)
+
+ deployments = [
+ _make_deployment("deploy-1"),
+ _make_deployment("deploy-2"),
+ _make_deployment("deploy-3"),
+ ]
+ result = router._filter_health_check_unhealthy_deployments(deployments)
+ assert len(result) == 2
+ assert all(d["model_info"]["id"] != "deploy-2" for d in result)
+
+ def test_filter_noop_when_disabled(self):
+ """When enable_health_check_routing=False, filter should be a no-op."""
+ health_cache = _make_health_cache(unhealthy_ids={"deploy-1"})
+ router = self._make_router_like(enable=False, health_cache=health_cache)
+
+ deployments = [
+ _make_deployment("deploy-1"),
+ _make_deployment("deploy-2"),
+ ]
+ result = router._filter_health_check_unhealthy_deployments(deployments)
+ assert len(result) == 2 # no filtering
+
+ def test_filter_returns_all_when_all_unhealthy(self):
+ """Safety net: if ALL deployments are unhealthy, return all (don't cause outage)."""
+ health_cache = _make_health_cache(
+ unhealthy_ids={"deploy-1", "deploy-2", "deploy-3"}
+ )
+ router = self._make_router_like(enable=True, health_cache=health_cache)
+
+ deployments = [
+ _make_deployment("deploy-1"),
+ _make_deployment("deploy-2"),
+ _make_deployment("deploy-3"),
+ ]
+ result = router._filter_health_check_unhealthy_deployments(deployments)
+ assert len(result) == 3 # all returned, safety net
+
+ def test_filter_returns_all_when_cache_empty(self):
+ """When cache is empty, all deployments should pass through."""
+ health_cache = _make_health_cache() # empty
+ router = self._make_router_like(enable=True, health_cache=health_cache)
+
+ deployments = [
+ _make_deployment("deploy-1"),
+ _make_deployment("deploy-2"),
+ ]
+ result = router._filter_health_check_unhealthy_deployments(deployments)
+ assert len(result) == 2
+
+
+class TestAsyncFilterHealthCheckUnhealthyDeployments:
+ """Test the async filter method."""
+
+ def _make_router_like(self, enable: bool, health_cache: DeploymentHealthCache):
+ from litellm.router import Router
+
+ class FakeRouter:
+ def __init__(self):
+ self.enable_health_check_routing = enable
+ self.health_state_cache = health_cache
+
+ fake = FakeRouter()
+ fake._async_filter_health_check_unhealthy_deployments = (
+ Router._async_filter_health_check_unhealthy_deployments.__get__(
+ fake, FakeRouter
+ )
+ )
+ return fake
+
+ @pytest.mark.asyncio
+ async def test_async_filter_removes_unhealthy(self):
+ """Async version: unhealthy deployments removed."""
+ health_cache = _make_health_cache(unhealthy_ids={"deploy-2"})
+ router = self._make_router_like(enable=True, health_cache=health_cache)
+
+ deployments = [
+ _make_deployment("deploy-1"),
+ _make_deployment("deploy-2"),
+ _make_deployment("deploy-3"),
+ ]
+ result = await router._async_filter_health_check_unhealthy_deployments(
+ healthy_deployments=deployments
+ )
+ assert len(result) == 2
+ assert all(d["model_info"]["id"] != "deploy-2" for d in result)
+
+ @pytest.mark.asyncio
+ async def test_async_filter_safety_net(self):
+ """Async version: safety net when all unhealthy."""
+ health_cache = _make_health_cache(unhealthy_ids={"deploy-1", "deploy-2"})
+ router = self._make_router_like(enable=True, health_cache=health_cache)
+
+ deployments = [
+ _make_deployment("deploy-1"),
+ _make_deployment("deploy-2"),
+ ]
+ result = await router._async_filter_health_check_unhealthy_deployments(
+ healthy_deployments=deployments
+ )
+ assert len(result) == 2 # safety net
+
+
+class TestBuildDeploymentHealthStates:
+ """Test the build_deployment_health_states function."""
+
+ def test_builds_states_from_endpoints(self):
+ from litellm.proxy.health_check import build_deployment_health_states
+
+ healthy = [{"model": "gpt-4", "model_id": "deploy-1"}]
+ unhealthy = [{"model": "gpt-4", "model_id": "deploy-2", "error": "timeout"}]
+
+ states = build_deployment_health_states(healthy, unhealthy)
+ assert states["deploy-1"]["is_healthy"] is True
+ assert states["deploy-2"]["is_healthy"] is False
+
+ def test_no_model_id_skipped(self):
+ from litellm.proxy.health_check import build_deployment_health_states
+
+ healthy = [{"model": "gpt-4"}] # no model_id
+ unhealthy = [{"model": "gpt-4", "model_id": "deploy-2"}]
+
+ states = build_deployment_health_states(healthy, unhealthy)
+ assert "deploy-1" not in states
+ assert states["deploy-2"]["is_healthy"] is False
+
+ def test_empty_endpoints(self):
+ from litellm.proxy.health_check import build_deployment_health_states
+
+ states = build_deployment_health_states([], [])
+ assert states == {}
diff --git a/tests/test_litellm/router_utils/test_health_state_cache.py b/tests/test_litellm/router_utils/test_health_state_cache.py
new file mode 100644
index 00000000000..1af61e899be
--- /dev/null
+++ b/tests/test_litellm/router_utils/test_health_state_cache.py
@@ -0,0 +1,113 @@
+"""
+Tests for DeploymentHealthCache - the cache layer for health-check-driven routing.
+"""
+
+import time
+
+import pytest
+
+from litellm.caching.caching import DualCache
+from litellm.router_utils.health_state_cache import DeploymentHealthCache
+
+
+@pytest.fixture
+def cache():
+ return DualCache()
+
+
+@pytest.fixture
+def health_cache(cache):
+ return DeploymentHealthCache(cache=cache, staleness_threshold=60.0)
+
+
+def test_set_and_get_unhealthy_ids(health_cache):
+ """Write states, verify unhealthy set is returned correctly."""
+ now = time.time()
+ states = {
+ "deploy-1": {"is_healthy": True, "timestamp": now, "reason": ""},
+ "deploy-2": {"is_healthy": False, "timestamp": now, "reason": "check_failed"},
+ "deploy-3": {"is_healthy": False, "timestamp": now, "reason": "timeout"},
+ }
+ health_cache.set_deployment_health_states(states)
+ result = health_cache.get_unhealthy_deployment_ids()
+ assert result == {"deploy-2", "deploy-3"}
+
+
+@pytest.mark.asyncio
+async def test_async_get_unhealthy_ids(health_cache):
+ """Async version of set and get."""
+ now = time.time()
+ states = {
+ "deploy-1": {"is_healthy": True, "timestamp": now, "reason": ""},
+ "deploy-2": {"is_healthy": False, "timestamp": now, "reason": "check_failed"},
+ }
+ health_cache.set_deployment_health_states(states)
+ result = await health_cache.async_get_unhealthy_deployment_ids()
+ assert result == {"deploy-2"}
+
+
+def test_staleness_filtering(health_cache):
+ """Entries older than staleness_threshold should be ignored."""
+ old_time = time.time() - 120 # 2 minutes ago, threshold is 60s
+ states = {
+ "deploy-1": {
+ "is_healthy": False,
+ "timestamp": old_time,
+ "reason": "check_failed",
+ },
+ }
+ health_cache.set_deployment_health_states(states)
+ result = health_cache.get_unhealthy_deployment_ids()
+ assert result == set() # stale entry should be ignored
+
+
+def test_empty_cache_returns_empty_set(health_cache):
+ """No data in cache should return empty set."""
+ result = health_cache.get_unhealthy_deployment_ids()
+ assert result == set()
+
+
+def test_all_healthy_returns_empty_set(health_cache):
+ """All healthy deployments should return empty set."""
+ now = time.time()
+ states = {
+ "deploy-1": {"is_healthy": True, "timestamp": now, "reason": ""},
+ "deploy-2": {"is_healthy": True, "timestamp": now, "reason": ""},
+ }
+ health_cache.set_deployment_health_states(states)
+ result = health_cache.get_unhealthy_deployment_ids()
+ assert result == set()
+
+
+def test_mixed_stale_and_fresh(health_cache):
+ """Only fresh unhealthy entries should be returned."""
+ now = time.time()
+ old_time = now - 120 # stale
+ states = {
+ "deploy-1": {
+ "is_healthy": False,
+ "timestamp": old_time,
+ "reason": "stale",
+ },
+ "deploy-2": {
+ "is_healthy": False,
+ "timestamp": now,
+ "reason": "fresh",
+ },
+ }
+ health_cache.set_deployment_health_states(states)
+ result = health_cache.get_unhealthy_deployment_ids()
+ assert result == {"deploy-2"}
+
+
+def test_malformed_state_entries_are_skipped(health_cache):
+ """Non-dict entries in the cache should be skipped safely."""
+ now = time.time()
+ states = {
+ "deploy-1": {"is_healthy": False, "timestamp": now, "reason": "bad"},
+ "deploy-2": "not_a_dict", # malformed
+ "deploy-3": None, # malformed
+ }
+ health_cache.set_deployment_health_states(states)
+ result = health_cache.get_unhealthy_deployment_ids()
+ assert result == {"deploy-1"}
diff --git a/tests/test_litellm/test_claude_opus_4_6_config.py b/tests/test_litellm/test_claude_opus_4_6_config.py
index 7ee2ea33957..654ef1b9771 100644
--- a/tests/test_litellm/test_claude_opus_4_6_config.py
+++ b/tests/test_litellm/test_claude_opus_4_6_config.py
@@ -24,70 +24,82 @@ def test_claude_4_6_australia_region_uses_au_prefix_not_apac():
Related: The 'apac.' prefix is valid for Asia-Pacific (Singapore) region models,
but should not be used for Australia which has its own 'au.' prefix.
"""
- json_path = os.path.join(os.path.dirname(__file__), "../../model_prices_and_context_window.json")
+ json_path = os.path.join(
+ os.path.dirname(__file__), "../../model_prices_and_context_window.json"
+ )
with open(json_path) as f:
model_data = json.load(f)
# Verify au.anthropic.claude-opus-4-6-v1 exists (correct)
- assert "au.anthropic.claude-opus-4-6-v1" in model_data, \
- "Missing Australia region model: au.anthropic.claude-opus-4-6-v1"
+ assert (
+ "au.anthropic.claude-opus-4-6-v1" in model_data
+ ), "Missing Australia region model: au.anthropic.claude-opus-4-6-v1"
# Verify apac.anthropic.claude-opus-4-6-v1 does NOT exist (incorrect)
- assert "apac.anthropic.claude-opus-4-6-v1" not in model_data, \
- "Incorrect model entry exists: apac.anthropic.claude-opus-4-6-v1 should be au.anthropic.claude-opus-4-6-v1"
+ assert (
+ "apac.anthropic.claude-opus-4-6-v1" not in model_data
+ ), "Incorrect model entry exists: apac.anthropic.claude-opus-4-6-v1 should be au.anthropic.claude-opus-4-6-v1"
# Verify au.anthropic.claude-sonnet-4-6 exists (correct)
- assert "au.anthropic.claude-sonnet-4-6" in model_data, \
- "Missing Australia region model: au.anthropic.claude-sonnet-4-6"
+ assert (
+ "au.anthropic.claude-sonnet-4-6" in model_data
+ ), "Missing Australia region model: au.anthropic.claude-sonnet-4-6"
# Verify apac.anthropic.claude-sonnet-4-6 does NOT exist (incorrect)
- assert "apac.anthropic.claude-sonnet-4-6" not in model_data, \
- "Incorrect model entry exists: apac.anthropic.claude-sonnet-4-6 should be au.anthropic.claude-sonnet-4-6"
+ assert (
+ "apac.anthropic.claude-sonnet-4-6" not in model_data
+ ), "Incorrect model entry exists: apac.anthropic.claude-sonnet-4-6 should be au.anthropic.claude-sonnet-4-6"
# Verify the au. model is registered in bedrock_converse_models
- assert "au.anthropic.claude-opus-4-6-v1" in litellm.bedrock_converse_models, \
- "au.anthropic.claude-opus-4-6-v1 not registered in bedrock_converse_models"
+ assert (
+ "au.anthropic.claude-opus-4-6-v1" in litellm.bedrock_converse_models
+ ), "au.anthropic.claude-opus-4-6-v1 not registered in bedrock_converse_models"
# Verify apac. is NOT registered for this model
- assert "apac.anthropic.claude-opus-4-6-v1" not in litellm.bedrock_converse_models, \
- "apac.anthropic.claude-opus-4-6-v1 should not be in bedrock_converse_models"
+ assert (
+ "apac.anthropic.claude-opus-4-6-v1" not in litellm.bedrock_converse_models
+ ), "apac.anthropic.claude-opus-4-6-v1 should not be in bedrock_converse_models"
# Verify the au. model is registered in bedrock_converse_models
- assert "au.anthropic.claude-sonnet-4-6" in litellm.bedrock_converse_models, \
- "au.anthropic.claude-sonnet-4-6 not registered in bedrock_converse_models"
+ assert (
+ "au.anthropic.claude-sonnet-4-6" in litellm.bedrock_converse_models
+ ), "au.anthropic.claude-sonnet-4-6 not registered in bedrock_converse_models"
# Verify apac. is NOT registered for this model
- assert "apac.anthropic.claude-sonnet-4-6" not in litellm.bedrock_converse_models, \
- "apac.anthropic.claude-sonnet-4-6 should not be in bedrock_converse_models"
+ assert (
+ "apac.anthropic.claude-sonnet-4-6" not in litellm.bedrock_converse_models
+ ), "apac.anthropic.claude-sonnet-4-6 should not be in bedrock_converse_models"
def test_opus_4_6_model_pricing_and_capabilities():
- json_path = os.path.join(os.path.dirname(__file__), "../../model_prices_and_context_window.json")
+ json_path = os.path.join(
+ os.path.dirname(__file__), "../../model_prices_and_context_window.json"
+ )
with open(json_path) as f:
model_data = json.load(f)
expected_models = {
"claude-opus-4-6": {
"provider": "anthropic",
- "has_long_context_pricing": True,
+ "has_long_context_pricing": False,
"tool_use_system_prompt_tokens": 346,
"max_input_tokens": 1000000,
},
"claude-opus-4-6-20260205": {
"provider": "anthropic",
- "has_long_context_pricing": True,
+ "has_long_context_pricing": False,
"tool_use_system_prompt_tokens": 346,
"max_input_tokens": 1000000,
},
"anthropic.claude-opus-4-6-v1": {
"provider": "bedrock_converse",
- "has_long_context_pricing": True,
+ "has_long_context_pricing": False,
"tool_use_system_prompt_tokens": 346,
"max_input_tokens": 1000000,
},
"vertex_ai/claude-opus-4-6": {
"provider": "vertex_ai-anthropic_models",
- "has_long_context_pricing": True,
+ "has_long_context_pricing": False,
"tool_use_system_prompt_tokens": 346,
"max_input_tokens": 1000000,
},
@@ -119,6 +131,11 @@ def test_opus_4_6_model_pricing_and_capabilities():
assert info["output_cost_per_token_above_200k_tokens"] == 3.75e-05
assert info["cache_creation_input_token_cost_above_200k_tokens"] == 1.25e-05
assert info["cache_read_input_token_cost_above_200k_tokens"] == 1e-06
+ else:
+ assert "input_cost_per_token_above_200k_tokens" not in info
+ assert "output_cost_per_token_above_200k_tokens" not in info
+ assert "cache_creation_input_token_cost_above_200k_tokens" not in info
+ assert "cache_read_input_token_cost_above_200k_tokens" not in info
assert info["supports_assistant_prefill"] is False
assert info["supports_function_calling"] is True
@@ -126,11 +143,16 @@ def test_opus_4_6_model_pricing_and_capabilities():
assert info["supports_reasoning"] is True
assert info["supports_tool_choice"] is True
assert info["supports_vision"] is True
- assert info["tool_use_system_prompt_tokens"] == config["tool_use_system_prompt_tokens"]
+ assert (
+ info["tool_use_system_prompt_tokens"]
+ == config["tool_use_system_prompt_tokens"]
+ )
def test_opus_4_6_bedrock_regional_model_pricing():
- json_path = os.path.join(os.path.dirname(__file__), "../../model_prices_and_context_window.json")
+ json_path = os.path.join(
+ os.path.dirname(__file__), "../../model_prices_and_context_window.json"
+ )
with open(json_path) as f:
model_data = json.load(f)
@@ -140,40 +162,24 @@ def test_opus_4_6_bedrock_regional_model_pricing():
"output_cost_per_token": 2.5e-05,
"cache_creation_input_token_cost": 6.25e-06,
"cache_read_input_token_cost": 5e-07,
- "input_cost_per_token_above_200k_tokens": 1e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
},
"us.anthropic.claude-opus-4-6-v1": {
"input_cost_per_token": 5.5e-06,
"output_cost_per_token": 2.75e-05,
"cache_creation_input_token_cost": 6.875e-06,
"cache_read_input_token_cost": 5.5e-07,
- "input_cost_per_token_above_200k_tokens": 1.1e-05,
- "output_cost_per_token_above_200k_tokens": 4.125e-05,
- "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
- "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
},
"eu.anthropic.claude-opus-4-6-v1": {
"input_cost_per_token": 5.5e-06,
"output_cost_per_token": 2.75e-05,
"cache_creation_input_token_cost": 6.875e-06,
"cache_read_input_token_cost": 5.5e-07,
- "input_cost_per_token_above_200k_tokens": 1.1e-05,
- "output_cost_per_token_above_200k_tokens": 4.125e-05,
- "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
- "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
},
"au.anthropic.claude-opus-4-6-v1": {
"input_cost_per_token": 5.5e-06,
"output_cost_per_token": 2.75e-05,
"cache_creation_input_token_cost": 6.875e-06,
"cache_read_input_token_cost": 5.5e-07,
- "input_cost_per_token_above_200k_tokens": 1.1e-05,
- "output_cost_per_token_above_200k_tokens": 4.125e-05,
- "cache_creation_input_token_cost_above_200k_tokens": 1.375e-05,
- "cache_read_input_token_cost_above_200k_tokens": 1.1e-06,
},
}
@@ -186,12 +192,18 @@ def test_opus_4_6_bedrock_regional_model_pricing():
assert info["max_tokens"] == 128000
assert info["supports_assistant_prefill"] is False
assert info["tool_use_system_prompt_tokens"] == 346
+ assert "input_cost_per_token_above_200k_tokens" not in info
+ assert "output_cost_per_token_above_200k_tokens" not in info
+ assert "cache_creation_input_token_cost_above_200k_tokens" not in info
+ assert "cache_read_input_token_cost_above_200k_tokens" not in info
for key, value in expected.items():
assert info[key] == value
def test_opus_4_6_alias_and_dated_metadata_match():
- json_path = os.path.join(os.path.dirname(__file__), "../../model_prices_and_context_window.json")
+ json_path = os.path.join(
+ os.path.dirname(__file__), "../../model_prices_and_context_window.json"
+ )
with open(json_path) as f:
model_data = json.load(f)
@@ -207,10 +219,6 @@ def test_opus_4_6_alias_and_dated_metadata_match():
"cache_creation_input_token_cost",
"cache_creation_input_token_cost_above_1hr",
"cache_read_input_token_cost",
- "input_cost_per_token_above_200k_tokens",
- "output_cost_per_token_above_200k_tokens",
- "cache_creation_input_token_cost_above_200k_tokens",
- "cache_read_input_token_cost_above_200k_tokens",
"supports_assistant_prefill",
"tool_use_system_prompt_tokens",
]
diff --git a/tests/test_litellm/test_constants.py b/tests/test_litellm/test_constants.py
index 8fff3ec40d4..b3c13c6e26e 100644
--- a/tests/test_litellm/test_constants.py
+++ b/tests/test_litellm/test_constants.py
@@ -1,3 +1,4 @@
+import ast
import inspect
import json
import os
@@ -17,57 +18,56 @@ import litellm
from litellm import constants
-def test_all_numeric_constants_can_be_overridden():
+def _build_constant_env_var_map() -> dict[str, str]:
"""
- Test that all integer and float constants in constants.py can be overridden with environment variables.
- This ensures that any new constants added in the future will be configurable via environment variables.
+ Build a mapping of CONSTANT_NAME -> ENV_VAR_NAME by parsing constants.py.
+
+ This keeps the test resilient when a constant name and env var name differ
+ (e.g., aliases like LITELLM_* env vars).
"""
- # Get all attributes from the constants module
- constants_attributes = inspect.getmembers(constants)
+ env_var_map: dict[str, str] = {}
+ constants_source = inspect.getsource(constants)
+ parsed = ast.parse(constants_source)
- # Filter for uppercase constants (by convention) that are integers or floats
- # Exclude booleans since bool is a subclass of int in Python
- numeric_constants = [
- (name, value)
- for name, value in constants_attributes
- if name.isupper() and isinstance(value, (int, float)) and not isinstance(value, bool)
- ]
-
- # Ensure we found some constants to test
- assert len(numeric_constants) > 0, "No numeric constants found to test"
-
- print("all numeric constants", json.dumps(numeric_constants, indent=4))
-
- # Constants that use a different env var name than the constant name
- constant_to_env_var = {
- "MAX_CALLBACKS": "LITELLM_MAX_CALLBACKS",
- "MCP_CLIENT_TIMEOUT": "LITELLM_MCP_CLIENT_TIMEOUT",
- "MCP_TOOL_LISTING_TIMEOUT": "LITELLM_MCP_TOOL_LISTING_TIMEOUT",
- "MCP_METADATA_TIMEOUT": "LITELLM_MCP_METADATA_TIMEOUT",
- "MCP_HEALTH_CHECK_TIMEOUT": "LITELLM_MCP_HEALTH_CHECK_TIMEOUT",
- }
-
- # Verify all numeric constants have environment variable support
- for name, value in numeric_constants:
- # Skip constants that are not meant to be overridden (if any)
- if name.startswith("_"):
+ for node in parsed.body:
+ if not isinstance(node, ast.Assign):
continue
- # Create a test value that's different from the default
- test_value = value + 1 if isinstance(value, int) else value + 0.1
+ if len(node.targets) != 1 or not isinstance(node.targets[0], ast.Name):
+ continue
- # Use the env var name that the constants module actually reads
- env_var_name = constant_to_env_var.get(name, name)
+ constant_name = node.targets[0].id
+ env_var_name = None
- # Set the environment variable
- with mock.patch.dict(os.environ, {env_var_name: str(test_value)}):
- print("overriding", name, "with", test_value)
- importlib.reload(constants)
+ for child in ast.walk(node.value):
+ if not isinstance(child, ast.Call):
+ continue
- # Get the new value after reload
- new_value = getattr(constants, name)
+ # os.getenv("ENV_NAME", default)
+ if (
+ isinstance(child.func, ast.Attribute)
+ and isinstance(child.func.value, ast.Name)
+ and child.func.value.id == "os"
+ and child.func.attr == "getenv"
+ and len(child.args) >= 1
+ and isinstance(child.args[0], ast.Constant)
+ and isinstance(child.args[0].value, str)
+ ):
+ env_var_name = child.args[0].value
+ break
- # Verify the value was overridden
- assert (
- new_value == test_value
- ), f"Failed to override {name} with environment variable. Expected {test_value}, got {new_value}"
+ # get_env_int("ENV_NAME", default)
+ if (
+ isinstance(child.func, ast.Name)
+ and child.func.id == "get_env_int"
+ and len(child.args) >= 1
+ and isinstance(child.args[0], ast.Constant)
+ and isinstance(child.args[0].value, str)
+ ):
+ env_var_name = child.args[0].value
+ break
+
+ if env_var_name:
+ env_var_map[constant_name] = env_var_name
+
+ return env_var_map
diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py
index f4a7c4f4990..a02464e946b 100644
--- a/tests/test_litellm/test_router.py
+++ b/tests/test_litellm/test_router.py
@@ -24,9 +24,9 @@ def test_update_kwargs_does_not_mutate_defaults_and_merges_metadata():
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
- "api_key": os.getenv("AZURE_API_KEY"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
- "api_base": os.getenv("AZURE_API_BASE"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
},
}
],
@@ -646,8 +646,15 @@ def test_arouter_responses_api_bridge():
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
- mock_response.json.return_value = {"id": "resp_test", "object": "response", "status": "completed", "output": []}
- mock_response.text = '{"id": "resp_test", "object": "response", "status": "completed", "output": []}'
+ mock_response.json.return_value = {
+ "id": "resp_test",
+ "object": "response",
+ "status": "completed",
+ "output": [],
+ }
+ mock_response.text = (
+ '{"id": "resp_test", "object": "response", "status": "completed", "output": []}'
+ )
with patch.object(client, "post", return_value=mock_response) as mock_post:
try:
@@ -2147,7 +2154,10 @@ def test_get_deployment_credentials_with_provider_aws_bedrock_runtime_endpoint()
)
assert credentials is not None
- assert credentials["aws_bedrock_runtime_endpoint"] == "https://bedrock-runtime.us-east-1.amazonaws.com"
+ assert (
+ credentials["aws_bedrock_runtime_endpoint"]
+ == "https://bedrock-runtime.us-east-1.amazonaws.com"
+ )
assert credentials["aws_access_key_id"] == "test-access-key"
assert credentials["aws_secret_access_key"] == "test-secret-key"
assert credentials["aws_region_name"] == "us-east-1"
@@ -2169,11 +2179,11 @@ def test_get_deployment_credentials_with_provider_resolves_credential_name():
credential_values={
"api_key": "resolved-api-key",
"api_base": "https://resolved.openai.azure.com",
- "api_version": "2024-02-01"
- }
+ "api_version": "2024-02-01",
+ },
)
]
-
+
router = litellm.Router(
model_list=[
{
@@ -2197,7 +2207,7 @@ def test_get_deployment_credentials_with_provider_resolves_credential_name():
assert credentials["custom_llm_provider"] == "azure"
# Ensure credential name is removed after resolution
assert "litellm_credential_name" not in credentials
-
+
# Cleanup
litellm.credential_list = []
@@ -2302,7 +2312,10 @@ async def test_aguardrail_helper():
# Mock the original function
async def mock_original_function(**kwargs):
- return {"result": "success", "selected_guardrail": kwargs.get("selected_guardrail")}
+ return {
+ "result": "success",
+ "selected_guardrail": kwargs.get("selected_guardrail"),
+ }
result = await router._aguardrail_helper(
model="content-filter",
@@ -2336,7 +2349,10 @@ async def test_aguardrail():
# Mock the original function
async def mock_original_function(**kwargs):
- return {"result": "success", "selected_guardrail": kwargs.get("selected_guardrail")}
+ return {
+ "result": "success",
+ "selected_guardrail": kwargs.get("selected_guardrail"),
+ }
result = await router.aguardrail(
guardrail_name="content-filter",
@@ -2346,6 +2362,7 @@ async def test_aguardrail():
assert result["result"] == "success"
assert result["selected_guardrail"]["id"] == "guardrail-1"
+
@pytest.mark.asyncio
async def test_anthropic_messages_call_type_is_cached():
"""
@@ -2417,36 +2434,33 @@ async def test_anthropic_messages_call_type_is_cached():
additional_headers=None,
),
)
-
+
cache = DualCache()
deployment_check = PromptCachingDeploymentCheck(cache=cache)
prompt_cache = PromptCachingCache(cache=cache)
-
+
# Create messages with enough tokens to pass the caching threshold
test_messages = [
{
- "role": "user",
+ "role": "user",
"content": [
{
- "type": "text",
+ "type": "text",
"text": "test long message here" * 1024,
- "cache_control": {
- "type": "ephemeral",
- "ttl": "5m"
- }
+ "cache_control": {"type": "ephemeral", "ttl": "5m"},
}
- ]
+ ],
}
]
test_model_id = "test-model-id-123"
-
+
# Create a payload with anthropic_messages call type
payload = create_standard_logging_payload()
payload["call_type"] = CallTypes.anthropic_messages.value
payload["messages"] = test_messages
payload["model"] = "anthropic/claude-3-5-sonnet-20240620"
payload["model_id"] = test_model_id
-
+
# Log the success event (should cache the model_id)
await deployment_check.async_log_success_event(
kwargs={"standard_logging_object": payload},
@@ -2454,19 +2468,23 @@ async def test_anthropic_messages_call_type_is_cached():
start_time=1234567890.0,
end_time=1234567891.0,
)
-
+
# Small delay to ensure cache write completes
await asyncio.sleep(0.1)
-
+
# Verify that the model_id was actually cached
cached_result = await prompt_cache.async_get_model_id(
messages=test_messages,
tools=None,
)
-
+
# This assertion will FAIL if anthropic_messages is filtered out
- assert cached_result is not None, "Model ID should be cached for anthropic_messages call type"
- assert cached_result["model_id"] == test_model_id, f"Expected {test_model_id}, got {cached_result['model_id']}"
+ assert (
+ cached_result is not None
+ ), "Model ID should be cached for anthropic_messages call type"
+ assert (
+ cached_result["model_id"] == test_model_id
+ ), f"Expected {test_model_id}, got {cached_result['model_id']}"
def test_update_kwargs_with_deployment_propagates_model_tags():
@@ -2682,9 +2700,7 @@ def test_credential_name_injected_as_tag():
)
kwargs: dict = {"metadata": {"tags": ["A.101"]}}
- deployment = router.get_deployment_by_model_group_name(
- model_group_name="xai-model"
- )
+ deployment = router.get_deployment_by_model_group_name(model_group_name="xai-model")
router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
assert "Credential: xAI" in kwargs["metadata"]["tags"]
@@ -2709,9 +2725,7 @@ def test_credential_name_not_duplicated_in_tags():
)
kwargs: dict = {"metadata": {"tags": ["Credential: xAI", "A.101"]}}
- deployment = router.get_deployment_by_model_group_name(
- model_group_name="xai-model"
- )
+ deployment = router.get_deployment_by_model_group_name(model_group_name="xai-model")
router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
assert kwargs["metadata"]["tags"].count("Credential: xAI") == 1
@@ -2733,9 +2747,7 @@ def test_credential_name_not_injected_when_absent():
)
kwargs: dict = {"metadata": {"tags": ["A.101"]}}
- deployment = router.get_deployment_by_model_group_name(
- model_group_name="gpt-model"
- )
+ deployment = router.get_deployment_by_model_group_name(model_group_name="gpt-model")
router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
assert kwargs["metadata"]["tags"] == ["A.101"]
diff --git a/tests/test_litellm/test_router_order_fallback.py b/tests/test_litellm/test_router_order_fallback.py
new file mode 100644
index 00000000000..760766a7461
--- /dev/null
+++ b/tests/test_litellm/test_router_order_fallback.py
@@ -0,0 +1,331 @@
+"""
+Tests for order-based fallback routing.
+
+When deployments have `order` set in litellm_params, lower order deployments
+should be tried first, and higher order deployments should be used as fallbacks
+when lower order deployments fail.
+"""
+
+from typing import Optional
+
+import pytest
+
+from litellm import Router
+from litellm.utils import _get_order_filtered_deployments
+
+# ---------------------------------------------------------------------------
+# Unit tests for _get_order_filtered_deployments
+# ---------------------------------------------------------------------------
+
+
+class TestGetOrderFilteredDeployments:
+ def _make_deployment(self, order: Optional[int], dep_id: str) -> dict:
+ params: dict = {"model": "gpt-4o", "api_key": "key"}
+ if order is not None:
+ params["order"] = order
+ return {
+ "model_name": "test-model",
+ "litellm_params": params,
+ "model_info": {"id": dep_id},
+ }
+
+ def test_returns_min_order_group(self):
+ deps = [
+ self._make_deployment(1, "a"),
+ self._make_deployment(2, "b"),
+ self._make_deployment(1, "c"),
+ ]
+ result = _get_order_filtered_deployments(deps)
+ assert len(result) == 2
+ assert all(d["model_info"]["id"] in ("a", "c") for d in result)
+
+ def test_target_order_filters_to_exact_level(self):
+ deps = [
+ self._make_deployment(1, "a"),
+ self._make_deployment(2, "b"),
+ self._make_deployment(3, "c"),
+ ]
+ result = _get_order_filtered_deployments(deps, target_order=2)
+ assert len(result) == 1
+ assert result[0]["model_info"]["id"] == "b"
+
+ def test_target_order_no_match_returns_all(self):
+ deps = [
+ self._make_deployment(1, "a"),
+ self._make_deployment(2, "b"),
+ ]
+ result = _get_order_filtered_deployments(deps, target_order=99)
+ assert len(result) == 2
+
+ def test_no_order_set_returns_all(self):
+ deps = [
+ self._make_deployment(None, "a"),
+ self._make_deployment(None, "b"),
+ ]
+ result = _get_order_filtered_deployments(deps)
+ assert len(result) == 2
+
+ def test_empty_list(self):
+ result = _get_order_filtered_deployments([])
+ assert result == []
+
+ def test_single_order_returns_all_with_that_order(self):
+ deps = [
+ self._make_deployment(1, "a"),
+ self._make_deployment(1, "b"),
+ ]
+ result = _get_order_filtered_deployments(deps)
+ assert len(result) == 2
+
+
+# ---------------------------------------------------------------------------
+# Integration tests for order-based fallback in Router
+# ---------------------------------------------------------------------------
+
+
+def test_router_order_without_pre_call_checks():
+ """Order filtering should work even when enable_pre_call_checks=False (default)."""
+ router = Router(
+ model_list=[
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "key",
+ "mock_response": "from order 1",
+ "order": 1,
+ },
+ "model_info": {"id": "1"},
+ },
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "key",
+ "mock_response": "from order 2",
+ "order": 2,
+ },
+ "model_info": {"id": "2"},
+ },
+ ],
+ num_retries=0,
+ enable_pre_call_checks=False,
+ )
+
+ for _ in range(20):
+ response = router.completion(
+ model="test-model",
+ messages=[{"role": "user", "content": "hi"}],
+ )
+ assert response._hidden_params["model_id"] == "1"
+
+
+def test_router_order_no_fallback_when_healthy():
+ """When order=1 is healthy, order=2 should never be used."""
+ router = Router(
+ model_list=[
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "key",
+ "mock_response": "from order 1",
+ "order": 1,
+ },
+ "model_info": {"id": "1"},
+ },
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "key",
+ "mock_response": "from order 2",
+ "order": 2,
+ },
+ "model_info": {"id": "2"},
+ },
+ ],
+ num_retries=0,
+ )
+
+ for _ in range(50):
+ response = router.completion(
+ model="test-model",
+ messages=[{"role": "user", "content": "hi"}],
+ )
+ assert response._hidden_params["model_id"] == "1"
+
+
+@pytest.mark.asyncio
+async def test_router_order_fallback_on_failure():
+ """When order=1 fails, order=2 should be tried as fallback."""
+ router = Router(
+ model_list=[
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "bad-key",
+ "mock_response": Exception("connection error"),
+ "order": 1,
+ },
+ "model_info": {"id": "1"},
+ },
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "good-key",
+ "mock_response": "success from order 2",
+ "order": 2,
+ },
+ "model_info": {"id": "2"},
+ },
+ ],
+ num_retries=0,
+ )
+
+ response = await router.acompletion(
+ model="test-model",
+ messages=[{"role": "user", "content": "hi"}],
+ )
+ assert response._hidden_params["model_id"] == "2"
+
+
+@pytest.mark.asyncio
+async def test_router_order_fallback_three_levels():
+ """When order=1 and order=2 both fail, order=3 should be tried."""
+ router = Router(
+ model_list=[
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "bad",
+ "mock_response": Exception("fail 1"),
+ "order": 1,
+ },
+ "model_info": {"id": "1"},
+ },
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "bad",
+ "mock_response": Exception("fail 2"),
+ "order": 2,
+ },
+ "model_info": {"id": "2"},
+ },
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "good",
+ "mock_response": "success from order 3",
+ "order": 3,
+ },
+ "model_info": {"id": "3"},
+ },
+ ],
+ num_retries=0,
+ )
+
+ response = await router.acompletion(
+ model="test-model",
+ messages=[{"role": "user", "content": "hi"}],
+ )
+ assert response._hidden_params["model_id"] == "3"
+
+
+@pytest.mark.asyncio
+async def test_router_order_fallback_then_external_fallback():
+ """When all order levels fail, external fallbacks should be tried."""
+ router = Router(
+ model_list=[
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "bad",
+ "mock_response": Exception("fail order 1"),
+ "order": 1,
+ },
+ "model_info": {"id": "1"},
+ },
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "bad",
+ "mock_response": Exception("fail order 2"),
+ "order": 2,
+ },
+ "model_info": {"id": "2"},
+ },
+ {
+ "model_name": "fallback-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "good",
+ "mock_response": "success from external fallback",
+ },
+ "model_info": {"id": "fallback"},
+ },
+ ],
+ fallbacks=[{"test-model": ["fallback-model"]}],
+ num_retries=0,
+ )
+
+ response = await router.acompletion(
+ model="test-model",
+ messages=[{"role": "user", "content": "hi"}],
+ )
+ assert response._hidden_params["model_id"] == "fallback"
+
+
+@pytest.mark.asyncio
+async def test_router_order_fallback_with_non_standard_fallbacks():
+ """Non-standard fallback formats (e.g. fallbacks=["model-name"]) passed
+ per-request should still be tried after all order levels are exhausted."""
+ router = Router(
+ model_list=[
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "bad",
+ "mock_response": Exception("fail order 1"),
+ "order": 1,
+ },
+ "model_info": {"id": "1"},
+ },
+ {
+ "model_name": "test-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "bad",
+ "mock_response": Exception("fail order 2"),
+ "order": 2,
+ },
+ "model_info": {"id": "2"},
+ },
+ {
+ "model_name": "fallback-model",
+ "litellm_params": {
+ "model": "gpt-4o",
+ "api_key": "good",
+ "mock_response": "success from non-standard fallback",
+ },
+ "model_info": {"id": "fallback"},
+ },
+ ],
+ num_retries=0,
+ )
+
+ response = await router.acompletion(
+ model="test-model",
+ messages=[{"role": "user", "content": "hi"}],
+ fallbacks=["fallback-model"], # non-standard format, passed per-request
+ )
+ assert response._hidden_params["model_id"] == "fallback"
diff --git a/tests/test_litellm/test_setup_wizard.py b/tests/test_litellm/test_setup_wizard.py
index e10bd893e31..c96d6d7ed6c 100644
--- a/tests/test_litellm/test_setup_wizard.py
+++ b/tests/test_litellm/test_setup_wizard.py
@@ -58,7 +58,7 @@ _ANTHROPIC = {
_AZURE = {
"id": "azure",
"name": "Azure OpenAI",
- "env_key": "AZURE_API_KEY",
+ "env_key": "AZURE_AI_API_KEY",
"models": [],
"test_model": None,
"needs_api_base": True,
@@ -123,8 +123,8 @@ def test_build_config_master_key_quoted():
def test_build_config_does_not_mutate_env_vars():
"""_build_config must not modify the caller's env_vars dict."""
env_vars = {
- "AZURE_API_KEY": "az-key",
- "_LITELLM_AZURE_API_BASE_AZURE": "https://my.azure.com",
+ "AZURE_AI_API_KEY": "az-key",
+ "_LITELLM_AZURE_AI_API_BASE_AZURE": "https://my.azure.com",
"_LITELLM_AZURE_DEPLOYMENT_AZURE": "my-deployment",
}
original_keys = set(env_vars.keys())
@@ -134,8 +134,8 @@ def test_build_config_does_not_mutate_env_vars():
def test_build_config_azure_uses_deployment_name():
env_vars = {
- "AZURE_API_KEY": "az-key",
- "_LITELLM_AZURE_API_BASE_AZURE": "https://my.azure.com",
+ "AZURE_AI_API_KEY": "az-key",
+ "_LITELLM_AZURE_AI_API_BASE_AZURE": "https://my.azure.com",
"_LITELLM_AZURE_DEPLOYMENT_AZURE": "my-gpt4o",
}
config = SetupWizard._build_config([_AZURE], env_vars, "sk-master")
@@ -147,7 +147,7 @@ def test_build_config_azure_uses_deployment_name():
def test_build_config_azure_no_deployment_skipped():
"""Azure without a deployment name should emit nothing (not fallback to gpt-4o)."""
- env_vars = {"AZURE_API_KEY": "az-key"} # no deployment sentinel
+ env_vars = {"AZURE_AI_API_KEY": "az-key"} # no deployment sentinel
config = SetupWizard._build_config([_AZURE], env_vars, "sk-master")
# No azure model entry should be emitted when deployment name is absent
assert "model: azure/" not in config
@@ -157,7 +157,7 @@ def test_build_config_no_display_name_collision_openai_and_azure():
"""OpenAI gpt-4o and azure gpt-4o should get distinct model_name values."""
env_vars = {
"OPENAI_API_KEY": "sk-openai",
- "AZURE_API_KEY": "az-key",
+ "AZURE_AI_API_KEY": "az-key",
"_LITELLM_AZURE_DEPLOYMENT_AZURE": "gpt-4o",
}
config = SetupWizard._build_config([_OPENAI, _AZURE], env_vars, "sk-master")
@@ -182,7 +182,7 @@ def test_build_config_internal_sentinel_keys_excluded():
"""_LITELLM_ prefixed sentinel keys must not appear in environment_variables."""
env_vars = {
"OPENAI_API_KEY": "sk-real",
- "_LITELLM_AZURE_API_BASE_AZURE": "https://x.azure.com",
+ "_LITELLM_AZURE_AI_API_BASE_AZURE": "https://x.azure.com",
}
config = SetupWizard._build_config([_OPENAI], env_vars, "sk-master")
assert "_LITELLM_" not in config
diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py
index 38b7b576d4f..e984403a05b 100644
--- a/tests/test_litellm/test_utils.py
+++ b/tests/test_litellm/test_utils.py
@@ -831,6 +831,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
},
},
"supports_native_streaming": {"type": "boolean"},
+ "supports_native_structured_output": {"type": "boolean"},
"tiered_pricing": {
"type": "array",
"items": {
@@ -961,6 +962,7 @@ def test_get_model_info_gemini():
and not "learnlm" in model
and not "imagen" in model
and not "veo" in model
+ and not "lyria" in model
and not "robotics" in model
):
assert info.get("tpm") is not None, f"{model} does not have tpm"
@@ -2788,6 +2790,22 @@ def test_model_info_for_openrouter_kimi_k2_5():
print("openrouter kimi-k2.5 model info", model_info)
+def test_gemini_lyria_3_preview_models_in_cost_map():
+ import json
+ from pathlib import Path
+
+ json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
+ with open(json_path) as f:
+ model_cost = json.load(f)
+
+ clip = model_cost.get("gemini/lyria-3-clip-preview")
+ pro = model_cost.get("gemini/lyria-3-pro-preview")
+ assert clip is not None and pro is not None
+ assert clip["litellm_provider"] == "gemini" and pro["litellm_provider"] == "gemini"
+ assert clip["max_input_tokens"] == 131072 == pro["max_input_tokens"]
+ assert clip["output_cost_per_image"] == 0.04
+
+
def test_model_info_for_fireworks_short_form_models():
"""
Test that fireworks_ai short-form model entries (fireworks_ai/)
diff --git a/tests/test_litellm/types/llms/test_types_llms_openai.py b/tests/test_litellm/types/llms/test_types_llms_openai.py
index 94221bd0efc..569743269a5 100644
--- a/tests/test_litellm/types/llms/test_types_llms_openai.py
+++ b/tests/test_litellm/types/llms/test_types_llms_openai.py
@@ -219,11 +219,13 @@ class TestAssistantMessageImageUrlContent:
# convert to list to consume it — this must not raise ValidationError.
content_blocks = list(raw_content) if raw_content is not None else []
- assert len(content_blocks) == 2, (
- f"Expected 2 content blocks (text + image_url), got {len(content_blocks)}: {content_blocks}"
- )
+ assert (
+ len(content_blocks) == 2
+ ), f"Expected 2 content blocks (text + image_url), got {len(content_blocks)}: {content_blocks}"
types = [b.get("type") for b in content_blocks if isinstance(b, dict)]
- assert "image_url" in types, f"image_url block was silently dropped; blocks: {content_blocks}"
+ assert (
+ "image_url" in types
+ ), f"image_url block was silently dropped; blocks: {content_blocks}"
def test_assistant_message_image_url_preserved_in_all_message_values(self):
"""
@@ -255,14 +257,16 @@ class TestAssistantMessageImageUrlContent:
assert assistant is not None, "Assistant message missing after serialisation"
content = assistant.get("content", [])
- assert isinstance(content, list), f"content should be a list, got {type(content)}"
- assert len(content) == 2, (
- f"Expected 2 content blocks (text + image_url), got {len(content)}: {content}"
- )
+ assert isinstance(
+ content, list
+ ), f"content should be a list, got {type(content)}"
+ assert (
+ len(content) == 2
+ ), f"Expected 2 content blocks (text + image_url), got {len(content)}: {content}"
types = [b.get("type") for b in content if isinstance(b, dict)]
- assert "image_url" in types, (
- f"image_url block was silently dropped during AllMessageValues serialisation; blocks: {content}"
- )
+ assert (
+ "image_url" in types
+ ), f"image_url block was silently dropped during AllMessageValues serialisation; blocks: {content}"
class TestResponsesAPIReasoningNullFields:
@@ -379,10 +383,14 @@ class TestResponsesAPIReasoningNullFields:
)
dumped = response.model_dump()
reasoning = [
- o for o in dumped["output"] if isinstance(o, dict) and o.get("type") == "reasoning"
+ o
+ for o in dumped["output"]
+ if isinstance(o, dict) and o.get("type") == "reasoning"
][0]
message = [
- o for o in dumped["output"] if isinstance(o, dict) and o.get("type") == "message"
+ o
+ for o in dumped["output"]
+ if isinstance(o, dict) and o.get("type") == "message"
][0]
assert "status" not in reasoning
assert "content" not in reasoning
@@ -410,3 +418,38 @@ class TestResponsesAPIReasoningNullFields:
assert dumped["error"] is None
assert "instructions" in dumped
assert dumped["instructions"] is None
+
+
+def test_normalize_fine_tuning_job_dict_maps_azure_pending():
+ from litellm.llms.openai.fine_tuning.handler import _normalize_fine_tuning_job_dict
+
+ out = _normalize_fine_tuning_job_dict(
+ {"organization_id": None, "result_files": None, "status": "pending"},
+ is_azure=True,
+ )
+ assert out["organization_id"] == ""
+ assert out["result_files"] == []
+ assert out["status"] == "queued"
+
+
+def test_normalize_fine_tuning_job_dict_openai_unchanged():
+ from litellm.llms.openai.fine_tuning.handler import _normalize_fine_tuning_job_dict
+
+ data = {"organization_id": None, "result_files": None, "status": "pending"}
+ out = _normalize_fine_tuning_job_dict(data, is_azure=False)
+ assert out is data
+
+
+def test_openai_file_object_accepts_pending_status():
+ from litellm.types.llms.openai import OpenAIFileObject
+
+ file_obj = OpenAIFileObject(
+ id="file-123",
+ bytes=1024,
+ created_at=1677610602,
+ filename="train.jsonl",
+ object="file",
+ purpose="fine-tune",
+ status="pending",
+ )
+ assert file_obj.status == "pending"
diff --git a/tests/test_team_logging.py b/tests/test_team_logging.py
index 913b1e19496..9e89d945eda 100644
--- a/tests/test_team_logging.py
+++ b/tests/test_team_logging.py
@@ -59,137 +59,3 @@ async def chat_completion(session, key, model="azure-gpt-3.5", request_metadata=
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
-
-
-@pytest.mark.skip(reason="flaky test - covered by simpler unit testing.")
-@pytest.mark.asyncio
-@pytest.mark.flaky(retries=12, delay=2)
-async def test_aaateam_logging():
- """
- -> Team 1 logs to project 1
- -> Create Key
- -> Make chat/completions call
- -> Fetch logs from langfuse
- """
- try:
- async with aiohttp.ClientSession() as session:
-
- key = await generate_key(
- session, models=["fake-openai-endpoint"], team_id="team-1"
- ) # team-1 logs to project 1
-
- from litellm._uuid import uuid
-
- _trace_id = f"trace-{uuid.uuid4()}"
- _request_metadata = {
- "trace_id": _trace_id,
- }
-
- await chat_completion(
- session,
- key["key"],
- model="fake-openai-endpoint",
- request_metadata=_request_metadata,
- )
-
- # Test - if the logs were sent to the correct team on langfuse
- import langfuse
-
- print(f"langfuse_public_key: {os.getenv('LANGFUSE_PROJECT1_PUBLIC')}")
- print(f"langfuse_secret_key: {os.getenv('LANGFUSE_HOST')}")
- langfuse_client = langfuse.Langfuse(
- public_key=os.getenv("LANGFUSE_PROJECT1_PUBLIC"),
- secret_key=os.getenv("LANGFUSE_PROJECT1_SECRET"),
- host="https://us.cloud.langfuse.com",
- )
-
- await asyncio.sleep(30)
-
- print(f"searching for trace_id={_trace_id} on langfuse")
-
- generations = langfuse_client.get_generations(trace_id=_trace_id).data
- print(generations)
- assert len(generations) == 1
- except Exception as e:
- pytest.fail(f"Unexpected error: {str(e)}")
-
-
-@pytest.mark.skip(reason="todo fix langfuse credential error")
-@pytest.mark.asyncio
-async def test_team_2logging():
- """
- -> Team 1 logs to project 2
- -> Create Key
- -> Make chat/completions call
- -> Fetch logs from langfuse
- """
- langfuse_public_key = os.getenv("LANGFUSE_PROJECT2_PUBLIC")
-
- print(f"langfuse_public_key: {langfuse_public_key}")
- langfuse_secret_key = os.getenv("LANGFUSE_PROJECT2_SECRET")
- print(f"langfuse_secret_key: {langfuse_secret_key}")
- langfuse_host = "https://us.cloud.langfuse.com"
-
- try:
- assert langfuse_public_key is not None
- assert langfuse_secret_key is not None
- except Exception as e:
- # skip test if langfuse credentials are not set
- return
-
- try:
- async with aiohttp.ClientSession() as session:
-
- key = await generate_key(
- session, models=["fake-openai-endpoint"], team_id="team-2"
- ) # team-1 logs to project 1
-
- from litellm._uuid import uuid
-
- _trace_id = f"trace-{uuid.uuid4()}"
- _request_metadata = {
- "trace_id": _trace_id,
- }
-
- await chat_completion(
- session,
- key["key"],
- model="fake-openai-endpoint",
- request_metadata=_request_metadata,
- )
-
- # Test - if the logs were sent to the correct team on langfuse
- import langfuse
-
- langfuse_client = langfuse.Langfuse(
- public_key=langfuse_public_key,
- secret_key=langfuse_secret_key,
- host=langfuse_host,
- )
-
- await asyncio.sleep(30)
-
- print(f"searching for trace_id={_trace_id} on langfuse")
-
- generations = langfuse_client.get_generations(trace_id=_trace_id).data
- print("Team 2 generations", generations)
-
- # team-2 should have 1 generation with this trace id
- assert len(generations) == 1
-
- # team-1 should have 0 generations with this trace id
- langfuse_client_1 = langfuse.Langfuse(
- public_key=os.getenv("LANGFUSE_PROJECT1_PUBLIC"),
- secret_key=os.getenv("LANGFUSE_PROJECT1_SECRET"),
- host="https://us.cloud.langfuse.com",
- )
-
- generations_team_1 = langfuse_client_1.get_generations(
- trace_id=_trace_id
- ).data
- print("Team 1 generations", generations_team_1)
-
- assert len(generations_team_1) == 0
-
- except Exception as e:
- pytest.fail("Team 2 logging failed: " + str(e))
diff --git a/tests/unified_google_tests/vertex_key.json b/tests/unified_google_tests/vertex_key.json
index 45ca6acc010..800969fb305 100644
--- a/tests/unified_google_tests/vertex_key.json
+++ b/tests/unified_google_tests/vertex_key.json
@@ -1,13 +1,13 @@
{
"type": "service_account",
- "project_id": "pathrise-convert-1606954137718",
+ "project_id": "litellm-ci-cd",
"private_key_id": "",
"private_key": "",
- "client_email": "ci-cd-723@pathrise-convert-1606954137718.iam.gserviceaccount.com",
- "client_id": "109577393201924326488",
+ "client_email": "test-litellm-ci-cd@litellm-ci-cd.iam.gserviceaccount.com",
+ "client_id": "116563532503305622785",
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
"token_uri": "https://oauth2.googleapis.com/token",
"auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs",
- "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/ci-cd-723%40pathrise-convert-1606954137718.iam.gserviceaccount.com",
+ "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/test-litellm-ci-cd%40litellm-ci-cd.iam.gserviceaccount.com",
"universe_domain": "googleapis.com"
}
diff --git a/tests/vector_store_tests/test_azure_ai_vector_store.py b/tests/vector_store_tests/test_azure_ai_vector_store.py
index 52eb6635a98..58e45f259ab 100644
--- a/tests/vector_store_tests/test_azure_ai_vector_store.py
+++ b/tests/vector_store_tests/test_azure_ai_vector_store.py
@@ -17,10 +17,10 @@ async def test_basic_search_vector_store(sync_mode):
"vector_store_id": "my-vector-index",
"custom_llm_provider": "azure_ai",
"azure_search_service_name": "azure-kb-search",
- "litellm_embedding_model": "azure/text-embedding-3-large",
+ "litellm_embedding_model": "azure_ai/text-embedding-3-large",
"litellm_embedding_config": {
- "api_base": os.getenv("AZURE_AI_SEARCH_EMBEDDING_API_BASE"),
- "api_key": os.getenv("AZURE_AI_SEARCH_EMBEDDING_API_KEY"),
+ "api_base": os.getenv("AZURE_AI_API_BASE"),
+ "api_key": os.getenv("AZURE_AI_API_KEY"),
},
"api_key": os.getenv("AZURE_SEARCH_API_KEY"),
}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/budgets/useBudgets.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/budgets/useBudgets.ts
new file mode 100644
index 00000000000..99c170b6791
--- /dev/null
+++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/budgets/useBudgets.ts
@@ -0,0 +1,70 @@
+import { useQuery, useMutation, useQueryClient, UseQueryResult } from "@tanstack/react-query";
+import { createQueryKeys } from "../common/queryKeysFactory";
+import { getBudgetList, budgetCreateCall, budgetUpdateCall, budgetDeleteCall } from "@/components/networking";
+import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
+import { budgetItem } from "@/components/budgets/budget_panel";
+
+export const budgetKeys = createQueryKeys("budgets");
+
+export const useBudgets = (): UseQueryResult => {
+ const { accessToken } = useAuthorized();
+ return useQuery({
+ queryKey: budgetKeys.list({}),
+ queryFn: async () => {
+ const data = await getBudgetList(accessToken!);
+ return (data ?? []).filter((item: budgetItem | null): item is budgetItem => item != null);
+ },
+ enabled: Boolean(accessToken),
+ });
+};
+
+export const useCreateBudget = () => {
+ const { accessToken } = useAuthorized();
+ const queryClient = useQueryClient();
+
+ return useMutation>({
+ mutationFn: async (formValues) => {
+ if (!accessToken) {
+ throw new Error("Access token is required");
+ }
+ return budgetCreateCall(accessToken, formValues);
+ },
+ onSuccess: () => {
+ queryClient.invalidateQueries({ queryKey: budgetKeys.all });
+ },
+ });
+};
+
+export const useUpdateBudget = () => {
+ const { accessToken } = useAuthorized();
+ const queryClient = useQueryClient();
+
+ return useMutation>({
+ mutationFn: async (formValues) => {
+ if (!accessToken) {
+ throw new Error("Access token is required");
+ }
+ return budgetUpdateCall(accessToken, formValues);
+ },
+ onSuccess: () => {
+ queryClient.invalidateQueries({ queryKey: budgetKeys.all });
+ },
+ });
+};
+
+export const useDeleteBudget = () => {
+ const { accessToken } = useAuthorized();
+ const queryClient = useQueryClient();
+
+ return useMutation({
+ mutationFn: async (budgetId) => {
+ if (!accessToken) {
+ throw new Error("Access token is required");
+ }
+ return budgetDeleteCall(accessToken, budgetId);
+ },
+ onSuccess: () => {
+ queryClient.invalidateQueries({ queryKey: budgetKeys.all });
+ },
+ });
+};
diff --git a/ui/litellm-dashboard/src/components/budgets/budget_modal.tsx b/ui/litellm-dashboard/src/components/budgets/budget_modal.tsx
index 490613de254..b5ad8aaff34 100644
--- a/ui/litellm-dashboard/src/components/budgets/budget_modal.tsx
+++ b/ui/litellm-dashboard/src/components/budgets/budget_modal.tsx
@@ -1,17 +1,17 @@
import React from "react";
import { TextInput, Accordion, AccordionHeader, AccordionBody } from "@tremor/react";
import { Button as Button2, Modal, Form, InputNumber, Select } from "antd";
-import { budgetCreateCall } from "../networking";
+import { useCreateBudget } from "@/app/(dashboard)/hooks/budgets/useBudgets";
import NotificationsManager from "../molecules/notifications_manager";
interface BudgetModalProps {
isModalVisible: boolean;
- accessToken: string | null;
setIsModalVisible: React.Dispatch>;
- setBudgetList: React.Dispatch>;
}
-const BudgetModal: React.FC = ({ isModalVisible, accessToken, setIsModalVisible, setBudgetList }) => {
+const BudgetModal: React.FC = ({ isModalVisible, setIsModalVisible }) => {
const [form] = Form.useForm();
+ const createBudget = useCreateBudget();
+
const handleOk = () => {
setIsModalVisible(false);
form.resetFields();
@@ -23,20 +23,15 @@ const BudgetModal: React.FC = ({ isModalVisible, accessToken,
};
const handleCreate = async (formValues: Record) => {
- if (accessToken == null || accessToken == undefined) {
- return;
- }
try {
NotificationsManager.info("Making API Call");
- // setIsModalVisible(true);
- const response = await budgetCreateCall(accessToken, formValues);
- console.log("key create Response:", response);
- setBudgetList((prevData) => (prevData ? [...prevData, response] : [response])); // Check if prevData is null
+ await createBudget.mutateAsync(formValues);
NotificationsManager.success("Budget Created");
form.resetFields();
+ setIsModalVisible(false);
} catch (error) {
- console.error("Error creating the key:", error);
- NotificationsManager.fromBackend(`Error creating the key: ${error}`);
+ console.error("Error creating the budget:", error);
+ NotificationsManager.fromBackend(`Error creating the budget: ${error}`);
}
};
diff --git a/ui/litellm-dashboard/src/components/budgets/budget_panel.test.tsx b/ui/litellm-dashboard/src/components/budgets/budget_panel.test.tsx
index 534693d3984..ecae379c9f1 100644
--- a/ui/litellm-dashboard/src/components/budgets/budget_panel.test.tsx
+++ b/ui/litellm-dashboard/src/components/budgets/budget_panel.test.tsx
@@ -1,31 +1,50 @@
-import * as networking from "../networking";
import { fireEvent, render, waitFor, screen } from "@testing-library/react";
import { act } from "@testing-library/react";
+import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { afterEach, describe, expect, it, vi } from "vitest";
import BudgetPanel from "./budget_panel";
-vi.mock("../networking", () => ({
- getBudgetList: vi.fn(),
- budgetDeleteCall: vi.fn(),
+const mockBudgets = [
+ {
+ budget_id: "budget-1",
+ max_budget: 100,
+ rpm_limit: 10,
+ tpm_limit: 1000,
+ updated_at: "2024-01-01T00:00:00Z",
+ },
+];
+
+vi.mock("@/app/(dashboard)/hooks/budgets/useBudgets", () => ({
+ useBudgets: vi.fn().mockReturnValue({ data: [], isLoading: false }),
+ useDeleteBudget: vi.fn().mockReturnValue({ mutateAsync: vi.fn(), isPending: false }),
+ useCreateBudget: vi.fn().mockReturnValue({ mutateAsync: vi.fn() }),
+ useUpdateBudget: vi.fn().mockReturnValue({ mutateAsync: vi.fn() }),
}));
+import { useBudgets, useDeleteBudget, useCreateBudget, useUpdateBudget } from "@/app/(dashboard)/hooks/budgets/useBudgets";
+
+const createQueryClient = () =>
+ new QueryClient({
+ defaultOptions: { queries: { retry: false, gcTime: 0 } },
+ });
+
+function renderWithProviders(ui: React.ReactElement) {
+ const qc = createQueryClient();
+ return render({ui});
+}
+
describe("Budget Panel", () => {
afterEach(() => {
vi.clearAllMocks();
});
it("should render the budget panel and load budgets", async () => {
- vi.mocked(networking.getBudgetList).mockResolvedValue([
- {
- budget_id: "budget-1",
- max_budget: "100",
- rpm_limit: 10,
- tpm_limit: 1000,
- updated_at: "2024-01-01T00:00:00Z",
- },
- ]);
+ vi.mocked(useBudgets).mockReturnValue({
+ data: mockBudgets,
+ isLoading: false,
+ } as any);
- render();
+ renderWithProviders();
await waitFor(() => {
expect(screen.getByText("Create a budget to assign to customers.")).toBeInTheDocument();
@@ -34,17 +53,20 @@ describe("Budget Panel", () => {
});
it("should open delete modal when clicking delete icon", async () => {
- vi.mocked(networking.getBudgetList).mockResolvedValue([
- {
- budget_id: "budget-to-delete",
- max_budget: "200",
- rpm_limit: 20,
- tpm_limit: 2000,
- updated_at: "2024-01-02T00:00:00Z",
- },
- ]);
+ vi.mocked(useBudgets).mockReturnValue({
+ data: [
+ {
+ budget_id: "budget-to-delete",
+ max_budget: 200,
+ rpm_limit: 20,
+ tpm_limit: 2000,
+ updated_at: "2024-01-02T00:00:00Z",
+ },
+ ],
+ isLoading: false,
+ } as any);
- render();
+ renderWithProviders();
await waitFor(() => {
expect(screen.getByText("budget-to-delete")).toBeInTheDocument();
@@ -62,18 +84,25 @@ describe("Budget Panel", () => {
});
it("should successfully delete a budget", async () => {
- vi.mocked(networking.getBudgetList).mockResolvedValue([
- {
- budget_id: "budget-to-delete",
- max_budget: "200",
- rpm_limit: 20,
- tpm_limit: 2000,
- updated_at: "2024-01-02T00:00:00Z",
- },
- ]);
- vi.mocked(networking.budgetDeleteCall).mockResolvedValue(undefined);
+ const deleteMutateAsync = vi.fn().mockResolvedValue(undefined);
+ vi.mocked(useBudgets).mockReturnValue({
+ data: [
+ {
+ budget_id: "budget-to-delete",
+ max_budget: 200,
+ rpm_limit: 20,
+ tpm_limit: 2000,
+ updated_at: "2024-01-02T00:00:00Z",
+ },
+ ],
+ isLoading: false,
+ } as any);
+ vi.mocked(useDeleteBudget).mockReturnValue({
+ mutateAsync: deleteMutateAsync,
+ isPending: false,
+ } as any);
- render();
+ renderWithProviders();
await waitFor(() => {
expect(screen.getByText("budget-to-delete")).toBeInTheDocument();
@@ -96,24 +125,43 @@ describe("Budget Panel", () => {
});
await waitFor(() => {
- expect(networking.budgetDeleteCall).toHaveBeenCalledWith("token-123", "budget-to-delete");
- expect(networking.getBudgetList).toHaveBeenCalledTimes(2); // Initial load + refresh after delete
+ expect(deleteMutateAsync).toHaveBeenCalledWith("budget-to-delete");
+ });
+ });
+
+ it("should render empty state without crashing", async () => {
+ vi.mocked(useBudgets).mockReturnValue({
+ data: [],
+ isLoading: false,
+ } as any);
+
+ renderWithProviders();
+
+ await waitFor(() => {
+ expect(screen.getByText("Create a budget to assign to customers.")).toBeInTheDocument();
});
});
it("should handle delete error", async () => {
- vi.mocked(networking.getBudgetList).mockResolvedValue([
- {
- budget_id: "budget-to-delete",
- max_budget: "200",
- rpm_limit: 20,
- tpm_limit: 2000,
- updated_at: "2024-01-02T00:00:00Z",
- },
- ]);
- vi.mocked(networking.budgetDeleteCall).mockRejectedValue(new Error("Delete failed"));
+ const deleteMutateAsync = vi.fn().mockRejectedValue(new Error("Delete failed"));
+ vi.mocked(useBudgets).mockReturnValue({
+ data: [
+ {
+ budget_id: "budget-to-delete",
+ max_budget: 200,
+ rpm_limit: 20,
+ tpm_limit: 2000,
+ updated_at: "2024-01-02T00:00:00Z",
+ },
+ ],
+ isLoading: false,
+ } as any);
+ vi.mocked(useDeleteBudget).mockReturnValue({
+ mutateAsync: deleteMutateAsync,
+ isPending: false,
+ } as any);
- render();
+ renderWithProviders();
await waitFor(() => {
expect(screen.getByText("budget-to-delete")).toBeInTheDocument();
@@ -136,10 +184,38 @@ describe("Budget Panel", () => {
});
await waitFor(() => {
- expect(networking.budgetDeleteCall).toHaveBeenCalledWith("token-123", "budget-to-delete");
+ expect(deleteMutateAsync).toHaveBeenCalledWith("budget-to-delete");
+ });
+ });
+
+ it("should open edit modal when clicking edit icon", async () => {
+ vi.mocked(useBudgets).mockReturnValue({
+ data: [
+ {
+ budget_id: "budget-to-edit",
+ max_budget: 300,
+ rpm_limit: 30,
+ tpm_limit: 3000,
+ updated_at: "2024-01-03T00:00:00Z",
+ },
+ ],
+ isLoading: false,
+ } as any);
+
+ renderWithProviders();
+
+ await waitFor(() => {
+ expect(screen.getByText("budget-to-edit")).toBeInTheDocument();
});
- // Modal should still be open (error handling)
- expect(screen.getByText("Delete Budget?")).toBeInTheDocument();
+ const editButton = screen.getByTestId("edit-budget-button");
+
+ act(() => {
+ fireEvent.click(editButton);
+ });
+
+ await waitFor(() => {
+ expect(screen.getByText("Edit Budget")).toBeInTheDocument();
+ });
});
});
diff --git a/ui/litellm-dashboard/src/components/budgets/budget_panel.tsx b/ui/litellm-dashboard/src/components/budgets/budget_panel.tsx
index b52ef5ab947..e42d0569652 100644
--- a/ui/litellm-dashboard/src/components/budgets/budget_panel.tsx
+++ b/ui/litellm-dashboard/src/components/budgets/budget_panel.tsx
@@ -19,12 +19,12 @@ import {
TabPanels,
Text,
} from "@tremor/react";
-import React, { useEffect, useState } from "react";
+import React, { useState } from "react";
import { Prism as SyntaxHighlighter } from "react-syntax-highlighter";
import DeleteResourceModal from "../common_components/DeleteResourceModal";
import TableIconActionButton from "../common_components/IconActionButton/TableIconActionButtons/TableIconActionButton";
import NotificationsManager from "../molecules/notifications_manager";
-import { budgetDeleteCall, getBudgetList } from "../networking";
+import { useBudgets, useDeleteBudget } from "@/app/(dashboard)/hooks/budgets/useBudgets";
import BudgetModal from "./budget_modal";
import EditBudgetModal from "./edit_budget_modal";
import { CREATE_END_USER_CURL_COMMAND, CHAT_COMPLETIONS_CURL_COMMAND, OPENAI_SDK_PYTHON_CODE } from "./constants";
@@ -35,7 +35,7 @@ interface BudgetSettingsPageProps {
export interface budgetItem {
budget_id: string;
- max_budget: string | null;
+ max_budget: number | null;
rpm_limit: number | null;
tpm_limit: number | null;
updated_at: string;
@@ -45,17 +45,10 @@ const BudgetPanel: React.FC = ({ accessToken }) => {
const [isCreateModelVisible, setIsCreateModelVisible] = useState(false);
const [isEditModalVisible, setIsEditModalVisible] = useState(false);
const [selectedBudget, setSelectedBudget] = useState(null);
- const [budgetList, setBudgetList] = useState([]);
- const [isDeleting, setIsDeleting] = useState(false);
const [isDeleteModalVisible, setIsDeleteModalVisible] = useState(false);
- useEffect(() => {
- if (!accessToken) {
- return;
- }
- getBudgetList(accessToken).then((data) => {
- setBudgetList(data);
- });
- }, [accessToken]);
+
+ const { data: budgetList = [] } = useBudgets();
+ const deleteBudget = useDeleteBudget();
const handleEditCall = async (budget: budgetItem) => {
if (accessToken == null) {
@@ -74,11 +67,9 @@ const BudgetPanel: React.FC = ({ accessToken }) => {
if (!selectedBudget || accessToken == null) {
return;
}
- setIsDeleting(true);
try {
- await budgetDeleteCall(accessToken, selectedBudget.budget_id);
+ await deleteBudget.mutateAsync(selectedBudget.budget_id);
NotificationsManager.success("Budget deleted.");
- await handleUpdateCall();
} catch (error) {
console.error("Error deleting budget:", error);
if (typeof NotificationsManager.fromBackend === "function") {
@@ -87,7 +78,6 @@ const BudgetPanel: React.FC = ({ accessToken }) => {
NotificationsManager.info("Failed to delete budget");
}
} finally {
- setIsDeleting(false);
setIsDeleteModalVisible(false);
setSelectedBudget(null);
}
@@ -97,15 +87,6 @@ const BudgetPanel: React.FC = ({ accessToken }) => {
setIsDeleteModalVisible(false);
};
- const handleUpdateCall = async () => {
- if (accessToken == null) {
- return;
- }
- getBudgetList(accessToken).then((data) => {
- setBudgetList(data);
- });
- };
-
return (