Merge branch 'BerriAI:main' into fix_vertex_expired_tokens
|
|
@ -839,6 +839,52 @@ jobs:
|
|||
paths:
|
||||
- guardrails_coverage.xml
|
||||
- guardrails_coverage
|
||||
|
||||
google_generate_content_endpoint_testing:
|
||||
docker:
|
||||
- image: cimg/python:3.11
|
||||
auth:
|
||||
username: ${DOCKERHUB_USERNAME}
|
||||
password: ${DOCKERHUB_PASSWORD}
|
||||
working_directory: ~/project
|
||||
|
||||
steps:
|
||||
- checkout
|
||||
- setup_google_dns
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
command: |
|
||||
python -m pip install --upgrade pip
|
||||
python -m pip install -r requirements.txt
|
||||
pip install "pytest==7.3.1"
|
||||
pip install "pytest-retry==1.6.3"
|
||||
pip install "pytest-cov==5.0.0"
|
||||
pip install "pytest-asyncio==0.21.1"
|
||||
pip install "respx==0.22.0"
|
||||
pip install "pydantic==2.10.2"
|
||||
# Run pytest and generate JUnit XML report
|
||||
- run:
|
||||
name: Run tests
|
||||
command: |
|
||||
pwd
|
||||
ls
|
||||
python -m pytest -vv tests/unified_google_tests --cov=litellm --cov-report=xml -x -s -v --junitxml=test-results/junit.xml --durations=5
|
||||
no_output_timeout: 120m
|
||||
- run:
|
||||
name: Rename the coverage files
|
||||
command: |
|
||||
mv coverage.xml google_generate_content_endpoint_coverage.xml
|
||||
mv .coverage google_generate_content_endpoint_coverage
|
||||
|
||||
# Store test results
|
||||
- store_test_results:
|
||||
path: test-results
|
||||
- persist_to_workspace:
|
||||
root: .
|
||||
paths:
|
||||
- google_generate_content_endpoint_coverage.xml
|
||||
- google_generate_content_endpoint_coverage
|
||||
|
||||
llm_responses_api_testing:
|
||||
docker:
|
||||
- image: cimg/python:3.11
|
||||
|
|
@ -925,6 +971,7 @@ jobs:
|
|||
command: |
|
||||
pwd
|
||||
ls
|
||||
prisma generate
|
||||
python -m pytest -vv tests/enterprise --cov=litellm --cov-report=xml -x -s -v --junitxml=test-results/junit-enterprise.xml --durations=10 -n 8
|
||||
no_output_timeout: 120m
|
||||
- run:
|
||||
|
|
@ -1333,7 +1380,6 @@ jobs:
|
|||
- run: python ./tests/code_coverage_tests/recursive_detector.py
|
||||
- run: python ./tests/code_coverage_tests/test_router_strategy_async.py
|
||||
- run: python ./tests/code_coverage_tests/litellm_logging_code_coverage.py
|
||||
- run: python ./tests/code_coverage_tests/bedrock_pricing.py
|
||||
- run: python ./tests/documentation_tests/test_env_keys.py
|
||||
- run: python ./tests/documentation_tests/test_router_settings.py
|
||||
- run: python ./tests/documentation_tests/test_api_docs.py
|
||||
|
|
@ -1342,6 +1388,7 @@ jobs:
|
|||
- run: python ./tests/documentation_tests/test_circular_imports.py
|
||||
- run: python ./tests/code_coverage_tests/prevent_key_leaks_in_exceptions.py
|
||||
- run: python ./tests/code_coverage_tests/check_unsafe_enterprise_import.py
|
||||
- run: python ./tests/code_coverage_tests/ban_copy_deepcopy_kwargs.py
|
||||
- run: helm lint ./deploy/charts/litellm-helm
|
||||
|
||||
db_migration_disable_update_check:
|
||||
|
|
@ -1475,6 +1522,25 @@ jobs:
|
|||
pip install "asyncio==3.4.3"
|
||||
pip install "PyGithub==1.59.1"
|
||||
pip install "openai==1.81.0"
|
||||
- run:
|
||||
name: Install dockerize
|
||||
command: |
|
||||
wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start PostgreSQL Database
|
||||
command: |
|
||||
docker run -d \
|
||||
--name postgres-db \
|
||||
-e POSTGRES_USER=postgres \
|
||||
-e POSTGRES_PASSWORD=postgres \
|
||||
-e POSTGRES_DB=circle_test \
|
||||
-p 5432:5432 \
|
||||
postgres:14
|
||||
- run:
|
||||
name: Wait for PostgreSQL to be ready
|
||||
command: dockerize -wait tcp://localhost:5432 -timeout 1m
|
||||
- run:
|
||||
name: Install Grype
|
||||
command: |
|
||||
|
|
@ -1485,13 +1551,13 @@ jobs:
|
|||
# Build and scan Dockerfile.database
|
||||
echo "Building and scanning Dockerfile.database..."
|
||||
docker build -t litellm-database:latest -f ./docker/Dockerfile.database .
|
||||
grype litellm-database:latest --fail-on high
|
||||
grype litellm-database:latest --fail-on critical
|
||||
|
||||
|
||||
# Build and scan main Dockerfile
|
||||
echo "Building and scanning main Dockerfile..."
|
||||
docker build -t litellm:latest .
|
||||
grype litellm:latest --fail-on high
|
||||
grype litellm:latest --fail-on critical
|
||||
- run:
|
||||
name: Build Docker image
|
||||
command: docker build -t my-app:latest -f ./docker/Dockerfile.database .
|
||||
|
|
@ -1500,7 +1566,7 @@ jobs:
|
|||
command: |
|
||||
docker run -d \
|
||||
-p 4000:4000 \
|
||||
-e DATABASE_URL=$PROXY_DATABASE_URL \
|
||||
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
|
||||
-e USE_PRISMA_MIGRATE=True \
|
||||
-e AZURE_API_KEY=$AZURE_API_KEY \
|
||||
-e REDIS_HOST=$REDIS_HOST \
|
||||
|
|
@ -1525,6 +1591,7 @@ jobs:
|
|||
-e LANGFUSE_PROJECT2_PUBLIC=$LANGFUSE_PROJECT2_PUBLIC \
|
||||
-e LANGFUSE_PROJECT1_SECRET=$LANGFUSE_PROJECT1_SECRET \
|
||||
-e LANGFUSE_PROJECT2_SECRET=$LANGFUSE_PROJECT2_SECRET \
|
||||
--add-host host.docker.internal:host-gateway \
|
||||
--name my-app \
|
||||
-v $(pwd)/proxy_server_config.yaml:/app/config.yaml \
|
||||
my-app:latest \
|
||||
|
|
@ -1532,13 +1599,10 @@ jobs:
|
|||
--port 4000 \
|
||||
--detailed_debug \
|
||||
- run:
|
||||
name: Install curl and dockerize
|
||||
name: Install curl
|
||||
command: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y curl
|
||||
sudo wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start outputting logs
|
||||
command: docker logs -f my-app
|
||||
|
|
@ -1616,6 +1680,25 @@ jobs:
|
|||
pip install "PyGithub==1.59.1"
|
||||
pip install "openai==1.81.0"
|
||||
# Run pytest and generate JUnit XML report
|
||||
- run:
|
||||
name: Install dockerize
|
||||
command: |
|
||||
wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start PostgreSQL Database
|
||||
command: |
|
||||
docker run -d \
|
||||
--name postgres-db \
|
||||
-e POSTGRES_USER=postgres \
|
||||
-e POSTGRES_PASSWORD=postgres \
|
||||
-e POSTGRES_DB=circle_test \
|
||||
-p 5432:5432 \
|
||||
postgres:14
|
||||
- run:
|
||||
name: Wait for PostgreSQL to be ready
|
||||
command: dockerize -wait tcp://localhost:5432 -timeout 1m
|
||||
- run:
|
||||
name: Build Docker image
|
||||
command: docker build -t my-app:latest -f ./docker/Dockerfile.database .
|
||||
|
|
@ -1624,7 +1707,7 @@ jobs:
|
|||
command: |
|
||||
docker run -d \
|
||||
-p 4000:4000 \
|
||||
-e DATABASE_URL=$PROXY_DATABASE_URL \
|
||||
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
|
||||
-e AZURE_API_KEY=$AZURE_BATCHES_API_KEY \
|
||||
-e AZURE_API_BASE=$AZURE_BATCHES_API_BASE \
|
||||
-e AZURE_API_VERSION="2024-05-01-preview" \
|
||||
|
|
@ -1650,6 +1733,7 @@ jobs:
|
|||
-e LANGFUSE_PROJECT2_PUBLIC=$LANGFUSE_PROJECT2_PUBLIC \
|
||||
-e LANGFUSE_PROJECT1_SECRET=$LANGFUSE_PROJECT1_SECRET \
|
||||
-e LANGFUSE_PROJECT2_SECRET=$LANGFUSE_PROJECT2_SECRET \
|
||||
--add-host host.docker.internal:host-gateway \
|
||||
--name my-app \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/oai_misc_config.yaml:/app/config.yaml \
|
||||
my-app:latest \
|
||||
|
|
@ -1657,13 +1741,10 @@ jobs:
|
|||
--port 4000 \
|
||||
--detailed_debug \
|
||||
- run:
|
||||
name: Install curl and dockerize
|
||||
name: Install curl
|
||||
command: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y curl
|
||||
sudo wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start outputting logs
|
||||
command: docker logs -f my-app
|
||||
|
|
@ -1738,6 +1819,25 @@ jobs:
|
|||
pip install "asyncio==3.4.3"
|
||||
pip install "PyGithub==1.59.1"
|
||||
pip install "openai==1.81.0"
|
||||
- run:
|
||||
name: Install dockerize
|
||||
command: |
|
||||
wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start PostgreSQL Database
|
||||
command: |
|
||||
docker run -d \
|
||||
--name postgres-db \
|
||||
-e POSTGRES_USER=postgres \
|
||||
-e POSTGRES_PASSWORD=postgres \
|
||||
-e POSTGRES_DB=circle_test \
|
||||
-p 5432:5432 \
|
||||
postgres:14
|
||||
- run:
|
||||
name: Wait for PostgreSQL to be ready
|
||||
command: dockerize -wait tcp://localhost:5432 -timeout 1m
|
||||
- run:
|
||||
name: Build Docker image
|
||||
command: docker build -t my-app:latest -f ./docker/Dockerfile.database .
|
||||
|
|
@ -1748,7 +1848,7 @@ jobs:
|
|||
command: |
|
||||
docker run -d \
|
||||
-p 4000:4000 \
|
||||
-e DATABASE_URL=$PROXY_DATABASE_URL \
|
||||
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
|
||||
-e REDIS_HOST=$REDIS_HOST \
|
||||
-e REDIS_PASSWORD=$REDIS_PASSWORD \
|
||||
-e REDIS_PORT=$REDIS_PORT \
|
||||
|
|
@ -1768,6 +1868,7 @@ jobs:
|
|||
-e APORIA_API_KEY_1=$APORIA_API_KEY_1 \
|
||||
-e COHERE_API_KEY=$COHERE_API_KEY \
|
||||
-e GCS_FLUSH_INTERVAL="1" \
|
||||
--add-host host.docker.internal:host-gateway \
|
||||
--name my-app \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/otel_test_config.yaml:/app/config.yaml \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/custom_guardrail.py:/app/custom_guardrail.py \
|
||||
|
|
@ -1812,13 +1913,14 @@ jobs:
|
|||
command: |
|
||||
docker run -d \
|
||||
-p 4000:4000 \
|
||||
-e DATABASE_URL=$PROXY_DATABASE_URL \
|
||||
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
|
||||
-e REDIS_HOST=$REDIS_HOST \
|
||||
-e REDIS_PASSWORD=$REDIS_PASSWORD \
|
||||
-e REDIS_PORT=$REDIS_PORT \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
-e OPENAI_API_KEY=$OPENAI_API_KEY \
|
||||
-e LITELLM_LICENSE="bad-license" \
|
||||
--add-host host.docker.internal:host-gateway \
|
||||
--name my-app-3 \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/enterprise_config.yaml:/app/config.yaml \
|
||||
my-app:latest \
|
||||
|
|
@ -1876,6 +1978,25 @@ jobs:
|
|||
pip install aiohttp
|
||||
python -m pip install --upgrade pip
|
||||
python -m pip install -r requirements.txt
|
||||
- run:
|
||||
name: Install dockerize
|
||||
command: |
|
||||
wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start PostgreSQL Database
|
||||
command: |
|
||||
docker run -d \
|
||||
--name postgres-db \
|
||||
-e POSTGRES_USER=postgres \
|
||||
-e POSTGRES_PASSWORD=postgres \
|
||||
-e POSTGRES_DB=circle_test \
|
||||
-p 5432:5432 \
|
||||
postgres:14
|
||||
- run:
|
||||
name: Wait for PostgreSQL to be ready
|
||||
command: dockerize -wait tcp://localhost:5432 -timeout 1m
|
||||
- run:
|
||||
name: Build Docker image
|
||||
command: docker build -t my-app:latest -f ./docker/Dockerfile.database .
|
||||
|
|
@ -1886,7 +2007,7 @@ jobs:
|
|||
command: |
|
||||
docker run -d \
|
||||
-p 4000:4000 \
|
||||
-e DATABASE_URL=$PROXY_DATABASE_URL \
|
||||
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
|
||||
-e REDIS_HOST=$REDIS_HOST \
|
||||
-e REDIS_PASSWORD=$REDIS_PASSWORD \
|
||||
-e REDIS_PORT=$REDIS_PORT \
|
||||
|
|
@ -1899,6 +2020,7 @@ jobs:
|
|||
-e DD_API_KEY=$DD_API_KEY \
|
||||
-e DD_SITE=$DD_SITE \
|
||||
-e AWS_REGION_NAME=$AWS_REGION_NAME \
|
||||
--add-host host.docker.internal:host-gateway \
|
||||
--name my-app \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/spend_tracking_config.yaml:/app/config.yaml \
|
||||
my-app:latest \
|
||||
|
|
@ -1906,13 +2028,10 @@ jobs:
|
|||
--port 4000 \
|
||||
--detailed_debug \
|
||||
- run:
|
||||
name: Install curl and dockerize
|
||||
name: Install curl
|
||||
command: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y curl
|
||||
sudo wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start outputting logs
|
||||
command: docker logs -f my-app
|
||||
|
|
@ -1971,6 +2090,25 @@ jobs:
|
|||
pip install "pytest-retry==1.6.3"
|
||||
pip install "pytest-mock==3.12.0"
|
||||
pip install "pytest-asyncio==0.21.1"
|
||||
- run:
|
||||
name: Install dockerize
|
||||
command: |
|
||||
wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start PostgreSQL Database
|
||||
command: |
|
||||
docker run -d \
|
||||
--name postgres-db \
|
||||
-e POSTGRES_USER=postgres \
|
||||
-e POSTGRES_PASSWORD=postgres \
|
||||
-e POSTGRES_DB=circle_test \
|
||||
-p 5432:5432 \
|
||||
postgres:14
|
||||
- run:
|
||||
name: Wait for PostgreSQL to be ready
|
||||
command: dockerize -wait tcp://localhost:5432 -timeout 1m
|
||||
- run:
|
||||
name: Build Docker image
|
||||
command: docker build -t my-app:latest -f ./docker/Dockerfile.database .
|
||||
|
|
@ -1981,7 +2119,7 @@ jobs:
|
|||
command: |
|
||||
docker run -d \
|
||||
-p 4000:4000 \
|
||||
-e DATABASE_URL=$PROXY_DATABASE_URL \
|
||||
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
|
||||
-e REDIS_HOST=$REDIS_HOST \
|
||||
-e REDIS_PASSWORD=$REDIS_PASSWORD \
|
||||
-e REDIS_PORT=$REDIS_PORT \
|
||||
|
|
@ -1990,6 +2128,7 @@ jobs:
|
|||
-e USE_DDTRACE=True \
|
||||
-e DD_API_KEY=$DD_API_KEY \
|
||||
-e DD_SITE=$DD_SITE \
|
||||
--add-host host.docker.internal:host-gateway \
|
||||
--name my-app \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/multi_instance_simple_config.yaml:/app/config.yaml \
|
||||
my-app:latest \
|
||||
|
|
@ -2001,7 +2140,7 @@ jobs:
|
|||
command: |
|
||||
docker run -d \
|
||||
-p 4001:4001 \
|
||||
-e DATABASE_URL=$PROXY_DATABASE_URL \
|
||||
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
|
||||
-e REDIS_HOST=$REDIS_HOST \
|
||||
-e REDIS_PASSWORD=$REDIS_PASSWORD \
|
||||
-e REDIS_PORT=$REDIS_PORT \
|
||||
|
|
@ -2010,6 +2149,7 @@ jobs:
|
|||
-e USE_DDTRACE=True \
|
||||
-e DD_API_KEY=$DD_API_KEY \
|
||||
-e DD_SITE=$DD_SITE \
|
||||
--add-host host.docker.internal:host-gateway \
|
||||
--name my-app-2 \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/multi_instance_simple_config.yaml:/app/config.yaml \
|
||||
my-app:latest \
|
||||
|
|
@ -2281,6 +2421,25 @@ jobs:
|
|||
pip install "langchain_mcp_adapters==0.0.5"
|
||||
pip install "langchain_openai==0.2.1"
|
||||
pip install "langgraph==0.3.18"
|
||||
- run:
|
||||
name: Install dockerize
|
||||
command: |
|
||||
wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start PostgreSQL Database
|
||||
command: |
|
||||
docker run -d \
|
||||
--name postgres-db \
|
||||
-e POSTGRES_USER=postgres \
|
||||
-e POSTGRES_PASSWORD=postgres \
|
||||
-e POSTGRES_DB=circle_test \
|
||||
-p 5432:5432 \
|
||||
postgres:14
|
||||
- run:
|
||||
name: Wait for PostgreSQL to be ready
|
||||
command: dockerize -wait tcp://localhost:5432 -timeout 1m
|
||||
# Run pytest and generate JUnit XML report
|
||||
- run:
|
||||
name: Build Docker image
|
||||
|
|
@ -2290,7 +2449,7 @@ jobs:
|
|||
command: |
|
||||
docker run -d \
|
||||
-p 4000:4000 \
|
||||
-e DATABASE_URL=$PROXY_DATABASE_URL \
|
||||
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
-e OPENAI_API_KEY=$OPENAI_API_KEY \
|
||||
-e GEMINI_API_KEY=$GEMINI_API_KEY \
|
||||
|
|
@ -2300,6 +2459,7 @@ jobs:
|
|||
-e DD_API_KEY=$DD_API_KEY \
|
||||
-e DD_SITE=$DD_SITE \
|
||||
-e LITELLM_LICENSE=$LITELLM_LICENSE \
|
||||
--add-host host.docker.internal:host-gateway \
|
||||
--name my-app \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/pass_through_config.yaml:/app/config.yaml \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/custom_auth_basic.py:/app/custom_auth_basic.py \
|
||||
|
|
@ -2307,14 +2467,6 @@ jobs:
|
|||
--config /app/config.yaml \
|
||||
--port 4000 \
|
||||
--detailed_debug \
|
||||
- run:
|
||||
name: Install curl and dockerize
|
||||
command: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y curl
|
||||
sudo wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start outputting logs
|
||||
command: docker logs -f my-app
|
||||
|
|
@ -2383,6 +2535,7 @@ jobs:
|
|||
ls
|
||||
python -m pytest -vv tests/pass_through_tests/ -x --junitxml=test-results/junit.xml --durations=5
|
||||
no_output_timeout: 120m
|
||||
|
||||
# Store test results
|
||||
- store_test_results:
|
||||
path: test-results
|
||||
|
|
@ -2429,16 +2582,6 @@ jobs:
|
|||
command: |
|
||||
cp model_prices_and_context_window.json litellm/model_prices_and_context_window_backup.json
|
||||
|
||||
- run:
|
||||
name: Check if litellm dir, tests dir, or pyproject.toml was modified
|
||||
command: |
|
||||
if [ -n "$(git diff --name-only $CIRCLE_SHA1^..$CIRCLE_SHA1 | grep -E 'pyproject\.toml|litellm/|tests/')" ]; then
|
||||
echo "litellm, tests, or pyproject.toml updated"
|
||||
else
|
||||
echo "No changes to litellm, tests, or pyproject.toml. Skipping PyPI publish."
|
||||
circleci step halt
|
||||
fi
|
||||
|
||||
- run:
|
||||
name: Checkout code
|
||||
command: git checkout $CIRCLE_SHA1
|
||||
|
|
@ -2740,6 +2883,25 @@ jobs:
|
|||
steps:
|
||||
- checkout
|
||||
- setup_google_dns
|
||||
- run:
|
||||
name: Install dockerize
|
||||
command: |
|
||||
wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start PostgreSQL Database
|
||||
command: |
|
||||
docker run -d \
|
||||
--name postgres-db \
|
||||
-e POSTGRES_USER=postgres \
|
||||
-e POSTGRES_PASSWORD=postgres \
|
||||
-e POSTGRES_DB=circle_test \
|
||||
-p 5432:5432 \
|
||||
postgres:14
|
||||
- run:
|
||||
name: Wait for PostgreSQL to be ready
|
||||
command: dockerize -wait tcp://localhost:5432 -timeout 1m
|
||||
- run:
|
||||
name: Build Docker image
|
||||
command: |
|
||||
|
|
@ -2759,7 +2921,6 @@ jobs:
|
|||
name: Check for expected error
|
||||
command: |
|
||||
if grep -q "Error: P1001: Can't reach database server at" docker_output.log && \
|
||||
grep -q "prisma.engine.errors.NotConnectedError: Not connected to the query engine" docker_output.log && \
|
||||
grep -q "ERROR: Application startup failed. Exiting." docker_output.log; then
|
||||
echo "Expected error found. Test passed."
|
||||
else
|
||||
|
|
@ -2904,6 +3065,12 @@ workflows:
|
|||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- google_generate_content_endpoint_testing:
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- llm_responses_api_testing:
|
||||
filters:
|
||||
branches:
|
||||
|
|
@ -2950,6 +3117,7 @@ workflows:
|
|||
requires:
|
||||
- llm_translation_testing
|
||||
- mcp_testing
|
||||
- google_generate_content_endpoint_testing
|
||||
- guardrails_testing
|
||||
- llm_responses_api_testing
|
||||
- litellm_mapped_tests
|
||||
|
|
@ -3009,6 +3177,7 @@ workflows:
|
|||
- test_bad_database_url
|
||||
- llm_translation_testing
|
||||
- mcp_testing
|
||||
- google_generate_content_endpoint_testing
|
||||
- llm_responses_api_testing
|
||||
- litellm_mapped_tests
|
||||
- batches_testing
|
||||
|
|
|
|||
|
|
@ -110,6 +110,22 @@ data:
|
|||
|
||||
Source: [GitHub Gist from troyharvey](https://gist.github.com/troyharvey/4506472732157221e04c6b15e3b3f094)
|
||||
|
||||
### Migration Job Settings
|
||||
|
||||
The migration job supports both ArgoCD and Helm hooks to ensure database migrations run at the appropriate time during deployments.
|
||||
|
||||
| Name | Description | Value |
|
||||
| ---------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | ----- |
|
||||
| `migrationJob.enabled` | Enable or disable the schema migration Job | `true` |
|
||||
| `migrationJob.backoffLimit` | Backoff limit for Job restarts | `4` |
|
||||
| `migrationJob.ttlSecondsAfterFinished` | TTL for completed migration jobs | `120` |
|
||||
| `migrationJob.annotations` | Additional annotations for the migration job pod | `{}` |
|
||||
| `migrationJob.extraContainers` | Additional containers to run alongside the migration job | `[]` |
|
||||
| `migrationJob.hooks.argocd.enabled` | Enable ArgoCD hooks for the migration job (uses PreSync hook with BeforeHookCreation delete policy) | `true` |
|
||||
| `migrationJob.hooks.helm.enabled` | Enable Helm hooks for the migration job (uses pre-install,pre-upgrade hooks with before-hook-creation delete policy) | `false` |
|
||||
| `migrationJob.hooks.helm.weight` | Helm hook execution order (lower weights executed first). Optional - defaults to "1" if not specified. | N/A |
|
||||
|
||||
|
||||
## Accessing the Admin UI
|
||||
When browsing to the URL published per the settings in `ingress.*`, you will
|
||||
be prompted for **Admin Configuration**. The **Proxy Endpoint** is the internal
|
||||
|
|
|
|||
|
|
@ -99,6 +99,12 @@ spec:
|
|||
value: {{ $val | quote }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
{{- if .Values.separateHealthApp }}
|
||||
- name: SEPARATE_HEALTH_APP
|
||||
value: "1"
|
||||
- name: SEPARATE_HEALTH_PORT
|
||||
value: {{ .Values.separateHealthPort | default "8081" | quote }}
|
||||
{{- end }}
|
||||
{{- with .Values.extraEnvVars }}
|
||||
{{- toYaml . | nindent 12 }}
|
||||
{{- end }}
|
||||
|
|
@ -118,19 +124,23 @@ spec:
|
|||
- name: http
|
||||
containerPort: {{ .Values.service.port }}
|
||||
protocol: TCP
|
||||
{{- if .Values.separateHealthApp }}
|
||||
- name: health
|
||||
containerPort: {{ .Values.separateHealthPort | default 8081 }}
|
||||
protocol: TCP
|
||||
{{- end }}
|
||||
livenessProbe:
|
||||
httpGet:
|
||||
path: /health/liveliness
|
||||
port: http
|
||||
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
|
||||
readinessProbe:
|
||||
httpGet:
|
||||
path: /health/readiness
|
||||
port: http
|
||||
# Give the container time to start up. Up to 5 minutes (10 * 30 seconds)
|
||||
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
|
||||
startupProbe:
|
||||
httpGet:
|
||||
path: /health/readiness
|
||||
port: http
|
||||
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
|
||||
failureThreshold: 30
|
||||
periodSeconds: 10
|
||||
resources:
|
||||
|
|
|
|||
|
|
@ -5,8 +5,15 @@ kind: Job
|
|||
metadata:
|
||||
name: {{ include "litellm.fullname" . }}-migrations
|
||||
annotations:
|
||||
{{- if .Values.migrationJob.hooks.argocd.enabled }}
|
||||
argocd.argoproj.io/hook: PreSync
|
||||
argocd.argoproj.io/hook-delete-policy: BeforeHookCreation # delete old migration on a new deploy in case the migration needs to make updates
|
||||
argocd.argoproj.io/hook-delete-policy: BeforeHookCreation
|
||||
{{- end }}
|
||||
{{- if .Values.migrationJob.hooks.helm.enabled }}
|
||||
helm.sh/hook: "pre-install,pre-upgrade"
|
||||
helm.sh/hook-delete-policy: "before-hook-creation"
|
||||
helm.sh/hook-weight: {{ .Values.migrationJob.hooks.helm.weight | default "1" | quote }}
|
||||
{{- end }}
|
||||
checksum/config: {{ toYaml .Values | sha256sum }}
|
||||
spec:
|
||||
template:
|
||||
|
|
@ -47,8 +54,6 @@ spec:
|
|||
- name: DATABASE_URL
|
||||
value: postgresql://{{ .Values.postgresql.auth.username }}:{{ .Values.postgresql.auth.password }}@{{ .Release.Name }}-postgresql/{{ .Values.postgresql.auth.database }}
|
||||
{{- end }}
|
||||
- name: DISABLE_SCHEMA_UPDATE
|
||||
value: "false" # always run the migration from the Helm PreSync hook, override the value set
|
||||
{{- if .Values.envVars }}
|
||||
{{- range $key, $val := .Values.envVars }}
|
||||
- name: {{ $key }}
|
||||
|
|
@ -58,6 +63,8 @@ spec:
|
|||
{{- with .Values.extraEnvVars }}
|
||||
{{- toYaml . | nindent 12 }}
|
||||
{{- end }}
|
||||
- name: DISABLE_SCHEMA_UPDATE
|
||||
value: "false" # always run the migration from the Helm PreSync hook, override the value set
|
||||
{{- with .Values.volumeMounts }}
|
||||
volumeMounts:
|
||||
{{- toYaml . | nindent 12 }}
|
||||
|
|
|
|||
|
|
@ -63,6 +63,12 @@ service:
|
|||
# optionally specify loadBalancerClass
|
||||
# loadBalancerClass: tailscale
|
||||
|
||||
# Separate health app configuration
|
||||
# When enabled, health checks will use a separate port and the application
|
||||
# will receive SEPARATE_HEALTH_APP=1 and SEPARATE_HEALTH_PORT from environment variables
|
||||
separateHealthApp: false
|
||||
separateHealthPort: 8081
|
||||
|
||||
ingress:
|
||||
enabled: false
|
||||
className: "nginx"
|
||||
|
|
@ -201,6 +207,13 @@ migrationJob:
|
|||
annotations: {}
|
||||
ttlSecondsAfterFinished: 120
|
||||
extraContainers: []
|
||||
|
||||
# Hook configuration
|
||||
hooks:
|
||||
argocd:
|
||||
enabled: true
|
||||
helm:
|
||||
enabled: false
|
||||
|
||||
# Additional environment variables to be added to the deployment as a map of key-value pairs
|
||||
envVars: {
|
||||
|
|
|
|||
|
|
@ -1,68 +1,66 @@
|
|||
version: "3.11"
|
||||
services:
|
||||
litellm:
|
||||
build:
|
||||
context: .
|
||||
args:
|
||||
target: runtime
|
||||
image: ghcr.io/berriai/litellm:main-stable
|
||||
#########################################
|
||||
## Uncomment these lines to start proxy with a config.yaml file ##
|
||||
# volumes:
|
||||
# - ./config.yaml:/app/config.yaml <<- this is missing in the docker-compose file currently
|
||||
# command:
|
||||
# - "--config=/app/config.yaml"
|
||||
##############################################
|
||||
ports:
|
||||
- "4000:4000" # Map the container port to the host, change the host port if necessary
|
||||
environment:
|
||||
DATABASE_URL: "postgresql://llmproxy:dbpassword9090@db:5432/litellm"
|
||||
STORE_MODEL_IN_DB: "True" # allows adding models to proxy via UI
|
||||
env_file:
|
||||
- .env # Load local .env file
|
||||
depends_on:
|
||||
- db # Indicates that this service depends on the 'db' service, ensuring 'db' starts first
|
||||
healthcheck: # Defines the health check configuration for the container
|
||||
test: [ "CMD-SHELL", "wget --no-verbose --tries=1 http://localhost:4000/health/liveliness || exit 1" ] # Command to execute for health check
|
||||
interval: 30s # Perform health check every 30 seconds
|
||||
timeout: 10s # Health check command times out after 10 seconds
|
||||
retries: 3 # Retry up to 3 times if health check fails
|
||||
start_period: 40s # Wait 40 seconds after container start before beginning health checks
|
||||
|
||||
db:
|
||||
image: postgres:16
|
||||
restart: always
|
||||
container_name: litellm_db
|
||||
environment:
|
||||
POSTGRES_DB: litellm
|
||||
POSTGRES_USER: llmproxy
|
||||
POSTGRES_PASSWORD: dbpassword9090
|
||||
ports:
|
||||
- "5432:5432"
|
||||
volumes:
|
||||
- postgres_data:/var/lib/postgresql/data # Persists Postgres data across container restarts
|
||||
healthcheck:
|
||||
test: ["CMD-SHELL", "pg_isready -d litellm -U llmproxy"]
|
||||
interval: 1s
|
||||
timeout: 5s
|
||||
retries: 10
|
||||
|
||||
prometheus:
|
||||
image: prom/prometheus
|
||||
volumes:
|
||||
- prometheus_data:/prometheus
|
||||
- ./prometheus.yml:/etc/prometheus/prometheus.yml
|
||||
ports:
|
||||
- "9090:9090"
|
||||
command:
|
||||
- "--config.file=/etc/prometheus/prometheus.yml"
|
||||
- "--storage.tsdb.path=/prometheus"
|
||||
- "--storage.tsdb.retention.time=15d"
|
||||
restart: always
|
||||
|
||||
volumes:
|
||||
prometheus_data:
|
||||
driver: local
|
||||
postgres_data:
|
||||
name: litellm_postgres_data # Named volume for Postgres data persistence
|
||||
|
||||
services:
|
||||
litellm:
|
||||
build:
|
||||
context: .
|
||||
args:
|
||||
target: runtime
|
||||
image: ghcr.io/berriai/litellm:main-stable
|
||||
#########################################
|
||||
## Uncomment these lines to start proxy with a config.yaml file ##
|
||||
# volumes:
|
||||
# - ./config.yaml:/app/config.yaml <<- this is missing in the docker-compose file currently
|
||||
# command:
|
||||
# - "--config=/app/config.yaml"
|
||||
##############################################
|
||||
ports:
|
||||
- "4000:4000" # Map the container port to the host, change the host port if necessary
|
||||
environment:
|
||||
DATABASE_URL: "postgresql://llmproxy:dbpassword9090@db:5432/litellm"
|
||||
STORE_MODEL_IN_DB: "True" # allows adding models to proxy via UI
|
||||
env_file:
|
||||
- .env # Load local .env file
|
||||
depends_on:
|
||||
- db # Indicates that this service depends on the 'db' service, ensuring 'db' starts first
|
||||
healthcheck: # Defines the health check configuration for the container
|
||||
test: [ "CMD-SHELL", "wget --no-verbose --tries=1 http://localhost:4000/health/liveliness || exit 1" ] # Command to execute for health check
|
||||
interval: 30s # Perform health check every 30 seconds
|
||||
timeout: 10s # Health check command times out after 10 seconds
|
||||
retries: 3 # Retry up to 3 times if health check fails
|
||||
start_period: 40s # Wait 40 seconds after container start before beginning health checks
|
||||
|
||||
db:
|
||||
image: postgres:16
|
||||
restart: always
|
||||
container_name: litellm_db
|
||||
environment:
|
||||
POSTGRES_DB: litellm
|
||||
POSTGRES_USER: llmproxy
|
||||
POSTGRES_PASSWORD: dbpassword9090
|
||||
ports:
|
||||
- "5432:5432"
|
||||
volumes:
|
||||
- postgres_data:/var/lib/postgresql/data # Persists Postgres data across container restarts
|
||||
healthcheck:
|
||||
test: ["CMD-SHELL", "pg_isready -d litellm -U llmproxy"]
|
||||
interval: 1s
|
||||
timeout: 5s
|
||||
retries: 10
|
||||
|
||||
prometheus:
|
||||
image: prom/prometheus
|
||||
volumes:
|
||||
- prometheus_data:/prometheus
|
||||
- ./prometheus.yml:/etc/prometheus/prometheus.yml
|
||||
ports:
|
||||
- "9090:9090"
|
||||
command:
|
||||
- "--config.file=/etc/prometheus/prometheus.yml"
|
||||
- "--storage.tsdb.path=/prometheus"
|
||||
- "--storage.tsdb.retention.time=15d"
|
||||
restart: always
|
||||
|
||||
volumes:
|
||||
prometheus_data:
|
||||
driver: local
|
||||
postgres_data:
|
||||
name: litellm_postgres_data # Named volume for Postgres data persistence
|
||||
|
|
|
|||
|
|
@ -33,7 +33,7 @@ WORKDIR /app
|
|||
# Install runtime dependencies
|
||||
USER root
|
||||
RUN apk upgrade --no-cache && \
|
||||
apk add --no-cache bash
|
||||
apk add --no-cache bash libstdc++ ca-certificates openssl
|
||||
|
||||
# Copy only necessary artifacts from builder stage for runtime
|
||||
COPY --from=builder /app/docker/entrypoint.sh /app/docker/prod_entrypoint.sh /app/docker/
|
||||
|
|
@ -71,6 +71,20 @@ RUN mkdir -p /nonexistent /.npm && \
|
|||
PRISMA_PATH=$(python -c "import os, prisma; print(os.path.dirname(prisma.__file__))") && \
|
||||
chown -R nobody:nogroup $PRISMA_PATH
|
||||
|
||||
# --- OpenShift Compatibility: Apply Red Hat recommended pattern ---
|
||||
# Get paths for directories that need write access at runtime
|
||||
RUN PRISMA_PATH=$(python -c "import os, prisma; print(os.path.dirname(prisma.__file__))") && \
|
||||
LITELLM_PROXY_EXTRAS_PATH=$(python -c "import os, litellm_proxy_extras; print(os.path.dirname(litellm_proxy_extras.__file__))" 2>/dev/null || echo "") && \
|
||||
# Set group ownership to 0 (root group) for OpenShift compatibility && \
|
||||
chgrp -R 0 $PRISMA_PATH && \
|
||||
[ -n "$LITELLM_PROXY_EXTRAS_PATH" ] && chgrp -R 0 $LITELLM_PROXY_EXTRAS_PATH || true && \
|
||||
# Mirror owner permissions to group (g=u) as recommended by Red Hat && \
|
||||
chmod -R g=u $PRISMA_PATH && \
|
||||
[ -n "$LITELLM_PROXY_EXTRAS_PATH" ] && chmod -R g=u $LITELLM_PROXY_EXTRAS_PATH || true && \
|
||||
# Ensure directories are writable by group && \
|
||||
chmod -R g+w $PRISMA_PATH && \
|
||||
[ -n "$LITELLM_PROXY_EXTRAS_PATH" ] && chmod -R g+w $LITELLM_PROXY_EXTRAS_PATH || true
|
||||
|
||||
# Switch to non-root user
|
||||
USER nobody
|
||||
|
||||
|
|
@ -86,4 +100,4 @@ ENTRYPOINT ["/app/docker/prod_entrypoint.sh"]
|
|||
|
||||
# Append "--detailed_debug" to the end of CMD to view detailed debug logs
|
||||
# CMD ["--port", "4000", "--detailed_debug"]
|
||||
CMD ["--port", "4000"]
|
||||
CMD ["--port", "4000"]
|
||||
|
|
@ -1,3 +1,65 @@
|
|||
# LiteLLM Docker
|
||||
# Docker Development Guide
|
||||
|
||||
This is a minimal Docker Compose setup for self-hosting LiteLLM.
|
||||
This guide provides instructions for building and running the LiteLLM application using Docker and Docker Compose.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- Docker
|
||||
- Docker Compose
|
||||
|
||||
## Building and Running the Application
|
||||
|
||||
To build and run the application, you will use the `docker-compose.yml` file located in the root of the project. This file is configured to use the `Dockerfile.non_root` for a secure, non-root container environment.
|
||||
|
||||
### 1. Set the Master Key
|
||||
|
||||
The application requires a `MASTER_KEY` for signing and validating tokens. You must set this key as an environment variable before running the application.
|
||||
|
||||
Create a `.env` file in the root of the project and add the following line:
|
||||
|
||||
```
|
||||
MASTER_KEY=your-secret-key
|
||||
```
|
||||
|
||||
Replace `your-secret-key` with a strong, randomly generated secret.
|
||||
|
||||
### 2. Build and Run the Containers
|
||||
|
||||
Once you have set the `MASTER_KEY`, you can build and run the containers using the following command:
|
||||
|
||||
```bash
|
||||
docker-compose up -d --build
|
||||
```
|
||||
|
||||
This command will:
|
||||
|
||||
- Build the Docker image using `Dockerfile.non_root`.
|
||||
- Start the `litellm`, `litellm_db`, and `prometheus` services in detached mode (`-d`).
|
||||
- The `--build` flag ensures that the image is rebuilt if there are any changes to the Dockerfile or the application code.
|
||||
|
||||
### 3. Verifying the Application is Running
|
||||
|
||||
You can check the status of the running containers with the following command:
|
||||
|
||||
```bash
|
||||
docker-compose ps
|
||||
```
|
||||
|
||||
To view the logs of the `litellm` container, run:
|
||||
|
||||
```bash
|
||||
docker-compose logs -f litellm
|
||||
```
|
||||
|
||||
### 4. Stopping the Application
|
||||
|
||||
To stop the running containers, use the following command:
|
||||
|
||||
```bash
|
||||
docker-compose down
|
||||
```
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
- **`build_admin_ui.sh: not found`**: This error can occur if the Docker build context is not set correctly. Ensure that you are running the `docker-compose` command from the root of the project.
|
||||
- **`Master key is not initialized`**: This error means the `MASTER_key` environment variable is not set. Make sure you have created a `.env` file in the project root with the `MASTER_KEY` defined.
|
||||
|
|
|
|||
446
docs/my-website/docs/completion/computer_use.md
Normal file
|
|
@ -0,0 +1,446 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Computer Use
|
||||
|
||||
Computer use allows models to interact with computer interfaces by taking screenshots and performing actions like clicking, typing, and scrolling. This enables AI models to autonomously operate desktop environments.
|
||||
|
||||
**Supported Providers:**
|
||||
- Anthropic API (`anthropic/`)
|
||||
- Bedrock (Anthropic) (`bedrock/`)
|
||||
- Vertex AI (Anthropic) (`vertex_ai/`)
|
||||
|
||||
**Supported Tool Types:**
|
||||
- `computer` - Computer interaction tool with display parameters
|
||||
- `bash` - Bash shell tool
|
||||
- `text_editor` - Text editor tool
|
||||
- `web_search` - Web search tool
|
||||
|
||||
LiteLLM will standardize the computer use tools across all supported providers.
|
||||
|
||||
## Quick Start
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="LiteLLM Python SDK">
|
||||
|
||||
```python
|
||||
import os
|
||||
from litellm import completion
|
||||
|
||||
os.environ["ANTHROPIC_API_KEY"] = "your-api-key"
|
||||
|
||||
# Computer use tool
|
||||
tools = [
|
||||
{
|
||||
"type": "computer_20241022",
|
||||
"name": "computer",
|
||||
"display_height_px": 768,
|
||||
"display_width_px": 1024,
|
||||
"display_number": 0,
|
||||
}
|
||||
]
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Take a screenshot and tell me what you see"
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg=="
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
response = completion(
|
||||
model="anthropic/claude-3-5-sonnet-latest",
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="LiteLLM Proxy Server">
|
||||
|
||||
1. Define computer use models on config.yaml
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: claude-3-5-sonnet-latest # Anthropic claude-3-5-sonnet-latest
|
||||
litellm_params:
|
||||
model: anthropic/claude-3-5-sonnet-latest
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
- model_name: claude-bedrock # Bedrock Anthropic model
|
||||
litellm_params:
|
||||
model: bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0
|
||||
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
|
||||
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
|
||||
aws_region_name: us-west-2
|
||||
model_info:
|
||||
supports_computer_use: True # set supports_computer_use to True so /model/info returns this attribute as True
|
||||
```
|
||||
|
||||
2. Run proxy server
|
||||
|
||||
```bash
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
3. Test it using the OpenAI Python SDK
|
||||
|
||||
```python
|
||||
import os
|
||||
from openai import OpenAI
|
||||
|
||||
client = OpenAI(
|
||||
api_key="sk-1234", # your litellm proxy api key
|
||||
base_url="http://0.0.0.0:4000"
|
||||
)
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model="claude-3-5-sonnet-latest",
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Take a screenshot and tell me what you see"
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg=="
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
tools=[
|
||||
{
|
||||
"type": "computer_20241022",
|
||||
"name": "computer",
|
||||
"display_height_px": 768,
|
||||
"display_width_px": 1024,
|
||||
"display_number": 0,
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Checking if a model supports `computer use`
|
||||
|
||||
<Tabs>
|
||||
<TabItem label="LiteLLM Python SDK" value="Python">
|
||||
|
||||
Use `litellm.supports_computer_use(model="")` -> returns `True` if model supports computer use and `False` if not
|
||||
|
||||
```python
|
||||
import litellm
|
||||
|
||||
assert litellm.supports_computer_use(model="anthropic/claude-3-5-sonnet-latest") == True
|
||||
assert litellm.supports_computer_use(model="anthropic/claude-3-7-sonnet-20250219") == True
|
||||
assert litellm.supports_computer_use(model="bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0") == True
|
||||
assert litellm.supports_computer_use(model="vertex_ai/claude-3-5-sonnet") == True
|
||||
assert litellm.supports_computer_use(model="openai/gpt-4") == False
|
||||
```
|
||||
</TabItem>
|
||||
|
||||
<TabItem label="LiteLLM Proxy Server" value="proxy">
|
||||
|
||||
1. Define computer use models on config.yaml
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: claude-3-5-sonnet-latest # Anthropic claude-3-5-sonnet-latest
|
||||
litellm_params:
|
||||
model: anthropic/claude-3-5-sonnet-latest
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
- model_name: claude-bedrock # Bedrock Anthropic model
|
||||
litellm_params:
|
||||
model: bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0
|
||||
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
|
||||
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
|
||||
aws_region_name: us-west-2
|
||||
model_info:
|
||||
supports_computer_use: True # set supports_computer_use to True so /model/info returns this attribute as True
|
||||
```
|
||||
|
||||
2. Run proxy server
|
||||
|
||||
```bash
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
3. Call `/model_group/info` to check if your model supports `computer use`
|
||||
|
||||
```shell
|
||||
curl -X 'GET' \
|
||||
'http://localhost:4000/model_group/info' \
|
||||
-H 'accept: application/json' \
|
||||
-H 'x-api-key: sk-1234'
|
||||
```
|
||||
|
||||
Expected Response
|
||||
|
||||
```json
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"model_group": "claude-3-5-sonnet-latest",
|
||||
"providers": ["anthropic"],
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"supports_computer_use": true, # 👈 supports_computer_use is true
|
||||
"supports_vision": true,
|
||||
"supports_function_calling": true
|
||||
},
|
||||
{
|
||||
"model_group": "claude-bedrock",
|
||||
"providers": ["bedrock"],
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"supports_computer_use": true, # 👈 supports_computer_use is true
|
||||
"supports_vision": true,
|
||||
"supports_function_calling": true
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Different Tool Types
|
||||
|
||||
Computer use supports several different tool types for various interaction modes:
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="computer" label="Computer Tool">
|
||||
|
||||
The `computer_20241022` tool provides direct screen interaction capabilities.
|
||||
|
||||
```python
|
||||
import os
|
||||
from litellm import completion
|
||||
|
||||
os.environ["ANTHROPIC_API_KEY"] = "your-api-key"
|
||||
|
||||
tools = [
|
||||
{
|
||||
"type": "computer_20241022",
|
||||
"name": "computer",
|
||||
"display_height_px": 768,
|
||||
"display_width_px": 1024,
|
||||
"display_number": 0,
|
||||
}
|
||||
]
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Click on the search button in the screenshot"
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg=="
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
response = completion(
|
||||
model="anthropic/claude-3-5-sonnet-latest",
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="bash" label="Bash Tool">
|
||||
|
||||
The `bash_20241022` tool provides command line interface access.
|
||||
|
||||
```python
|
||||
import os
|
||||
from litellm import completion
|
||||
|
||||
os.environ["ANTHROPIC_API_KEY"] = "your-api-key"
|
||||
|
||||
tools = [
|
||||
{
|
||||
"type": "bash_20241022",
|
||||
"name": "bash"
|
||||
}
|
||||
]
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "List the files in the current directory using bash"
|
||||
}
|
||||
]
|
||||
|
||||
response = completion(
|
||||
model="anthropic/claude-3-5-sonnet-latest",
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="text_editor" label="Text Editor Tool">
|
||||
|
||||
The `text_editor_20250124` tool provides text file editing capabilities.
|
||||
|
||||
```python
|
||||
import os
|
||||
from litellm import completion
|
||||
|
||||
os.environ["ANTHROPIC_API_KEY"] = "your-api-key"
|
||||
|
||||
tools = [
|
||||
{
|
||||
"type": "text_editor_20250124",
|
||||
"name": "str_replace_editor"
|
||||
}
|
||||
]
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Create a simple Python hello world script"
|
||||
}
|
||||
]
|
||||
|
||||
response = completion(
|
||||
model="anthropic/claude-3-5-sonnet-latest",
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Advanced Usage with Multiple Tools
|
||||
|
||||
You can combine different computer use tools in a single request:
|
||||
|
||||
```python
|
||||
import os
|
||||
from litellm import completion
|
||||
|
||||
os.environ["ANTHROPIC_API_KEY"] = "your-api-key"
|
||||
|
||||
tools = [
|
||||
{
|
||||
"type": "computer_20241022",
|
||||
"name": "computer",
|
||||
"display_height_px": 768,
|
||||
"display_width_px": 1024,
|
||||
"display_number": 0,
|
||||
},
|
||||
{
|
||||
"type": "bash_20241022",
|
||||
"name": "bash"
|
||||
},
|
||||
{
|
||||
"type": "text_editor_20250124",
|
||||
"name": "str_replace_editor"
|
||||
}
|
||||
]
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Take a screenshot, then create a file describing what you see, and finally use bash to show the file contents"
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg=="
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
response = completion(
|
||||
model="anthropic/claude-3-5-sonnet-latest",
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
## Spec
|
||||
|
||||
### Computer Tool (`computer_20241022`)
|
||||
|
||||
```json
|
||||
{
|
||||
"type": "computer_20241022",
|
||||
"name": "computer",
|
||||
"display_height_px": 768, // Required: Screen height in pixels
|
||||
"display_width_px": 1024, // Required: Screen width in pixels
|
||||
"display_number": 0 // Optional: Display number (default: 0)
|
||||
}
|
||||
```
|
||||
|
||||
### Bash Tool (`bash_20241022`)
|
||||
|
||||
```json
|
||||
{
|
||||
"type": "bash_20241022",
|
||||
"name": "bash" // Required: Tool name
|
||||
}
|
||||
```
|
||||
|
||||
### Text Editor Tool (`text_editor_20250124`)
|
||||
|
||||
```json
|
||||
{
|
||||
"type": "text_editor_20250124",
|
||||
"name": "str_replace_editor" // Required: Tool name
|
||||
}
|
||||
```
|
||||
|
||||
### Web Search Tool (`web_search_20250305`)
|
||||
|
||||
```json
|
||||
{
|
||||
"type": "web_search_20250305",
|
||||
"name": "web_search" // Required: Tool name
|
||||
}
|
||||
```
|
||||
|
|
@ -9,8 +9,14 @@ LiteLLM Supports logging to the following Datdog Integrations:
|
|||
- `datadog_llm_observability` [Datadog LLM Observability](https://www.datadoghq.com/product/llm-observability/)
|
||||
- `ddtrace-run` [Datadog Tracing](#datadog-tracing)
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="datadog" label="Datadog Logs">
|
||||
## Datadog Logs
|
||||
|
||||
| Feature | Details |
|
||||
|---------|---------|
|
||||
| **What is logged** | [StandardLoggingPayload](../proxy/logging_spec) |
|
||||
| **Events** | Success + Failure |
|
||||
| **Product Link** | [Datadog Logs](https://docs.datadoghq.com/logs/) |
|
||||
|
||||
|
||||
We will use the `--config` to set `litellm.callbacks = ["datadog"]` this will log all successful LLM calls to DataDog
|
||||
|
||||
|
|
@ -26,8 +32,16 @@ litellm_settings:
|
|||
service_callback: ["datadog"] # logs redis, postgres failures on datadog
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="datadog_llm_observability" label="Datadog LLM Observability">
|
||||
|
||||
## Datadog LLM Observability
|
||||
|
||||
**Overview**
|
||||
|
||||
| Feature | Details |
|
||||
|---------|---------|
|
||||
| **What is logged** | [StandardLoggingPayload](../proxy/logging_spec) |
|
||||
| **Events** | Success + Failure |
|
||||
| **Product Link** | [Datadog LLM Observability](https://www.datadoghq.com/product/llm-observability/) |
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
|
|
@ -38,8 +52,7 @@ litellm_settings:
|
|||
callbacks: ["datadog_llm_observability"] # logs llm success logs on datadog
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
|
||||
**Step 2**: Set Required env variables for datadog
|
||||
|
||||
|
|
@ -80,7 +93,53 @@ Expected output on Datadog
|
|||
|
||||
<Image img={require('../../img/dd_small1.png')} />
|
||||
|
||||
#### Datadog Tracing
|
||||
### Redacting Messages and Responses
|
||||
|
||||
This section covers how to redact sensitive data from messages and responses in the logged payload on Datadog LLM Observability.
|
||||
|
||||
|
||||
When redaction is enabled, the actual message content and response text will be excluded from Datadog logs while preserving metadata like token counts, latency, and model information.
|
||||
|
||||
**Step 1**: Configure redaction in your `config.yaml`
|
||||
|
||||
```yaml showLineNumbers title="config.yaml"
|
||||
model_list:
|
||||
- model_name: gpt-3.5-turbo
|
||||
litellm_params:
|
||||
model: gpt-3.5-turbo
|
||||
litellm_settings:
|
||||
callbacks: ["datadog_llm_observability"] # logs llm success logs on datadog
|
||||
|
||||
# Params to apply only for "datadog_llm_observability" callback
|
||||
datadog_llm_observability_params:
|
||||
turn_off_message_logging: true # redacts input messages and output responses
|
||||
```
|
||||
|
||||
**Step 2**: Send a chat completion request
|
||||
|
||||
```shell
|
||||
curl --location 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "what llm are you"
|
||||
}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
**Step 3**: Verify redaction in Datadog LLM Observability
|
||||
|
||||
On the Datadog LLM Observability page, you should see that both input messages and output responses are redacted, while metadata (token counts, timing, model info) remains visible.
|
||||
|
||||
<Image img={require('../../img/dd_llm_obs.png')} />
|
||||
|
||||
|
||||
|
||||
### Datadog Tracing
|
||||
|
||||
Use `ddtrace-run` to enable [Datadog Tracing](https://ddtrace.readthedocs.io/en/stable/installation_quickstart.html) on litellm proxy
|
||||
|
||||
|
|
@ -104,7 +163,7 @@ docker run \
|
|||
--config /app/config.yaml --detailed_debug
|
||||
```
|
||||
|
||||
### Set DD variables (`DD_SERVICE` etc)
|
||||
## Set DD variables (`DD_SERVICE` etc)
|
||||
|
||||
LiteLLM supports customizing the following Datadog environment variables
|
||||
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ MLflow’s integration with LiteLLM supports advanced observability compatible w
|
|||
Install MLflow:
|
||||
|
||||
```shell
|
||||
pip install mlflow
|
||||
pip install "litellm[mlflow]"
|
||||
```
|
||||
|
||||
To enable MLflow auto tracing for LiteLLM:
|
||||
|
|
@ -160,6 +160,102 @@ class CustomAgent:
|
|||
|
||||
This approach generates a unified trace, combining your custom Python code with LiteLLM calls.
|
||||
|
||||
## LiteLLM Proxy Server
|
||||
|
||||
### Dependencies
|
||||
|
||||
For using `mlflow` on LiteLLM Proxy Server, you need to install the `mlflow` package on your docker container.
|
||||
|
||||
```shell
|
||||
pip install "mlflow>=3.1.4"
|
||||
```
|
||||
|
||||
### Configuration
|
||||
|
||||
Configure MLflow in your LiteLLM proxy configuration file:
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: openai/*
|
||||
litellm_params:
|
||||
model: openai/*
|
||||
|
||||
litellm_settings:
|
||||
success_callback: ["mlflow"]
|
||||
failure_callback: ["mlflow"]
|
||||
```
|
||||
|
||||
### Environment Variables
|
||||
|
||||
For MLflow with Databricks service, set these required environment variables:
|
||||
|
||||
```shell
|
||||
DATABRICKS_TOKEN="dapixxxxx"
|
||||
DATABRICKS_HOST="https://dbc-xxxx.cloud.databricks.com"
|
||||
MLFLOW_TRACKING_URI="databricks"
|
||||
MLFLOW_REGISTRY_URI="databricks-uc"
|
||||
MLFLOW_EXPERIMENT_ID="xxxx"
|
||||
```
|
||||
|
||||
### Adding Tags for Better Tracing
|
||||
|
||||
You can add custom tags to your requests for improved trace organization and filtering in MLflow. Tags help you categorize and search your traces by job ID, task name, or any custom metadata.
|
||||
|
||||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="curl" label="curl">
|
||||
|
||||
```shell
|
||||
curl --location 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--data '{
|
||||
"model": "gemini-2.5-flash",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "what llm are you"
|
||||
}
|
||||
],
|
||||
"litellm_metadata": {
|
||||
"tags": ["jobID:214590dsff09fds", "taskName:run_page_classification"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="openai-python" label="OpenAI Python SDK">
|
||||
|
||||
```python
|
||||
from openai import OpenAI
|
||||
|
||||
# Initialize the OpenAI client pointing to your LiteLLM proxy
|
||||
client = OpenAI(
|
||||
api_key="sk-1234", # Your LiteLLM proxy API key
|
||||
base_url="http://0.0.0.0:4000" # Your LiteLLM proxy URL
|
||||
)
|
||||
|
||||
# Make a request with tags in metadata
|
||||
response = client.chat.completions.create(
|
||||
model="gemini-2.5-flash",
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": "what llm are you"
|
||||
}
|
||||
],
|
||||
extra_body={
|
||||
"litellm_metadata": {
|
||||
"tags": ["jobID:214590dsff09fds", "taskName:run_page_classification"]
|
||||
}
|
||||
}
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Support
|
||||
|
||||
|
|
|
|||
|
|
@ -33,7 +33,7 @@ Supports **ALL** Bedrock Endpoints (including streaming).
|
|||
|
||||
Let's call the Bedrock [`/converse` endpoint](https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_Converse.html)
|
||||
|
||||
1. Add AWS Keyss to your environment
|
||||
1. Add AWS Keys to your environment
|
||||
|
||||
```bash
|
||||
export AWS_ACCESS_KEY_ID="" # Access key
|
||||
|
|
@ -295,4 +295,4 @@ for event in response.get("completion"):
|
|||
|
||||
print(completion)
|
||||
|
||||
```
|
||||
```
|
||||
|
|
|
|||
|
|
@ -175,6 +175,25 @@ print(response)
|
|||
</Tabs>
|
||||
|
||||
|
||||
### Setting API Version
|
||||
|
||||
You can set the `api_version` for Azure OpenAI in your proxy config.yaml in the following ways
|
||||
|
||||
#### Option 1: Per Model Configuration
|
||||
|
||||
```yaml showLineNumbers title="config.yaml"
|
||||
model_list:
|
||||
- model_name: gpt-4
|
||||
litellm_params:
|
||||
model: azure/my-gpt4-deployment
|
||||
api_base: https://your-resource.openai.azure.com/
|
||||
api_version: "2024-08-01-preview" # Set version per model
|
||||
api_key: os.environ/AZURE_API_KEY
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
## Azure OpenAI Chat Completion Models
|
||||
|
||||
|
|
|
|||
|
|
@ -14,29 +14,7 @@ import TabItem from '@theme/TabItem';
|
|||
|
||||
:::
|
||||
|
||||
### SSO for UI
|
||||
|
||||
#### Step 1: Set upperbounds for keys
|
||||
Control the upperbound that users can use for `max_budget`, `budget_duration` or any `key/generate` param per key.
|
||||
|
||||
```yaml
|
||||
litellm_settings:
|
||||
upperbound_key_generate_params:
|
||||
max_budget: 100 # Optional[float], optional): upperbound of $100, for all /key/generate requests
|
||||
budget_duration: "10d" # Optional[str], optional): upperbound of 10 days for budget_duration values
|
||||
duration: "30d" # Optional[str], optional): upperbound of 30 days for all /key/generate requests
|
||||
max_parallel_requests: 1000 # (Optional[int], optional): Max number of requests that can be made in parallel. Defaults to None.
|
||||
tpm_limit: 1000 #(Optional[int], optional): Tpm limit. Defaults to None.
|
||||
rpm_limit: 1000 #(Optional[int], optional): Rpm limit. Defaults to None.
|
||||
|
||||
```
|
||||
|
||||
** Expected Behavior **
|
||||
|
||||
- Send a `/key/generate` request with `max_budget=200`
|
||||
- Key will be created with `max_budget=100` since 100 is the upper bound
|
||||
|
||||
#### Step 2: Setup Oauth Client
|
||||
### Usage (Google, Microsoft, Okta, etc.)
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="okta" label="Okta SSO">
|
||||
|
|
|
|||
|
|
@ -38,8 +38,7 @@ litellm_settings:
|
|||
context_window_fallbacks: [{"gpt-3.5-turbo-small": ["gpt-3.5-turbo-large", "claude-opus"]}] # fallbacks for ContextWindowExceededErrors
|
||||
|
||||
# MCP Aliases - Map aliases to MCP server names for easier tool access
|
||||
mcp_aliases: { "github": "github_mcp_server", "zapier": "zapier_mcp_server", "deepwiki": "deepwiki_mcp_server" } # Maps friendly aliases to MCP server names. Only the first alias for each server is used.
|
||||
|
||||
mcp_aliases: { "github": "github_mcp_server", "zapier": "zapier_mcp_server", "deepwiki": "deepwiki_mcp_server" } # Maps friendly aliases to MCP server names. Only the first alias for each server is used
|
||||
|
||||
# Caching settings
|
||||
cache: true
|
||||
|
|
@ -327,6 +326,7 @@ router_settings:
|
|||
| ATHINA_BASE_URL | Base URL for Athina service (defaults to `https://log.athina.ai`)
|
||||
| AUTH_STRATEGY | Strategy used for authentication (e.g., OAuth, API key)
|
||||
| ANTHROPIC_API_KEY | API key for Anthropic service
|
||||
| ANTHROPIC_API_BASE | Base URL for Anthropic API. Default is https://api.anthropic.com
|
||||
| AWS_ACCESS_KEY_ID | Access Key ID for AWS services
|
||||
| AWS_PROFILE_NAME | AWS CLI profile name to be used
|
||||
| AWS_REGION_NAME | Default AWS region for service interactions
|
||||
|
|
@ -336,6 +336,7 @@ router_settings:
|
|||
| AWS_WEB_IDENTITY_TOKEN | Web identity token for AWS
|
||||
| AZURE_API_VERSION | Version of the Azure API being used
|
||||
| AZURE_AUTHORITY_HOST | Azure authority host URL
|
||||
| AZURE_CERTIFICATE_PASSWORD | Password for Azure OpenAI certificate
|
||||
| AZURE_CLIENT_ID | Client ID for Azure services
|
||||
| AZURE_CLIENT_SECRET | Client secret for Azure services
|
||||
| AZURE_CODE_INTERPRETER_COST_PER_SESSION | Cost per session for Azure Code Interpreter service
|
||||
|
|
@ -371,6 +372,7 @@ router_settings:
|
|||
| CONFIDENT_API_KEY | API key for DeepEval integration
|
||||
| CUSTOM_TIKTOKEN_CACHE_DIR | Custom directory for Tiktoken cache
|
||||
| CONFIDENT_API_KEY | API key for Confident AI (Deepeval) Logging service
|
||||
| COHERE_API_BASE | Base URL for Cohere API. Default is https://api.cohere.com
|
||||
| DATABASE_HOST | Hostname for the database server
|
||||
| DATABASE_NAME | Name of the database
|
||||
| DATABASE_PASSWORD | Password for the database user
|
||||
|
|
@ -481,6 +483,7 @@ router_settings:
|
|||
| GENERIC_USER_PROVIDER_ATTRIBUTE | Attribute specifying the user's provider
|
||||
| GENERIC_USER_ROLE_ATTRIBUTE | Attribute specifying the user's role
|
||||
| GENERIC_USERINFO_ENDPOINT | Endpoint to fetch user information in generic OAuth
|
||||
| GEMINI_API_BASE | Base URL for Gemini API. Default is https://generativelanguage.googleapis.com
|
||||
| GALILEO_BASE_URL | Base URL for Galileo platform
|
||||
| GALILEO_PASSWORD | Password for Galileo authentication
|
||||
| GALILEO_PROJECT_ID | Project ID for Galileo usage
|
||||
|
|
@ -580,7 +583,7 @@ router_settings:
|
|||
| MAX_LANGFUSE_INITIALIZED_CLIENTS | Maximum number of Langfuse clients to initialize on proxy. Default is 20. This is set since langfuse initializes 1 thread everytime a client is initialized. We've had an incident in the past where we reached 100% cpu utilization because Langfuse was initialized several times.
|
||||
| MIN_NON_ZERO_TEMPERATURE | Minimum non-zero temperature value. Default is 0.0001
|
||||
| MINIMUM_PROMPT_CACHE_TOKEN_COUNT | Minimum token count for caching a prompt. Default is 1024
|
||||
| MISTRAL_API_BASE | Base URL for Mistral API
|
||||
| MISTRAL_API_BASE | Base URL for Mistral API. Default is https://api.mistral.ai
|
||||
| MISTRAL_API_KEY | API key for Mistral API
|
||||
| MICROSOFT_CLIENT_ID | Client ID for Microsoft services
|
||||
| MICROSOFT_CLIENT_SECRET | Client secret for Microsoft services
|
||||
|
|
@ -592,7 +595,7 @@ router_settings:
|
|||
| NON_LLM_CONNECTION_TIMEOUT | Timeout in seconds for non-LLM service connections. Default is 15
|
||||
| OAUTH_TOKEN_INFO_ENDPOINT | Endpoint for OAuth token info retrieval
|
||||
| OPENAI_BASE_URL | Base URL for OpenAI API
|
||||
| OPENAI_API_BASE | Base URL for OpenAI API
|
||||
| OPENAI_API_BASE | Base URL for OpenAI API. Default is https://api.openai.com/
|
||||
| OPENAI_API_KEY | API key for OpenAI services
|
||||
| OPENAI_FILE_SEARCH_COST_PER_1K_CALLS | Cost per 1000 calls for OpenAI file search. Default is 0.0025
|
||||
| OPENAI_ORGANIZATION | Organization identifier for OpenAI
|
||||
|
|
|
|||
|
|
@ -60,9 +60,114 @@ Supported from v1.72.2+
|
|||
[Get free 7-day trial key](https://www.litellm.ai/enterprise#trial)
|
||||
:::
|
||||
|
||||
### Usage
|
||||
|
||||
1. Setup custom auth file
|
||||
|
||||
```python
|
||||
"""
|
||||
Example custom auth function.
|
||||
|
||||
This will allow all keys starting with "my-custom-key" to pass through.
|
||||
"""
|
||||
from typing import Union
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
async def user_api_key_auth(
|
||||
request: Request, api_key: str
|
||||
) -> Union[UserAPIKeyAuth, str]:
|
||||
try:
|
||||
if api_key.startswith("my-custom-key"):
|
||||
return "sk-P1zJMdsqCPNN54alZd_ETw"
|
||||
else:
|
||||
raise Exception("Invalid API key")
|
||||
except Exception:
|
||||
raise Exception("Invalid API key")
|
||||
|
||||
```
|
||||
|
||||
2. Setup config.yaml
|
||||
|
||||
Key change set `mode: auto`. This will check both litellm api key auth + custom auth.
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: "openai-model"
|
||||
litellm_params:
|
||||
model: "gpt-3.5-turbo"
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
general_settings:
|
||||
custom_auth: custom_auth_auto.user_api_key_auth
|
||||
custom_auth_settings:
|
||||
mode: "auto" # can be 'on', 'off', 'auto' - 'auto' checks both litellm api key auth + custom auth
|
||||
```
|
||||
|
||||
Flow:
|
||||
1. Checks custom auth first
|
||||
2. If custom auth fails, checks litellm api key auth
|
||||
3. If both fail, returns 401
|
||||
|
||||
|
||||
3. Test it!
|
||||
|
||||
```bash
|
||||
curl -L -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-H 'Authorization: Bearer sk-P1zJMdsqCPNN54alZd_ETw' \
|
||||
-d '{
|
||||
"model": "openai-model",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hey! My name is John"
|
||||
}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
#### Bubble up custom exceptions
|
||||
|
||||
If you want to bubble up custom exceptions, you can do so by raising a `ProxyException`.
|
||||
|
||||
```python
|
||||
"""
|
||||
Example custom auth function.
|
||||
|
||||
This will allow all keys starting with "my-custom-key" to pass through.
|
||||
"""
|
||||
|
||||
from typing import Union
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth, ProxyException
|
||||
|
||||
|
||||
async def user_api_key_auth(
|
||||
request: Request, api_key: str
|
||||
) -> Union[UserAPIKeyAuth, str]:
|
||||
try:
|
||||
if api_key.startswith("my-custom-key"):
|
||||
return "sk-P1zJMdsqCPNN54alZd_ETw"
|
||||
if api_key == "invalid-api-key":
|
||||
# raise a custom exception back to the client
|
||||
raise ProxyException(
|
||||
message="Invalid API key",
|
||||
type="invalid_request_error",
|
||||
param="api_key",
|
||||
code=401,
|
||||
)
|
||||
else:
|
||||
raise Exception("Invalid API key")
|
||||
except Exception:
|
||||
raise Exception("Invalid API key")
|
||||
|
||||
```
|
||||
|
|
@ -491,6 +491,47 @@ guardrails:
|
|||
default_on: true # run on every request
|
||||
```
|
||||
|
||||
|
||||
### ✨ Model-level Guardrails
|
||||
|
||||
:::info
|
||||
|
||||
✨ This is an Enterprise only feature [Get a free trial](https://www.litellm.ai/enterprise#trial)
|
||||
|
||||
:::
|
||||
|
||||
|
||||
This is great for cases when you have an on-prem and hosted model, and just want to run prevent sending PII to the hosted model.
|
||||
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: claude-sonnet-4
|
||||
litellm_params:
|
||||
model: anthropic/claude-sonnet-4-20250514
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
api_base: https://api.anthropic.com/v1
|
||||
guardrails: ["azure-text-moderation"]
|
||||
- model_name: openai-gpt-4o
|
||||
litellm_params:
|
||||
model: openai/gpt-4o
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "presidio-pii"
|
||||
litellm_params:
|
||||
guardrail: presidio # supported values: "aporia", "bedrock", "lakera", "presidio"
|
||||
mode: "pre_call"
|
||||
presidio_language: "en" # optional: set default language for PII analysis
|
||||
pii_entities_config:
|
||||
PERSON: "BLOCK" # Will mask credit card numbers
|
||||
- guardrail_name: azure-text-moderation
|
||||
litellm_params:
|
||||
guardrail: azure/text_moderations
|
||||
mode: "post_call"
|
||||
api_key: os.environ/AZURE_GUARDRAIL_API_KEY
|
||||
api_base: os.environ/AZURE_GUARDRAIL_API_BASE
|
||||
```
|
||||
|
||||
### ✨ Disable team from turning on/off guardrails
|
||||
|
||||
:::info
|
||||
|
|
|
|||
|
|
@ -1,6 +1,15 @@
|
|||
# Health Checks
|
||||
Use this to health check all LLMs defined in your config.yaml
|
||||
|
||||
## When to Use Each Endpoint
|
||||
|
||||
| Endpoint | Use Case | Purpose |
|
||||
|----------|----------|---------|
|
||||
| `/health/liveliness` | **Container liveness probes** | Basic alive check - use for container restart decisions |
|
||||
| `/health/readiness` | **Load balancer health checks** | Ready to accept traffic - includes DB connection status |
|
||||
| `/health` | **Model health monitoring** | Comprehensive LLM model health - makes actual API calls |
|
||||
| `/health/services` | **Service debugging** | Check specific integrations (datadog, langfuse, etc.) |
|
||||
|
||||
## Summary
|
||||
|
||||
The proxy exposes:
|
||||
|
|
@ -219,7 +228,7 @@ Here's how to use it:
|
|||
```
|
||||
general_settings:
|
||||
background_health_checks: True # enable background health checks
|
||||
health_check_interval: 300 # frequency of background health checks
|
||||
health_check_interval: 300 # frequency of background health checks
|
||||
```
|
||||
|
||||
2. Start server
|
||||
|
|
@ -229,7 +238,24 @@ $ litellm /path/to/config.yaml
|
|||
|
||||
3. Query health endpoint:
|
||||
```
|
||||
curl --location 'http://0.0.0.0:4000/health'
|
||||
curl --location 'http://0.0.0.0:4000/health'
|
||||
```
|
||||
|
||||
### Disable Background Health Checks For Specific Models
|
||||
|
||||
Use this if you want to disable background health checks for specific models.
|
||||
|
||||
If `background_health_checks` is enabled you can skip individual models by
|
||||
setting `disable_background_health_check: true` in the model's `model_info`.
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: openai/gpt-4o
|
||||
litellm_params:
|
||||
model: openai/gpt-4o
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
model_info:
|
||||
disable_background_health_check: true
|
||||
```
|
||||
|
||||
### Hide details
|
||||
|
|
|
|||
|
|
@ -1539,6 +1539,9 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \
|
|||
|
||||
## [Datadog](../observability/datadog)
|
||||
|
||||
👉 Go here for using [Datadog LLM Observability](../observability/datadog) with LiteLLM Proxy
|
||||
|
||||
|
||||
## Lunary
|
||||
#### Step1: Install dependencies and set your environment variables
|
||||
Install the dependencies
|
||||
|
|
@ -1590,54 +1593,7 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
|||
|
||||
## MLflow
|
||||
|
||||
|
||||
#### Step1: Install dependencies
|
||||
Install the dependencies.
|
||||
|
||||
```shell
|
||||
pip install litellm mlflow
|
||||
```
|
||||
|
||||
#### Step 2: Create a `config.yaml` with `mlflow` callback
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: "*"
|
||||
litellm_params:
|
||||
model: "*"
|
||||
litellm_settings:
|
||||
success_callback: ["mlflow"]
|
||||
failure_callback: ["mlflow"]
|
||||
```
|
||||
|
||||
#### Step 3: Start the LiteLLM proxy
|
||||
```shell
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
#### Step 4: Make a request
|
||||
|
||||
```shell
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What is the capital of France?"
|
||||
}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
#### Step 5: Review traces
|
||||
|
||||
Run the following command to start MLflow UI and review recorded traces.
|
||||
|
||||
```shell
|
||||
mlflow ui
|
||||
```
|
||||
👉 Follow the tutorial [here](../observability/mlflow) to get started with mlflow on LiteLLM Proxy Server
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -271,16 +271,7 @@ Or [watch on Loom](https://www.loom.com/share/b08be303331246b88fdc053940d03281?s
|
|||
## Extras
|
||||
### Expected Performance in Production
|
||||
|
||||
1 LiteLLM Uvicorn Worker on Kubernetes
|
||||
|
||||
| Description | Value |
|
||||
|--------------|-------|
|
||||
| Avg latency | `50ms` |
|
||||
| Median latency | `51ms` |
|
||||
| `/chat/completions` Requests/second | `100` |
|
||||
| `/chat/completions` Requests/minute | `6000` |
|
||||
| `/chat/completions` Requests/hour | `360K` |
|
||||
|
||||
See benchmarks [here](../benchmarks#performance-metrics)
|
||||
|
||||
### Verifying Debugging logs are off
|
||||
|
||||
|
|
|
|||
|
|
@ -130,28 +130,57 @@ general_settings:
|
|||
|
||||
Set the field in the jwt token, which corresponds to a litellm user / team / org.
|
||||
|
||||
**Note:** All JWT fields support dot notation to access nested claims (e.g., `"user.sub"`, `"resource_access.client.roles"`).
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
enable_jwt_auth: True
|
||||
litellm_jwtauth:
|
||||
admin_jwt_scope: "litellm-proxy-admin"
|
||||
team_id_jwt_field: "client_id" # 👈 CAN BE ANY FIELD
|
||||
user_id_jwt_field: "sub" # 👈 CAN BE ANY FIELD
|
||||
org_id_jwt_field: "org_id" # 👈 CAN BE ANY FIELD
|
||||
end_user_id_jwt_field: "customer_id" # 👈 CAN BE ANY FIELD
|
||||
team_id_jwt_field: "client_id" # 👈 CAN BE ANY FIELD (supports dot notation for nested claims)
|
||||
user_id_jwt_field: "sub" # 👈 CAN BE ANY FIELD (supports dot notation for nested claims)
|
||||
org_id_jwt_field: "org_id" # 👈 CAN BE ANY FIELD (supports dot notation for nested claims)
|
||||
end_user_id_jwt_field: "customer_id" # 👈 CAN BE ANY FIELD (supports dot notation for nested claims)
|
||||
```
|
||||
|
||||
Expected JWT:
|
||||
Expected JWT (flat structure):
|
||||
|
||||
```
|
||||
```json
|
||||
{
|
||||
"client_id": "my-unique-team",
|
||||
"sub": "my-unique-user",
|
||||
"org_id": "my-unique-org",
|
||||
"org_id": "my-unique-org"
|
||||
}
|
||||
```
|
||||
|
||||
**Or with nested structure using dot notation:**
|
||||
|
||||
```json
|
||||
{
|
||||
"user": {
|
||||
"sub": "my-unique-user",
|
||||
"email": "user@example.com"
|
||||
},
|
||||
"tenant": {
|
||||
"team_id": "my-unique-team"
|
||||
},
|
||||
"organization": {
|
||||
"id": "my-unique-org"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Configuration for nested example:**
|
||||
|
||||
```yaml
|
||||
litellm_jwtauth:
|
||||
user_id_jwt_field: "user.sub"
|
||||
user_email_jwt_field: "user.email"
|
||||
team_id_jwt_field: "tenant.team_id"
|
||||
org_id_jwt_field: "organization.id"
|
||||
```
|
||||
|
||||
Now litellm will automatically update the spend for the user/team/org in the db for each call.
|
||||
|
||||
### JWT Scopes
|
||||
|
|
@ -407,9 +436,15 @@ environment_variables:
|
|||
JWT_AUDIENCE: "api://LiteLLM_Proxy" # ensures audience is validated
|
||||
```
|
||||
|
||||
- `object_id_jwt_field`: The field in the JWT token that contains the object id. This id can be either a user id or a team id. Use this instead of `user_id_jwt_field` and `team_id_jwt_field`. If the same field could be both.
|
||||
- `object_id_jwt_field`: The field in the JWT token that contains the object id. This id can be either a user id or a team id. Use this instead of `user_id_jwt_field` and `team_id_jwt_field`. If the same field could be both. **Supports dot notation** for nested claims (e.g., `"profile.object_id"`).
|
||||
|
||||
- `roles_jwt_field`: The field in the JWT token that contains the roles. This field is a list of roles that the user has. To index into a nested field, use dot notation - eg. `resource_access.litellm-test-client-id.roles`.
|
||||
- `roles_jwt_field`: The field in the JWT token that contains the roles. This field is a list of roles that the user has. **Supports dot notation** for nested fields - e.g., `resource_access.litellm-test-client-id.roles`.
|
||||
|
||||
**Additional JWT Field Configuration Options:**
|
||||
|
||||
- `team_ids_jwt_field`: Field containing team IDs (as a list). **Supports dot notation** (e.g., `"groups"`, `"teams.ids"`).
|
||||
- `user_email_jwt_field`: Field containing user email. **Supports dot notation** (e.g., `"email"`, `"user.email"`).
|
||||
- `end_user_id_jwt_field`: Field containing end-user ID for cost tracking. **Supports dot notation** (e.g., `"customer_id"`, `"customer.id"`).
|
||||
|
||||
- `role_mappings`: A list of role mappings. Map the received role in the JWT token to an internal role on LiteLLM.
|
||||
|
||||
|
|
|
|||
91
docs/my-website/docs/tutorials/cost_tracking_coding.md
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
import Image from '@theme/IdealImage';
|
||||
|
||||
# Track Usage for Coding Tools
|
||||
|
||||
Track usage and costs for AI-powered coding tools like Claude Code, Roo Code, Gemini CLI, and OpenAI Codex through LiteLLM.
|
||||
|
||||
Monitor requests, costs, and user engagement metrics for each coding tool using User-Agent headers.
|
||||
|
||||
<Image
|
||||
img={require('../../img/agent_1.png')}
|
||||
style={{width: '100%', display: 'block', margin: '2rem auto'}}
|
||||
/>
|
||||
|
||||
|
||||
## Who This Is For
|
||||
|
||||
Central AI Platform teams providing developers access to coding tools through LiteLLM. Monitor tool engagement and track individual user usage patterns.
|
||||
|
||||
## What You Can Track
|
||||
|
||||
### Summary Metrics
|
||||
- Cost per coding tool
|
||||
- Successful requests and token usage per tool
|
||||
|
||||
### User Engagement Metrics
|
||||
- Daily, weekly, and monthly active users for each User-Agent
|
||||
|
||||
## Quick Start
|
||||
|
||||
### 1. Connect Your Coding Tool to LiteLLM
|
||||
|
||||
Configure your coding tool to send requests through the LiteLLM proxy with appropriate User-Agent headers.
|
||||
|
||||
**Setup guides:**
|
||||
- [Use LiteLLM with Claude Code](../../docs/tutorials/claude_responses_api)
|
||||
- [Use LiteLLM with Gemini CLI](../../docs/tutorials/litellm_gemini_cli)
|
||||
- [Use LiteLLM with OpenAI Codex](../../docs/tutorials/openai_codex)
|
||||
|
||||
### 2. Send Requests with User-Agent Headers
|
||||
|
||||
Ensure your coding tool includes identifying User-Agent headers in API requests.
|
||||
|
||||
### 3. Verify Tracking in LiteLLM Logs
|
||||
|
||||
Confirm LiteLLM is properly tracking requests by checking logs for the expected User-Agent values.
|
||||
|
||||
<Image
|
||||
img={require('../../img/agent_2.png')}
|
||||
style={{width: '100%', display: 'block', margin: '2rem auto'}}
|
||||
/>
|
||||
|
||||
### 4. View Usage Dashboard
|
||||
|
||||
Access the LiteLLM dashboard to view aggregated usage metrics and user engagement data.
|
||||
|
||||
#### Summary Metrics
|
||||
|
||||
View total cost and successful requests for each coding tool.
|
||||
|
||||
<Image
|
||||
img={require('../../img/agent_3.png')}
|
||||
style={{width: '100%', display: 'block', margin: '2rem auto'}}
|
||||
/>
|
||||
|
||||
#### Daily, Weekly, and Monthly Active Users
|
||||
|
||||
View active user metrics for each coding tool.
|
||||
|
||||
<Image
|
||||
img={require('../../img/agent_4.png')}
|
||||
style={{width: '80%', display: 'block', margin: '2rem auto'}}
|
||||
/>
|
||||
|
||||
## How LiteLLM Identifies Coding Tools
|
||||
|
||||
LiteLLM tracks coding tools by monitoring the `User-Agent` header in incoming API requests (`/chat/completions`, `/responses`, etc.). Each unique User-Agent is tracked separately for usage analytics.
|
||||
|
||||
### Example Request
|
||||
|
||||
Example using `claude-cli` as the User-Agent:
|
||||
|
||||
```shell
|
||||
curl -X POST \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "User-Agent: claude-cli/1.0" \
|
||||
-d '{"model": "claude-3-5-sonnet-latest", "messages": [{"role": "user", "content": "Hello, how are you?"}]}' \
|
||||
http://localhost:4000/chat/completions
|
||||
```
|
||||
BIN
docs/my-website/img/agent_1.png
Normal file
|
After Width: | Height: | Size: 470 KiB |
BIN
docs/my-website/img/agent_2.png
Normal file
|
After Width: | Height: | Size: 218 KiB |
BIN
docs/my-website/img/agent_3.png
Normal file
|
After Width: | Height: | Size: 211 KiB |
BIN
docs/my-website/img/agent_4.png
Normal file
|
After Width: | Height: | Size: 130 KiB |
BIN
docs/my-website/img/dd_llm_obs.png
Normal file
|
After Width: | Height: | Size: 232 KiB |
BIN
docs/my-website/img/release_notes/auto_router.png
Normal file
|
After Width: | Height: | Size: 468 KiB |
BIN
docs/my-website/img/release_notes/mcp_header_propogation.png
Normal file
|
After Width: | Height: | Size: 867 KiB |
BIN
docs/my-website/img/release_notes/model_level_guardrails.jpg
Normal file
|
After Width: | Height: | Size: 310 KiB |
291
docs/my-website/release_notes/v1.74.15-stable/index.md
Normal file
|
|
@ -0,0 +1,291 @@
|
|||
---
|
||||
title: "[Pre-Release] v1.74.15-stable"
|
||||
slug: "v1-74-15"
|
||||
date: 2025-08-02T10:00:00
|
||||
authors:
|
||||
- name: Krrish Dholakia
|
||||
title: CEO, LiteLLM
|
||||
url: https://www.linkedin.com/in/krish-d/
|
||||
image_url: https://pbs.twimg.com/profile_images/1298587542745358340/DZv3Oj-h_400x400.jpg
|
||||
- name: Ishaan Jaffer
|
||||
title: CTO, LiteLLM
|
||||
url: https://www.linkedin.com/in/reffajnaahsi/
|
||||
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
|
||||
|
||||
hide_table_of_contents: false
|
||||
---
|
||||
|
||||
import Image from '@theme/IdealImage';
|
||||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
## Deploy this version
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="docker" label="Docker">
|
||||
|
||||
``` showLineNumbers title="docker run litellm"
|
||||
docker run \
|
||||
-e STORE_MODEL_IN_DB=True \
|
||||
-p 4000:4000 \
|
||||
ghcr.io/berriai/litellm:v1.74.15.rc.1
|
||||
```
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="pip" label="Pip">
|
||||
|
||||
``` showLineNumbers title="pip install litellm"
|
||||
pip install litellm==1.74.15.post1
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
---
|
||||
|
||||
## Key Highlights
|
||||
|
||||
- **User Agent Activity Tracking** - Track how much usage each coding tool gets.
|
||||
- **Prompt Management** - Use Git-Ops style prompt management with prompt templates.
|
||||
- **MCP Gateway: Guardrails** - Support for using Guardrails with MCP servers.
|
||||
- **Google AI Studio Imagen4** - Support for using Imagen4 models on Google AI Studio.
|
||||
|
||||
---
|
||||
|
||||
## User Agent Activity Tracking
|
||||
|
||||
<Image
|
||||
img={require('../../img/agent_1.png')}
|
||||
style={{width: '100%', display: 'block', margin: '2rem auto'}}
|
||||
/>
|
||||
|
||||
<br/>
|
||||
|
||||
This release brings support for tracking usage and costs for AI-powered coding tools like Claude Code, Roo Code, Gemini CLI through LiteLLM. You can now track LLM cost, total tokens used, and DAU/WAU/MAU for each coding tool.
|
||||
|
||||
This is great to central AI Platform teams looking to track how they are helping developer productivity.
|
||||
|
||||
[Read More](https://docs.litellm.ai/docs/tutorials/cost_tracking_coding)
|
||||
|
||||
---
|
||||
|
||||
## Prompt Management
|
||||
|
||||
<br/>
|
||||
|
||||
|
||||
|
||||
[Read More](../../docs/proxy/prompt_management)
|
||||
|
||||
---
|
||||
|
||||
## New Models / Updated Models
|
||||
|
||||
#### New Model Support
|
||||
|
||||
| Provider | Model | Context Window | Input ($/1M tokens) | Output ($/1M tokens) | Cost per Image |
|
||||
| ----------- | -------------------------------------- | -------------- | ------------------- | -------------------- | -------------- |
|
||||
| OpenRouter | `openrouter/x-ai/grok-4` | 256k | $3 | $15 | N/A |
|
||||
| Google AI Studio | `gemini/imagen-4.0-generate-preview-06-06` | N/A | N/A | N/A | $0.04 |
|
||||
| Google AI Studio | `gemini/imagen-4.0-ultra-generate-preview-06-06` | N/A | N/A | N/A | $0.06 |
|
||||
| Google AI Studio | `gemini/imagen-4.0-fast-generate-preview-06-06` | N/A | N/A | N/A | $0.02 |
|
||||
| Google AI Studio | `gemini/imagen-3.0-generate-002` | N/A | N/A | N/A | $0.04 |
|
||||
| Google AI Studio | `gemini/imagen-3.0-generate-001` | N/A | N/A | N/A | $0.04 |
|
||||
| Google AI Studio | `gemini/imagen-3.0-fast-generate-001` | N/A | N/A | N/A | $0.02 |
|
||||
|
||||
#### Features
|
||||
|
||||
- **[Google AI Studio](../../docs/providers/gemini)**
|
||||
- Added Google AI Studio Imagen4 model family support - [PR #13065](https://github.com/BerriAI/litellm/pull/13065), [Get Started](../../docs/providers/google_ai_studio/image_gen)
|
||||
- **[Azure OpenAI](../../docs/providers/azure/azure)**
|
||||
- Azure `api_version="preview"` support - [PR #13072](https://github.com/BerriAI/litellm/pull/13072), [Get Started](../../docs/providers/azure/azure#setting-api-version)
|
||||
- Password protected certificate files support - [PR #12995](https://github.com/BerriAI/litellm/pull/12995), [Get Started](../../docs/providers/azure/azure#authentication)
|
||||
- **[AWS Bedrock](../../docs/providers/bedrock)**
|
||||
- Cost tracking via Anthropic `/v1/messages` - [PR #13072](https://github.com/BerriAI/litellm/pull/13072)
|
||||
- Computer use support - [PR #13150](https://github.com/BerriAI/litellm/pull/13150)
|
||||
- **[OpenRouter](../../docs/providers/openrouter)**
|
||||
- Added Grok4 model support - [PR #13018](https://github.com/BerriAI/litellm/pull/13018)
|
||||
- **[Anthropic](../../docs/providers/anthropic)**
|
||||
- Auto Cache Control Injection - Improved cache_control_injection_points with negative index support - [PR #13187](https://github.com/BerriAI/litellm/pull/13187), [Get Started](../../docs/tutorials/prompt_caching)
|
||||
- Working mid-stream fallbacks with token usage tracking - [PR #13149](https://github.com/BerriAI/litellm/pull/13149), [PR #13170](https://github.com/BerriAI/litellm/pull/13170)
|
||||
- **[Perplexity](../../docs/providers/perplexity)**
|
||||
- Citation annotations support - [PR #13225](https://github.com/BerriAI/litellm/pull/13225)
|
||||
|
||||
#### Bugs
|
||||
|
||||
- **[Gemini](../../docs/providers/gemini)**
|
||||
- Fix merge_reasoning_content_in_choices parameter issue - [PR #13066](https://github.com/BerriAI/litellm/pull/13066), [Get Started](../../docs/tutorials/openweb_ui#render-thinking-content-on-open-webui)
|
||||
- Added support for using `GOOGLE_API_KEY` environment variable for Google AI Studio - [PR #12507](https://github.com/BerriAI/litellm/pull/12507)
|
||||
- **[vLLM/OpenAI-like](../../docs/providers/vllm)**
|
||||
- Fix missing extra_headers support for embeddings - [PR #13198](https://github.com/BerriAI/litellm/pull/13198)
|
||||
|
||||
---
|
||||
|
||||
## LLM API Endpoints
|
||||
|
||||
#### Bugs
|
||||
|
||||
- **[/generateContent](../../docs/generateContent)**
|
||||
- Support for query_params in generateContent routes for API Key setting - [PR #13100](https://github.com/BerriAI/litellm/pull/13100)
|
||||
- Ensure "x-goog-api-key" is used for auth to google ai studio when using /generateContent on LiteLLM - [PR #13098](https://github.com/BerriAI/litellm/pull/13098)
|
||||
- Ensure tool calling works as expected on generateContent - [PR #13189](https://github.com/BerriAI/litellm/pull/13189)
|
||||
- **[/vertex_ai (Passthrough)](../../docs/pass_through/vertex_ai)**
|
||||
- Ensure multimodal embedding responses are logged properly - [PR #13050](https://github.com/BerriAI/litellm/pull/13050)
|
||||
|
||||
---
|
||||
|
||||
## [MCP Gateway](../../docs/mcp)
|
||||
|
||||
#### Features
|
||||
|
||||
- **Health Check Improvements**
|
||||
- Add health check endpoints for MCP servers - [PR #13106](https://github.com/BerriAI/litellm/pull/13106)
|
||||
- **Guardrails Integration**
|
||||
- Add pre and during call hooks initialization - [PR #13067](https://github.com/BerriAI/litellm/pull/13067)
|
||||
- Move pre and during hooks to ProxyLogging - [PR #13109](https://github.com/BerriAI/litellm/pull/13109)
|
||||
- MCP pre and during guardrails implementation - [PR #13188](https://github.com/BerriAI/litellm/pull/13188)
|
||||
- **Protocol & Header Support**
|
||||
- Add protocol headers support - [PR #13062](https://github.com/BerriAI/litellm/pull/13062)
|
||||
- **URL & Namespacing**
|
||||
- Improve MCP server URL validation for internal/Kubernetes URLs - [PR #13099](https://github.com/BerriAI/litellm/pull/13099)
|
||||
|
||||
|
||||
#### Bugs
|
||||
|
||||
- **UI**
|
||||
- Fix scrolling issue with MCP tools - [PR #13015](https://github.com/BerriAI/litellm/pull/13015)
|
||||
- Fix MCP client list failure - [PR #13114](https://github.com/BerriAI/litellm/pull/13114)
|
||||
|
||||
|
||||
[Read More](../../docs/mcp)
|
||||
|
||||
|
||||
---
|
||||
|
||||
## Management Endpoints / UI
|
||||
|
||||
#### Features
|
||||
|
||||
- **Usage Analytics**
|
||||
- New tab for user agent activity tracking - [PR #13146](https://github.com/BerriAI/litellm/pull/13146)
|
||||
- Daily usage per user analytics - [PR #13147](https://github.com/BerriAI/litellm/pull/13147)
|
||||
- Default usage chart date range set to last 7 days - [PR #12917](https://github.com/BerriAI/litellm/pull/12917)
|
||||
- New advanced date range picker component - [PR #13141](https://github.com/BerriAI/litellm/pull/13141), [PR #13221](https://github.com/BerriAI/litellm/pull/13221)
|
||||
- Show loader on usage cost charts after date selection - [PR #13113](https://github.com/BerriAI/litellm/pull/13113)
|
||||
- **Models**
|
||||
- Added Voyage, Jinai, Deepinfra and VolcEngine providers on UI - [PR #13131](https://github.com/BerriAI/litellm/pull/13131)
|
||||
- Added Sagemaker on UI - [PR #13117](https://github.com/BerriAI/litellm/pull/13117)
|
||||
- Preserve model order in `/v1/models` and `/model_group/info` endpoints - [PR #13178](https://github.com/BerriAI/litellm/pull/13178)
|
||||
|
||||
- **Key Management**
|
||||
- Properly parse JSON options for key generation in UI - [PR #12989](https://github.com/BerriAI/litellm/pull/12989)
|
||||
- **Authentication**
|
||||
- **JWT Fields**
|
||||
- Add dot notation support for all JWT fields - [PR #13013](https://github.com/BerriAI/litellm/pull/13013)
|
||||
|
||||
#### Bugs
|
||||
|
||||
- **Permissions**
|
||||
- Fix object permission for organizations - [PR #13142](https://github.com/BerriAI/litellm/pull/13142)
|
||||
- Fix list team v2 security check - [PR #13094](https://github.com/BerriAI/litellm/pull/13094)
|
||||
- **Models**
|
||||
- Fix model reload on model update - [PR #13216](https://github.com/BerriAI/litellm/pull/13216)
|
||||
- **Router Settings**
|
||||
- Fix displaying models for fallbacks in UI - [PR #13191](https://github.com/BerriAI/litellm/pull/13191)
|
||||
- Fix wildcard model name handling with custom values - [PR #13116](https://github.com/BerriAI/litellm/pull/13116)
|
||||
- Fix fallback delete functionality - [PR #12606](https://github.com/BerriAI/litellm/pull/12606)
|
||||
|
||||
---
|
||||
|
||||
## Logging / Guardrail Integrations
|
||||
|
||||
#### Features
|
||||
|
||||
- **[MLFlow](../../docs/proxy/logging#mlflow)**
|
||||
- Allow adding tags for MLFlow logging requests - [PR #13108](https://github.com/BerriAI/litellm/pull/13108)
|
||||
- **[Langfuse OTEL](../../docs/proxy/logging#langfuse)**
|
||||
- Add comprehensive metadata support to Langfuse OpenTelemetry integration - [PR #12956](https://github.com/BerriAI/litellm/pull/12956)
|
||||
- **[Datadog LLM Observability](../../docs/proxy/logging#datadog)**
|
||||
- Allow redacting message/response content for specific logging integrations - [PR #13158](https://github.com/BerriAI/litellm/pull/13158)
|
||||
|
||||
#### Bugs
|
||||
|
||||
- **API Key Logging**
|
||||
- Fix API Key being logged inappropriately - [PR #12978](https://github.com/BerriAI/litellm/pull/12978)
|
||||
- **MCP Spend Tracking**
|
||||
- Set default value for MCP namespace tool name in spend table - [PR #12894](https://github.com/BerriAI/litellm/pull/12894)
|
||||
|
||||
---
|
||||
|
||||
## Performance / Loadbalancing / Reliability improvements
|
||||
|
||||
#### Features
|
||||
|
||||
- **Background Health Checks**
|
||||
- Allow disabling background health checks for specific deployments - [PR #13186](https://github.com/BerriAI/litellm/pull/13186)
|
||||
- **Database Connection Management**
|
||||
- Ensure stale Prisma clients disconnect DB connections properly - [PR #13140](https://github.com/BerriAI/litellm/pull/13140)
|
||||
- **Jitter Improvements**
|
||||
- Fix jitter calculation (should be added not multiplied) - [PR #12901](https://github.com/BerriAI/litellm/pull/12901)
|
||||
|
||||
#### Bugs
|
||||
|
||||
- **Anthropic Streaming**
|
||||
- Always use choice index=0 for Anthropic streaming responses - [PR #12666](https://github.com/BerriAI/litellm/pull/12666)
|
||||
- **Custom Auth**
|
||||
- Bubble up custom exceptions properly - [PR #13093](https://github.com/BerriAI/litellm/pull/13093)
|
||||
- **OTEL with Managed Files**
|
||||
- Fix using managed files with OTEL integration - [PR #13171](https://github.com/BerriAI/litellm/pull/13171)
|
||||
|
||||
---
|
||||
|
||||
## General Proxy Improvements
|
||||
|
||||
#### Features
|
||||
|
||||
- **Database Migration**
|
||||
- Move to use_prisma_migrate by default - [PR #13117](https://github.com/BerriAI/litellm/pull/13117)
|
||||
- Resolve team-only models on auth checks - [PR #13117](https://github.com/BerriAI/litellm/pull/13117)
|
||||
- **Infrastructure**
|
||||
- Loosened MCP Python version restrictions - [PR #13102](https://github.com/BerriAI/litellm/pull/13102)
|
||||
- Migrate build_and_test to CI/CD Postgres DB - [PR #13166](https://github.com/BerriAI/litellm/pull/13166)
|
||||
- **Helm Charts**
|
||||
- Allow Helm hooks for migration jobs - [PR #13174](https://github.com/BerriAI/litellm/pull/13174)
|
||||
- Fix Helm migration job schema updates - [PR #12809](https://github.com/BerriAI/litellm/pull/12809)
|
||||
|
||||
#### Bugs
|
||||
|
||||
- **Docker**
|
||||
- Remove obsolete `version` attribute in docker-compose - [PR #13172](https://github.com/BerriAI/litellm/pull/13172)
|
||||
- Add openssl in runtime stage for non-root Dockerfile - [PR #13168](https://github.com/BerriAI/litellm/pull/13168)
|
||||
- **Database Configuration**
|
||||
- Fix DB config through environment variables - [PR #13111](https://github.com/BerriAI/litellm/pull/13111)
|
||||
- **Logging**
|
||||
- Suppress httpx logging - [PR #13217](https://github.com/BerriAI/litellm/pull/13217)
|
||||
- **Token Counting**
|
||||
- Ignore unsupported keys like prefix in token counter - [PR #11954](https://github.com/BerriAI/litellm/pull/11954)
|
||||
---
|
||||
|
||||
## New Contributors
|
||||
* @5731la made their first contribution in https://github.com/BerriAI/litellm/pull/12989
|
||||
* @restato made their first contribution in https://github.com/BerriAI/litellm/pull/12980
|
||||
* @strickvl made their first contribution in https://github.com/BerriAI/litellm/pull/12956
|
||||
* @Ne0-1 made their first contribution in https://github.com/BerriAI/litellm/pull/12995
|
||||
* @maxrabin made their first contribution in https://github.com/BerriAI/litellm/pull/13079
|
||||
* @lvuna made their first contribution in https://github.com/BerriAI/litellm/pull/12894
|
||||
* @Maximgitman made their first contribution in https://github.com/BerriAI/litellm/pull/12666
|
||||
* @pathikrit made their first contribution in https://github.com/BerriAI/litellm/pull/12901
|
||||
* @huetterma made their first contribution in https://github.com/BerriAI/litellm/pull/12809
|
||||
* @betterthanbreakfast made their first contribution in https://github.com/BerriAI/litellm/pull/13029
|
||||
* @phosae made their first contribution in https://github.com/BerriAI/litellm/pull/12606
|
||||
* @sahusiddharth made their first contribution in https://github.com/BerriAI/litellm/pull/12507
|
||||
* @Amit-kr26 made their first contribution in https://github.com/BerriAI/litellm/pull/11954
|
||||
* @kowyo made their first contribution in https://github.com/BerriAI/litellm/pull/13172
|
||||
* @AnandKhinvasara made their first contribution in https://github.com/BerriAI/litellm/pull/13187
|
||||
* @unique-jakub made their first contribution in https://github.com/BerriAI/litellm/pull/13174
|
||||
* @tyumentsev4 made their first contribution in https://github.com/BerriAI/litellm/pull/13134
|
||||
* @aayush-malviya-acquia made their first contribution in https://github.com/BerriAI/litellm/pull/12978
|
||||
* @kankute-sameer made their first contribution in https://github.com/BerriAI/litellm/pull/13225
|
||||
* @AlexanderYastrebov made their first contribution in https://github.com/BerriAI/litellm/pull/13178
|
||||
|
||||
## **[Full Changelog](https://github.com/BerriAI/litellm/compare/v1.74.9-stable...v1.74.15.rc)**
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
---
|
||||
title: "[PRE-RELEASE] v1.74.9-stable"
|
||||
title: "v1.74.9-stable - Auto-Router"
|
||||
slug: "v1-74-9"
|
||||
date: 2025-07-27T10:00:00
|
||||
authors:
|
||||
|
|
@ -21,21 +21,104 @@ import TabItem from '@theme/TabItem';
|
|||
|
||||
## Deploy this version
|
||||
|
||||
:::info
|
||||
<Tabs>
|
||||
<TabItem value="docker" label="Docker">
|
||||
|
||||
This release is not live yet.
|
||||
``` showLineNumbers title="docker run litellm"
|
||||
docker run \
|
||||
-e STORE_MODEL_IN_DB=True \
|
||||
-p 4000:4000 \
|
||||
ghcr.io/berriai/litellm:v1.74.9-stable.patch.1
|
||||
```
|
||||
</TabItem>
|
||||
|
||||
:::
|
||||
<TabItem value="pip" label="Pip">
|
||||
|
||||
``` showLineNumbers title="pip install litellm"
|
||||
pip install litellm==1.74.9.post2
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
---
|
||||
|
||||
## Key Highlights
|
||||
|
||||
- **Auto-Router** - Automatically route requests to specific models based on request content.
|
||||
- **Model-level Guardrails** - Only run guardrails when specific models are used.
|
||||
- **MCP Header Propagation** - Propagate headers from client to backend MCP.
|
||||
- **New LLM Providers** - Added Bedrock inpainting support and Recraft API image generation / image edits support.
|
||||
|
||||
---
|
||||
|
||||
## Auto-Router
|
||||
|
||||
<Image img={require('../../img/release_notes/auto_router.png')} />
|
||||
|
||||
<br/>
|
||||
|
||||
This release introduces auto-routing to models based on request content. This means **Proxy Admins** can define a set of keywords that always routes to specific models when **users** opt in to using the auto-router.
|
||||
|
||||
This is great for internal use cases where you don't want **users** to think about which model to use - for example, use Claude models for coding vs GPT models for generating ad copy.
|
||||
|
||||
|
||||
[Read More](../../docs/proxy/auto_routing)
|
||||
|
||||
---
|
||||
|
||||
## Model-level Guardrails
|
||||
|
||||
<Image img={require('../../img/release_notes/model_level_guardrails.jpg')} />
|
||||
|
||||
<br/>
|
||||
|
||||
This release brings model-level guardrails support to your config.yaml + UI. This is great for cases when you have an on-prem and hosted model, and just want to run prevent sending PII to the hosted model.
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: claude-sonnet-4
|
||||
litellm_params:
|
||||
model: anthropic/claude-sonnet-4-20250514
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
api_base: https://api.anthropic.com/v1
|
||||
guardrails: ["azure-text-moderation"] # 👈 KEY CHANGE
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: azure-text-moderation
|
||||
litellm_params:
|
||||
guardrail: azure/text_moderations
|
||||
mode: "post_call"
|
||||
api_key: os.environ/AZURE_GUARDRAIL_API_KEY
|
||||
api_base: os.environ/AZURE_GUARDRAIL_API_BASE
|
||||
```
|
||||
|
||||
|
||||
[Read More](../../docs/proxy/guardrails/quick_start#model-level-guardrails)
|
||||
|
||||
---
|
||||
## MCP Header Propagation
|
||||
|
||||
<Image img={require('../../img/release_notes/mcp_header_propogation.png')} />
|
||||
|
||||
<br/>
|
||||
|
||||
v1.74.9-stable allows you to propagate MCP server specific authentication headers via LiteLLM
|
||||
|
||||
- Allowing users to specify which `header_name` is to be propagated to which `mcp_server` via headers
|
||||
- Allows adding of different deployments of same MCP server type to use different authentication headers
|
||||
|
||||
|
||||
[Read More](https://docs.litellm.ai/docs/mcp#new-server-specific-auth-headers-recommended)
|
||||
|
||||
---
|
||||
## New Models / Updated Models
|
||||
|
||||
#### Pricing / Context Window Updates
|
||||
|
||||
| Provider | Model | Context Window | Input ($/1M tokens) | Output ($/1M tokens) |
|
||||
| ----------- | -------------------------------------- | -------------- | ------------------- | -------------------- |
|
||||
| Fireworks AI | `fireworks/models/kimi-k2-instruct | 131k | $0.6 | $2.5 |
|
||||
| Fireworks AI | `fireworks/models/kimi-k2-instruct` | 131k | $0.6 | $2.5 |
|
||||
| OpenRouter | `openrouter/qwen/qwen-vl-plus` | 8192 | $0.21 | $0.63 |
|
||||
| OpenRouter | `openrouter/qwen/qwen3-coder` | 8192 | $1 | $5 |
|
||||
| OpenRouter | `openrouter/bytedance/ui-tars-1.5-7b` | 128k | $0.10 | $0.20 |
|
||||
|
|
|
|||
|
|
@ -78,6 +78,7 @@ const sidebars = {
|
|||
"tutorials/litellm_qwen_code_cli",
|
||||
"tutorials/github_copilot_integration",
|
||||
"tutorials/claude_responses_api",
|
||||
"tutorials/cost_tracking_coding",
|
||||
]
|
||||
},
|
||||
|
||||
|
|
@ -485,6 +486,7 @@ const sidebars = {
|
|||
"completion/vision",
|
||||
"completion/json_mode",
|
||||
"reasoning_content",
|
||||
"completion/computer_use",
|
||||
"completion/prompt_caching",
|
||||
"completion/predict_outputs",
|
||||
"completion/knowledgebase",
|
||||
|
|
|
|||
|
|
@ -173,6 +173,7 @@ class AporiaGuardrail(CustomGuardrail):
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
):
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
|
|
|
|||
|
|
@ -95,6 +95,7 @@ class _ENTERPRISE_GoogleTextModeration(CustomLogger):
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -42,6 +42,7 @@ class _ENTERPRISE_OpenAI_Moderation(CustomLogger):
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
):
|
||||
text = ""
|
||||
|
|
|
|||
|
|
@ -105,6 +105,7 @@ class _ENTERPRISE_LlamaGuard(CustomLogger):
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -127,6 +127,7 @@ class _ENTERPRISE_LLMGuard(CustomLogger):
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -147,6 +147,7 @@ class PagerDutyAlerting(SlackAlerting):
|
|||
"audio_transcription",
|
||||
"pass_through_endpoint",
|
||||
"rerank",
|
||||
"mcp_call",
|
||||
],
|
||||
) -> Optional[Union[Exception, str, dict]]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ from typing import Any, Optional
|
|||
from fastapi import Request
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
|
||||
|
||||
async def enterprise_custom_auth(
|
||||
|
|
@ -24,6 +24,8 @@ async def enterprise_custom_auth(
|
|||
elif custom_auth_settings["mode"] == "auto":
|
||||
try:
|
||||
return await user_custom_auth(request, api_key)
|
||||
except ProxyException as e:
|
||||
raise e
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Error in custom auth, checking litellm auth: {e}"
|
||||
|
|
|
|||
|
|
@ -290,6 +290,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
"aretrieve_fine_tuning_job",
|
||||
"alist_fine_tuning_jobs",
|
||||
"acancel_fine_tuning_job",
|
||||
"mcp_call",
|
||||
],
|
||||
) -> Union[Exception, str, Dict, None]:
|
||||
"""
|
||||
|
|
|
|||
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.14-py3-none-any.whl
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.14.tar.gz
vendored
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
-- Add health check fields to MCP server table
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN "status" TEXT DEFAULT 'unknown';
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN "last_health_check" TIMESTAMP(3);
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN "health_check_error" TEXT;
|
||||
|
|
@ -0,0 +1,15 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE "LiteLLM_PromptTable" (
|
||||
"id" TEXT NOT NULL,
|
||||
"prompt_id" TEXT NOT NULL,
|
||||
"litellm_params" JSONB NOT NULL,
|
||||
"prompt_info" JSONB,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"updated_at" TIMESTAMP(3) NOT NULL,
|
||||
|
||||
CONSTRAINT "LiteLLM_PromptTable_pkey" PRIMARY KEY ("id")
|
||||
);
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX "LiteLLM_PromptTable_prompt_id_key" ON "LiteLLM_PromptTable"("prompt_id");
|
||||
|
||||
|
|
@ -179,6 +179,10 @@ model LiteLLM_MCPServerTable {
|
|||
updated_by String?
|
||||
mcp_info Json? @default("{}")
|
||||
mcp_access_groups String[]
|
||||
// Health check status
|
||||
status String? @default("unknown")
|
||||
last_health_check DateTime?
|
||||
health_check_error String?
|
||||
// Stdio-specific fields
|
||||
command String?
|
||||
args String[] @default([])
|
||||
|
|
@ -516,6 +520,16 @@ model LiteLLM_GuardrailsTable {
|
|||
updated_at DateTime @updatedAt
|
||||
}
|
||||
|
||||
// Prompt table for storing prompt configurations
|
||||
model LiteLLM_PromptTable {
|
||||
id String @id @default(uuid())
|
||||
prompt_id String @unique
|
||||
litellm_params Json
|
||||
prompt_info Json?
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
}
|
||||
|
||||
model LiteLLM_HealthCheckTable {
|
||||
health_check_id String @id @default(uuid())
|
||||
model_name String
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[tool.poetry]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.2.12"
|
||||
version = "0.2.15"
|
||||
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.2.12"
|
||||
version = "0.2.15"
|
||||
version_files = [
|
||||
"pyproject.toml:version",
|
||||
"../requirements.txt:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -5,7 +5,18 @@ warnings.filterwarnings("ignore", message=".*conflict with protected namespace.*
|
|||
### INIT VARIABLES ####################
|
||||
import threading
|
||||
import os
|
||||
from typing import Callable, List, Optional, Dict, Union, Any, Literal, get_args
|
||||
from typing import (
|
||||
Callable,
|
||||
List,
|
||||
Optional,
|
||||
Dict,
|
||||
Union,
|
||||
Any,
|
||||
Literal,
|
||||
get_args,
|
||||
TYPE_CHECKING,
|
||||
)
|
||||
from litellm.types.integrations.datadog_llm_obs import DatadogLLMObsInitParams
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.caching.caching import Cache, DualCache, RedisCache, InMemoryCache
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
|
|
@ -60,6 +71,11 @@ from litellm.constants import (
|
|||
DEFAULT_SOFT_BUDGET,
|
||||
DEFAULT_ALLOWED_FAILS,
|
||||
)
|
||||
from litellm.integrations.dotprompt import (
|
||||
global_prompt_manager,
|
||||
global_prompt_directory,
|
||||
set_global_prompt_directory,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailItem
|
||||
from litellm.types.secret_managers.main import (
|
||||
KeyManagementSystem,
|
||||
|
|
@ -82,7 +98,6 @@ if litellm_mode == "DEV":
|
|||
|
||||
# Register async client cleanup to prevent resource leaks
|
||||
register_async_client_cleanup()
|
||||
|
||||
####################################################
|
||||
if set_verbose == True:
|
||||
_turn_on_debug()
|
||||
|
|
@ -129,6 +144,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
|
|||
"s3_v2",
|
||||
"aws_sqs",
|
||||
"vector_store_pre_call_hook",
|
||||
"dotprompt",
|
||||
]
|
||||
logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None
|
||||
_known_custom_logger_compatible_callbacks: List = list(
|
||||
|
|
@ -144,22 +160,22 @@ prometheus_initialize_budget_metrics: Optional[bool] = False
|
|||
require_auth_for_metrics_endpoint: Optional[bool] = False
|
||||
argilla_batch_size: Optional[int] = None
|
||||
datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged payload.
|
||||
gcs_pub_sub_use_v1: Optional[
|
||||
bool
|
||||
] = False # if you want to use v1 gcs pubsub logged payload
|
||||
generic_api_use_v1: Optional[
|
||||
bool
|
||||
] = False # if you want to use v1 generic api logged payload
|
||||
gcs_pub_sub_use_v1: Optional[bool] = (
|
||||
False # if you want to use v1 gcs pubsub logged payload
|
||||
)
|
||||
generic_api_use_v1: Optional[bool] = (
|
||||
False # if you want to use v1 generic api logged payload
|
||||
)
|
||||
argilla_transformation_object: Optional[Dict[str, Any]] = None
|
||||
_async_input_callback: List[
|
||||
Union[str, Callable, CustomLogger]
|
||||
] = [] # internal variable - async custom callbacks are routed here.
|
||||
_async_success_callback: List[
|
||||
Union[str, Callable, CustomLogger]
|
||||
] = [] # internal variable - async custom callbacks are routed here.
|
||||
_async_failure_callback: List[
|
||||
Union[str, Callable, CustomLogger]
|
||||
] = [] # internal variable - async custom callbacks are routed here.
|
||||
_async_input_callback: List[Union[str, Callable, CustomLogger]] = (
|
||||
[]
|
||||
) # internal variable - async custom callbacks are routed here.
|
||||
_async_success_callback: List[Union[str, Callable, CustomLogger]] = (
|
||||
[]
|
||||
) # internal variable - async custom callbacks are routed here.
|
||||
_async_failure_callback: List[Union[str, Callable, CustomLogger]] = (
|
||||
[]
|
||||
) # internal variable - async custom callbacks are routed here.
|
||||
pre_call_rules: List[Callable] = []
|
||||
post_call_rules: List[Callable] = []
|
||||
turn_off_message_logging: Optional[bool] = False
|
||||
|
|
@ -167,18 +183,18 @@ log_raw_request_response: bool = False
|
|||
redact_messages_in_exceptions: Optional[bool] = False
|
||||
redact_user_api_key_info: Optional[bool] = False
|
||||
filter_invalid_headers: Optional[bool] = False
|
||||
add_user_information_to_llm_headers: Optional[
|
||||
bool
|
||||
] = None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers
|
||||
add_user_information_to_llm_headers: Optional[bool] = (
|
||||
None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers
|
||||
)
|
||||
store_audit_logs = False # Enterprise feature, allow users to see audit logs
|
||||
### end of callbacks #############
|
||||
|
||||
email: Optional[
|
||||
str
|
||||
] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
|
||||
token: Optional[
|
||||
str
|
||||
] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
|
||||
email: Optional[str] = (
|
||||
None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
|
||||
)
|
||||
token: Optional[str] = (
|
||||
None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
|
||||
)
|
||||
telemetry = True
|
||||
max_tokens: int = DEFAULT_MAX_TOKENS # OpenAI Defaults
|
||||
drop_params = bool(os.getenv("LITELLM_DROP_PARAMS", False))
|
||||
|
|
@ -253,6 +269,11 @@ blocked_user_list: Optional[Union[str, List]] = None
|
|||
banned_keywords_list: Optional[Union[str, List]] = None
|
||||
llm_guard_mode: Literal["all", "key-specific", "request-specific"] = "all"
|
||||
guardrail_name_config_map: Dict[str, GuardrailItem] = {}
|
||||
### PROMPTS ###
|
||||
from litellm.types.prompts.init_prompts import PromptSpec
|
||||
|
||||
prompt_name_config_map: Dict[str, PromptSpec] = {}
|
||||
|
||||
##################
|
||||
### PREVIEW FEATURES ###
|
||||
enable_preview_features: bool = False
|
||||
|
|
@ -266,11 +287,15 @@ enable_loadbalancing_on_batch_endpoints: Optional[bool] = None
|
|||
enable_caching_on_provider_specific_optional_params: bool = (
|
||||
False # feature-flag for caching on optional params - e.g. 'top_k'
|
||||
)
|
||||
caching: bool = False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
|
||||
caching_with_models: bool = False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
|
||||
cache: Optional[
|
||||
Cache
|
||||
] = None # cache object <- use this - https://docs.litellm.ai/docs/caching
|
||||
caching: bool = (
|
||||
False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
|
||||
)
|
||||
caching_with_models: bool = (
|
||||
False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
|
||||
)
|
||||
cache: Optional[Cache] = (
|
||||
None # cache object <- use this - https://docs.litellm.ai/docs/caching
|
||||
)
|
||||
default_in_memory_ttl: Optional[float] = None
|
||||
default_redis_ttl: Optional[float] = None
|
||||
default_redis_batch_cache_expiry: Optional[float] = None
|
||||
|
|
@ -278,9 +303,9 @@ model_alias_map: Dict[str, str] = {}
|
|||
model_group_alias_map: Dict[str, str] = {}
|
||||
model_group_settings: Optional["ModelGroupSettings"] = None
|
||||
max_budget: float = 0.0 # set the max budget across all providers
|
||||
budget_duration: Optional[
|
||||
str
|
||||
] = None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
|
||||
budget_duration: Optional[str] = (
|
||||
None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
|
||||
)
|
||||
default_soft_budget: float = (
|
||||
DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0
|
||||
)
|
||||
|
|
@ -289,14 +314,19 @@ forward_traceparent_to_llm_provider: bool = False
|
|||
|
||||
_current_cost = 0.0 # private variable, used if max budget is set
|
||||
error_logs: Dict = {}
|
||||
add_function_to_prompt: bool = False # if function calling not supported by api, append function call details to system prompt
|
||||
add_function_to_prompt: bool = (
|
||||
False # if function calling not supported by api, append function call details to system prompt
|
||||
)
|
||||
client_session: Optional[httpx.Client] = None
|
||||
aclient_session: Optional[httpx.AsyncClient] = None
|
||||
model_fallbacks: Optional[List] = None # Deprecated for 'litellm.fallbacks'
|
||||
model_cost_map_url: str = "https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json"
|
||||
model_cost_map_url: str = (
|
||||
"https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json"
|
||||
)
|
||||
suppress_debug_info = False
|
||||
dynamodb_table_name: Optional[str] = None
|
||||
s3_callback_params: Optional[Dict] = None
|
||||
datadog_llm_observability_params: Optional[Union[DatadogLLMObsInitParams, Dict]] = None
|
||||
aws_sqs_callback_params: Optional[Dict] = None
|
||||
generic_logger_headers: Optional[Dict] = None
|
||||
default_key_generate_params: Optional[Dict] = None
|
||||
|
|
@ -321,7 +351,9 @@ prometheus_metrics_config: Optional[List] = None
|
|||
disable_add_prefix_to_prompt: bool = (
|
||||
False # used by anthropic, to disable adding prefix to prompt
|
||||
)
|
||||
disable_copilot_system_to_assistant: bool = False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior.
|
||||
disable_copilot_system_to_assistant: bool = (
|
||||
False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior.
|
||||
)
|
||||
public_model_groups: Optional[List[str]] = None
|
||||
public_model_groups_links: Dict[str, str] = {}
|
||||
#### REQUEST PRIORITIZATION #####
|
||||
|
|
@ -329,13 +361,17 @@ priority_reservation: Optional[Dict[str, float]] = None
|
|||
|
||||
|
||||
######## Networking Settings ########
|
||||
use_aiohttp_transport: bool = True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead.
|
||||
use_aiohttp_transport: bool = (
|
||||
True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead.
|
||||
)
|
||||
aiohttp_trust_env: bool = False # set to true to use HTTP_ Proxy settings
|
||||
disable_aiohttp_transport: bool = False # Set this to true to use httpx instead
|
||||
disable_aiohttp_trust_env: bool = (
|
||||
False # When False, aiohttp will respect HTTP(S)_PROXY env vars
|
||||
)
|
||||
force_ipv4: bool = False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6.
|
||||
force_ipv4: bool = (
|
||||
False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6.
|
||||
)
|
||||
module_level_aclient = AsyncHTTPHandler(
|
||||
timeout=request_timeout, client_alias="module level aclient"
|
||||
)
|
||||
|
|
@ -349,13 +385,13 @@ fallbacks: Optional[List] = None
|
|||
context_window_fallbacks: Optional[List] = None
|
||||
content_policy_fallbacks: Optional[List] = None
|
||||
allowed_fails: int = 3
|
||||
num_retries_per_request: Optional[
|
||||
int
|
||||
] = None # for the request overall (incl. fallbacks + model retries)
|
||||
num_retries_per_request: Optional[int] = (
|
||||
None # for the request overall (incl. fallbacks + model retries)
|
||||
)
|
||||
####### SECRET MANAGERS #####################
|
||||
secret_manager_client: Optional[
|
||||
Any
|
||||
] = None # list of instantiated key management clients - e.g. azure kv, infisical, etc.
|
||||
secret_manager_client: Optional[Any] = (
|
||||
None # list of instantiated key management clients - e.g. azure kv, infisical, etc.
|
||||
)
|
||||
_google_kms_resource_name: Optional[str] = None
|
||||
_key_management_system: Optional[KeyManagementSystem] = None
|
||||
_key_management_settings: KeyManagementSettings = KeyManagementSettings()
|
||||
|
|
@ -494,6 +530,7 @@ lambda_ai_models: List = []
|
|||
hyperbolic_models: List = []
|
||||
recraft_models: List = []
|
||||
|
||||
|
||||
def is_bedrock_pricing_only_model(key: str) -> bool:
|
||||
"""
|
||||
Excludes keys with the pattern 'bedrock/<region>/<model>'. These are in the model_prices_and_context_window.json file for pricing purposes only.
|
||||
|
|
@ -1223,12 +1260,12 @@ from .types.llms.custom_llm import CustomLLMItem
|
|||
from .types.utils import GenericStreamingChunk
|
||||
|
||||
custom_provider_map: List[CustomLLMItem] = []
|
||||
_custom_providers: List[
|
||||
str
|
||||
] = [] # internal helper util, used to track names of custom providers
|
||||
disable_hf_tokenizer_download: Optional[
|
||||
bool
|
||||
] = None # disable huggingface tokenizer download. Defaults to openai clk100
|
||||
_custom_providers: List[str] = (
|
||||
[]
|
||||
) # internal helper util, used to track names of custom providers
|
||||
disable_hf_tokenizer_download: Optional[bool] = (
|
||||
None # disable huggingface tokenizer download. Defaults to openai clk100
|
||||
)
|
||||
global_disable_no_log_param: bool = False
|
||||
|
||||
### PASSTHROUGH ###
|
||||
|
|
|
|||
|
|
@ -108,6 +108,10 @@ verbose_router_logger.addHandler(handler)
|
|||
verbose_proxy_logger.addHandler(handler)
|
||||
verbose_logger.addHandler(handler)
|
||||
|
||||
# Suppress httpx request logging at INFO level
|
||||
httpx_logger = logging.getLogger("httpx")
|
||||
httpx_logger.setLevel(logging.WARNING)
|
||||
|
||||
ALL_LOGGERS = [
|
||||
logging.getLogger(),
|
||||
verbose_logger,
|
||||
|
|
|
|||
|
|
@ -789,6 +789,7 @@ def completion_cost( # noqa: PLR0915
|
|||
from litellm.llms.recraft.cost_calculator import (
|
||||
cost_calculator as recraft_image_cost_calculator,
|
||||
)
|
||||
|
||||
return recraft_image_cost_calculator(
|
||||
model=model,
|
||||
image_response=completion_response,
|
||||
|
|
@ -797,6 +798,7 @@ def completion_cost( # noqa: PLR0915
|
|||
from litellm.llms.gemini.image_generation.cost_calculator import (
|
||||
cost_calculator as gemini_image_cost_calculator,
|
||||
)
|
||||
|
||||
return gemini_image_cost_calculator(
|
||||
model=model,
|
||||
image_response=completion_response,
|
||||
|
|
@ -867,7 +869,10 @@ def completion_cost( # noqa: PLR0915
|
|||
from litellm.proxy._experimental.mcp_server.cost_calculator import (
|
||||
MCPCostCalculator,
|
||||
)
|
||||
return MCPCostCalculator.calculate_mcp_tool_call_cost(litellm_logging_obj=litellm_logging_obj)
|
||||
|
||||
return MCPCostCalculator.calculate_mcp_tool_call_cost(
|
||||
litellm_logging_obj=litellm_logging_obj
|
||||
)
|
||||
# Calculate cost based on prompt_tokens, completion_tokens
|
||||
if (
|
||||
"togethercomputer" in model
|
||||
|
|
@ -1318,7 +1323,7 @@ class BaseTokenUsageProcessor:
|
|||
combined.completion_tokens_details = CompletionTokensDetails()
|
||||
|
||||
# Check what keys exist in the model's completion_tokens_details
|
||||
for attr in dir(usage.completion_tokens_details):
|
||||
for attr in usage.completion_tokens_details.model_fields:
|
||||
if not attr.startswith("_") and not callable(
|
||||
getattr(usage.completion_tokens_details, attr)
|
||||
):
|
||||
|
|
@ -1326,7 +1331,8 @@ class BaseTokenUsageProcessor:
|
|||
combined.completion_tokens_details, attr, 0
|
||||
)
|
||||
new_val = getattr(usage.completion_tokens_details, attr, 0)
|
||||
if new_val is not None:
|
||||
|
||||
if new_val is not None and current_val is not None:
|
||||
setattr(
|
||||
combined.completion_tokens_details,
|
||||
attr,
|
||||
|
|
|
|||
|
|
@ -829,3 +829,65 @@ class BlockedPiiEntityError(Exception):
|
|||
self.guardrail_name = guardrail_name
|
||||
self.message = f"Blocked entity detected: {entity_type} by Guardrail: {guardrail_name}. This entity is not allowed to be used in this request."
|
||||
super().__init__(self.message)
|
||||
|
||||
|
||||
class MidStreamFallbackError(ServiceUnavailableError): # type: ignore
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
model: str,
|
||||
llm_provider: str,
|
||||
original_exception: Optional[Exception] = None,
|
||||
response: Optional[httpx.Response] = None,
|
||||
litellm_debug_info: Optional[str] = None,
|
||||
max_retries: Optional[int] = None,
|
||||
num_retries: Optional[int] = None,
|
||||
generated_content: str = "",
|
||||
is_pre_first_chunk: bool = False,
|
||||
):
|
||||
self.status_code = 503 # Service Unavailable
|
||||
self.message = f"litellm.MidStreamFallbackError: {message}"
|
||||
self.model = model
|
||||
self.llm_provider = llm_provider
|
||||
self.original_exception = original_exception
|
||||
self.litellm_debug_info = litellm_debug_info
|
||||
self.max_retries = max_retries
|
||||
self.num_retries = num_retries
|
||||
self.generated_content = generated_content
|
||||
self.is_pre_first_chunk = is_pre_first_chunk
|
||||
|
||||
# Create a response if one wasn't provided
|
||||
if response is None:
|
||||
self.response = httpx.Response(
|
||||
status_code=self.status_code,
|
||||
request=httpx.Request(
|
||||
method="POST",
|
||||
url=f"https://{llm_provider}.com/v1/",
|
||||
),
|
||||
)
|
||||
else:
|
||||
self.response = response
|
||||
|
||||
# Call the parent constructor
|
||||
super().__init__(
|
||||
message=self.message,
|
||||
llm_provider=llm_provider,
|
||||
model=model,
|
||||
response=self.response,
|
||||
litellm_debug_info=self.litellm_debug_info,
|
||||
max_retries=self.max_retries,
|
||||
num_retries=self.num_retries,
|
||||
)
|
||||
|
||||
def __str__(self):
|
||||
_message = self.message
|
||||
if self.num_retries:
|
||||
_message += f" LiteLLM Retried: {self.num_retries} times"
|
||||
if self.max_retries:
|
||||
_message += f", LiteLLM Max Retries: {self.max_retries}"
|
||||
if self.original_exception:
|
||||
_message += f" Original exception: {type(self.original_exception).__name__}: {str(self.original_exception)}"
|
||||
return _message
|
||||
|
||||
def __repr__(self):
|
||||
return self.__str__()
|
||||
|
|
|
|||
|
|
@ -12,11 +12,15 @@ from mcp.client.stdio import stdio_client
|
|||
from mcp.client.streamable_http import streamablehttp_client
|
||||
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
|
||||
from mcp.types import CallToolResult as MCPCallToolResult
|
||||
from mcp.types import TextContent
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.mcp import (
|
||||
MCPAuth,
|
||||
MCPAuthType,
|
||||
MCPSpecVersion,
|
||||
MCPSpecVersionType,
|
||||
MCPStdioConfig,
|
||||
MCPTransport,
|
||||
MCPTransportType,
|
||||
|
|
@ -44,6 +48,7 @@ class MCPClient:
|
|||
auth_value: Optional[str] = None,
|
||||
timeout: float = 60.0,
|
||||
stdio_config: Optional[MCPStdioConfig] = None,
|
||||
protocol_version: MCPSpecVersionType = MCPSpecVersion.jun_2025,
|
||||
):
|
||||
self.server_url: str = server_url
|
||||
self.transport_type: MCPTransport = transport_type
|
||||
|
|
@ -57,6 +62,7 @@ class MCPClient:
|
|||
self._session_ctx = None
|
||||
self._task: Optional[asyncio.Task] = None
|
||||
self.stdio_config: Optional[MCPStdioConfig] = stdio_config
|
||||
self.protocol_version: MCPSpecVersionType = protocol_version
|
||||
|
||||
# handle the basic auth value if provided
|
||||
if auth_value:
|
||||
|
|
@ -118,9 +124,18 @@ class MCPClient:
|
|||
self._session_ctx = ClientSession(self._transport[0], self._transport[1])
|
||||
self._session = await self._session_ctx.__aenter__()
|
||||
await self._session.initialize()
|
||||
except Exception:
|
||||
except ValueError as e:
|
||||
# Re-raise ValueError exceptions (like missing stdio_config)
|
||||
verbose_logger.warning(f"MCP client connection failed: {str(e)}")
|
||||
await self.disconnect()
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"MCP client connection failed: {str(e)}")
|
||||
await self.disconnect()
|
||||
# Don't raise other exceptions, let the calling code handle it gracefully
|
||||
# This allows the server manager to continue with other servers
|
||||
# Instead of raising, we'll let the calling code handle the failure
|
||||
pass
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
"""Cleanup when exiting context manager."""
|
||||
|
|
@ -169,23 +184,40 @@ class MCPClient:
|
|||
|
||||
def _get_auth_headers(self) -> dict:
|
||||
"""Generate authentication headers based on auth type."""
|
||||
if not self._mcp_auth_value:
|
||||
return {}
|
||||
headers = {}
|
||||
|
||||
if self._mcp_auth_value:
|
||||
if self.auth_type == MCPAuth.bearer_token:
|
||||
headers["Authorization"] = f"Bearer {self._mcp_auth_value}"
|
||||
elif self.auth_type == MCPAuth.basic:
|
||||
headers["Authorization"] = f"Basic {self._mcp_auth_value}"
|
||||
elif self.auth_type == MCPAuth.api_key:
|
||||
headers["X-API-Key"] = self._mcp_auth_value
|
||||
|
||||
# Handle protocol version - it might be a string or enum
|
||||
if hasattr(self.protocol_version, 'value'):
|
||||
# It's an enum
|
||||
protocol_version_str = self.protocol_version.value
|
||||
else:
|
||||
# It's a string
|
||||
protocol_version_str = str(self.protocol_version)
|
||||
|
||||
headers["MCP-Protocol-Version"] = protocol_version_str
|
||||
return headers
|
||||
|
||||
if self.auth_type == MCPAuth.bearer_token:
|
||||
return {"Authorization": f"Bearer {self._mcp_auth_value}"}
|
||||
elif self.auth_type == MCPAuth.basic:
|
||||
return {"Authorization": f"Basic {self._mcp_auth_value}"}
|
||||
elif self.auth_type == MCPAuth.api_key:
|
||||
return {"X-API-Key": self._mcp_auth_value}
|
||||
return {}
|
||||
|
||||
async def list_tools(self) -> List[MCPTool]:
|
||||
"""List available tools from the server."""
|
||||
if not self._session:
|
||||
await self.connect()
|
||||
try:
|
||||
await self.connect()
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"MCP client connection failed: {str(e)}")
|
||||
return []
|
||||
|
||||
if self._session is None:
|
||||
raise ValueError("Session is not initialized")
|
||||
verbose_logger.warning("MCP client session is not initialized")
|
||||
return []
|
||||
|
||||
try:
|
||||
result = await self._session.list_tools()
|
||||
|
|
@ -193,9 +225,11 @@ class MCPClient:
|
|||
except asyncio.CancelledError:
|
||||
await self.disconnect()
|
||||
raise
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"MCP client list_tools failed: {str(e)}")
|
||||
await self.disconnect()
|
||||
raise
|
||||
# Return empty list instead of raising to allow graceful degradation
|
||||
return []
|
||||
|
||||
async def call_tool(
|
||||
self, call_tool_request_params: MCPCallToolRequestParams
|
||||
|
|
@ -204,10 +238,21 @@ class MCPClient:
|
|||
Call an MCP Tool.
|
||||
"""
|
||||
if not self._session:
|
||||
await self.connect()
|
||||
try:
|
||||
await self.connect()
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"MCP client connection failed: {str(e)}")
|
||||
return MCPCallToolResult(
|
||||
content=[TextContent(type="text", text=f"{str(e)}")],
|
||||
isError=True
|
||||
)
|
||||
|
||||
if self._session is None:
|
||||
raise ValueError("Session is not initialized")
|
||||
verbose_logger.warning("MCP client session is not initialized")
|
||||
return MCPCallToolResult(
|
||||
content=[TextContent(type="text", text="MCP client session is not initialized")],
|
||||
isError=True,
|
||||
)
|
||||
|
||||
try:
|
||||
tool_result = await self._session.call_tool(
|
||||
|
|
@ -218,8 +263,13 @@ class MCPClient:
|
|||
except asyncio.CancelledError:
|
||||
await self.disconnect()
|
||||
raise
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"MCP client call_tool failed: {str(e)}")
|
||||
await self.disconnect()
|
||||
raise
|
||||
# Return a default error result instead of raising
|
||||
return MCPCallToolResult(
|
||||
content=[TextContent(type="text", text=f"{str(e)}")], # Empty content for error case
|
||||
isError=True,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -24,11 +24,14 @@ if TYPE_CHECKING:
|
|||
GenerateContentConfigDict,
|
||||
GenerateContentContentListUnionDict,
|
||||
GenerateContentResponse,
|
||||
ToolConfigDict,
|
||||
)
|
||||
else:
|
||||
GenerateContentConfigDict = Any
|
||||
GenerateContentContentListUnionDict = Any
|
||||
GenerateContentResponse = Any
|
||||
ToolConfigDict = Any
|
||||
|
||||
|
||||
####### ENVIRONMENT VARIABLES ###################
|
||||
# Initialize any necessary instances or variables here
|
||||
|
|
@ -83,6 +86,7 @@ class GenerateContentHelper:
|
|||
config: Optional[GenerateContentConfigDict] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
stream: bool = False,
|
||||
tools: Optional[ToolConfigDict] = None,
|
||||
**kwargs,
|
||||
) -> GenerateContentSetupResult:
|
||||
"""
|
||||
|
|
@ -166,6 +170,7 @@ class GenerateContentHelper:
|
|||
generate_content_provider_config.transform_generate_content_request(
|
||||
model=model,
|
||||
contents=contents,
|
||||
tools=tools,
|
||||
generate_content_config_dict=generate_content_config_dict,
|
||||
)
|
||||
)
|
||||
|
|
@ -200,6 +205,7 @@ async def agenerate_content(
|
|||
model: str,
|
||||
contents: GenerateContentContentListUnionDict,
|
||||
config: Optional[GenerateContentConfigDict] = None,
|
||||
tools: Optional[ToolConfigDict] = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
|
|
@ -235,6 +241,7 @@ async def agenerate_content(
|
|||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
tools=tools,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
@ -263,6 +270,7 @@ def generate_content(
|
|||
model: str,
|
||||
contents: GenerateContentContentListUnionDict,
|
||||
config: Optional[GenerateContentConfigDict] = None,
|
||||
tools: Optional[ToolConfigDict] = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
|
|
@ -296,6 +304,7 @@ def generate_content(
|
|||
config=config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
stream=False,
|
||||
tools=tools,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
@ -316,6 +325,7 @@ def generate_content(
|
|||
response = base_llm_http_handler.generate_content_handler(
|
||||
model=setup_result.model,
|
||||
contents=contents,
|
||||
tools=tools,
|
||||
generate_content_provider_config=setup_result.generate_content_provider_config,
|
||||
generate_content_config_dict=setup_result.generate_content_config_dict,
|
||||
custom_llm_provider=setup_result.custom_llm_provider,
|
||||
|
|
@ -346,6 +356,7 @@ async def agenerate_content_stream(
|
|||
model: str,
|
||||
contents: GenerateContentContentListUnionDict,
|
||||
config: Optional[GenerateContentConfigDict] = None,
|
||||
tools: Optional[ToolConfigDict] = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
|
|
@ -377,6 +388,7 @@ async def agenerate_content_stream(
|
|||
"config": config,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"stream": True,
|
||||
"tools": tools,
|
||||
**kwargs,
|
||||
}
|
||||
)
|
||||
|
|
@ -402,6 +414,7 @@ async def agenerate_content_stream(
|
|||
contents=contents,
|
||||
generate_content_provider_config=setup_result.generate_content_provider_config,
|
||||
generate_content_config_dict=setup_result.generate_content_config_dict,
|
||||
tools=tools,
|
||||
custom_llm_provider=setup_result.custom_llm_provider,
|
||||
litellm_params=setup_result.litellm_params,
|
||||
logging_obj=setup_result.litellm_logging_obj,
|
||||
|
|
@ -429,6 +442,7 @@ def generate_content_stream(
|
|||
model: str,
|
||||
contents: GenerateContentContentListUnionDict,
|
||||
config: Optional[GenerateContentConfigDict] = None,
|
||||
tools: Optional[ToolConfigDict] = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
|
|
@ -454,6 +468,7 @@ def generate_content_stream(
|
|||
config=config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
stream=True,
|
||||
tools=tools,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
@ -476,6 +491,7 @@ def generate_content_stream(
|
|||
contents=contents,
|
||||
generate_content_provider_config=setup_result.generate_content_provider_config,
|
||||
generate_content_config_dict=setup_result.generate_content_config_dict,
|
||||
tools=tools,
|
||||
custom_llm_provider=setup_result.custom_llm_provider,
|
||||
litellm_params=setup_result.litellm_params,
|
||||
logging_obj=setup_result.litellm_logging_obj,
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ Users can define
|
|||
import copy
|
||||
from typing import Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.integrations.custom_prompt_management import CustomPromptManagement
|
||||
from litellm.types.integrations.anthropic_cache_control_hook import (
|
||||
|
|
@ -80,12 +81,22 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
# Case 1: Target by specific index
|
||||
if targetted_index is not None:
|
||||
original_index = targetted_index
|
||||
# Handle negative indices (convert to positive)
|
||||
if targetted_index < 0:
|
||||
targetted_index += len(messages)
|
||||
|
||||
if 0 <= targetted_index < len(messages):
|
||||
messages[targetted_index] = (
|
||||
AnthropicCacheControlHook._safe_insert_cache_control_in_message(
|
||||
messages[targetted_index], control
|
||||
)
|
||||
)
|
||||
else:
|
||||
verbose_logger.warning(
|
||||
f"AnthropicCacheControlHook: Provided index {original_index} is out of bounds for message list of length {len(messages)}. "
|
||||
f"Targeted index was {targetted_index}. Skipping cache control injection for this point."
|
||||
)
|
||||
# Case 2: Target by role
|
||||
elif targetted_role is not None:
|
||||
for msg in messages:
|
||||
|
|
|
|||
|
|
@ -234,7 +234,6 @@ class CustomGuardrail(CustomLogger):
|
|||
Returns True if the guardrail should be run on the event_type
|
||||
"""
|
||||
requested_guardrails = self.get_guardrail_from_metadata(data)
|
||||
|
||||
verbose_logger.debug(
|
||||
"inside should_run_guardrail for guardrail=%s event_type= %s guardrail_supported_event_hooks= %s requested_guardrails= %s self.default_on= %s",
|
||||
self.guardrail_name,
|
||||
|
|
@ -243,7 +242,6 @@ class CustomGuardrail(CustomLogger):
|
|||
requested_guardrails,
|
||||
self.default_on,
|
||||
)
|
||||
|
||||
if self.default_on is True:
|
||||
if self._event_hook_is_event_type(event_type):
|
||||
if isinstance(self.event_hook, Mode):
|
||||
|
|
@ -287,7 +285,6 @@ class CustomGuardrail(CustomLogger):
|
|||
)
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
return True
|
||||
|
||||
def _event_hook_is_event_type(self, event_type: GuardrailEventHooks) -> bool:
|
||||
|
|
|
|||
|
|
@ -33,7 +33,13 @@ if TYPE_CHECKING:
|
|||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
from litellm.types.mcp import (
|
||||
MCPDuringCallRequestObject,
|
||||
MCPDuringCallResponseObject,
|
||||
MCPPostCallResponseObject,
|
||||
MCPPreCallRequestObject,
|
||||
MCPPreCallResponseObject,
|
||||
)
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
|
|
@ -42,13 +48,30 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
UserAPIKeyAuth = Any
|
||||
MCPPostCallResponseObject = Any
|
||||
MCPPreCallRequestObject = Any
|
||||
MCPPreCallResponseObject = Any
|
||||
MCPDuringCallRequestObject = Any
|
||||
MCPDuringCallResponseObject = Any
|
||||
PreRoutingHookResponse = Any
|
||||
|
||||
|
||||
class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callback#callback-class
|
||||
# Class variables or attributes
|
||||
def __init__(self, message_logging: bool = True, **kwargs) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
turn_off_message_logging: bool = False,
|
||||
|
||||
# deprecated param, use `turn_off_message_logging` instead
|
||||
message_logging: bool = True,
|
||||
**kwargs
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
turn_off_message_logging: bool - if True, the message logging will be turned off. Message and response will be redacted from StandardLoggingPayload.
|
||||
message_logging: bool - deprecated param, use `turn_off_message_logging` instead
|
||||
"""
|
||||
self.message_logging = message_logging
|
||||
self.turn_off_message_logging = turn_off_message_logging
|
||||
pass
|
||||
|
||||
def log_pre_api_call(self, model, messages, kwargs):
|
||||
|
|
@ -258,6 +281,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
"audio_transcription",
|
||||
"pass_through_endpoint",
|
||||
"rerank",
|
||||
"mcp_call",
|
||||
],
|
||||
) -> Optional[
|
||||
Union[Exception, str, dict]
|
||||
|
|
@ -304,6 +328,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
) -> Any:
|
||||
pass
|
||||
|
|
@ -387,6 +412,60 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
#########################################################
|
||||
# MCP TOOL CALL HOOKS
|
||||
#########################################################
|
||||
async def async_pre_mcp_tool_call_hook(
|
||||
self,
|
||||
kwargs,
|
||||
request_obj: MCPPreCallRequestObject,
|
||||
start_time,
|
||||
end_time
|
||||
) -> Optional[MCPPreCallResponseObject]:
|
||||
"""
|
||||
This hook gets called before the MCP tool call is made.
|
||||
|
||||
Useful for:
|
||||
- Validating tool calls before execution
|
||||
- Modifying arguments before they are sent to the MCP server
|
||||
- Implementing access control and rate limiting
|
||||
- Adding custom metadata or tracking information
|
||||
|
||||
Args:
|
||||
kwargs: The logging kwargs containing model call details
|
||||
request_obj: MCPPreCallRequestObject containing tool name, arguments, and metadata
|
||||
start_time: Start time of the request
|
||||
end_time: End time of the request
|
||||
|
||||
Returns:
|
||||
MCPPreCallResponseObject with validation results and any modifications
|
||||
"""
|
||||
return None
|
||||
|
||||
async def async_during_mcp_tool_call_hook(
|
||||
self,
|
||||
kwargs,
|
||||
request_obj: MCPDuringCallRequestObject,
|
||||
start_time,
|
||||
end_time
|
||||
) -> Optional[MCPDuringCallResponseObject]:
|
||||
"""
|
||||
This hook gets called during the MCP tool call execution.
|
||||
|
||||
Useful for:
|
||||
- Concurrent monitoring and validation during tool execution
|
||||
- Implementing timeouts and cancellation logic
|
||||
- Real-time cost tracking and billing
|
||||
- Performance monitoring and metrics collection
|
||||
|
||||
Args:
|
||||
kwargs: The logging kwargs containing model call details
|
||||
request_obj: MCPDuringCallRequestObject containing tool execution context
|
||||
start_time: Start time of the request
|
||||
end_time: End time of the request
|
||||
|
||||
Returns:
|
||||
MCPDuringCallResponseObject with execution control decisions
|
||||
"""
|
||||
return None
|
||||
|
||||
async def async_post_mcp_tool_call_hook(
|
||||
self, kwargs, response_obj: MCPPostCallResponseObject, start_time, end_time
|
||||
) -> Optional[MCPPostCallResponseObject]:
|
||||
|
|
@ -470,3 +549,49 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
if LITELLM_METADATA_FIELD in request_kwargs:
|
||||
return LITELLM_METADATA_FIELD
|
||||
return OLD_LITELLM_METADATA_FIELD
|
||||
|
||||
def redact_standard_logging_payload_from_model_call_details(
|
||||
self, model_call_details: Dict
|
||||
) -> Dict:
|
||||
"""
|
||||
Only redacts messages and responses when self.turn_off_message_logging is True
|
||||
|
||||
|
||||
By default, self.turn_off_message_logging is False and this does nothing.
|
||||
|
||||
Return a redacted deepcopy of the provided logging payload.
|
||||
|
||||
This is useful for logging payloads that contain sensitive information.
|
||||
"""
|
||||
from copy import copy
|
||||
|
||||
from litellm import Choices, Message, ModelResponse
|
||||
from litellm.types.utils import LiteLLMCommonStrings
|
||||
turn_off_message_logging: bool = getattr(self, "turn_off_message_logging", False)
|
||||
|
||||
if turn_off_message_logging is False:
|
||||
return model_call_details
|
||||
|
||||
# Only make a shallow copy of the top-level dict to avoid deepcopy issues
|
||||
# with complex objects like AuthenticationError that may be present
|
||||
model_call_details_copy = copy(model_call_details)
|
||||
redacted_str = LiteLLMCommonStrings.redacted_by_litellm.value
|
||||
standard_logging_object = model_call_details.get("standard_logging_object")
|
||||
if standard_logging_object is None:
|
||||
return model_call_details_copy
|
||||
|
||||
# Make a copy of just the standard_logging_object to avoid modifying the original
|
||||
standard_logging_object_copy = copy(standard_logging_object)
|
||||
|
||||
if standard_logging_object_copy.get("messages") is not None:
|
||||
standard_logging_object_copy["messages"] = [Message(content=redacted_str).model_dump()]
|
||||
|
||||
if standard_logging_object_copy.get("response") is not None:
|
||||
model_response = ModelResponse(
|
||||
choices=[Choices(message=Message(content=redacted_str))]
|
||||
)
|
||||
model_response_dict = model_response.model_dump()
|
||||
standard_logging_object_copy["response"] = model_response_dict
|
||||
|
||||
model_call_details_copy["standard_logging_object"] = standard_logging_object_copy
|
||||
return model_call_details_copy
|
||||
|
|
|
|||
|
|
@ -58,18 +58,40 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
|
|||
asyncio.create_task(self.periodic_flush())
|
||||
self.flush_lock = asyncio.Lock()
|
||||
self.log_queue: List[LLMObsPayload] = []
|
||||
|
||||
#########################################################
|
||||
# Handle datadog_llm_observability_params set as litellm.datadog_llm_observability_params
|
||||
#########################################################
|
||||
dict_datadog_llm_obs_params = self._get_datadog_llm_obs_params()
|
||||
kwargs.update(dict_datadog_llm_obs_params)
|
||||
CustomBatchLogger.__init__(self, **kwargs, flush_lock=self.flush_lock)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"DataDogLLMObs: Error initializing - {str(e)}")
|
||||
raise e
|
||||
|
||||
def _get_datadog_llm_obs_params(self) -> Dict:
|
||||
"""
|
||||
Get the datadog_llm_observability_params from litellm.datadog_llm_observability_params
|
||||
|
||||
These are params specific to initializing the DataDogLLMObsLogger e.g. turn_off_message_logging
|
||||
"""
|
||||
dict_datadog_llm_obs_params: Dict = {}
|
||||
if litellm.datadog_llm_observability_params is not None:
|
||||
if isinstance(litellm.datadog_llm_observability_params, DatadogLLMObsInitParams):
|
||||
dict_datadog_llm_obs_params = litellm.datadog_llm_observability_params.model_dump()
|
||||
elif isinstance(litellm.datadog_llm_observability_params, Dict):
|
||||
# only allow params that are of DatadogLLMObsInitParams
|
||||
dict_datadog_llm_obs_params = DatadogLLMObsInitParams(**litellm.datadog_llm_observability_params).model_dump()
|
||||
return dict_datadog_llm_obs_params
|
||||
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
verbose_logger.debug(
|
||||
f"DataDogLLMObs: Logging success event for model {kwargs.get('model', 'unknown')}"
|
||||
)
|
||||
payload = self.create_llm_obs_payload(
|
||||
kwargs, response_obj, start_time, end_time
|
||||
kwargs, start_time, end_time
|
||||
)
|
||||
verbose_logger.debug(f"DataDogLLMObs: Payload: {payload}")
|
||||
self.log_queue.append(payload)
|
||||
|
|
@ -128,7 +150,7 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
|
|||
verbose_logger.exception(f"DataDogLLMObs: Error sending batch - {str(e)}")
|
||||
|
||||
def create_llm_obs_payload(
|
||||
self, kwargs: Dict, response_obj: Any, start_time: datetime, end_time: datetime
|
||||
self, kwargs: Dict, start_time: datetime, end_time: datetime
|
||||
) -> LLMObsPayload:
|
||||
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object"
|
||||
|
|
@ -138,6 +160,7 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
|
|||
|
||||
messages = standard_logging_payload["messages"]
|
||||
messages = self._ensure_string_content(messages=messages)
|
||||
response_obj = standard_logging_payload.get("response")
|
||||
|
||||
metadata = kwargs.get("litellm_params", {}).get("metadata", {})
|
||||
|
||||
|
|
@ -146,7 +169,10 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
|
|||
messages
|
||||
)
|
||||
)
|
||||
output_meta = OutputMeta(messages=self._get_response_messages(response_obj))
|
||||
output_meta = OutputMeta(messages=self._get_response_messages(
|
||||
response_obj=response_obj,
|
||||
call_type=standard_logging_payload.get("call_type")
|
||||
))
|
||||
|
||||
meta = Meta(
|
||||
kind=self._get_datadog_span_kind(standard_logging_payload.get("call_type")),
|
||||
|
|
@ -198,14 +224,16 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
|
|||
return 0.0
|
||||
|
||||
|
||||
def _get_response_messages(self, response_obj: Any) -> List[Any]:
|
||||
def _get_response_messages(
|
||||
self, response_obj: Any, call_type: Optional[str]
|
||||
) -> List[Any]:
|
||||
"""
|
||||
Get the messages from the response object
|
||||
|
||||
for now this handles logging /chat/completions responses
|
||||
"""
|
||||
if isinstance(response_obj, litellm.ModelResponse):
|
||||
return [response_obj["choices"][0]["message"].json()]
|
||||
if call_type in [CallTypes.completion.value, CallTypes.acompletion.value]:
|
||||
return [response_obj["choices"][0]["message"]]
|
||||
return []
|
||||
|
||||
def _get_datadog_span_kind(self, call_type: Optional[str]) -> Literal["llm", "tool", "task", "embedding", "retrieval"]:
|
||||
|
|
|
|||
316
litellm/integrations/dotprompt/README.md
Normal file
|
|
@ -0,0 +1,316 @@
|
|||
# LiteLLM Dotprompt Manager
|
||||
|
||||
A powerful prompt management system for LiteLLM that supports [Google's Dotprompt specification](https://google.github.io/dotprompt/getting-started/). This allows you to manage your AI prompts in organized `.prompt` files with YAML frontmatter, Handlebars templating, and full integration with LiteLLM's completion API.
|
||||
|
||||
## Features
|
||||
|
||||
- **📁 File-based prompt management**: Organize prompts in `.prompt` files
|
||||
- **🎯 YAML frontmatter**: Define model, parameters, and schemas in file headers
|
||||
- **🔧 Handlebars templating**: Use `{{variable}}` syntax with Jinja2 backend
|
||||
- **✅ Input validation**: Automatic validation against defined schemas
|
||||
- **🔗 LiteLLM integration**: Works seamlessly with `litellm.completion()`
|
||||
- **💬 Smart message parsing**: Converts prompts to proper chat messages
|
||||
- **⚙️ Parameter extraction**: Automatically applies model settings from prompts
|
||||
|
||||
## Quick Start
|
||||
|
||||
### 1. Create a `.prompt` file
|
||||
|
||||
Create a file called `chat_assistant.prompt`:
|
||||
|
||||
```yaml
|
||||
---
|
||||
model: gpt-4
|
||||
temperature: 0.7
|
||||
max_tokens: 150
|
||||
input:
|
||||
schema:
|
||||
user_message: string
|
||||
system_context?: string
|
||||
---
|
||||
|
||||
{% if system_context %}System: {{system_context}}
|
||||
|
||||
{% endif %}User: {{user_message}}
|
||||
```
|
||||
|
||||
### 2. Use with LiteLLM
|
||||
|
||||
```python
|
||||
import litellm
|
||||
|
||||
litellm.set_global_prompt_directory("path/to/your/prompts")
|
||||
|
||||
# Use with completion - the model prefix 'dotprompt/' tells LiteLLM to use prompt management
|
||||
response = litellm.completion(
|
||||
model="dotprompt/gpt-4", # The actual model comes from the .prompt file
|
||||
prompt_id="chat_assistant",
|
||||
prompt_variables={
|
||||
"user_message": "What is machine learning?",
|
||||
"system_context": "You are a helpful AI tutor."
|
||||
},
|
||||
# Any additional messages will be appended after the prompt
|
||||
messages=[{"role": "user", "content": "Please explain it simply."}]
|
||||
)
|
||||
|
||||
print(response.choices[0].message.content)
|
||||
```
|
||||
|
||||
## Prompt File Format
|
||||
|
||||
### Basic Structure
|
||||
|
||||
```yaml
|
||||
---
|
||||
# Model configuration
|
||||
model: gpt-4
|
||||
temperature: 0.7
|
||||
max_tokens: 500
|
||||
|
||||
# Input schema (optional)
|
||||
input:
|
||||
schema:
|
||||
name: string
|
||||
age: integer
|
||||
preferences?: array
|
||||
---
|
||||
|
||||
# Template content using Handlebars syntax
|
||||
Hello {{name}}!
|
||||
|
||||
{% if age >= 18 %}
|
||||
You're an adult, so here are some mature recommendations:
|
||||
{% else %}
|
||||
Here are some age-appropriate suggestions:
|
||||
{% endif %}
|
||||
|
||||
{% for pref in preferences %}
|
||||
- Based on your interest in {{pref}}, I recommend...
|
||||
{% endfor %}
|
||||
```
|
||||
|
||||
### Supported Frontmatter Fields
|
||||
|
||||
- **`model`**: The LLM model to use (e.g., `gpt-4`, `claude-3-sonnet`)
|
||||
- **`input.schema`**: Define expected input variables and their types
|
||||
- **`output.format`**: Expected output format (`json`, `text`, etc.)
|
||||
- **`output.schema`**: Structure of expected output
|
||||
|
||||
### Additional Parameters
|
||||
|
||||
- **`temperature`**: Model temperature (0.0 to 1.0)
|
||||
- **`max_tokens`**: Maximum tokens to generate
|
||||
- **`top_p`**: Nucleus sampling parameter (0.0 to 1.0)
|
||||
- **`frequency_penalty`**: Frequency penalty (0.0 to 1.0)
|
||||
- **`presence_penalty`**: Presence penalty (0.0 to 1.0)
|
||||
- any other parameters that are not model or schema-related will be treated as optional parameters to the model.
|
||||
|
||||
### Input Schema Types
|
||||
|
||||
- `string` or `str`: Text values
|
||||
- `integer` or `int`: Whole numbers
|
||||
- `float`: Decimal numbers
|
||||
- `boolean` or `bool`: True/false values
|
||||
- `array` or `list`: Lists of values
|
||||
- `object` or `dict`: Key-value objects
|
||||
|
||||
Use `?` suffix for optional fields: `name?: string`
|
||||
|
||||
## Message Format Conversion
|
||||
|
||||
The dotprompt manager intelligently converts your rendered prompts into proper chat messages:
|
||||
|
||||
### Simple Text → User Message
|
||||
```yaml
|
||||
---
|
||||
model: gpt-4
|
||||
---
|
||||
Tell me about {{topic}}.
|
||||
```
|
||||
Becomes: `[{"role": "user", "content": "Tell me about AI."}]`
|
||||
|
||||
### Role-Based Format → Multiple Messages
|
||||
```yaml
|
||||
---
|
||||
model: gpt-4
|
||||
---
|
||||
System: You are a {{role}}.
|
||||
|
||||
User: {{question}}
|
||||
```
|
||||
|
||||
Becomes:
|
||||
```python
|
||||
[
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "What is AI?"}
|
||||
]
|
||||
```
|
||||
|
||||
|
||||
## Example Prompts
|
||||
|
||||
### Data Extraction
|
||||
```yaml
|
||||
# extract_info.prompt
|
||||
---
|
||||
model: gemini/gemini-1.5-pro
|
||||
input:
|
||||
schema:
|
||||
text: string
|
||||
output:
|
||||
format: json
|
||||
schema:
|
||||
title?: string
|
||||
summary: string
|
||||
tags: array
|
||||
---
|
||||
|
||||
Extract the requested information from the given text. Return JSON format.
|
||||
|
||||
Text: {{text}}
|
||||
```
|
||||
|
||||
### Code Assistant
|
||||
```yaml
|
||||
# code_helper.prompt
|
||||
---
|
||||
model: claude-3-5-sonnet-20241022
|
||||
temperature: 0.2
|
||||
max_tokens: 2000
|
||||
input:
|
||||
schema:
|
||||
language: string
|
||||
task: string
|
||||
code?: string
|
||||
---
|
||||
|
||||
You are an expert {{language}} programmer.
|
||||
|
||||
Task: {{task}}
|
||||
|
||||
{% if code %}
|
||||
Current code:
|
||||
```{{language}}
|
||||
{{code}}
|
||||
```
|
||||
{% endif %}
|
||||
|
||||
Please provide a complete, well-documented solution.
|
||||
```
|
||||
|
||||
### Multi-turn Conversation
|
||||
```yaml
|
||||
# conversation.prompt
|
||||
---
|
||||
model: gpt-4
|
||||
temperature: 0.8
|
||||
input:
|
||||
schema:
|
||||
personality: string
|
||||
context: string
|
||||
---
|
||||
|
||||
System: You are a {{personality}}. {{context}}
|
||||
|
||||
User: Let's start our conversation.
|
||||
```
|
||||
|
||||
## API Reference
|
||||
|
||||
### PromptManager
|
||||
|
||||
The core class for managing `.prompt` files.
|
||||
|
||||
#### Methods
|
||||
|
||||
- **`__init__(prompt_directory: str)`**: Initialize with directory path
|
||||
- **`render(prompt_id: str, variables: dict) -> str`**: Render prompt with variables
|
||||
- **`list_prompts() -> List[str]`**: Get all available prompt IDs
|
||||
- **`get_prompt(prompt_id: str) -> PromptTemplate`**: Get prompt template object
|
||||
- **`get_prompt_metadata(prompt_id: str) -> dict`**: Get prompt metadata
|
||||
- **`reload_prompts() -> None`**: Reload all prompts from directory
|
||||
- **`add_prompt(prompt_id: str, content: str, metadata: dict)`**: Add prompt programmatically
|
||||
|
||||
### DotpromptManager
|
||||
|
||||
LiteLLM integration class extending `PromptManagementBase`.
|
||||
|
||||
#### Methods
|
||||
|
||||
- **`__init__(prompt_directory: str)`**: Initialize with directory path
|
||||
- **`should_run_prompt_management(prompt_id: str, params: dict) -> bool`**: Check if prompt exists
|
||||
- **`set_prompt_directory(directory: str)`**: Change prompt directory
|
||||
- **`reload_prompts()`**: Reload prompts from directory
|
||||
|
||||
### PromptTemplate
|
||||
|
||||
Represents a single prompt with metadata.
|
||||
|
||||
#### Properties
|
||||
|
||||
- **`content: str`**: The prompt template content
|
||||
- **`metadata: dict`**: Full metadata from frontmatter
|
||||
- **`model: str`**: Specified model name
|
||||
- **`temperature: float`**: Model temperature
|
||||
- **`max_tokens: int`**: Token limit
|
||||
- **`input_schema: dict`**: Input validation schema
|
||||
- **`output_format: str`**: Expected output format
|
||||
- **`output_schema: dict`**: Output structure schema
|
||||
|
||||
## Best Practices
|
||||
|
||||
1. **Organize by purpose**: Group related prompts in subdirectories
|
||||
2. **Use descriptive names**: `extract_user_info.prompt` vs `prompt1.prompt`
|
||||
3. **Define schemas**: Always specify input schemas for validation
|
||||
4. **Version control**: Store `.prompt` files in git for change tracking
|
||||
5. **Test prompts**: Use the test framework to validate prompt behavior
|
||||
6. **Keep templates focused**: One prompt should do one thing well
|
||||
7. **Use includes**: Break complex prompts into reusable components
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Common Issues
|
||||
|
||||
**Prompt not found**: Ensure the `.prompt` file exists and has correct extension
|
||||
```python
|
||||
# Check available prompts
|
||||
from litellm.integrations.dotprompt import get_dotprompt_manager
|
||||
manager = get_dotprompt_manager()
|
||||
print(manager.prompt_manager.list_prompts())
|
||||
```
|
||||
|
||||
**Template errors**: Verify Handlebars syntax and variable names
|
||||
```python
|
||||
# Test rendering directly
|
||||
manager.prompt_manager.render("my_prompt", {"test": "value"})
|
||||
```
|
||||
|
||||
**Model not working**: Check that model name in frontmatter is correct
|
||||
```python
|
||||
# Check prompt metadata
|
||||
metadata = manager.prompt_manager.get_prompt_metadata("my_prompt")
|
||||
print(metadata)
|
||||
```
|
||||
|
||||
### Validation Errors
|
||||
|
||||
Input validation failures show helpful error messages:
|
||||
```
|
||||
ValueError: Invalid type for field 'age': expected int, got str
|
||||
```
|
||||
|
||||
Make sure your variables match the defined schema types.
|
||||
|
||||
## Contributing
|
||||
|
||||
The LiteLLM Dotprompt manager follows the [Dotprompt specification](https://google.github.io/dotprompt/) for maximum compatibility. When contributing:
|
||||
|
||||
1. Ensure compatibility with existing `.prompt` files
|
||||
2. Add tests for new features
|
||||
3. Update documentation
|
||||
4. Follow the existing code style
|
||||
|
||||
## License
|
||||
|
||||
This prompt management system is part of LiteLLM and follows the same license terms.
|
||||
71
litellm/integrations/dotprompt/__init__.py
Normal file
|
|
@ -0,0 +1,71 @@
|
|||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .prompt_manager import PromptManager, PromptTemplate
|
||||
from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec
|
||||
from litellm.integrations.custom_prompt_management import CustomPromptManagement
|
||||
|
||||
from litellm.types.prompts.init_prompts import SupportedPromptIntegrations
|
||||
|
||||
from .dotprompt_manager import DotpromptManager
|
||||
|
||||
# Global instances
|
||||
global_prompt_directory: Optional[str] = None
|
||||
global_prompt_manager: Optional["PromptManager"] = None
|
||||
|
||||
|
||||
def set_global_prompt_directory(directory: str) -> None:
|
||||
"""
|
||||
Set the global prompt directory for dotprompt files.
|
||||
|
||||
Args:
|
||||
directory: Path to directory containing .prompt files
|
||||
"""
|
||||
import litellm
|
||||
|
||||
litellm.global_prompt_directory = directory # type: ignore
|
||||
|
||||
|
||||
def prompt_initializer(
|
||||
litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec"
|
||||
) -> "CustomPromptManagement":
|
||||
"""
|
||||
Initialize a prompt from a .prompt file.
|
||||
"""
|
||||
prompt_directory = getattr(litellm_params, "prompt_directory", None)
|
||||
prompt_data = getattr(litellm_params, "prompt_data", None)
|
||||
prompt_id = getattr(litellm_params, "prompt_id", None)
|
||||
if prompt_directory:
|
||||
raise ValueError(
|
||||
"Cannot set prompt_directory when working with prompt_initializer. Needs to be a specific dotprompt file"
|
||||
)
|
||||
|
||||
prompt_file = getattr(litellm_params, "prompt_file", None)
|
||||
|
||||
try:
|
||||
dot_prompt_manager = DotpromptManager(
|
||||
prompt_directory=prompt_directory,
|
||||
prompt_data=prompt_data,
|
||||
prompt_file=prompt_file,
|
||||
prompt_id=prompt_id,
|
||||
)
|
||||
|
||||
return dot_prompt_manager
|
||||
except Exception as e:
|
||||
|
||||
raise e
|
||||
|
||||
|
||||
prompt_initializer_registry = {
|
||||
SupportedPromptIntegrations.DOT_PROMPT.value: prompt_initializer,
|
||||
}
|
||||
|
||||
# Export public API
|
||||
__all__ = [
|
||||
"PromptManager",
|
||||
"DotpromptManager",
|
||||
"PromptTemplate",
|
||||
"set_global_prompt_directory",
|
||||
"global_prompt_directory",
|
||||
"global_prompt_manager",
|
||||
]
|
||||
291
litellm/integrations/dotprompt/dotprompt_manager.py
Normal file
|
|
@ -0,0 +1,291 @@
|
|||
"""
|
||||
Dotprompt manager that integrates with LiteLLM's prompt management system.
|
||||
Builds on top of PromptManagementBase to provide .prompt file support.
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from litellm.integrations.custom_prompt_management import CustomPromptManagement
|
||||
from litellm.integrations.prompt_management_base import PromptManagementClient
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
from .prompt_manager import PromptManager, PromptTemplate
|
||||
|
||||
|
||||
class DotpromptManager(CustomPromptManagement):
|
||||
"""
|
||||
Dotprompt manager that integrates with LiteLLM's prompt management system.
|
||||
|
||||
This class enables using .prompt files with the litellm completion() function
|
||||
by implementing the PromptManagementBase interface.
|
||||
|
||||
Usage:
|
||||
# Set global prompt directory
|
||||
litellm.prompt_directory = "path/to/prompts"
|
||||
|
||||
# Use with completion
|
||||
response = litellm.completion(
|
||||
model="dotprompt/gpt-4",
|
||||
prompt_id="my_prompt",
|
||||
prompt_variables={"variable": "value"},
|
||||
messages=[{"role": "user", "content": "This will be combined with the prompt"}]
|
||||
)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
prompt_directory: Optional[str] = None,
|
||||
prompt_file: Optional[str] = None,
|
||||
prompt_data: Optional[Union[dict, str]] = None,
|
||||
prompt_id: Optional[str] = None,
|
||||
):
|
||||
import litellm
|
||||
|
||||
self.prompt_directory = prompt_directory or litellm.global_prompt_directory
|
||||
# Support for JSON-based prompts stored in memory/database
|
||||
if isinstance(prompt_data, str):
|
||||
self.prompt_data = json.loads(prompt_data)
|
||||
else:
|
||||
self.prompt_data = prompt_data or {}
|
||||
|
||||
self._prompt_manager: Optional[PromptManager] = None
|
||||
self.prompt_file = prompt_file
|
||||
self.prompt_id = prompt_id
|
||||
|
||||
@property
|
||||
def integration_name(self) -> str:
|
||||
"""Integration name used in model names like 'dotprompt/gpt-4'."""
|
||||
return "dotprompt"
|
||||
|
||||
@property
|
||||
def prompt_manager(self) -> PromptManager:
|
||||
"""Lazy-load the prompt manager."""
|
||||
if self._prompt_manager is None:
|
||||
if (
|
||||
self.prompt_directory is None
|
||||
and not self.prompt_data
|
||||
and not self.prompt_file
|
||||
):
|
||||
raise ValueError(
|
||||
"Either prompt_directory or prompt_data must be set before using dotprompt manager. "
|
||||
"Set litellm.global_prompt_directory, initialize with prompt_directory parameter, or provide prompt_data."
|
||||
)
|
||||
self._prompt_manager = PromptManager(
|
||||
prompt_directory=self.prompt_directory,
|
||||
prompt_data=self.prompt_data,
|
||||
prompt_file=self.prompt_file,
|
||||
prompt_id=self.prompt_id,
|
||||
)
|
||||
return self._prompt_manager
|
||||
|
||||
def should_run_prompt_management(
|
||||
self,
|
||||
prompt_id: str,
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
) -> bool:
|
||||
"""
|
||||
Determine if prompt management should run based on the prompt_id.
|
||||
|
||||
Returns True if the prompt_id exists in our prompt manager.
|
||||
"""
|
||||
try:
|
||||
return prompt_id in self.prompt_manager.list_prompts()
|
||||
except Exception:
|
||||
# If there's any error accessing prompts, don't run prompt management
|
||||
return False
|
||||
|
||||
def _compile_prompt_helper(
|
||||
self,
|
||||
prompt_id: str,
|
||||
prompt_variables: Optional[dict],
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
) -> PromptManagementClient:
|
||||
"""
|
||||
Compile a .prompt file into a PromptManagementClient structure.
|
||||
|
||||
This method:
|
||||
1. Loads the prompt template from the .prompt file
|
||||
2. Renders it with the provided variables
|
||||
3. Converts the rendered text into chat messages
|
||||
4. Extracts model and optional parameters from metadata
|
||||
"""
|
||||
|
||||
try:
|
||||
|
||||
# Get the prompt template
|
||||
template = self.prompt_manager.get_prompt(prompt_id)
|
||||
if template is None:
|
||||
raise ValueError(f"Prompt '{prompt_id}' not found in prompt directory")
|
||||
|
||||
# Render the template with variables
|
||||
rendered_content = self.prompt_manager.render(prompt_id, prompt_variables)
|
||||
|
||||
# Convert rendered content to chat messages
|
||||
messages = self._convert_to_messages(rendered_content)
|
||||
|
||||
# Extract model from metadata (if specified)
|
||||
template_model = template.model
|
||||
|
||||
# Extract optional parameters from metadata
|
||||
optional_params = self._extract_optional_params(template)
|
||||
|
||||
return PromptManagementClient(
|
||||
prompt_id=prompt_id,
|
||||
prompt_template=messages,
|
||||
prompt_template_model=template_model,
|
||||
prompt_template_optional_params=optional_params,
|
||||
completed_messages=None,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise ValueError(f"Error compiling prompt '{prompt_id}': {e}")
|
||||
|
||||
def get_chat_completion_prompt(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
non_default_params: dict,
|
||||
prompt_id: Optional[str],
|
||||
prompt_variables: Optional[dict],
|
||||
dynamic_callback_params: StandardCallbackDynamicParams,
|
||||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
) -> Tuple[str, List[AllMessageValues], dict]:
|
||||
|
||||
from litellm.integrations.prompt_management_base import PromptManagementBase
|
||||
|
||||
return PromptManagementBase.get_chat_completion_prompt(
|
||||
self,
|
||||
model,
|
||||
messages,
|
||||
non_default_params,
|
||||
prompt_id,
|
||||
prompt_variables,
|
||||
dynamic_callback_params,
|
||||
prompt_label,
|
||||
prompt_version,
|
||||
)
|
||||
|
||||
def _convert_to_messages(self, rendered_content: str) -> List[AllMessageValues]:
|
||||
"""
|
||||
Convert rendered prompt content to chat messages.
|
||||
|
||||
This method supports multiple formats:
|
||||
1. Simple text -> converted to user message
|
||||
2. Text with role prefixes (System:, User:, Assistant:) -> parsed into separate messages
|
||||
3. Already formatted as a single message
|
||||
"""
|
||||
# Clean up the content
|
||||
content = rendered_content.strip()
|
||||
|
||||
# Try to parse role-based format (System: ..., User: ..., etc.)
|
||||
messages = []
|
||||
current_role = None
|
||||
current_content = []
|
||||
|
||||
lines = content.split("\n")
|
||||
|
||||
for line in lines:
|
||||
line = line.strip()
|
||||
|
||||
# Check for role prefixes
|
||||
if line.startswith("System:"):
|
||||
if current_role and current_content:
|
||||
messages.append(
|
||||
self._create_message(
|
||||
current_role, "\n".join(current_content).strip()
|
||||
)
|
||||
)
|
||||
current_role = "system"
|
||||
current_content = [line[7:].strip()] # Remove "System:" prefix
|
||||
elif line.startswith("User:"):
|
||||
if current_role and current_content:
|
||||
messages.append(
|
||||
self._create_message(
|
||||
current_role, "\n".join(current_content).strip()
|
||||
)
|
||||
)
|
||||
current_role = "user"
|
||||
current_content = [line[5:].strip()] # Remove "User:" prefix
|
||||
elif line.startswith("Assistant:"):
|
||||
if current_role and current_content:
|
||||
messages.append(
|
||||
self._create_message(
|
||||
current_role, "\n".join(current_content).strip()
|
||||
)
|
||||
)
|
||||
current_role = "assistant"
|
||||
current_content = [line[10:].strip()] # Remove "Assistant:" prefix
|
||||
else:
|
||||
# Continue current message content
|
||||
if current_role:
|
||||
current_content.append(line)
|
||||
else:
|
||||
# No role prefix found, treat as user message
|
||||
current_role = "user"
|
||||
current_content = [line]
|
||||
|
||||
# Add the last message
|
||||
if current_role and current_content:
|
||||
content_text = "\n".join(current_content).strip()
|
||||
if content_text: # Only add if there's actual content
|
||||
messages.append(self._create_message(current_role, content_text))
|
||||
|
||||
# If no messages were created, treat the entire content as a user message
|
||||
if not messages and content:
|
||||
messages.append(self._create_message("user", content))
|
||||
|
||||
return messages
|
||||
|
||||
def _create_message(self, role: str, content: str) -> AllMessageValues:
|
||||
"""Create a message with the specified role and content."""
|
||||
return {
|
||||
"role": role, # type: ignore
|
||||
"content": content,
|
||||
}
|
||||
|
||||
def _extract_optional_params(self, template: PromptTemplate) -> dict:
|
||||
"""
|
||||
Extract optional parameters from the prompt template metadata.
|
||||
|
||||
Includes parameters like temperature, max_tokens, etc.
|
||||
"""
|
||||
optional_params = {}
|
||||
|
||||
# Extract common parameters from metadata
|
||||
if template.optional_params is not None:
|
||||
optional_params.update(template.optional_params)
|
||||
|
||||
return optional_params
|
||||
|
||||
def set_prompt_directory(self, prompt_directory: str) -> None:
|
||||
"""Set the prompt directory and reload prompts."""
|
||||
self.prompt_directory = prompt_directory
|
||||
self._prompt_manager = None # Reset to force reload
|
||||
|
||||
def reload_prompts(self) -> None:
|
||||
"""Reload all prompts from the directory."""
|
||||
if self._prompt_manager:
|
||||
self._prompt_manager.reload_prompts()
|
||||
|
||||
def add_prompt_from_json(self, prompt_id: str, json_data: Dict[str, Any]) -> None:
|
||||
"""Add a prompt from JSON data."""
|
||||
content = json_data.get("content", "")
|
||||
metadata = json_data.get("metadata", {})
|
||||
self.prompt_manager.add_prompt(prompt_id, content, metadata)
|
||||
|
||||
def load_prompts_from_json(self, prompts_data: Dict[str, Dict[str, Any]]) -> None:
|
||||
"""Load multiple prompts from JSON data."""
|
||||
self.prompt_manager.load_prompts_from_json_data(prompts_data)
|
||||
|
||||
def get_prompts_as_json(self) -> Dict[str, Dict[str, Any]]:
|
||||
"""Get all prompts in JSON format."""
|
||||
return self.prompt_manager.get_all_prompts_as_json()
|
||||
|
||||
def convert_prompt_file_to_json(self, file_path: str) -> Dict[str, Any]:
|
||||
"""Convert a .prompt file to JSON format."""
|
||||
return self.prompt_manager.prompt_file_to_json(file_path)
|
||||
343
litellm/integrations/dotprompt/prompt_manager.py
Normal file
|
|
@ -0,0 +1,343 @@
|
|||
"""
|
||||
Based on Google's GenAI Kit dotprompt implementation: https://google.github.io/dotprompt/reference/frontmatter/
|
||||
"""
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import yaml
|
||||
from jinja2 import DictLoader, Environment, select_autoescape
|
||||
|
||||
|
||||
class PromptTemplate:
|
||||
"""Represents a single prompt template with metadata and content."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
content: str,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
template_id: Optional[str] = None,
|
||||
):
|
||||
self.content = content
|
||||
self.metadata = metadata or {}
|
||||
self.template_id = template_id
|
||||
|
||||
# Extract common metadata fields
|
||||
restricted_keys = ["model", "input", "output"]
|
||||
self.model = self.metadata.get("model")
|
||||
self.input_schema = self.metadata.get("input", {}).get("schema", {})
|
||||
self.output_format = self.metadata.get("output", {}).get("format")
|
||||
self.output_schema = self.metadata.get("output", {}).get("schema", {})
|
||||
self.optional_params = {}
|
||||
for key in self.metadata.keys():
|
||||
if key not in restricted_keys:
|
||||
self.optional_params[key] = self.metadata[key]
|
||||
|
||||
def __repr__(self):
|
||||
return f"PromptTemplate(id='{self.template_id}', model='{self.model}')"
|
||||
|
||||
|
||||
class PromptManager:
|
||||
"""
|
||||
Manager for loading and rendering .prompt files following the Dotprompt specification.
|
||||
|
||||
Supports:
|
||||
- YAML frontmatter for metadata
|
||||
- Handlebars-style templating (using Jinja2)
|
||||
- Input/output schema validation
|
||||
- Model configuration
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
prompt_id: Optional[str] = None,
|
||||
prompt_directory: Optional[str] = None,
|
||||
prompt_data: Optional[Dict[str, Dict[str, Any]]] = None,
|
||||
prompt_file: Optional[str] = None,
|
||||
):
|
||||
self.prompt_directory = Path(prompt_directory) if prompt_directory else None
|
||||
self.prompts: Dict[str, PromptTemplate] = {}
|
||||
self.prompt_file = prompt_file
|
||||
self.jinja_env = Environment(
|
||||
loader=DictLoader({}),
|
||||
autoescape=select_autoescape(["html", "xml"]),
|
||||
# Use Handlebars-style delimiters to match Dotprompt spec
|
||||
variable_start_string="{{",
|
||||
variable_end_string="}}",
|
||||
block_start_string="{%",
|
||||
block_end_string="%}",
|
||||
comment_start_string="{#",
|
||||
comment_end_string="#}",
|
||||
)
|
||||
|
||||
# Load prompts from directory if provided
|
||||
if self.prompt_directory:
|
||||
self._load_prompts()
|
||||
|
||||
if self.prompt_file:
|
||||
if not prompt_id:
|
||||
raise ValueError("prompt_id is required when prompt_file is provided")
|
||||
|
||||
template = self._load_prompt_file(self.prompt_file, prompt_id)
|
||||
self.prompts[prompt_id] = template
|
||||
|
||||
# Load prompts from JSON data if provided
|
||||
if prompt_data:
|
||||
self._load_prompts_from_json(prompt_data, prompt_id)
|
||||
|
||||
def _load_prompts(self) -> None:
|
||||
"""Load all .prompt files from the prompt directory."""
|
||||
if not self.prompt_directory or not self.prompt_directory.exists():
|
||||
raise ValueError(
|
||||
f"Prompt directory does not exist: {self.prompt_directory}"
|
||||
)
|
||||
|
||||
prompt_files = list(self.prompt_directory.glob("*.prompt"))
|
||||
|
||||
for prompt_file in prompt_files:
|
||||
try:
|
||||
prompt_id = prompt_file.stem # filename without extension
|
||||
template = self._load_prompt_file(prompt_file, prompt_id)
|
||||
self.prompts[prompt_id] = template
|
||||
# Optional: print(f"Loaded prompt: {prompt_id}")
|
||||
except Exception:
|
||||
# Optional: print(f"Error loading prompt file {prompt_file}")
|
||||
pass
|
||||
|
||||
def _load_prompts_from_json(
|
||||
self, prompt_data: Dict[str, Dict[str, Any]], prompt_id: Optional[str] = None
|
||||
) -> None:
|
||||
"""Load prompts from JSON data structure.
|
||||
|
||||
Expected format:
|
||||
{
|
||||
"prompt_id": {
|
||||
"content": "template content",
|
||||
"metadata": {"model": "gpt-4", "temperature": 0.7, ...}
|
||||
}
|
||||
}
|
||||
|
||||
or
|
||||
|
||||
{
|
||||
"content": "template content",
|
||||
"metadata": {"model": "gpt-4", "temperature": 0.7, ...}
|
||||
} + prompt_id
|
||||
"""
|
||||
if prompt_id:
|
||||
prompt_data = {prompt_id: prompt_data}
|
||||
|
||||
for prompt_id, prompt_info in prompt_data.items():
|
||||
try:
|
||||
content = prompt_info.get("content", "")
|
||||
metadata = prompt_info.get("metadata", {})
|
||||
|
||||
template = PromptTemplate(
|
||||
content=content,
|
||||
metadata=metadata,
|
||||
template_id=prompt_id,
|
||||
)
|
||||
self.prompts[prompt_id] = template
|
||||
except Exception:
|
||||
# Optional: print(f"Error loading prompt from JSON: {prompt_id}")
|
||||
pass
|
||||
|
||||
def _load_prompt_file(
|
||||
self, file_path: Union[str, Path], prompt_id: str
|
||||
) -> PromptTemplate:
|
||||
"""Load and parse a single .prompt file."""
|
||||
if isinstance(file_path, str):
|
||||
file_path = Path(file_path)
|
||||
|
||||
content = file_path.read_text(encoding="utf-8")
|
||||
|
||||
# Split frontmatter and content
|
||||
frontmatter, template_content = self._parse_frontmatter(content)
|
||||
|
||||
return PromptTemplate(
|
||||
content=template_content.strip(),
|
||||
metadata=frontmatter,
|
||||
template_id=prompt_id,
|
||||
)
|
||||
|
||||
def _parse_frontmatter(self, content: str) -> Tuple[Dict[str, Any], str]:
|
||||
"""Parse YAML frontmatter from prompt content."""
|
||||
# Match YAML frontmatter between --- delimiters
|
||||
frontmatter_pattern = r"^---\s*\n(.*?)\n---\s*\n(.*)$"
|
||||
match = re.match(frontmatter_pattern, content, re.DOTALL)
|
||||
|
||||
if match:
|
||||
frontmatter_yaml = match.group(1)
|
||||
template_content = match.group(2)
|
||||
|
||||
try:
|
||||
frontmatter = yaml.safe_load(frontmatter_yaml) or {}
|
||||
except yaml.YAMLError as e:
|
||||
raise ValueError(f"Invalid YAML frontmatter: {e}")
|
||||
else:
|
||||
# No frontmatter found, treat entire content as template
|
||||
frontmatter = {}
|
||||
template_content = content
|
||||
|
||||
return frontmatter, template_content
|
||||
|
||||
def render(
|
||||
self, prompt_id: str, prompt_variables: Optional[Dict[str, Any]] = None
|
||||
) -> str:
|
||||
"""
|
||||
Render a prompt template with the given variables.
|
||||
|
||||
Args:
|
||||
prompt_id: The ID of the prompt template to render
|
||||
prompt_variables: Variables to substitute in the template
|
||||
|
||||
Returns:
|
||||
The rendered prompt string
|
||||
|
||||
Raises:
|
||||
KeyError: If prompt_id is not found
|
||||
ValueError: If template rendering fails
|
||||
"""
|
||||
if prompt_id not in self.prompts:
|
||||
available_prompts = list(self.prompts.keys())
|
||||
raise KeyError(
|
||||
f"Prompt '{prompt_id}' not found. Available prompts: {available_prompts}"
|
||||
)
|
||||
|
||||
template = self.prompts[prompt_id]
|
||||
variables = prompt_variables or {}
|
||||
|
||||
# Validate input variables against schema if defined
|
||||
if template.input_schema:
|
||||
self._validate_input(variables, template.input_schema)
|
||||
|
||||
try:
|
||||
# Create Jinja2 template and render
|
||||
jinja_template = self.jinja_env.from_string(template.content)
|
||||
rendered = jinja_template.render(**variables)
|
||||
return rendered
|
||||
except Exception as e:
|
||||
raise ValueError(f"Error rendering template '{prompt_id}': {e}")
|
||||
|
||||
def _validate_input(
|
||||
self, variables: Dict[str, Any], schema: Dict[str, Any]
|
||||
) -> None:
|
||||
"""Basic validation of input variables against schema."""
|
||||
for field_name, field_type in schema.items():
|
||||
if field_name in variables:
|
||||
value = variables[field_name]
|
||||
expected_type = self._get_python_type(field_type)
|
||||
|
||||
if not isinstance(value, expected_type):
|
||||
raise ValueError(
|
||||
f"Invalid type for field '{field_name}': "
|
||||
f"expected {getattr(expected_type, '__name__', str(expected_type))}, got {type(value).__name__}"
|
||||
)
|
||||
|
||||
def _get_python_type(self, schema_type: str) -> Union[type, tuple]:
|
||||
"""Convert schema type string to Python type."""
|
||||
type_mapping: Dict[str, Union[type, tuple]] = {
|
||||
"string": str,
|
||||
"str": str,
|
||||
"number": (int, float),
|
||||
"integer": int,
|
||||
"int": int,
|
||||
"float": float,
|
||||
"boolean": bool,
|
||||
"bool": bool,
|
||||
"array": list,
|
||||
"list": list,
|
||||
"object": dict,
|
||||
"dict": dict,
|
||||
}
|
||||
|
||||
return type_mapping.get(schema_type.lower(), str) # type: ignore
|
||||
|
||||
def get_prompt(self, prompt_id: str) -> Optional[PromptTemplate]:
|
||||
"""Get a prompt template by ID."""
|
||||
return self.prompts.get(prompt_id)
|
||||
|
||||
def list_prompts(self) -> List[str]:
|
||||
"""Get a list of all available prompt IDs."""
|
||||
return list(self.prompts.keys())
|
||||
|
||||
def get_prompt_metadata(self, prompt_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""Get metadata for a specific prompt."""
|
||||
template = self.prompts.get(prompt_id)
|
||||
return template.metadata if template else None
|
||||
|
||||
def reload_prompts(self) -> None:
|
||||
"""Reload all prompts from the directory (if directory was provided)."""
|
||||
self.prompts.clear()
|
||||
if self.prompt_directory:
|
||||
self._load_prompts()
|
||||
|
||||
def add_prompt(
|
||||
self, prompt_id: str, content: str, metadata: Optional[Dict[str, Any]] = None
|
||||
) -> None:
|
||||
"""Add a prompt template programmatically."""
|
||||
template = PromptTemplate(
|
||||
content=content, metadata=metadata or {}, template_id=prompt_id
|
||||
)
|
||||
self.prompts[prompt_id] = template
|
||||
|
||||
def prompt_file_to_json(self, file_path: Union[str, Path]) -> Dict[str, Any]:
|
||||
"""Convert a .prompt file to JSON format.
|
||||
|
||||
Args:
|
||||
file_path: Path to the .prompt file
|
||||
|
||||
Returns:
|
||||
Dictionary with 'content' and 'metadata' keys
|
||||
"""
|
||||
file_path = Path(file_path)
|
||||
content = file_path.read_text(encoding="utf-8")
|
||||
|
||||
# Parse frontmatter and content
|
||||
frontmatter, template_content = self._parse_frontmatter(content)
|
||||
|
||||
return {"content": template_content.strip(), "metadata": frontmatter}
|
||||
|
||||
def json_to_prompt_file(self, prompt_data: Dict[str, Any]) -> str:
|
||||
"""Convert JSON prompt data to .prompt file format.
|
||||
|
||||
Args:
|
||||
prompt_data: Dictionary with 'content' and 'metadata' keys
|
||||
|
||||
Returns:
|
||||
String content in .prompt file format
|
||||
"""
|
||||
content = prompt_data.get("content", "")
|
||||
metadata = prompt_data.get("metadata", {})
|
||||
|
||||
if not metadata:
|
||||
# No metadata, return just the content
|
||||
return content
|
||||
|
||||
# Convert metadata to YAML frontmatter
|
||||
import yaml
|
||||
|
||||
frontmatter_yaml = yaml.dump(metadata, default_flow_style=False)
|
||||
|
||||
return f"---\n{frontmatter_yaml}---\n{content}"
|
||||
|
||||
def get_all_prompts_as_json(self) -> Dict[str, Dict[str, Any]]:
|
||||
"""Get all loaded prompts in JSON format.
|
||||
|
||||
Returns:
|
||||
Dictionary mapping prompt_id to prompt data
|
||||
"""
|
||||
result = {}
|
||||
for prompt_id, template in self.prompts.items():
|
||||
result[prompt_id] = {
|
||||
"content": template.content,
|
||||
"metadata": template.metadata,
|
||||
}
|
||||
return result
|
||||
|
||||
def load_prompts_from_json_data(
|
||||
self, prompt_data: Dict[str, Dict[str, Any]]
|
||||
) -> None:
|
||||
"""Load additional prompts from JSON data (merges with existing prompts)."""
|
||||
self._load_prompts_from_json(prompt_data)
|
||||
|
|
@ -1,10 +1,15 @@
|
|||
import json
|
||||
import threading
|
||||
from typing import Optional
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
else:
|
||||
StandardLoggingPayload = Any
|
||||
|
||||
|
||||
class MlflowLogger(CustomLogger):
|
||||
def __init__(self):
|
||||
|
|
@ -178,7 +183,7 @@ class MlflowLogger(CustomLogger):
|
|||
"call_type": kwargs.get("call_type"),
|
||||
"model": kwargs.get("model"),
|
||||
}
|
||||
standard_obj = kwargs.get("standard_logging_object")
|
||||
standard_obj: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object")
|
||||
if standard_obj:
|
||||
attributes.update(
|
||||
{
|
||||
|
|
@ -192,6 +197,7 @@ class MlflowLogger(CustomLogger):
|
|||
"raw_llm_response": standard_obj.get("response"),
|
||||
"response_cost": standard_obj.get("response_cost"),
|
||||
"saved_cache_cost": standard_obj.get("saved_cache_cost"),
|
||||
"request_tags": standard_obj.get("request_tags"),
|
||||
}
|
||||
)
|
||||
else:
|
||||
|
|
@ -226,6 +232,7 @@ class MlflowLogger(CustomLogger):
|
|||
"""
|
||||
import mlflow
|
||||
|
||||
|
||||
call_type = kwargs.get("call_type", "completion")
|
||||
span_name = f"litellm-{call_type}"
|
||||
span_type = self._get_span_type(call_type)
|
||||
|
|
@ -237,7 +244,7 @@ class MlflowLogger(CustomLogger):
|
|||
if active_span := mlflow.get_current_active_span(): # type: ignore
|
||||
return self._client.start_span(
|
||||
name=span_name,
|
||||
request_id=active_span.request_id,
|
||||
trace_id=active_span.request_id,
|
||||
parent_id=active_span.span_id,
|
||||
span_type=span_type,
|
||||
inputs=inputs,
|
||||
|
|
@ -250,21 +257,24 @@ class MlflowLogger(CustomLogger):
|
|||
span_type=span_type,
|
||||
inputs=inputs,
|
||||
attributes=attributes,
|
||||
tags=self._transform_tag_list_to_dict(attributes.get("request_tags", [])),
|
||||
start_time_ns=start_time_ns,
|
||||
)
|
||||
def _transform_tag_list_to_dict(self, tag_list: list) -> dict:
|
||||
return {tag: "" for tag in tag_list}
|
||||
|
||||
def _end_span_or_trace(self, span, outputs, end_time_ns, status):
|
||||
"""End an MLflow span or a trace."""
|
||||
if span.parent_id is None:
|
||||
self._client.end_trace(
|
||||
request_id=span.request_id,
|
||||
trace_id=span.request_id,
|
||||
outputs=outputs,
|
||||
status=status,
|
||||
end_time_ns=end_time_ns,
|
||||
)
|
||||
else:
|
||||
self._client.end_span(
|
||||
request_id=span.request_id,
|
||||
trace_id=span.request_id,
|
||||
span_id=span.span_id,
|
||||
outputs=outputs,
|
||||
status=status,
|
||||
|
|
|
|||
|
|
@ -54,6 +54,7 @@ class PromptManagementBase(ABC):
|
|||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
) -> PromptManagementClient:
|
||||
|
||||
compiled_prompt_client = self._compile_prompt_helper(
|
||||
prompt_id=prompt_id,
|
||||
prompt_variables=prompt_variables,
|
||||
|
|
@ -91,6 +92,7 @@ class PromptManagementBase(ABC):
|
|||
prompt_label: Optional[str] = None,
|
||||
prompt_version: Optional[int] = None,
|
||||
) -> Tuple[str, List[AllMessageValues], dict]:
|
||||
|
||||
if prompt_id is None:
|
||||
raise ValueError("prompt_id is required for Prompt Management Base class")
|
||||
if not self.should_run_prompt_management(
|
||||
|
|
|
|||
|
|
@ -18,24 +18,22 @@ else:
|
|||
|
||||
|
||||
def safe_divide_seconds(
|
||||
seconds: float,
|
||||
denominator: float,
|
||||
default: Optional[float] = None
|
||||
seconds: float, denominator: float, default: Optional[float] = None
|
||||
) -> Optional[float]:
|
||||
"""
|
||||
Safely divide seconds by denominator, handling zero division.
|
||||
|
||||
|
||||
Args:
|
||||
seconds: Time duration in seconds
|
||||
denominator: The divisor (e.g., number of tokens)
|
||||
default: Value to return if division by zero (defaults to None)
|
||||
|
||||
|
||||
Returns:
|
||||
The result of the division as a float (seconds per unit), or default if denominator is zero
|
||||
"""
|
||||
if denominator <= 0:
|
||||
return default
|
||||
|
||||
|
||||
return float(seconds / denominator)
|
||||
|
||||
|
||||
|
|
@ -203,3 +201,50 @@ def preserve_upstream_non_openai_attributes(
|
|||
for key, value in original_chunk.model_dump().items():
|
||||
if key not in expected_keys:
|
||||
setattr(model_response, key, value)
|
||||
|
||||
|
||||
def safe_deep_copy(data):
|
||||
"""
|
||||
Safe Deep Copy
|
||||
|
||||
The LiteLLM Request has some object that can-not be pickled / deep copied
|
||||
|
||||
Use this function to safely deep copy the LiteLLM Request
|
||||
"""
|
||||
import copy
|
||||
|
||||
import litellm
|
||||
|
||||
if litellm.safe_memory_mode is True:
|
||||
return data
|
||||
|
||||
litellm_parent_otel_span: Optional[Any] = None
|
||||
# Step 1: Remove the litellm_parent_otel_span
|
||||
litellm_parent_otel_span = None
|
||||
if isinstance(data, dict):
|
||||
# remove litellm_parent_otel_span since this is not picklable
|
||||
if "metadata" in data and "litellm_parent_otel_span" in data["metadata"]:
|
||||
litellm_parent_otel_span = data["metadata"].pop("litellm_parent_otel_span")
|
||||
data["metadata"]["litellm_parent_otel_span"] = "placeholder"
|
||||
if (
|
||||
"litellm_metadata" in data
|
||||
and "litellm_parent_otel_span" in data["litellm_metadata"]
|
||||
):
|
||||
litellm_parent_otel_span = data["litellm_metadata"].pop(
|
||||
"litellm_parent_otel_span"
|
||||
)
|
||||
data["litellm_metadata"]["litellm_parent_otel_span"] = "placeholder"
|
||||
new_data = copy.deepcopy(data)
|
||||
|
||||
# Step 2: re-add the litellm_parent_otel_span after doing a deep copy
|
||||
if isinstance(data, dict) and litellm_parent_otel_span is not None:
|
||||
if "metadata" in data and "litellm_parent_otel_span" in data["metadata"]:
|
||||
data["metadata"]["litellm_parent_otel_span"] = litellm_parent_otel_span
|
||||
if (
|
||||
"litellm_metadata" in data
|
||||
and "litellm_parent_otel_span" in data["litellm_metadata"]
|
||||
):
|
||||
data["litellm_metadata"][
|
||||
"litellm_parent_otel_span"
|
||||
] = litellm_parent_otel_span
|
||||
return new_data
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ Example:
|
|||
"datadog" -> DataDogLogger
|
||||
"prometheus" -> PrometheusLogger
|
||||
"""
|
||||
|
||||
from typing import Union
|
||||
|
||||
from litellm.integrations.agentops import AgentOps
|
||||
|
|
@ -31,10 +32,12 @@ from litellm.integrations.mlflow import MlflowLogger
|
|||
from litellm.integrations.openmeter import OpenMeterLogger
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry
|
||||
from litellm.integrations.opik.opik import OpikLogger
|
||||
|
||||
try:
|
||||
from litellm_enterprise.integrations.prometheus import PrometheusLogger
|
||||
except Exception:
|
||||
PrometheusLogger = None
|
||||
from litellm.integrations.dotprompt import DotpromptManager
|
||||
from litellm.integrations.s3_v2 import S3Logger
|
||||
from litellm.integrations.sqs import SQSLogger
|
||||
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
|
||||
|
|
@ -47,6 +50,7 @@ class CustomLoggerRegistry:
|
|||
"""
|
||||
Registry mapping the callback class string to the class type.
|
||||
"""
|
||||
|
||||
CALLBACK_CLASS_STR_TO_CLASS_TYPE = {
|
||||
"lago": LagoLogger,
|
||||
"openmeter": OpenMeterLogger,
|
||||
|
|
@ -80,6 +84,7 @@ class CustomLoggerRegistry:
|
|||
"aws_sqs": SQSLogger,
|
||||
"dynamic_rate_limiter": _PROXY_DynamicRateLimitHandler,
|
||||
"vector_store_pre_call_hook": VectorStorePreCallHook,
|
||||
"dotprompt": DotpromptManager,
|
||||
}
|
||||
|
||||
try:
|
||||
|
|
@ -110,14 +115,17 @@ class CustomLoggerRegistry:
|
|||
def get_callback_str_from_class_type(cls, class_type: type) -> Union[str, None]:
|
||||
"""
|
||||
Get the callback string from the class type.
|
||||
|
||||
|
||||
Args:
|
||||
class_type: The class type to find the string for
|
||||
|
||||
|
||||
Returns:
|
||||
str: The callback string, or None if not found
|
||||
"""
|
||||
for callback_str, callback_class in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items():
|
||||
for (
|
||||
callback_str,
|
||||
callback_class,
|
||||
) in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items():
|
||||
if callback_class == class_type:
|
||||
return callback_str
|
||||
return None
|
||||
|
|
@ -127,15 +135,18 @@ class CustomLoggerRegistry:
|
|||
"""
|
||||
Get all callback strings that map to the same class type.
|
||||
Some class types (like OpenTelemetry) have multiple string mappings.
|
||||
|
||||
|
||||
Args:
|
||||
class_type: The class type to find all strings for
|
||||
|
||||
|
||||
Returns:
|
||||
list: List of callback strings that map to the class type
|
||||
"""
|
||||
callback_strs: list[str] = []
|
||||
for callback_str, callback_class in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items():
|
||||
for (
|
||||
callback_str,
|
||||
callback_class,
|
||||
) in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items():
|
||||
if callback_class == class_type:
|
||||
callback_strs.append(callback_str)
|
||||
return callback_strs
|
||||
return callback_strs
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
import uuid
|
||||
from copy import deepcopy
|
||||
from typing import Optional
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.core_helpers import safe_deep_copy
|
||||
|
||||
from .asyncify import run_async_function
|
||||
|
||||
|
|
@ -41,7 +41,7 @@ async def async_completion_with_fallbacks(**kwargs):
|
|||
most_recent_exception_str: Optional[str] = None
|
||||
for fallback in fallbacks:
|
||||
try:
|
||||
completion_kwargs = deepcopy(base_kwargs)
|
||||
completion_kwargs = safe_deep_copy(base_kwargs)
|
||||
# Handle dictionary fallback configurations
|
||||
if isinstance(fallback, dict):
|
||||
model = fallback.pop("model", original_model)
|
||||
|
|
|
|||
|
|
@ -120,6 +120,7 @@ from ..integrations.azure_storage.azure_storage import AzureBlobStorageLogger
|
|||
from ..integrations.custom_prompt_management import CustomPromptManagement
|
||||
from ..integrations.datadog.datadog import DataDogLogger
|
||||
from ..integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
|
||||
from ..integrations.dotprompt import DotpromptManager
|
||||
from ..integrations.dynamodb import DyanmoDBLogger
|
||||
from ..integrations.galileo import GalileoObserve
|
||||
from ..integrations.gcs_bucket.gcs_bucket import GCSBucketLogger
|
||||
|
|
@ -167,11 +168,10 @@ try:
|
|||
from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import (
|
||||
SMTPEmailLogger,
|
||||
)
|
||||
from litellm_enterprise.integrations.prometheus import PrometheusLogger
|
||||
from litellm_enterprise.litellm_core_utils.litellm_logging import (
|
||||
StandardLoggingPayloadSetup as EnterpriseStandardLoggingPayloadSetup,
|
||||
)
|
||||
from litellm_enterprise.integrations.prometheus import PrometheusLogger
|
||||
|
||||
|
||||
EnterpriseStandardLoggingPayloadSetupVAR: Optional[
|
||||
Type[EnterpriseStandardLoggingPayloadSetup]
|
||||
|
|
@ -504,6 +504,15 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if "custom_llm_provider" in self.model_call_details:
|
||||
self.custom_llm_provider = self.model_call_details["custom_llm_provider"]
|
||||
|
||||
def update_messages(self, messages: List[AllMessageValues]):
|
||||
"""
|
||||
Update the logged value of the messages in the model_call_details
|
||||
|
||||
Allows pre-call hooks to update the messages before the call is made
|
||||
"""
|
||||
self.messages = messages
|
||||
self.model_call_details["messages"] = messages
|
||||
|
||||
def should_run_prompt_management_hooks(
|
||||
self,
|
||||
non_default_params: Dict,
|
||||
|
|
@ -599,9 +608,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
custom_logger = (
|
||||
prompt_management_logger
|
||||
or self.get_custom_logger_for_prompt_management(
|
||||
model=model,
|
||||
tools=tools,
|
||||
non_default_params=non_default_params
|
||||
model=model, tools=tools, non_default_params=non_default_params
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -673,16 +680,16 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
# Vector Store / Knowledge Base hooks
|
||||
#########################################################
|
||||
if litellm.vector_store_registry is not None:
|
||||
|
||||
vector_store_custom_logger = _init_custom_logger_compatible_class(
|
||||
logging_integration="vector_store_pre_call_hook",
|
||||
internal_usage_cache=None,
|
||||
llm_router=None,
|
||||
)
|
||||
self.model_call_details["prompt_integration"] = (
|
||||
vector_store_custom_logger.__class__.__name__
|
||||
)
|
||||
return vector_store_custom_logger
|
||||
|
||||
vector_store_custom_logger = _init_custom_logger_compatible_class(
|
||||
logging_integration="vector_store_pre_call_hook",
|
||||
internal_usage_cache=None,
|
||||
llm_router=None,
|
||||
)
|
||||
self.model_call_details["prompt_integration"] = (
|
||||
vector_store_custom_logger.__class__.__name__
|
||||
)
|
||||
return vector_store_custom_logger
|
||||
|
||||
return None
|
||||
|
||||
|
|
@ -945,7 +952,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if additional_args.get("request_str", None) is not None:
|
||||
# print the sagemaker / bedrock client request
|
||||
curl_command = "\nRequest Sent from LiteLLM:\n"
|
||||
curl_command += additional_args.get("request_str", None)
|
||||
request_str = additional_args.get("request_str", "")
|
||||
curl_command += request_str
|
||||
elif api_base == "":
|
||||
curl_command = str(self.model_call_details)
|
||||
return curl_command
|
||||
|
|
@ -1314,9 +1322,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if (
|
||||
EnterpriseCallbackControls is not None
|
||||
and EnterpriseCallbackControls.is_callback_disabled_dynamically(
|
||||
callback=callback,
|
||||
callback=callback,
|
||||
litellm_params=litellm_params,
|
||||
standard_callback_dynamic_params = self.standard_callback_dynamic_params
|
||||
standard_callback_dynamic_params=self.standard_callback_dynamic_params,
|
||||
)
|
||||
):
|
||||
verbose_logger.debug(
|
||||
|
|
@ -2265,15 +2273,20 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
if isinstance(callback, CustomLogger): # custom logger class
|
||||
model_call_details: Dict = self.model_call_details
|
||||
##################################
|
||||
# call redaction hook for custom logger
|
||||
model_call_details = callback.redact_standard_logging_payload_from_model_call_details(
|
||||
model_call_details=model_call_details
|
||||
)
|
||||
##################################
|
||||
if self.stream is True:
|
||||
if (
|
||||
"async_complete_streaming_response"
|
||||
in self.model_call_details
|
||||
):
|
||||
if "async_complete_streaming_response" in model_call_details:
|
||||
await callback.async_log_success_event(
|
||||
kwargs=self.model_call_details,
|
||||
response_obj=self.model_call_details[
|
||||
kwargs=model_call_details,
|
||||
response_obj=model_call_details[
|
||||
"async_complete_streaming_response"
|
||||
],
|
||||
start_time=start_time,
|
||||
|
|
@ -2281,14 +2294,14 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
else:
|
||||
await callback.async_log_stream_event( # [TODO]: move this to being an async log stream event function
|
||||
kwargs=self.model_call_details,
|
||||
kwargs=model_call_details,
|
||||
response_obj=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
else:
|
||||
await callback.async_log_success_event(
|
||||
kwargs=self.model_call_details,
|
||||
kwargs=model_call_details,
|
||||
response_obj=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
|
|
@ -3209,6 +3222,8 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
|
|||
_in_memory_loggers.append(_literalai_logger)
|
||||
return _literalai_logger # type: ignore
|
||||
elif logging_integration == "prometheus":
|
||||
if PrometheusLogger is None:
|
||||
raise ValueError("PrometheusLogger is not initialized")
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, PrometheusLogger):
|
||||
return callback # type: ignore
|
||||
|
|
@ -3483,7 +3498,7 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
|
|||
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
|
||||
VectorStorePreCallHook,
|
||||
)
|
||||
|
||||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, VectorStorePreCallHook):
|
||||
return callback
|
||||
|
|
@ -3526,11 +3541,21 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
|
|||
humanloop_logger = HumanloopLogger()
|
||||
_in_memory_loggers.append(humanloop_logger)
|
||||
return humanloop_logger # type: ignore
|
||||
elif logging_integration == "dotprompt":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, DotpromptManager):
|
||||
return callback
|
||||
|
||||
dotprompt_logger = DotpromptManager()
|
||||
_in_memory_loggers.append(dotprompt_logger)
|
||||
return dotprompt_logger # type: ignore
|
||||
return None
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"[Non-Blocking Error] Error initializing custom logger: {e}"
|
||||
)
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def get_custom_logger_compatible_class( # noqa: PLR0915
|
||||
|
|
@ -3571,7 +3596,7 @@ def get_custom_logger_compatible_class( # noqa: PLR0915
|
|||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, LiteralAILogger):
|
||||
return callback
|
||||
elif logging_integration == "prometheus":
|
||||
elif logging_integration == "prometheus" and PrometheusLogger is not None:
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, PrometheusLogger):
|
||||
return callback
|
||||
|
|
@ -3674,7 +3699,7 @@ def get_custom_logger_compatible_class( # noqa: PLR0915
|
|||
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
|
||||
VectorStorePreCallHook,
|
||||
)
|
||||
|
||||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, VectorStorePreCallHook):
|
||||
return callback
|
||||
|
|
|
|||
|
|
@ -822,3 +822,41 @@ def set_last_user_message(
|
|||
messages.reverse()
|
||||
messages.append({"role": "user", "content": content})
|
||||
return messages
|
||||
|
||||
|
||||
def convert_prefix_message_to_non_prefix_messages(
|
||||
messages: List[AllMessageValues],
|
||||
) -> List[AllMessageValues]:
|
||||
"""
|
||||
For models that don't support {prefix: true} in messages, we need to convert the prefix message to a non-prefix message.
|
||||
|
||||
Use prompt:
|
||||
|
||||
{"role": "assistant", "content": "value", "prefix": true} -> [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are a helpful assistant. You are given a message and you need to respond to it. You are also given a generated content. You need to respond to the message in continuation of the generated content. Do not repeat the same content. Your response should be in continuation of this text: ",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": message["content"],
|
||||
},
|
||||
]
|
||||
|
||||
do this in place
|
||||
"""
|
||||
new_messages: List[AllMessageValues] = []
|
||||
for message in messages:
|
||||
if message.get("prefix"):
|
||||
new_messages.append(
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are a helpful assistant. You are given a message and you need to respond to it. You are also given a generated content. You need to respond to the message in continuation of the generated content. Do not repeat the same content. Your response should be in continuation of this text: ",
|
||||
}
|
||||
)
|
||||
new_messages.append(
|
||||
{**{k: v for k, v in message.items() if k != "prefix"}} # type: ignore
|
||||
)
|
||||
else:
|
||||
new_messages.append(message)
|
||||
return new_messages
|
||||
|
|
|
|||
|
|
@ -1121,13 +1121,14 @@ def convert_to_gemini_tool_call_result(
|
|||
}
|
||||
"""
|
||||
content_str: str = ""
|
||||
if isinstance(message["content"], str):
|
||||
content_str = message["content"]
|
||||
elif isinstance(message["content"], List):
|
||||
content_list = message["content"]
|
||||
for content in content_list:
|
||||
if content["type"] == "text":
|
||||
content_str += content["text"]
|
||||
if "content" in message:
|
||||
if isinstance(message["content"], str):
|
||||
content_str = message["content"]
|
||||
elif isinstance(message["content"], List):
|
||||
content_list = message["content"]
|
||||
for content in content_list:
|
||||
if content["type"] == "text":
|
||||
content_str += content["text"]
|
||||
name: Optional[str] = message.get("name", "") # type: ignore
|
||||
|
||||
# Recover name from last message with tool calls
|
||||
|
|
|
|||
|
|
@ -940,8 +940,8 @@ class CustomStreamWrapper:
|
|||
and not self.sent_last_thinking_block
|
||||
and model_response.choices[0].delta.content
|
||||
):
|
||||
model_response.choices[0].delta.content = (
|
||||
"</think>" + (model_response.choices[0].delta.content or "")
|
||||
model_response.choices[0].delta.content = "</think>" + (
|
||||
model_response.choices[0].delta.content or ""
|
||||
)
|
||||
self.sent_last_thinking_block = True
|
||||
|
||||
|
|
@ -1841,13 +1841,25 @@ class CustomStreamWrapper:
|
|||
self.logging_obj.async_failure_handler(e, traceback_exception) # type: ignore
|
||||
)
|
||||
## Map to OpenAI Exception
|
||||
raise exception_type(
|
||||
model=self.model,
|
||||
custom_llm_provider=self.custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs={},
|
||||
extra_kwargs={},
|
||||
)
|
||||
try:
|
||||
exception_type(
|
||||
model=self.model,
|
||||
custom_llm_provider=self.custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs={},
|
||||
extra_kwargs={},
|
||||
)
|
||||
except Exception as e:
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
|
||||
raise MidStreamFallbackError(
|
||||
message=str(e),
|
||||
model=self.model,
|
||||
llm_provider=self.custom_llm_provider or "anthropic",
|
||||
original_exception=e,
|
||||
generated_content=self.response_uptil_now,
|
||||
is_pre_first_chunk=not self.sent_first_chunk,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _strip_sse_data_from_chunk(chunk: Optional[str]) -> Optional[str]:
|
||||
|
|
|
|||
|
|
@ -462,9 +462,8 @@ def _count_messages(
|
|||
default_token_count,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported type {type(value)} for key {key} in message {message}"
|
||||
)
|
||||
# Skip unsupported keys instead of raising an error
|
||||
continue
|
||||
return num_tokens
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -640,7 +640,8 @@ class ModelResponseIterator:
|
|||
]
|
||||
] = None
|
||||
|
||||
index = int(chunk.get("index", 0))
|
||||
# Always use index=0 for OpenAI choice format (fixes multi-choice errors)
|
||||
index = 0
|
||||
if type_chunk == "content_block_delta":
|
||||
"""
|
||||
Anthropic content chunk
|
||||
|
|
|
|||
|
|
@ -159,9 +159,16 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
supported_openai_params = self.get_supported_openai_params(model)
|
||||
|
||||
api_version_times = api_version.split("-")
|
||||
api_version_year = api_version_times[0]
|
||||
api_version_month = api_version_times[1]
|
||||
api_version_day = api_version_times[2]
|
||||
|
||||
if len(api_version_times) >= 3:
|
||||
api_version_year = api_version_times[0]
|
||||
api_version_month = api_version_times[1]
|
||||
api_version_day = api_version_times[2]
|
||||
else:
|
||||
api_version_year = None
|
||||
api_version_month = None
|
||||
api_version_day = None
|
||||
|
||||
for param, value in non_default_params.items():
|
||||
if param == "tool_choice":
|
||||
"""
|
||||
|
|
@ -171,47 +178,57 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
"""
|
||||
## check if api version supports this param ##
|
||||
if (
|
||||
api_version_year < "2023"
|
||||
or (api_version_year == "2023" and api_version_month < "12")
|
||||
or (
|
||||
api_version_year == "2023"
|
||||
and api_version_month == "12"
|
||||
and api_version_day < "01"
|
||||
)
|
||||
api_version_year is None
|
||||
or api_version_month is None
|
||||
or api_version_day is None
|
||||
):
|
||||
if litellm.drop_params is True or (
|
||||
drop_params is not None and drop_params is True
|
||||
):
|
||||
pass
|
||||
else:
|
||||
raise UnsupportedParamsError(
|
||||
status_code=400,
|
||||
message=f"""Azure does not support 'tool_choice', for api_version={api_version}. Bump your API version to '2023-12-01-preview' or later. This parameter requires 'api_version="2023-12-01-preview"' or later. Azure API Reference: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#chat-completions""",
|
||||
)
|
||||
elif value == "required" and (
|
||||
api_version_year == "2024" and api_version_month <= "05"
|
||||
): ## check if tool_choice value is supported ##
|
||||
if litellm.drop_params is True or (
|
||||
drop_params is not None and drop_params is True
|
||||
):
|
||||
pass
|
||||
else:
|
||||
raise UnsupportedParamsError(
|
||||
status_code=400,
|
||||
message=f"Azure does not support '{value}' as a {param} param, for api_version={api_version}. To drop 'tool_choice=required' for calls with this Azure API version, set `litellm.drop_params=True` or for proxy:\n\n`litellm_settings:\n drop_params: true`\nAzure API Reference: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#chat-completions",
|
||||
)
|
||||
else:
|
||||
optional_params["tool_choice"] = value
|
||||
else:
|
||||
if (
|
||||
api_version_year < "2023"
|
||||
or (api_version_year == "2023" and api_version_month < "12")
|
||||
or (
|
||||
api_version_year == "2023"
|
||||
and api_version_month == "12"
|
||||
and api_version_day < "01"
|
||||
)
|
||||
):
|
||||
if litellm.drop_params is True or (
|
||||
drop_params is not None and drop_params is True
|
||||
):
|
||||
pass
|
||||
else:
|
||||
raise UnsupportedParamsError(
|
||||
status_code=400,
|
||||
message=f"""Azure does not support 'tool_choice', for api_version={api_version}. Bump your API version to '2023-12-01-preview' or later. This parameter requires 'api_version="2023-12-01-preview"' or later. Azure API Reference: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#chat-completions""",
|
||||
)
|
||||
elif value == "required" and (
|
||||
api_version_year == "2024" and api_version_month <= "05"
|
||||
): ## check if tool_choice value is supported ##
|
||||
if litellm.drop_params is True or (
|
||||
drop_params is not None and drop_params is True
|
||||
):
|
||||
pass
|
||||
else:
|
||||
raise UnsupportedParamsError(
|
||||
status_code=400,
|
||||
message=f"Azure does not support '{value}' as a {param} param, for api_version={api_version}. To drop 'tool_choice=required' for calls with this Azure API version, set `litellm.drop_params=True` or for proxy:\n\n`litellm_settings:\n drop_params: true`\nAzure API Reference: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#chat-completions",
|
||||
)
|
||||
else:
|
||||
optional_params["tool_choice"] = value
|
||||
elif param == "response_format" and isinstance(value, dict):
|
||||
_is_response_format_supported_model = (
|
||||
self._is_response_format_supported_model(model)
|
||||
)
|
||||
|
||||
is_response_format_supported_api_version = (
|
||||
self._is_response_format_supported_api_version(
|
||||
api_version_year, api_version_month
|
||||
if api_version_year is None or api_version_month is None:
|
||||
is_response_format_supported_api_version = True
|
||||
else:
|
||||
is_response_format_supported_api_version = (
|
||||
self._is_response_format_supported_api_version(
|
||||
api_version_year, api_version_month
|
||||
)
|
||||
)
|
||||
)
|
||||
is_response_format_supported = (
|
||||
is_response_format_supported_api_version
|
||||
and _is_response_format_supported_model
|
||||
|
|
|
|||
|
|
@ -10,12 +10,14 @@ if TYPE_CHECKING:
|
|||
GenerateContentConfigDict,
|
||||
GenerateContentContentListUnionDict,
|
||||
GenerateContentResponse,
|
||||
ToolConfigDict,
|
||||
)
|
||||
else:
|
||||
GenerateContentConfigDict = Any
|
||||
GenerateContentContentListUnionDict = Any
|
||||
GenerateContentResponse = Any
|
||||
LiteLLMLoggingObj = Any
|
||||
ToolConfigDict = Any
|
||||
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
|
|
@ -145,6 +147,7 @@ class BaseGoogleGenAIGenerateContentConfig(ABC):
|
|||
self,
|
||||
model: str,
|
||||
contents: GenerateContentContentListUnionDict,
|
||||
tools: Optional[ToolConfigDict],
|
||||
generate_content_config_dict: Dict,
|
||||
) -> dict:
|
||||
"""
|
||||
|
|
@ -153,6 +156,7 @@ class BaseGoogleGenAIGenerateContentConfig(ABC):
|
|||
Args:
|
||||
model: The model name
|
||||
contents: Input contents
|
||||
tools: Tools
|
||||
generate_content_request_params: Request parameters
|
||||
litellm_params: LiteLLM parameters
|
||||
headers: Request headers
|
||||
|
|
|
|||
|
|
@ -49,6 +49,14 @@ from litellm.utils import add_dummy_tool, has_tool_call_blocks, supports_reasoni
|
|||
|
||||
from ..common_utils import BedrockError, BedrockModelInfo, get_bedrock_tool_name
|
||||
|
||||
# Computer use tool prefixes supported by Bedrock
|
||||
BEDROCK_COMPUTER_USE_TOOLS = [
|
||||
"computer_use_preview",
|
||||
"computer_",
|
||||
"bash_",
|
||||
"text_editor_"
|
||||
]
|
||||
|
||||
|
||||
class AmazonConverseConfig(BaseConfig):
|
||||
"""
|
||||
|
|
@ -218,6 +226,98 @@ class AmazonConverseConfig(BaseConfig):
|
|||
+ self.get_supported_video_types()
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
for tool in tools:
|
||||
if "type" in tool:
|
||||
tool_type = tool["type"]
|
||||
for computer_use_prefix in BEDROCK_COMPUTER_USE_TOOLS:
|
||||
if tool_type.startswith(computer_use_prefix):
|
||||
return True
|
||||
return False
|
||||
|
||||
def _transform_computer_use_tools(
|
||||
self, computer_use_tools: List[OpenAIChatCompletionToolParam]
|
||||
) -> List[dict]:
|
||||
"""Transform computer use tools to Bedrock format."""
|
||||
transformed_tools: List[dict] = []
|
||||
|
||||
for tool in computer_use_tools:
|
||||
tool_type = tool.get("type", "")
|
||||
|
||||
# Check if this is a computer use tool with the startswith method
|
||||
is_computer_use_tool = False
|
||||
for computer_use_prefix in BEDROCK_COMPUTER_USE_TOOLS:
|
||||
if tool_type.startswith(computer_use_prefix):
|
||||
is_computer_use_tool = True
|
||||
break
|
||||
|
||||
transformed_tool: dict = {}
|
||||
if is_computer_use_tool:
|
||||
if tool_type.startswith("computer_") and "function" in tool:
|
||||
# Computer use tool with function format
|
||||
func = tool["function"]
|
||||
transformed_tool = {
|
||||
"type": tool_type,
|
||||
"name": func.get("name", "computer"),
|
||||
**func.get("parameters", {})
|
||||
}
|
||||
else:
|
||||
# Direct tools - just need to ensure name is present
|
||||
transformed_tool = dict(tool)
|
||||
if "name" not in transformed_tool:
|
||||
if tool_type.startswith("bash_"):
|
||||
transformed_tool["name"] = "bash"
|
||||
elif tool_type.startswith("text_editor_"):
|
||||
transformed_tool["name"] = "str_replace_editor"
|
||||
else:
|
||||
# Pass through other tools as-is
|
||||
transformed_tool = dict(tool)
|
||||
|
||||
transformed_tools.append(transformed_tool)
|
||||
|
||||
return transformed_tools
|
||||
|
||||
def _separate_computer_use_tools(
|
||||
self, tools: List[OpenAIChatCompletionToolParam], model: str
|
||||
) -> Tuple[List[OpenAIChatCompletionToolParam], List[OpenAIChatCompletionToolParam]]:
|
||||
"""
|
||||
Separate computer use tools from regular function tools.
|
||||
|
||||
Args:
|
||||
tools: List of tools to separate
|
||||
model: The model name to check if it supports computer use
|
||||
|
||||
Returns:
|
||||
Tuple of (computer_use_tools, regular_tools)
|
||||
"""
|
||||
computer_use_tools = []
|
||||
regular_tools = []
|
||||
|
||||
for tool in tools:
|
||||
if "type" in tool:
|
||||
tool_type = tool["type"]
|
||||
is_computer_use_tool = False
|
||||
for computer_use_prefix in BEDROCK_COMPUTER_USE_TOOLS:
|
||||
if tool_type.startswith(computer_use_prefix):
|
||||
is_computer_use_tool = True
|
||||
break
|
||||
if is_computer_use_tool:
|
||||
computer_use_tools.append(tool)
|
||||
else:
|
||||
regular_tools.append(tool)
|
||||
else:
|
||||
regular_tools.append(tool)
|
||||
|
||||
return computer_use_tools, regular_tools
|
||||
|
||||
|
||||
|
||||
def _create_json_tool_call_for_response_format(
|
||||
self,
|
||||
json_schema: Optional[dict] = None,
|
||||
|
|
@ -546,9 +646,31 @@ class AmazonConverseConfig(BaseConfig):
|
|||
self._handle_top_k_value(model, inference_params)
|
||||
)
|
||||
|
||||
bedrock_tools: List[ToolBlock] = _bedrock_tools_pt(
|
||||
inference_params.pop("tools", [])
|
||||
)
|
||||
original_tools = inference_params.pop("tools", [])
|
||||
|
||||
# Initialize bedrock_tools
|
||||
bedrock_tools: List[ToolBlock] = []
|
||||
|
||||
# Only separate tools if computer use tools are actually present
|
||||
if original_tools and self.is_computer_use_tool_used(original_tools, model):
|
||||
# Separate computer use tools from regular function tools
|
||||
computer_use_tools, regular_tools = self._separate_computer_use_tools(
|
||||
original_tools, model
|
||||
)
|
||||
|
||||
# Process regular function tools using existing logic
|
||||
bedrock_tools = _bedrock_tools_pt(regular_tools)
|
||||
|
||||
# Add computer use tools and anthropic_beta if needed (only when computer use tools are present)
|
||||
if computer_use_tools:
|
||||
additional_request_params["anthropic_beta"] = ["computer-use-2024-10-22"]
|
||||
# Transform computer use tools to proper Bedrock format
|
||||
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
|
||||
bedrock_tools = _bedrock_tools_pt(original_tools)
|
||||
|
||||
bedrock_tool_config: Optional[ToolConfigBlock] = None
|
||||
if len(bedrock_tools) > 0:
|
||||
tool_choice_values: ToolChoiceValuesBlock = inference_params.pop(
|
||||
|
|
|
|||
|
|
@ -1384,7 +1384,9 @@ class AWSEventStreamDecoder:
|
|||
"name": None,
|
||||
"arguments": "{}",
|
||||
},
|
||||
"index": chunk_data["contentBlockIndex"],
|
||||
"index": self.tool_calls_index
|
||||
if self.tool_calls_index is not None
|
||||
else index,
|
||||
}
|
||||
elif "stopReason" in chunk_data:
|
||||
finish_reason = map_finish_reason(chunk_data.get("stopReason", "stop"))
|
||||
|
|
|
|||
|
|
@ -3196,6 +3196,7 @@ class BaseLLMHTTPHandler:
|
|||
contents: Any,
|
||||
generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig,
|
||||
generate_content_config_dict: Dict,
|
||||
tools: Any,
|
||||
custom_llm_provider: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
|
|
@ -3221,6 +3222,7 @@ class BaseLLMHTTPHandler:
|
|||
contents=contents,
|
||||
generate_content_provider_config=generate_content_provider_config,
|
||||
generate_content_config_dict=generate_content_config_dict,
|
||||
tools=tools,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -3256,6 +3258,7 @@ class BaseLLMHTTPHandler:
|
|||
data = generate_content_provider_config.transform_generate_content_request(
|
||||
model=model,
|
||||
contents=contents,
|
||||
tools=tools,
|
||||
generate_content_config_dict=generate_content_config_dict,
|
||||
)
|
||||
|
||||
|
|
@ -3317,6 +3320,7 @@ class BaseLLMHTTPHandler:
|
|||
contents: Any,
|
||||
generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig,
|
||||
generate_content_config_dict: Dict,
|
||||
tools: Any,
|
||||
custom_llm_provider: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
|
|
@ -3360,6 +3364,7 @@ class BaseLLMHTTPHandler:
|
|||
data = generate_content_provider_config.transform_generate_content_request(
|
||||
model=model,
|
||||
contents=contents,
|
||||
tools=tools,
|
||||
generate_content_config_dict=generate_content_config_dict,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -44,7 +44,7 @@ class GeminiModelInfo(BaseLLMModelInfo):
|
|||
|
||||
@staticmethod
|
||||
def get_api_key(api_key: Optional[str] = None) -> Optional[str]:
|
||||
return api_key or (get_secret_str("GEMINI_API_KEY"))
|
||||
return api_key or (get_secret_str("GOOGLE_API_KEY")) or (get_secret_str("GEMINI_API_KEY"))
|
||||
|
||||
@staticmethod
|
||||
def get_base_model(model: str) -> Optional[str]:
|
||||
|
|
@ -66,7 +66,7 @@ class GeminiModelInfo(BaseLLMModelInfo):
|
|||
endpoint = f"/{self.api_version}/models"
|
||||
if api_base is None or api_key is None:
|
||||
raise ValueError(
|
||||
"GEMINI_API_BASE or GEMINI_API_KEY is not set. Please set the environment variable, to query Gemini's `/models` endpoint."
|
||||
"GEMINI_API_BASE or GEMINI_API_KEY/GOOGLE_API_KEY is not set. Please set the environment variable, to query Gemini's `/models` endpoint."
|
||||
)
|
||||
|
||||
response = litellm.module_level_client.get(
|
||||
|
|
@ -133,3 +133,7 @@ def encode_unserializable_types(
|
|||
else:
|
||||
processed_data[key] = value
|
||||
return processed_data
|
||||
|
||||
|
||||
def get_api_key_from_env() -> Optional[str]:
|
||||
return get_secret_str("GOOGLE_API_KEY") or get_secret_str("GEMINI_API_KEY")
|
||||
|
|
|
|||
|
|
@ -11,7 +11,6 @@ from litellm.llms.base_llm.google_genai.transformation import (
|
|||
BaseGoogleGenAIGenerateContentConfig,
|
||||
)
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -19,17 +18,26 @@ if TYPE_CHECKING:
|
|||
GenerateContentConfigDict,
|
||||
GenerateContentContentListUnionDict,
|
||||
GenerateContentResponse,
|
||||
ToolConfigDict,
|
||||
)
|
||||
else:
|
||||
GenerateContentConfigDict = Any
|
||||
GenerateContentContentListUnionDict = Any
|
||||
GenerateContentResponse = Any
|
||||
|
||||
ToolConfigDict = Any
|
||||
|
||||
from ..common_utils import get_api_key_from_env
|
||||
|
||||
class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
|
||||
"""
|
||||
Configuration for calling Google models in their native format.
|
||||
"""
|
||||
##############################
|
||||
# Constants
|
||||
##############################
|
||||
XGOOGLE_API_KEY = "x-goog-api-key"
|
||||
##############################
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Literal["gemini", "vertex_ai"]:
|
||||
return "gemini"
|
||||
|
|
@ -113,8 +121,9 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
|
|||
default_headers = {
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
if api_key is not None:
|
||||
default_headers["Authorization"] = f"Bearer {api_key}"
|
||||
gemini_api_key = self._get_google_ai_studio_api_key(dict(litellm_params or {}))
|
||||
if gemini_api_key is not None:
|
||||
default_headers[self.XGOOGLE_API_KEY] = gemini_api_key
|
||||
if headers is not None:
|
||||
default_headers.update(headers)
|
||||
|
||||
|
|
@ -124,7 +133,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
|
|||
return (
|
||||
litellm_params.pop("api_key", None)
|
||||
or litellm_params.pop("gemini_api_key", None)
|
||||
or get_secret_str("GEMINI_API_KEY")
|
||||
or get_api_key_from_env()
|
||||
or litellm.api_key
|
||||
)
|
||||
|
||||
|
|
@ -252,6 +261,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
|
|||
self,
|
||||
model: str,
|
||||
contents: GenerateContentContentListUnionDict,
|
||||
tools: Optional[ToolConfigDict],
|
||||
generate_content_config_dict: Dict,
|
||||
) -> dict:
|
||||
from litellm.types.google_genai.main import (
|
||||
|
|
@ -261,6 +271,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
|
|||
typed_generate_content_request = GenerateContentRequestDict(
|
||||
model=model,
|
||||
contents=contents,
|
||||
tools=tools,
|
||||
generationConfig=GenerateContentConfigDict(**generate_content_config_dict),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -190,11 +190,9 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
|
|||
predictions = response_data.get("predictions", [])
|
||||
for prediction in predictions:
|
||||
# Google AI returns base64 encoded images in the prediction
|
||||
generated_images = prediction.get("generatedImages", [])
|
||||
for image_data in generated_images:
|
||||
model_response.data.append(ImageObject(
|
||||
b64_json=image_data.get("bytesBase64Encoded", None),
|
||||
url=None, # Google AI returns base64, not URLs
|
||||
))
|
||||
model_response.data.append(ImageObject(
|
||||
b64_json=prediction.get("bytesBase64Encoded", None),
|
||||
url=None, # Google AI returns base64, not URLs
|
||||
))
|
||||
|
||||
return model_response
|
||||
|
|
@ -3,7 +3,6 @@ This file contains the transformation logic for the Gemini realtime API.
|
|||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional, Union, cast
|
||||
|
||||
|
|
@ -55,7 +54,7 @@ from litellm.types.realtime import (
|
|||
)
|
||||
from litellm.utils import get_empty_usage
|
||||
|
||||
from ..common_utils import encode_unserializable_types
|
||||
from ..common_utils import encode_unserializable_types, get_api_key_from_env
|
||||
|
||||
MAP_GEMINI_FIELD_TO_OPENAI_EVENT: Dict[str, OpenAIRealtimeEventTypes] = {
|
||||
"setupComplete": OpenAIRealtimeEventTypes.SESSION_CREATED,
|
||||
|
|
@ -81,7 +80,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
|
|||
if api_base is None:
|
||||
api_base = "wss://generativelanguage.googleapis.com"
|
||||
if api_key is None:
|
||||
api_key = os.environ.get("GEMINI_API_KEY")
|
||||
api_key = get_api_key_from_env()
|
||||
if api_key is None:
|
||||
raise ValueError("api_key is required for Gemini API calls")
|
||||
api_base = api_base.replace("https://", "wss://")
|
||||
|
|
@ -188,9 +187,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
|
|||
|
||||
vertex_gemini_config = VertexGeminiConfig()
|
||||
vertex_gemini_config._map_function(value)
|
||||
optional_params["generationConfig"][
|
||||
"tools"
|
||||
] = vertex_gemini_config._map_function(value)
|
||||
optional_params["generationConfig"]["tools"] = (
|
||||
vertex_gemini_config._map_function(value)
|
||||
)
|
||||
elif key == "input_audio_transcription" and value is not None:
|
||||
optional_params["inputAudioTranscription"] = {}
|
||||
elif key == "turn_detection":
|
||||
|
|
@ -201,10 +200,10 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
|
|||
if (
|
||||
len(transformed_audio_activity_config) > 0
|
||||
): # if the config is not empty, add it to the optional params
|
||||
optional_params[
|
||||
"realtimeInputConfig"
|
||||
] = BidiGenerateContentRealtimeInputConfig(
|
||||
automaticActivityDetection=transformed_audio_activity_config
|
||||
optional_params["realtimeInputConfig"] = (
|
||||
BidiGenerateContentRealtimeInputConfig(
|
||||
automaticActivityDetection=transformed_audio_activity_config
|
||||
)
|
||||
)
|
||||
if len(optional_params["generationConfig"]) == 0:
|
||||
optional_params.pop("generationConfig")
|
||||
|
|
@ -405,15 +404,17 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
|
|||
output_index=0,
|
||||
event_id="event_{}".format(uuid.uuid4()),
|
||||
item_id=output_item_id,
|
||||
part={
|
||||
"type": "text",
|
||||
"text": "",
|
||||
}
|
||||
if delta_type == "text"
|
||||
else {
|
||||
"type": "audio",
|
||||
"transcript": "",
|
||||
},
|
||||
part=(
|
||||
{
|
||||
"type": "text",
|
||||
"text": "",
|
||||
}
|
||||
if delta_type == "text"
|
||||
else {
|
||||
"type": "audio",
|
||||
"transcript": "",
|
||||
}
|
||||
),
|
||||
response_id=response_id,
|
||||
)
|
||||
response_items.append(response_content_part_added)
|
||||
|
|
@ -440,9 +441,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
|
|||
)
|
||||
|
||||
return OpenAIRealtimeResponseDelta(
|
||||
type="response.text.delta"
|
||||
if delta_type == "text"
|
||||
else "response.audio.delta",
|
||||
type=(
|
||||
"response.text.delta"
|
||||
if delta_type == "text"
|
||||
else "response.audio.delta"
|
||||
),
|
||||
content_index=0,
|
||||
event_id="event_{}".format(uuid.uuid4()),
|
||||
item_id=output_item_id,
|
||||
|
|
@ -513,12 +516,14 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
|
|||
event_id="event_{}".format(uuid.uuid4()),
|
||||
item_id=current_output_item_id,
|
||||
output_index=0,
|
||||
part={"type": "text", "text": delta_done_event_text}
|
||||
if delta_done_event_text and delta_type == "text"
|
||||
else {
|
||||
"type": "audio",
|
||||
"transcript": "", # gemini doesn't return transcript for audio
|
||||
},
|
||||
part=(
|
||||
{"type": "text", "text": delta_done_event_text}
|
||||
if delta_done_event_text and delta_type == "text"
|
||||
else {
|
||||
"type": "audio",
|
||||
"transcript": "", # gemini doesn't return transcript for audio
|
||||
}
|
||||
),
|
||||
response_id=current_response_id,
|
||||
)
|
||||
returned_items.append(response_content_part_done)
|
||||
|
|
@ -535,12 +540,14 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
|
|||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "text", "text": delta_done_event_text}
|
||||
if delta_done_event_text and delta_type == "text"
|
||||
else {
|
||||
"type": "audio",
|
||||
"transcript": "",
|
||||
}
|
||||
(
|
||||
{"type": "text", "text": delta_done_event_text}
|
||||
if delta_done_event_text and delta_type == "text"
|
||||
else {
|
||||
"type": "audio",
|
||||
"transcript": "",
|
||||
}
|
||||
)
|
||||
],
|
||||
},
|
||||
)
|
||||
|
|
@ -674,9 +681,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
|
|||
object="realtime.response",
|
||||
id=current_response_id,
|
||||
status="completed",
|
||||
output=[output_item["item"] for output_item in output_items]
|
||||
if output_items
|
||||
else [],
|
||||
output=(
|
||||
[output_item["item"] for output_item in output_items]
|
||||
if output_items
|
||||
else []
|
||||
),
|
||||
conversation_id=current_conversation_id,
|
||||
modalities=_modalities,
|
||||
usage=responses_api_usage.model_dump(),
|
||||
|
|
@ -828,9 +837,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
|
|||
"session_configuration_request"
|
||||
]
|
||||
current_item_chunks = realtime_response_transform_input["current_item_chunks"]
|
||||
current_delta_type: Optional[
|
||||
ALL_DELTA_TYPES
|
||||
] = realtime_response_transform_input["current_delta_type"]
|
||||
current_delta_type: Optional[ALL_DELTA_TYPES] = (
|
||||
realtime_response_transform_input["current_delta_type"]
|
||||
)
|
||||
returned_message: List[OpenAIRealtimeEvents] = []
|
||||
|
||||
for key, value in json_message.items():
|
||||
|
|
|
|||
|
|
@ -705,8 +705,14 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
if api_key is None:
|
||||
api_key = get_secret_str("OPENAI_API_KEY")
|
||||
|
||||
# Strip api_base to just the base URL (scheme + host + port)
|
||||
parsed_url = httpx.URL(api_base)
|
||||
base_url = f"{parsed_url.scheme}://{parsed_url.host}"
|
||||
if parsed_url.port:
|
||||
base_url += f":{parsed_url.port}"
|
||||
|
||||
response = litellm.module_level_client.get(
|
||||
url=f"{api_base}/v1/models",
|
||||
url=f"{base_url}/v1/models",
|
||||
headers={"Authorization": f"Bearer {api_key}"},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,8 @@ from litellm.types.utils import Usage, PromptTokensDetailsWrapper
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.types.llms.openai import ChatCompletionAnnotation
|
||||
from litellm.types.llms.openai import ChatCompletionAnnotationURLCitation
|
||||
|
||||
|
||||
class PerplexityChatConfig(OpenAIGPTConfig):
|
||||
|
|
@ -102,7 +104,10 @@ class PerplexityChatConfig(OpenAIGPTConfig):
|
|||
# Extract and enhance usage with Perplexity-specific fields
|
||||
try:
|
||||
raw_response_json = raw_response.json()
|
||||
self._enhance_usage_with_perplexity_fields(model_response, raw_response_json)
|
||||
self._enhance_usage_with_perplexity_fields(
|
||||
model_response, raw_response_json
|
||||
)
|
||||
self._add_citations_as_annotations(model_response, raw_response_json)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Error extracting Perplexity-specific usage fields: {e}")
|
||||
|
||||
|
|
@ -131,7 +136,9 @@ class PerplexityChatConfig(OpenAIGPTConfig):
|
|||
if citations:
|
||||
# Count total characters in citations as a proxy for citation tokens
|
||||
# This is an estimation - in practice, you might want to use proper tokenization
|
||||
total_citation_chars = sum(len(str(citation)) for citation in citations if citation)
|
||||
total_citation_chars = sum(
|
||||
len(str(citation)) for citation in citations if citation
|
||||
)
|
||||
# Rough estimation: ~4 characters per token (OpenAI's general rule)
|
||||
if total_citation_chars > 0:
|
||||
citation_tokens = max(1, total_citation_chars // 4)
|
||||
|
|
@ -150,7 +157,9 @@ class PerplexityChatConfig(OpenAIGPTConfig):
|
|||
num_search_queries = raw_response_json.get("search_queries")
|
||||
|
||||
# Create or update prompt_tokens_details to include web search requests and citation tokens
|
||||
if citation_tokens > 0 or (num_search_queries is not None and num_search_queries > 0):
|
||||
if citation_tokens > 0 or (
|
||||
num_search_queries is not None and num_search_queries > 0
|
||||
):
|
||||
if usage.prompt_tokens_details is None:
|
||||
usage.prompt_tokens_details = PromptTokensDetailsWrapper()
|
||||
|
||||
|
|
@ -161,3 +170,82 @@ class PerplexityChatConfig(OpenAIGPTConfig):
|
|||
# Store search queries count in the standard web_search_requests field
|
||||
if num_search_queries is not None and num_search_queries > 0:
|
||||
usage.prompt_tokens_details.web_search_requests = num_search_queries
|
||||
|
||||
def _add_citations_as_annotations(
|
||||
self, model_response: ModelResponse, raw_response_json: dict
|
||||
) -> None:
|
||||
"""
|
||||
Extract citations and search_results from Perplexity API response
|
||||
and add them as ChatCompletionAnnotation objects to the message.
|
||||
"""
|
||||
if not model_response.choices:
|
||||
return
|
||||
|
||||
# Get the first choice (assuming single response)
|
||||
choice = model_response.choices[0]
|
||||
if not hasattr(choice, "message") or choice.message is None:
|
||||
return
|
||||
|
||||
message = choice.message
|
||||
annotations = []
|
||||
|
||||
# Extract citations from the response
|
||||
citations = raw_response_json.get("citations", [])
|
||||
search_results = raw_response_json.get("search_results", [])
|
||||
|
||||
# Create a mapping of URLs to search result titles
|
||||
url_to_title = {}
|
||||
for result in search_results:
|
||||
if isinstance(result, dict) and "url" in result and "title" in result:
|
||||
url_to_title[result["url"]] = result["title"]
|
||||
|
||||
# Get the message content to find citation positions
|
||||
content = getattr(message, "content", "")
|
||||
if not content:
|
||||
return
|
||||
|
||||
# Find all citation markers like [1], [2], [3], [4] in the text
|
||||
import re
|
||||
|
||||
citation_pattern = r"\[(\d+)\]"
|
||||
citation_matches = list(re.finditer(citation_pattern, content))
|
||||
|
||||
# Create a mapping of citation numbers to URLs
|
||||
citation_number_to_url = {}
|
||||
for i, citation in enumerate(citations):
|
||||
if isinstance(citation, str):
|
||||
citation_number_to_url[i + 1] = citation # 1-indexed
|
||||
|
||||
# Create annotations for each citation match found in the text
|
||||
for match in citation_matches:
|
||||
citation_number = int(match.group(1))
|
||||
if citation_number in citation_number_to_url:
|
||||
url = citation_number_to_url[citation_number]
|
||||
title = url_to_title.get(url, "")
|
||||
|
||||
# Create the URL citation annotation with actual text positions
|
||||
url_citation: ChatCompletionAnnotationURLCitation = {
|
||||
"url": url,
|
||||
"title": title,
|
||||
"start_index": match.start(),
|
||||
"end_index": match.end(),
|
||||
}
|
||||
|
||||
annotation: ChatCompletionAnnotation = {
|
||||
"type": "url_citation",
|
||||
"url_citation": url_citation,
|
||||
}
|
||||
|
||||
annotations.append(annotation)
|
||||
|
||||
# Add annotations to the message if we have any
|
||||
if annotations:
|
||||
if not hasattr(message, "annotations") or message.annotations is None:
|
||||
message.annotations = []
|
||||
message.annotations.extend(annotations)
|
||||
|
||||
# Also add the raw citations and search_results as attributes for backward compatibility
|
||||
if citations:
|
||||
setattr(model_response, "citations", citations)
|
||||
if search_results:
|
||||
setattr(model_response, "search_results", search_results)
|
||||
|
|
@ -1,16 +1,39 @@
|
|||
"""
|
||||
Transformation for Calling Google models in their native format.
|
||||
"""
|
||||
from typing import Literal
|
||||
from typing import Literal, Optional, Union
|
||||
|
||||
from litellm.llms.gemini.google_genai.transformation import GoogleGenAIConfig
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
|
||||
class VertexAIGoogleGenAIConfig(GoogleGenAIConfig):
|
||||
"""
|
||||
Configuration for calling Google models in their native format.
|
||||
"""
|
||||
HEADER_NAME = "Authorization"
|
||||
BEARER_PREFIX = "Bearer"
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Literal["gemini", "vertex_ai"]:
|
||||
return "vertex_ai"
|
||||
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
api_key: Optional[str],
|
||||
headers: Optional[dict],
|
||||
model: str,
|
||||
litellm_params: Optional[Union[GenericLiteLLMParams, dict]]
|
||||
) -> dict:
|
||||
default_headers = {
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
if api_key is not None:
|
||||
default_headers[self.HEADER_NAME] = f"{self.BEARER_PREFIX} {api_key}"
|
||||
if headers is not None:
|
||||
default_headers.update(headers)
|
||||
|
||||
return default_headers
|
||||
|
||||
|
|
@ -147,6 +147,7 @@ from .llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
|||
from .llms.custom_llm import CustomLLM, custom_chat_llm_router
|
||||
from .llms.databricks.embed.handler import DatabricksEmbeddingHandler
|
||||
from .llms.deprecated_providers import aleph_alpha, palm
|
||||
from .llms.gemini.common_utils import get_api_key_from_env
|
||||
from .llms.groq.chat.handler import GroqChatCompletion
|
||||
from .llms.huggingface.embedding.handler import HuggingFaceEmbedding
|
||||
from .llms.nlp_cloud.chat.handler import completion as nlp_cloud_chat_completion
|
||||
|
|
@ -1048,11 +1049,13 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
non_default_params = get_non_default_completion_params(kwargs=kwargs)
|
||||
litellm_params = {} # used to prevent unbound var errors
|
||||
## PROMPT MANAGEMENT HOOKS ##
|
||||
|
||||
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and (
|
||||
litellm_logging_obj.should_run_prompt_management_hooks(
|
||||
prompt_id=prompt_id, non_default_params=non_default_params
|
||||
)
|
||||
):
|
||||
|
||||
(
|
||||
model,
|
||||
messages,
|
||||
|
|
@ -2595,7 +2598,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
|
||||
gemini_api_key = (
|
||||
api_key
|
||||
or get_secret("GEMINI_API_KEY")
|
||||
or get_api_key_from_env()
|
||||
or get_secret("PALM_API_KEY") # older palm api key should also work
|
||||
or litellm.api_key
|
||||
)
|
||||
|
|
@ -3895,6 +3898,9 @@ def embedding( # noqa: PLR0915
|
|||
or get_secret_str("OPENAI_LIKE_API_KEY")
|
||||
)
|
||||
|
||||
if extra_headers is not None:
|
||||
optional_params["extra_headers"] = extra_headers
|
||||
|
||||
## EMBEDDING CALL
|
||||
response = openai_like_embedding.embedding(
|
||||
model=model,
|
||||
|
|
@ -3998,9 +4004,7 @@ def embedding( # noqa: PLR0915
|
|||
litellm_params={},
|
||||
)
|
||||
elif custom_llm_provider == "gemini":
|
||||
gemini_api_key = (
|
||||
api_key or get_secret_str("GEMINI_API_KEY") or litellm.api_key
|
||||
)
|
||||
gemini_api_key = api_key or get_api_key_from_env() or litellm.api_key
|
||||
|
||||
api_base = api_base or litellm.api_base or get_secret_str("GEMINI_API_BASE")
|
||||
|
||||
|
|
@ -5430,6 +5434,7 @@ def speech( # noqa: PLR0915
|
|||
|
||||
##### Health Endpoints #######################
|
||||
|
||||
|
||||
async def ahealth_check(
|
||||
model_params: dict,
|
||||
mode: Optional[
|
||||
|
|
@ -5475,7 +5480,11 @@ async def ahealth_check(
|
|||
log_raw_request_response=True,
|
||||
)
|
||||
model_params["litellm_logging_obj"] = litellm_logging_obj
|
||||
model_params = HealthCheckHelpers._update_model_params_with_health_check_tracking_information(model_params=model_params)
|
||||
model_params = (
|
||||
HealthCheckHelpers._update_model_params_with_health_check_tracking_information(
|
||||
model_params=model_params
|
||||
)
|
||||
)
|
||||
#########################################################
|
||||
try:
|
||||
model: Optional[str] = model_params.get("model", None)
|
||||
|
|
|
|||
|
|
@ -10480,6 +10480,20 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_prompt_caching": true
|
||||
},
|
||||
"openrouter/x-ai/grok-4":{
|
||||
"max_tokens": 256000,
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 256000,
|
||||
"input_cost_per_token": 3e-06,
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"litellm_provider": "openrouter",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_reasoning": true,
|
||||
"source": "https://openrouter.ai/x-ai/grok-4",
|
||||
"supports_web_search": true
|
||||
},
|
||||
"openrouter/bytedance/ui-tars-1.5-7b":{
|
||||
"max_tokens": 2048,
|
||||
"max_input_tokens": 131072,
|
||||
|
|
|
|||
|
|
@ -16,8 +16,9 @@ class MCPAuthenticatedUser(AuthenticatedUser):
|
|||
4. Server-specific authentication headers
|
||||
"""
|
||||
|
||||
def __init__(self, user_api_key_auth: UserAPIKeyAuth, mcp_auth_header: Optional[str] = None, mcp_servers: Optional[List[str]] = None, mcp_server_auth_headers: Optional[Dict[str, str]] = None):
|
||||
def __init__(self, user_api_key_auth: UserAPIKeyAuth, mcp_auth_header: Optional[str] = None, mcp_servers: Optional[List[str]] = None, mcp_server_auth_headers: Optional[Dict[str, str]] = None, mcp_protocol_version: Optional[str] = None):
|
||||
self.user_api_key_auth = user_api_key_auth
|
||||
self.mcp_auth_header = mcp_auth_header
|
||||
self.mcp_servers = mcp_servers
|
||||
self.mcp_server_auth_headers = mcp_server_auth_headers or {}
|
||||
self.mcp_protocol_version = mcp_protocol_version
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import List, Optional, Tuple, Dict
|
||||
from typing import List, Optional, Tuple, Dict, Set
|
||||
|
||||
from starlette.datastructures import Headers
|
||||
from starlette.requests import Request
|
||||
|
|
@ -28,9 +28,12 @@ class MCPRequestHandler:
|
|||
LITELLM_MCP_SERVERS_HEADER_NAME = SpecialHeaders.mcp_servers.value
|
||||
|
||||
LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value
|
||||
|
||||
# MCP Protocol Version header
|
||||
MCP_PROTOCOL_VERSION_HEADER_NAME = "MCP-Protocol-Version"
|
||||
|
||||
@staticmethod
|
||||
async def process_mcp_request(scope: Scope) -> Tuple[UserAPIKeyAuth, Optional[str], Optional[List[str]], Optional[Dict[str, str]]]:
|
||||
async def process_mcp_request(scope: Scope) -> Tuple[UserAPIKeyAuth, Optional[str], Optional[List[str]], Optional[Dict[str, str]], Optional[str]]:
|
||||
"""
|
||||
Process and validate MCP request headers from the ASGI scope.
|
||||
This includes:
|
||||
|
|
@ -46,6 +49,7 @@ class MCPRequestHandler:
|
|||
mcp_auth_header: Optional[str] MCP auth header to be passed to the MCP server (deprecated)
|
||||
mcp_servers: Optional[List[str]] List of MCP servers and access groups to use
|
||||
mcp_server_auth_headers: Optional[Dict[str, str]] Server-specific auth headers in format {server_alias: auth_value}
|
||||
mcp_protocol_version: Optional[str] MCP protocol version from request header
|
||||
|
||||
Raises:
|
||||
HTTPException: If headers are invalid or missing required headers
|
||||
|
|
@ -61,6 +65,9 @@ class MCPRequestHandler:
|
|||
# Get the new server-specific auth headers
|
||||
mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
|
||||
|
||||
# Get MCP protocol version from header
|
||||
mcp_protocol_version = headers.get(MCPRequestHandler.MCP_PROTOCOL_VERSION_HEADER_NAME)
|
||||
|
||||
# Parse MCP servers from header
|
||||
mcp_servers_header = headers.get(MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME)
|
||||
verbose_logger.debug(f"Raw MCP servers header: {mcp_servers_header}")
|
||||
|
|
@ -82,7 +89,7 @@ class MCPRequestHandler:
|
|||
validated_user_api_key_auth = await user_api_key_auth(
|
||||
api_key=litellm_api_key, request=request
|
||||
)
|
||||
return validated_user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers
|
||||
return validated_user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, mcp_protocol_version
|
||||
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -219,25 +226,29 @@ class MCPRequestHandler:
|
|||
"""
|
||||
from typing import List
|
||||
|
||||
allowed_mcp_servers: List[str] = []
|
||||
allowed_mcp_servers_for_key = (
|
||||
await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
|
||||
)
|
||||
allowed_mcp_servers_for_team = (
|
||||
await MCPRequestHandler._get_allowed_mcp_servers_for_team(user_api_key_auth)
|
||||
)
|
||||
try:
|
||||
allowed_mcp_servers: List[str] = []
|
||||
allowed_mcp_servers_for_key = (
|
||||
await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
|
||||
)
|
||||
allowed_mcp_servers_for_team = (
|
||||
await MCPRequestHandler._get_allowed_mcp_servers_for_team(user_api_key_auth)
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# If team has mcp_servers, then key must have a subset of the team's mcp_servers
|
||||
#########################################################
|
||||
if len(allowed_mcp_servers_for_team) > 0:
|
||||
for _mcp_server in allowed_mcp_servers_for_key:
|
||||
if _mcp_server in allowed_mcp_servers_for_team:
|
||||
allowed_mcp_servers.append(_mcp_server)
|
||||
else:
|
||||
allowed_mcp_servers = allowed_mcp_servers_for_key
|
||||
#########################################################
|
||||
# If team has mcp_servers, then key must have a subset of the team's mcp_servers
|
||||
#########################################################
|
||||
if len(allowed_mcp_servers_for_team) > 0:
|
||||
for _mcp_server in allowed_mcp_servers_for_key:
|
||||
if _mcp_server in allowed_mcp_servers_for_team:
|
||||
allowed_mcp_servers.append(_mcp_server)
|
||||
else:
|
||||
allowed_mcp_servers = allowed_mcp_servers_for_key
|
||||
|
||||
return list(set(allowed_mcp_servers))
|
||||
return list(set(allowed_mcp_servers))
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_mcp_servers_for_key(
|
||||
|
|
@ -255,25 +266,29 @@ class MCPRequestHandler:
|
|||
verbose_logger.debug("prisma_client is None")
|
||||
return []
|
||||
|
||||
key_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": user_api_key_auth.object_permission_id},
|
||||
try:
|
||||
key_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": user_api_key_auth.object_permission_id},
|
||||
)
|
||||
)
|
||||
)
|
||||
if key_object_permission is None:
|
||||
return []
|
||||
if key_object_permission is None:
|
||||
return []
|
||||
|
||||
# Get direct MCP servers
|
||||
direct_mcp_servers = key_object_permission.mcp_servers or []
|
||||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
key_object_permission.mcp_access_groups or []
|
||||
)
|
||||
|
||||
# Combine both lists
|
||||
all_servers = direct_mcp_servers + access_group_servers
|
||||
return list(set(all_servers))
|
||||
# Get direct MCP servers
|
||||
direct_mcp_servers = key_object_permission.mcp_servers or []
|
||||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
key_object_permission.mcp_access_groups or []
|
||||
)
|
||||
|
||||
# Combine both lists
|
||||
all_servers = direct_mcp_servers + access_group_servers
|
||||
return list(set(all_servers))
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get allowed MCP servers for key: {str(e)}")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_mcp_servers_for_team(
|
||||
|
|
@ -297,37 +312,41 @@ class MCPRequestHandler:
|
|||
verbose_logger.debug("prisma_client is None")
|
||||
return []
|
||||
|
||||
team_obj: Optional[LiteLLM_TeamTable] = (
|
||||
await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": user_api_key_auth.team_id},
|
||||
try:
|
||||
team_obj: Optional[LiteLLM_TeamTable] = (
|
||||
await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": user_api_key_auth.team_id},
|
||||
)
|
||||
)
|
||||
)
|
||||
if team_obj is None:
|
||||
verbose_logger.debug("team_obj is None")
|
||||
return []
|
||||
if team_obj is None:
|
||||
verbose_logger.debug("team_obj is None")
|
||||
return []
|
||||
|
||||
object_permissions = team_obj.object_permission
|
||||
if object_permissions is None:
|
||||
return []
|
||||
object_permissions = team_obj.object_permission
|
||||
if object_permissions is None:
|
||||
return []
|
||||
|
||||
# Get direct MCP servers
|
||||
direct_mcp_servers = object_permissions.mcp_servers or []
|
||||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
object_permissions.mcp_access_groups or []
|
||||
)
|
||||
|
||||
# Combine both lists
|
||||
all_servers = direct_mcp_servers + access_group_servers
|
||||
return list(set(all_servers))
|
||||
# Get direct MCP servers
|
||||
direct_mcp_servers = object_permissions.mcp_servers or []
|
||||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
object_permissions.mcp_access_groups or []
|
||||
)
|
||||
|
||||
# Combine both lists
|
||||
all_servers = direct_mcp_servers + access_group_servers
|
||||
return list(set(all_servers))
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get allowed MCP servers for team: {str(e)}")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _get_config_server_ids_for_access_groups(config_mcp_servers, access_groups: List[str]) -> set:
|
||||
def _get_config_server_ids_for_access_groups(config_mcp_servers, access_groups: List[str]) -> Set[str]:
|
||||
"""
|
||||
Helper to get server_ids from config-loaded servers that match any of the given access groups.
|
||||
"""
|
||||
server_ids = set()
|
||||
server_ids: Set[str] = set()
|
||||
for server_id, server in config_mcp_servers.items():
|
||||
if server.access_groups:
|
||||
if any(group in server.access_groups for group in access_groups):
|
||||
|
|
@ -335,11 +354,11 @@ class MCPRequestHandler:
|
|||
return server_ids
|
||||
|
||||
@staticmethod
|
||||
async def _get_db_server_ids_for_access_groups(prisma_client, access_groups: List[str]) -> set:
|
||||
async def _get_db_server_ids_for_access_groups(prisma_client, access_groups: List[str]) -> Set[str]:
|
||||
"""
|
||||
Helper to get server_ids from DB servers that match any of the given access groups.
|
||||
"""
|
||||
server_ids = set()
|
||||
server_ids: Set[str] = set()
|
||||
if access_groups and prisma_client is not None:
|
||||
try:
|
||||
mcp_servers = await prisma_client.db.litellm_mcpservertable.find_many(
|
||||
|
|
@ -363,20 +382,26 @@ class MCPRequestHandler:
|
|||
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
# Use the new helper for config-loaded servers
|
||||
server_ids = MCPRequestHandler._get_config_server_ids_for_access_groups(
|
||||
global_mcp_server_manager.config_mcp_servers, access_groups
|
||||
)
|
||||
try:
|
||||
# Import here to avoid circular import
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
# Use the new helper for config-loaded servers
|
||||
server_ids = MCPRequestHandler._get_config_server_ids_for_access_groups(
|
||||
global_mcp_server_manager.config_mcp_servers, access_groups
|
||||
)
|
||||
|
||||
# Use the new helper for DB servers
|
||||
db_server_ids = await MCPRequestHandler._get_db_server_ids_for_access_groups(
|
||||
prisma_client, access_groups
|
||||
)
|
||||
server_ids.update(db_server_ids)
|
||||
# Use the new helper for DB servers
|
||||
db_server_ids = await MCPRequestHandler._get_db_server_ids_for_access_groups(
|
||||
prisma_client, access_groups
|
||||
)
|
||||
server_ids.update(db_server_ids)
|
||||
|
||||
return list(server_ids)
|
||||
return list(server_ids)
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get MCP servers from access groups: {str(e)}")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
async def get_mcp_access_groups(
|
||||
|
|
|
|||
|
|
@ -82,14 +82,20 @@ async def get_mcp_servers(
|
|||
"""
|
||||
Returns the matching mcp servers from the db with the server_ids
|
||||
"""
|
||||
mcp_servers: List[LiteLLM_MCPServerTable] = (
|
||||
_mcp_servers: List[LiteLLM_MCPServerTable] = (
|
||||
await prisma_client.db.litellm_mcpservertable.find_many(
|
||||
where={
|
||||
"server_id": {"in": server_ids},
|
||||
}
|
||||
)
|
||||
)
|
||||
return mcp_servers
|
||||
final_mcp_servers: List[LiteLLM_MCPServerTable] = []
|
||||
for _mcp_server in _mcp_servers:
|
||||
final_mcp_servers.append(
|
||||
LiteLLM_MCPServerTable(**_mcp_server.model_dump())
|
||||
)
|
||||
|
||||
return final_mcp_servers
|
||||
|
||||
|
||||
async def get_mcp_servers_by_verificationtoken(
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ router = APIRouter(
|
|||
if MCP_AVAILABLE:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
_convert_protocol_version_to_enum,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
ListMCPToolsRestAPIResponseObject,
|
||||
|
|
@ -140,15 +141,52 @@ if MCP_AVAILABLE:
|
|||
REST API to call a specific MCP tool with the provided arguments
|
||||
"""
|
||||
from litellm.proxy.proxy_server import add_litellm_data_to_request, proxy_config
|
||||
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
|
||||
from fastapi import HTTPException
|
||||
|
||||
data = await request.json()
|
||||
data = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
return await call_mcp_tool(**data)
|
||||
try:
|
||||
data = await request.json()
|
||||
data = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
return await call_mcp_tool(**data)
|
||||
except BlockedPiiEntityError as e:
|
||||
verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "blocked_pii_entity",
|
||||
"message": str(e),
|
||||
"entity_type": getattr(e, 'entity_type', None),
|
||||
"guardrail_name": getattr(e, 'guardrail_name', None)
|
||||
}
|
||||
)
|
||||
except GuardrailRaisedException as e:
|
||||
verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "guardrail_violation",
|
||||
"message": str(e),
|
||||
"guardrail_name": getattr(e, 'guardrail_name', None)
|
||||
}
|
||||
)
|
||||
except HTTPException as e:
|
||||
# Re-raise HTTPException as-is to preserve status code and detail
|
||||
verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}")
|
||||
raise e
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Unexpected error in MCP tool call: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": "internal_server_error",
|
||||
"message": f"An unexpected error occurred: {str(e)}"
|
||||
}
|
||||
)
|
||||
|
||||
########################################################
|
||||
# MCP Connection testing routes
|
||||
|
|
@ -174,7 +212,7 @@ if MCP_AVAILABLE:
|
|||
name=request.alias or request.server_name or "",
|
||||
url=request.url,
|
||||
transport=request.transport,
|
||||
spec_version=request.spec_version,
|
||||
spec_version=_convert_protocol_version_to_enum(request.spec_version),
|
||||
auth_type=request.auth_type,
|
||||
mcp_info=request.mcp_info,
|
||||
),
|
||||
|
|
@ -203,7 +241,7 @@ if MCP_AVAILABLE:
|
|||
name=request.alias or request.server_name or "",
|
||||
url=request.url,
|
||||
transport=request.transport,
|
||||
spec_version=request.spec_version,
|
||||
spec_version=_convert_protocol_version_to_enum(request.spec_version),
|
||||
auth_type=request.auth_type,
|
||||
mcp_info=request.mcp_info,
|
||||
),
|
||||
|
|
|
|||
|
|
@ -64,7 +64,7 @@ if MCP_AVAILABLE:
|
|||
global_mcp_tool_registry,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
get_server_name_prefix_tool_mcp,
|
||||
get_server_name_prefix_tool_mcp,
|
||||
)
|
||||
|
||||
######################################################
|
||||
|
|
@ -168,24 +168,34 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
List all available tools
|
||||
"""
|
||||
# Get user authentication from context variable
|
||||
user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers = get_auth_context()
|
||||
verbose_logger.debug(
|
||||
f"MCP list_tools - User API Key Auth from context: {user_api_key_auth}"
|
||||
)
|
||||
verbose_logger.debug(
|
||||
f"MCP list_tools - MCP servers from context: {mcp_servers}"
|
||||
)
|
||||
verbose_logger.debug(
|
||||
f"MCP list_tools - MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}"
|
||||
)
|
||||
# Get mcp_servers from context variable
|
||||
return await _list_mcp_tools(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
)
|
||||
try:
|
||||
# Get user authentication from context variable
|
||||
user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, mcp_protocol_version = get_auth_context()
|
||||
verbose_logger.debug(
|
||||
f"MCP list_tools - User API Key Auth from context: {user_api_key_auth}"
|
||||
)
|
||||
verbose_logger.debug(
|
||||
f"MCP list_tools - MCP servers from context: {mcp_servers}"
|
||||
)
|
||||
verbose_logger.debug(
|
||||
f"MCP list_tools - MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}"
|
||||
)
|
||||
# Get mcp_servers from context variable
|
||||
verbose_logger.debug("MCP list_tools - Calling _list_mcp_tools")
|
||||
tools = await _list_mcp_tools(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
mcp_protocol_version=mcp_protocol_version,
|
||||
)
|
||||
verbose_logger.info(f"MCP list_tools - Successfully returned {len(tools)} tools")
|
||||
return tools
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error in list_tools endpoint: {str(e)}")
|
||||
# Return empty list instead of failing completely
|
||||
# This prevents the HTTP stream from failing and allows the client to get a response
|
||||
return []
|
||||
|
||||
@server.call_tool()
|
||||
async def mcp_server_tool_call(
|
||||
|
|
@ -208,9 +218,10 @@ if MCP_AVAILABLE:
|
|||
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
from litellm.proxy.proxy_server import proxy_config
|
||||
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
|
||||
|
||||
# Validate arguments
|
||||
user_api_key_auth, mcp_auth_header, _, mcp_server_auth_headers = get_auth_context()
|
||||
user_api_key_auth, mcp_auth_header, _, mcp_server_auth_headers, mcp_protocol_version = get_auth_context()
|
||||
|
||||
verbose_logger.debug(
|
||||
f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}"
|
||||
|
|
@ -241,11 +252,37 @@ if MCP_AVAILABLE:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
mcp_protocol_version=mcp_protocol_version,
|
||||
**data, # for logging
|
||||
)
|
||||
except BlockedPiiEntityError as e:
|
||||
verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}")
|
||||
# Return error as text content for MCP protocol
|
||||
return [TextContent(
|
||||
text=f"Error: Blocked PII entity detected - {str(e)}",
|
||||
type="text"
|
||||
)]
|
||||
except GuardrailRaisedException as e:
|
||||
verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}")
|
||||
# Return error as text content for MCP protocol
|
||||
return [TextContent(
|
||||
text=f"Error: Guardrail violation - {str(e)}",
|
||||
type="text"
|
||||
)]
|
||||
except HTTPException as e:
|
||||
verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}")
|
||||
# Return error as text content for MCP protocol
|
||||
return [TextContent(
|
||||
text=f"Error: {str(e.detail)}",
|
||||
type="text"
|
||||
)]
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"MCP mcp_server_tool_call - error: {e}")
|
||||
raise e
|
||||
# Return error as text content for MCP protocol
|
||||
return [TextContent(
|
||||
text=f"Error: {str(e)}",
|
||||
type="text"
|
||||
)]
|
||||
|
||||
return response
|
||||
|
||||
|
|
@ -262,6 +299,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header: Optional[str],
|
||||
mcp_servers: Optional[List[str]],
|
||||
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
|
||||
mcp_protocol_version: Optional[str] = None,
|
||||
) -> List[MCPTool]:
|
||||
"""
|
||||
Helper method to fetch tools from MCP servers based on server filtering criteria.
|
||||
|
|
@ -325,6 +363,7 @@ if MCP_AVAILABLE:
|
|||
tools = await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
mcp_protocol_version=mcp_protocol_version,
|
||||
)
|
||||
all_tools.extend(tools)
|
||||
verbose_logger.debug(f"Successfully fetched {len(tools)} tools from server {server.name}")
|
||||
|
|
@ -342,6 +381,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header: Optional[str] = None,
|
||||
mcp_servers: Optional[List[str]] = None,
|
||||
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
|
||||
mcp_protocol_version: Optional[str] = None,
|
||||
) -> List[MCPTool]:
|
||||
"""
|
||||
List all available MCP tools.
|
||||
|
|
@ -357,28 +397,38 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
if not MCP_AVAILABLE:
|
||||
return []
|
||||
|
||||
# Get tools from managed MCP servers
|
||||
managed_tools = await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
)
|
||||
# Get tools from managed MCP servers with error handling
|
||||
managed_tools = []
|
||||
try:
|
||||
managed_tools = await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
mcp_protocol_version=mcp_protocol_version,
|
||||
)
|
||||
verbose_logger.debug(f"Successfully fetched {len(managed_tools)} tools from managed MCP servers")
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error getting tools from managed MCP servers: {str(e)}")
|
||||
# Continue with empty managed tools list instead of failing completely
|
||||
|
||||
# Get tools from local registry
|
||||
local_tools_raw = global_mcp_tool_registry.list_tools()
|
||||
|
||||
# Convert local tools to MCPTool format
|
||||
local_tools = []
|
||||
for tool in local_tools_raw:
|
||||
# Convert from litellm.types.mcp_server.tool_registry.MCPTool to mcp.types.Tool
|
||||
mcp_tool = MCPTool(
|
||||
name=tool.name,
|
||||
description=tool.description,
|
||||
inputSchema=tool.input_schema
|
||||
)
|
||||
local_tools.append(mcp_tool)
|
||||
try:
|
||||
local_tools_raw = global_mcp_tool_registry.list_tools()
|
||||
|
||||
# Convert local tools to MCPTool format
|
||||
for tool in local_tools_raw:
|
||||
# Convert from litellm.types.mcp_server.tool_registry.MCPTool to mcp.types.Tool
|
||||
mcp_tool = MCPTool(
|
||||
name=tool.name,
|
||||
description=tool.description,
|
||||
inputSchema=tool.input_schema
|
||||
)
|
||||
local_tools.append(mcp_tool)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error getting tools from local registry: {str(e)}")
|
||||
# Continue with empty local tools list instead of failing completely
|
||||
|
||||
# Combine all tools
|
||||
all_tools = managed_tools + local_tools
|
||||
|
|
@ -392,6 +442,7 @@ if MCP_AVAILABLE:
|
|||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
mcp_auth_header: Optional[str] = None,
|
||||
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
|
||||
mcp_protocol_version: Optional[str] = None,
|
||||
**kwargs: Any
|
||||
) -> List[Union[TextContent, ImageContent, EmbeddedResource]]:
|
||||
"""
|
||||
|
|
@ -439,6 +490,8 @@ if MCP_AVAILABLE:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
mcp_protocol_version=mcp_protocol_version,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
# Fall back to local tool registry (use original name)
|
||||
|
|
@ -490,14 +543,20 @@ if MCP_AVAILABLE:
|
|||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
mcp_auth_header: Optional[str] = None,
|
||||
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
|
||||
mcp_protocol_version: Optional[str] = None,
|
||||
litellm_logging_obj: Optional[Any] = None,
|
||||
) -> List[Union[TextContent, ImageContent, EmbeddedResource]]:
|
||||
"""Handle tool execution for managed server tools"""
|
||||
# Import here to avoid circular import
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
call_tool_result = await global_mcp_server_manager.call_tool(
|
||||
name=name,
|
||||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result)
|
||||
return call_tool_result.content # type: ignore[return-value]
|
||||
|
|
@ -533,15 +592,15 @@ if MCP_AVAILABLE:
|
|||
mcp_servers_from_path = [s.strip() for s in mcp_servers_str.split(",") if s.strip()]
|
||||
|
||||
if mcp_servers_from_path is not None:
|
||||
user_api_key_auth, mcp_auth_header, _, mcp_server_auth_headers = (
|
||||
user_api_key_auth, mcp_auth_header, _, mcp_server_auth_headers, mcp_protocol_version = (
|
||||
await MCPRequestHandler.process_mcp_request(scope)
|
||||
)
|
||||
mcp_servers = mcp_servers_from_path
|
||||
else:
|
||||
user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers = (
|
||||
user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, mcp_protocol_version = (
|
||||
await MCPRequestHandler.process_mcp_request(scope)
|
||||
)
|
||||
return user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers
|
||||
return user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, mcp_protocol_version
|
||||
|
||||
async def handle_streamable_http_mcp(
|
||||
scope: Scope, receive: Receive, send: Send
|
||||
|
|
@ -549,15 +608,17 @@ if MCP_AVAILABLE:
|
|||
"""Handle MCP requests through StreamableHTTP."""
|
||||
try:
|
||||
path = scope.get("path", "")
|
||||
user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers = await extract_mcp_auth_context(scope, path)
|
||||
user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, mcp_protocol_version = await extract_mcp_auth_context(scope, path)
|
||||
verbose_logger.debug(f"MCP request mcp_servers (header/path): {mcp_servers}")
|
||||
verbose_logger.debug(f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}")
|
||||
verbose_logger.debug(f"MCP protocol version: {mcp_protocol_version}")
|
||||
# Set the auth context variable for easy access in MCP functions
|
||||
set_auth_context(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
mcp_protocol_version=mcp_protocol_version,
|
||||
)
|
||||
|
||||
# Ensure session managers are initialized
|
||||
|
|
@ -569,20 +630,36 @@ if MCP_AVAILABLE:
|
|||
await session_manager.handle_request(scope, receive, send)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error handling MCP request: {e}")
|
||||
raise e
|
||||
# Instead of re-raising, try to send a graceful error response
|
||||
try:
|
||||
# Send a proper HTTP error response instead of letting the exception bubble up
|
||||
from starlette.responses import JSONResponse
|
||||
from starlette.status import HTTP_500_INTERNAL_SERVER_ERROR
|
||||
|
||||
error_response = JSONResponse(
|
||||
status_code=HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
content={"error": "MCP request failed", "details": str(e)}
|
||||
)
|
||||
await error_response(scope, receive, send)
|
||||
except Exception as response_error:
|
||||
verbose_logger.exception(f"Failed to send error response: {response_error}")
|
||||
# If we can't send a proper response, re-raise the original error
|
||||
raise e
|
||||
|
||||
async def handle_sse_mcp(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
"""Handle MCP requests through SSE."""
|
||||
try:
|
||||
path = scope.get("path", "")
|
||||
user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers = await extract_mcp_auth_context(scope, path)
|
||||
user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, mcp_protocol_version = await extract_mcp_auth_context(scope, path)
|
||||
verbose_logger.debug(f"MCP request mcp_servers (header/path): {mcp_servers}")
|
||||
verbose_logger.debug(f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}")
|
||||
verbose_logger.debug(f"MCP protocol version: {mcp_protocol_version}")
|
||||
set_auth_context(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
mcp_protocol_version=mcp_protocol_version,
|
||||
)
|
||||
|
||||
if not _SESSION_MANAGERS_INITIALIZED:
|
||||
|
|
@ -592,7 +669,21 @@ if MCP_AVAILABLE:
|
|||
await sse_session_manager.handle_request(scope, receive, send)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error handling MCP request: {e}")
|
||||
raise e
|
||||
# Instead of re-raising, try to send a graceful error response
|
||||
try:
|
||||
# Send a proper HTTP error response instead of letting the exception bubble up
|
||||
from starlette.responses import JSONResponse
|
||||
from starlette.status import HTTP_500_INTERNAL_SERVER_ERROR
|
||||
|
||||
error_response = JSONResponse(
|
||||
status_code=HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
content={"error": "MCP request failed", "details": str(e)}
|
||||
)
|
||||
await error_response(scope, receive, send)
|
||||
except Exception as response_error:
|
||||
verbose_logger.exception(f"Failed to send error response: {response_error}")
|
||||
# If we can't send a proper response, re-raise the original error
|
||||
raise e
|
||||
|
||||
app = FastAPI(
|
||||
title=LITELLM_MCP_SERVER_NAME,
|
||||
|
|
@ -626,6 +717,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header: Optional[str] = None,
|
||||
mcp_servers: Optional[List[str]] = None,
|
||||
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
|
||||
mcp_protocol_version: Optional[str] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Set the UserAPIKeyAuth in the auth context variable.
|
||||
|
|
@ -641,11 +733,12 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
mcp_protocol_version=mcp_protocol_version,
|
||||
)
|
||||
auth_context_var.set(auth_user)
|
||||
|
||||
def get_auth_context() -> (
|
||||
Tuple[Optional[UserAPIKeyAuth], Optional[str], Optional[List[str]], Optional[Dict[str, str]]]
|
||||
Tuple[Optional[UserAPIKeyAuth], Optional[str], Optional[List[str]], Optional[Dict[str, str]], Optional[str]]
|
||||
):
|
||||
"""
|
||||
Get the UserAPIKeyAuth from the auth context variable.
|
||||
|
|
@ -661,8 +754,9 @@ if MCP_AVAILABLE:
|
|||
auth_user.mcp_auth_header,
|
||||
auth_user.mcp_servers,
|
||||
auth_user.mcp_server_auth_headers,
|
||||
auth_user.mcp_protocol_version,
|
||||
)
|
||||
return None, None, None, None
|
||||
return None, None, None, None, None
|
||||
|
||||
########################################################
|
||||
############ End of Auth Context Functions #############
|
||||
|
|
|
|||