mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge branch 'litellm_internal_staging' of github.com:BerriAI/litellm into litellm_bedrock_mantle_claude_5_models
This commit is contained in:
commit
740bc96ad4
1588 changed files with 44377 additions and 17261 deletions
|
|
@ -24,9 +24,12 @@ jobs:
|
|||
- name: Update JSON Data
|
||||
run: |
|
||||
uv run --frozen --with 'aiohttp==3.13.3' python ".github/workflows/auto_update_price_and_context_window_file.py"
|
||||
- name: Regenerate JSON Schema
|
||||
run: |
|
||||
uv run --frozen python ci_cd/generate_model_prices_schema.py
|
||||
- name: Create Pull Request
|
||||
run: |
|
||||
git add model_prices_and_context_window.json
|
||||
git add model_prices_and_context_window.json model_prices_and_context_window.schema.json
|
||||
git commit -m "Update model_prices_and_context_window.json file: $(date +'%Y-%m-%d')"
|
||||
gh pr create --title "Update model_prices_and_context_window.json file" \
|
||||
--body "Automated update for model_prices_and_context_window.json" \
|
||||
|
|
|
|||
1
.github/workflows/test-linting.yml
vendored
1
.github/workflows/test-linting.yml
vendored
|
|
@ -104,6 +104,7 @@ jobs:
|
|||
- name: Check basedpyright budget (delta vs base)
|
||||
env:
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
NODE_OPTIONS: --max-old-space-size=12288
|
||||
run: |
|
||||
(uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py --base "$BASE_SHA"
|
||||
|
||||
|
|
|
|||
10
.github/workflows/test-litellm-ui-lint.yml
vendored
10
.github/workflows/test-litellm-ui-lint.yml
vendored
|
|
@ -29,7 +29,15 @@ jobs:
|
|||
id: changed
|
||||
env:
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
run: |
|
||||
# base.sha is the base branch tip from when the PR was opened, while
|
||||
# actions/checkout leaves HEAD on a merge of the PR into the *current*
|
||||
# base tip. "$BASE_SHA"...HEAD therefore spans every base-branch commit
|
||||
# landed since, so a PR that touches no UI file still gets linted
|
||||
# against hundreds of other people's files. Diff the PR head against its
|
||||
# own merge base instead, which is exactly what this PR changed.
|
||||
merge_base=$(git merge-base "$BASE_SHA" "$HEAD_SHA")
|
||||
: > "$RUNNER_TEMP/prettier_files.txt"
|
||||
: > "$RUNNER_TEMP/eslint_files.txt"
|
||||
while IFS= read -r f; do
|
||||
|
|
@ -41,7 +49,7 @@ jobs:
|
|||
*.json | *.css | *.scss | *.md | *.mdx | *.yml | *.yaml | *.html)
|
||||
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt" ;;
|
||||
esac
|
||||
done < <(git diff --name-only --diff-filter=ACMR --relative "$BASE_SHA"...HEAD -- .)
|
||||
done < <(git diff --name-only --diff-filter=ACMR --relative "$merge_base" "$HEAD_SHA" -- .)
|
||||
if [ -s "$RUNNER_TEMP/prettier_files.txt" ] || [ -s "$RUNNER_TEMP/eslint_files.txt" ]; then
|
||||
echo "has_files=true" >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
|
|
|
|||
9
.github/workflows/test-model-map.yaml
vendored
9
.github/workflows/test-model-map.yaml
vendored
|
|
@ -22,3 +22,12 @@ jobs:
|
|||
- name: Validate model_prices_and_context_window.json
|
||||
run: |
|
||||
jq empty model_prices_and_context_window.json
|
||||
|
||||
- name: Set up uv
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Check model_prices_and_context_window.schema.json is in sync
|
||||
run: |
|
||||
uv run --frozen python ci_cd/generate_model_prices_schema.py --check
|
||||
|
|
|
|||
151
.github/workflows/test_server_root_path.yml
vendored
151
.github/workflows/test_server_root_path.yml
vendored
|
|
@ -1,151 +0,0 @@
|
|||
name: Test Proxy SERVER_ROOT_PATH Routing
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
jobs:
|
||||
test-server-root-path:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
root_path: ["/api/v1", "/llmproxy"]
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
sudo rm -rf /usr/local/lib/android /usr/share/dotnet /opt/ghc /usr/local/share/boost
|
||||
sudo apt-get clean
|
||||
df -h /
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3.12.0
|
||||
|
||||
- name: Build Docker image
|
||||
uses: docker/build-push-action@0adf9959216b96bec444f325f1e493d4aa344497 # v6.14.0
|
||||
with:
|
||||
context: .
|
||||
file: ./docker/Dockerfile.non_root
|
||||
tags: litellm-test:${{ github.sha }}
|
||||
load: true
|
||||
push: false
|
||||
|
||||
- name: Start LiteLLM container with SERVER_ROOT_PATH
|
||||
run: |
|
||||
docker run -d \
|
||||
--name litellm-test \
|
||||
-p 4000:4000 \
|
||||
-e SERVER_ROOT_PATH="${{ matrix.root_path }}" \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
litellm-test:${{ github.sha }} \
|
||||
--detailed_debug
|
||||
|
||||
- name: Wait for container to be healthy
|
||||
run: |
|
||||
echo "Waiting for LiteLLM to start..."
|
||||
max_attempts=30
|
||||
attempt=0
|
||||
|
||||
while [ $attempt -lt $max_attempts ]; do
|
||||
if docker logs litellm-test 2>&1 | grep -q "Uvicorn running"; then
|
||||
echo "LiteLLM started successfully"
|
||||
break
|
||||
fi
|
||||
attempt=$((attempt + 1))
|
||||
echo "Attempt $attempt/$max_attempts - waiting for server to start..."
|
||||
sleep 2
|
||||
done
|
||||
|
||||
if [ $attempt -eq $max_attempts ]; then
|
||||
echo "Server failed to start within timeout"
|
||||
docker logs litellm-test
|
||||
exit 1
|
||||
fi
|
||||
|
||||
sleep 5
|
||||
|
||||
- name: Show container logs
|
||||
if: always()
|
||||
run: docker logs litellm-test
|
||||
|
||||
- name: Test UI endpoint with root path
|
||||
run: |
|
||||
ROOT_PATH="${{ matrix.root_path }}"
|
||||
echo "Testing UI at: http://localhost:4000${ROOT_PATH}/ui/"
|
||||
|
||||
for i in 1 2 3; do
|
||||
content=$(curl -sL --max-time 5 -H "Authorization: Bearer sk-1234" "http://localhost:4000${ROOT_PATH}/ui/")
|
||||
if echo "$content" | grep -q -E "(html|<!DOCTYPE|<head|<body)"; then
|
||||
echo "UI page contains valid HTML content"
|
||||
exit 0
|
||||
fi
|
||||
echo "Attempt $i/3 - no valid HTML, retrying in 5s..."
|
||||
sleep 5
|
||||
done
|
||||
echo "UI page does not contain expected HTML content"
|
||||
echo "Response: $content"
|
||||
docker logs litellm-test
|
||||
exit 1
|
||||
|
||||
- name: Setup Node for Playwright
|
||||
uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4.4.0
|
||||
with:
|
||||
node-version: "20"
|
||||
|
||||
- name: Install e2e deps and Chromium
|
||||
working-directory: tests/e2e/ui
|
||||
run: |
|
||||
retry() {
|
||||
local attempt=1
|
||||
local max_attempts=4
|
||||
until "$@"; do
|
||||
if [ "$attempt" -ge "$max_attempts" ]; then
|
||||
echo "Command failed after $attempt attempts: $*"
|
||||
return 1
|
||||
fi
|
||||
echo "Attempt $attempt failed: $*. Retrying in $((attempt * 15))s..."
|
||||
sleep $((attempt * 15))
|
||||
attempt=$((attempt + 1))
|
||||
done
|
||||
}
|
||||
|
||||
npm config set fetch-retries 5
|
||||
npm config set fetch-retry-mintimeout 20000
|
||||
npm config set fetch-retry-maxtimeout 120000
|
||||
|
||||
retry npm ci
|
||||
retry npx playwright install --with-deps chromium
|
||||
|
||||
- name: Run SERVER_ROOT_PATH redirect e2e
|
||||
working-directory: tests/e2e/ui
|
||||
env:
|
||||
SERVER_ROOT_PATH: ${{ matrix.root_path }}
|
||||
run: npx playwright test --config=serverRootPath.config.ts
|
||||
|
||||
- name: Upload Playwright artifacts on failure
|
||||
if: failure()
|
||||
uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4.6.2
|
||||
with:
|
||||
name: playwright-trace-${{ strategy.job-index }}
|
||||
path: tests/e2e/ui/test-results/
|
||||
retention-days: 7
|
||||
|
||||
- name: Cleanup
|
||||
if: always()
|
||||
run: |
|
||||
docker stop litellm-test || true
|
||||
docker rm litellm-test || true
|
||||
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -141,3 +141,4 @@ crash.*.log
|
|||
.coverage
|
||||
|
||||
ui/litellm-dashboard/out/
|
||||
litellm.log
|
||||
|
|
|
|||
2
Makefile
2
Makefile
|
|
@ -176,6 +176,8 @@ lint-ruff-FULL-dev: install-dev
|
|||
if [ -n "$$files" ]; then echo "$$files" | xargs $(UV_RUN) ruff check; \
|
||||
else echo "No changed .py files to check."; fi
|
||||
|
||||
lint-basedpyright lint-basedpyright-budget-update: export NODE_OPTIONS := --max-old-space-size=12288
|
||||
|
||||
lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging
|
||||
|
||||
|
|
|
|||
|
|
@ -70,6 +70,10 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/project/",
|
||||
"/memory/",
|
||||
"/mcp/",
|
||||
# Control plane (see the List Endpoints + Tables standard). Every resource
|
||||
# eventually moves under this prefix, so allowlist it once rather than
|
||||
# per-resource.
|
||||
"/management/v1/",
|
||||
# Spend / analytics
|
||||
"/spend/",
|
||||
"/analytics/",
|
||||
|
|
|
|||
|
|
@ -1,30 +1,30 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 37484
|
||||
"limit": 31903
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2704
|
||||
"limit": 2645
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 330
|
||||
"limit": 329
|
||||
},
|
||||
"reportAttributeAccessIssue": {
|
||||
"limit": 516
|
||||
},
|
||||
"reportCallIssue": {
|
||||
"limit": 124
|
||||
"limit": 123
|
||||
},
|
||||
"reportConstantRedefinition": {
|
||||
"limit": 59
|
||||
},
|
||||
"reportDeprecated": {
|
||||
"limit": 326
|
||||
"limit": 325
|
||||
},
|
||||
"reportDuplicateImport": {
|
||||
"limit": 42
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 10389
|
||||
"limit": 10214
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 11
|
||||
|
|
@ -54,10 +54,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5900
|
||||
"limit": 5869
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15903
|
||||
"limit": 15861
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 41
|
||||
|
|
@ -72,7 +72,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportOptionalMemberAccess": {
|
||||
"limit": 1085
|
||||
"limit": 1079
|
||||
},
|
||||
"reportOptionalOperand": {
|
||||
"limit": 0
|
||||
|
|
@ -84,13 +84,13 @@
|
|||
"limit": 77
|
||||
},
|
||||
"reportPrivateUsage": {
|
||||
"limit": 2438
|
||||
"limit": 2437
|
||||
},
|
||||
"reportRedeclaration": {
|
||||
"limit": 12
|
||||
},
|
||||
"reportReturnType": {
|
||||
"limit": 225
|
||||
"limit": 219
|
||||
},
|
||||
"reportTypedDictNotRequiredAccess": {
|
||||
"limit": 27
|
||||
|
|
@ -99,31 +99,31 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 45894
|
||||
"limit": 45366
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 113
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 40539
|
||||
"limit": 40477
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 20403
|
||||
"limit": 20338
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 32141
|
||||
"limit": 32047
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 177
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 1025
|
||||
"limit": 1021
|
||||
},
|
||||
"reportUnnecessaryContains": {
|
||||
"limit": 7
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 1209
|
||||
"limit": 1205
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 165
|
||||
|
|
@ -135,10 +135,10 @@
|
|||
"limit": 33
|
||||
},
|
||||
"reportUnusedFunction": {
|
||||
"limit": 206
|
||||
"limit": 204
|
||||
},
|
||||
"reportUnusedImport": {
|
||||
"limit": 1005
|
||||
"limit": 1003
|
||||
},
|
||||
"reportUnusedVariable": {
|
||||
"limit": 1297
|
||||
|
|
|
|||
325
ci_cd/generate_model_prices_schema.py
Normal file
325
ci_cd/generate_model_prices_schema.py
Normal file
|
|
@ -0,0 +1,325 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import jsonschema
|
||||
|
||||
REPO_ROOT = Path(__file__).parent.parent
|
||||
PRICES_PATH = REPO_ROOT / "model_prices_and_context_window.json"
|
||||
SCHEMA_PATH = REPO_ROOT / "model_prices_and_context_window.schema.json"
|
||||
|
||||
SPECIAL_ROOT_KEYS = frozenset({"sample_spec", "fallback_generalizations"})
|
||||
|
||||
JsonSchema = dict
|
||||
|
||||
NONNEG_NUMBER: JsonSchema = {"type": "number", "minimum": 0}
|
||||
NONNEG_INTEGER: JsonSchema = {"type": "integer", "minimum": 0}
|
||||
BOOLEAN: JsonSchema = {"type": "boolean"}
|
||||
STRING: JsonSchema = {"type": "string"}
|
||||
|
||||
EXTRA_BOOLEAN_KEYS = frozenset(
|
||||
{
|
||||
"gemini_native_audio",
|
||||
"gemini_audio_only_live",
|
||||
"uses_embed_content",
|
||||
"use_openai_responses_path",
|
||||
"bedrock_converse_supports_strict_tools",
|
||||
}
|
||||
)
|
||||
|
||||
OBJECT_KEYS: dict[str, JsonSchema] = {
|
||||
"search_context_cost_per_query": {
|
||||
"type": "object",
|
||||
"description": "USD cost per web search query, keyed by search context size.",
|
||||
"properties": {
|
||||
"search_context_size_low": NONNEG_NUMBER,
|
||||
"search_context_size_medium": NONNEG_NUMBER,
|
||||
"search_context_size_high": NONNEG_NUMBER,
|
||||
},
|
||||
"additionalProperties": False,
|
||||
},
|
||||
"metadata": {
|
||||
"type": "object",
|
||||
"description": "Free-form notes about the entry (e.g. pricing derivation).",
|
||||
},
|
||||
"provider_specific_entry": {
|
||||
"type": "object",
|
||||
"description": "Provider-internal routing hints (e.g. bedrock_invocation_schema).",
|
||||
},
|
||||
}
|
||||
|
||||
ARRAY_KEYS: dict[str, JsonSchema] = {
|
||||
"supported_endpoints": {
|
||||
"type": "array",
|
||||
"description": "OpenAI-style API routes this model can be called through, e.g. /v1/chat/completions.",
|
||||
"items": STRING,
|
||||
},
|
||||
"supported_modalities": {
|
||||
"type": "array",
|
||||
"description": "Input modalities the model accepts.",
|
||||
"items": {"type": "string", "enum": ["text", "image", "audio", "video"]},
|
||||
},
|
||||
"supported_output_modalities": {
|
||||
"type": "array",
|
||||
"description": "Output modalities the model can produce.",
|
||||
"items": {"type": "string", "enum": ["text", "image", "audio", "video", "code"]},
|
||||
},
|
||||
"supported_regions": {
|
||||
"type": "array",
|
||||
"description": "Cloud regions the model is available in ('global' or region ids).",
|
||||
"items": STRING,
|
||||
},
|
||||
"tiered_pricing": {
|
||||
"type": "array",
|
||||
"description": "Context-length or result-count tiered rates; each tier's costs apply within its range.",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"range": {
|
||||
"type": "array",
|
||||
"description": "[min, max] prompt-token span this tier applies to.",
|
||||
"items": NONNEG_NUMBER,
|
||||
"minItems": 2,
|
||||
"maxItems": 2,
|
||||
},
|
||||
"max_results_range": {
|
||||
"type": "array",
|
||||
"description": "[min, max] result-count span this tier applies to (search models).",
|
||||
"items": NONNEG_NUMBER,
|
||||
"minItems": 2,
|
||||
"maxItems": 2,
|
||||
},
|
||||
"input_cost_per_token": NONNEG_NUMBER,
|
||||
"output_cost_per_token": NONNEG_NUMBER,
|
||||
"output_cost_per_reasoning_token": NONNEG_NUMBER,
|
||||
"cache_read_input_token_cost": NONNEG_NUMBER,
|
||||
"input_cost_per_query": NONNEG_NUMBER,
|
||||
},
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
INTEGER_KEYS: dict[str, JsonSchema] = {
|
||||
"max_tokens": {
|
||||
**NONNEG_INTEGER,
|
||||
"description": "Legacy field: max output tokens if the provider specifies it, else max input tokens.",
|
||||
},
|
||||
"max_input_tokens": {
|
||||
**NONNEG_INTEGER,
|
||||
"description": "Maximum prompt/context tokens the model accepts.",
|
||||
},
|
||||
"max_output_tokens": {
|
||||
**NONNEG_INTEGER,
|
||||
"description": "Maximum tokens the model can generate in one response.",
|
||||
},
|
||||
"output_vector_size": {
|
||||
**NONNEG_INTEGER,
|
||||
"description": "Embedding dimension for embedding models.",
|
||||
},
|
||||
"prompt_cache_min_tokens": {
|
||||
**NONNEG_INTEGER,
|
||||
"description": "Smallest prefix the provider will actually cache; absent means the provider default applies.",
|
||||
},
|
||||
"tpm": {**NONNEG_INTEGER, "description": "Provider default tokens-per-minute limit."},
|
||||
"rpm": {**NONNEG_INTEGER, "description": "Provider default requests-per-minute limit."},
|
||||
}
|
||||
|
||||
NUMBER_KEYS: dict[str, JsonSchema] = {
|
||||
"regional_processing_uplift_multiplier_eu": {
|
||||
"type": "number",
|
||||
"minimum": 1,
|
||||
"description": "Multiplier applied to all token costs for EU data residency (e.g. 1.10 = +10%).",
|
||||
},
|
||||
"regional_processing_uplift_multiplier_us": {
|
||||
"type": "number",
|
||||
"minimum": 1,
|
||||
"description": "Multiplier applied to all token costs for US data residency (e.g. 1.10 = +10%).",
|
||||
},
|
||||
}
|
||||
|
||||
COST_DESCRIPTIONS: dict[str, str] = {
|
||||
"input_cost_per_token": "USD per prompt token.",
|
||||
"output_cost_per_token": "USD per generated token.",
|
||||
"output_cost_per_reasoning_token": "USD per reasoning/thinking token, when billed separately.",
|
||||
"cache_creation_input_token_cost": "USD per token written to the provider's prompt cache.",
|
||||
"cache_read_input_token_cost": "USD per prompt token served from the provider's prompt cache.",
|
||||
"input_cost_per_token_batches": "USD per prompt token via the provider's batch API.",
|
||||
"output_cost_per_token_batches": "USD per generated token via the provider's batch API.",
|
||||
}
|
||||
|
||||
|
||||
def cost_description(key: str) -> Optional[str]:
|
||||
if key in COST_DESCRIPTIONS:
|
||||
return COST_DESCRIPTIONS[key]
|
||||
if key.endswith("_flex"):
|
||||
return "Flex service-tier rate for the same-named base field."
|
||||
if key.endswith("_priority"):
|
||||
return "Priority service-tier rate for the same-named base field."
|
||||
if "_above_" in key:
|
||||
return "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
return None
|
||||
|
||||
|
||||
def cost_schema(key: str) -> JsonSchema:
|
||||
description = cost_description(key)
|
||||
return {**NONNEG_NUMBER, "description": description} if description else dict(NONNEG_NUMBER)
|
||||
|
||||
|
||||
def string_key_schemas(modes: tuple) -> dict[str, JsonSchema]:
|
||||
return {
|
||||
"litellm_provider": {
|
||||
"type": "string",
|
||||
"description": "LiteLLM provider slug; one of https://docs.litellm.ai/docs/providers.",
|
||||
},
|
||||
"mode": {
|
||||
"type": "string",
|
||||
"description": "Primary API surface / task type of the model.",
|
||||
"enum": list(modes),
|
||||
},
|
||||
"source": {
|
||||
"type": "string",
|
||||
"description": "URL of the provider pricing/model page this entry was taken from.",
|
||||
},
|
||||
"deprecation_date": {
|
||||
"type": "string",
|
||||
"description": "Date the provider deprecates the model, YYYY-MM-DD.",
|
||||
"format": "date",
|
||||
"pattern": "^\\d{4}-(0[1-9]|1[0-2])-(0[1-9]|[12]\\d|3[01])$",
|
||||
},
|
||||
"web_search_billing_unit": {
|
||||
"type": "string",
|
||||
"description": "Whether web search is billed per query or per prompt.",
|
||||
"enum": ["per_query", "per_prompt"],
|
||||
},
|
||||
"bedrock_output_config_effort_ceiling": {
|
||||
"type": "string",
|
||||
"description": "Highest reasoning effort the Bedrock output_config accepts for this model.",
|
||||
"enum": ["low", "medium", "high", "max", "xhigh"],
|
||||
},
|
||||
"comment": STRING,
|
||||
"audio_transcription_config": STRING,
|
||||
}
|
||||
|
||||
|
||||
def classify(key: str, modes: tuple) -> Optional[JsonSchema]:
|
||||
curated = {**OBJECT_KEYS, **ARRAY_KEYS, **string_key_schemas(modes), **INTEGER_KEYS, **NUMBER_KEYS}
|
||||
if key in curated:
|
||||
return curated[key]
|
||||
if key.startswith("supports_") or key in EXTRA_BOOLEAN_KEYS:
|
||||
return BOOLEAN
|
||||
if "cost" in key:
|
||||
return cost_schema(key)
|
||||
return None
|
||||
|
||||
|
||||
def build_schema(prices: dict) -> JsonSchema:
|
||||
entries = {name: entry for name, entry in prices.items() if name not in SPECIAL_ROOT_KEYS}
|
||||
all_keys = tuple(sorted({key for entry in entries.values() for key in entry}))
|
||||
modes = tuple(sorted({entry["mode"] for entry in entries.values() if "mode" in entry}))
|
||||
unclassified = tuple(key for key in all_keys if classify(key, modes) is None)
|
||||
if unclassified:
|
||||
raise SystemExit(
|
||||
f"Unclassified keys in {PRICES_PATH.name}: {', '.join(unclassified)}. "
|
||||
f"Add them to the key tables in {Path(__file__).name} and rerun it."
|
||||
)
|
||||
entry_properties = {key: classify(key, modes) for key in all_keys}
|
||||
return {
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"title": "LiteLLM model_prices_and_context_window.json",
|
||||
"description": (
|
||||
"Schema for LiteLLM's model price and context window registry "
|
||||
"(https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json). "
|
||||
"Every top-level key except 'sample_spec' and 'fallback_generalizations' is a model id, "
|
||||
"optionally prefixed with its provider (e.g. 'azure/gpt-5.4'), mapping to a model entry. "
|
||||
"All costs are USD per unit. New optional fields are added regularly, so consumers should "
|
||||
"ignore unknown fields rather than reject them."
|
||||
),
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"sample_spec": {
|
||||
"type": "object",
|
||||
"description": (
|
||||
"Documentation placeholder illustrating the entry shape; not a real model and not "
|
||||
"schema-conformant (several values are prose)."
|
||||
),
|
||||
},
|
||||
"fallback_generalizations": {
|
||||
"type": "object",
|
||||
"description": "Regex rules that generalize unknown model ids to known families; not a model entry.",
|
||||
"properties": {
|
||||
"rules": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": STRING,
|
||||
"pattern": STRING,
|
||||
"description": STRING,
|
||||
},
|
||||
"required": ["name", "pattern"],
|
||||
"additionalProperties": True,
|
||||
},
|
||||
}
|
||||
},
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
"additionalProperties": {"$ref": "#/$defs/modelEntry"},
|
||||
"$defs": {
|
||||
"modelEntry": {
|
||||
"type": "object",
|
||||
"description": (
|
||||
"Pricing, limits, and capability flags for one model. Fields other than litellm_provider "
|
||||
"are optional; boolean capability flags are simply omitted when unknown or false."
|
||||
),
|
||||
"required": ["litellm_provider"],
|
||||
"properties": entry_properties,
|
||||
"additionalProperties": True,
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def render(schema: JsonSchema) -> str:
|
||||
return json.dumps(schema, indent=2) + "\n"
|
||||
|
||||
|
||||
def validation_errors(prices: dict, schema: JsonSchema) -> tuple:
|
||||
validator = jsonschema.Draft202012Validator(
|
||||
schema, format_checker=jsonschema.Draft202012Validator.FORMAT_CHECKER
|
||||
)
|
||||
return tuple(
|
||||
f"{'.'.join(str(part) for part in error.absolute_path)}: {error.message}"
|
||||
for error in validator.iter_errors(prices)
|
||||
)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
check = "--check" in sys.argv[1:]
|
||||
prices = json.loads(PRICES_PATH.read_text())
|
||||
rendered = render(build_schema(prices))
|
||||
errors = validation_errors(prices, json.loads(rendered))
|
||||
if errors:
|
||||
print(f"{PRICES_PATH.name} does not validate against the generated schema:")
|
||||
print("\n".join(errors[:20]))
|
||||
return 1
|
||||
if not check:
|
||||
SCHEMA_PATH.write_text(rendered)
|
||||
print(f"wrote {SCHEMA_PATH}")
|
||||
return 0
|
||||
if not SCHEMA_PATH.exists() or SCHEMA_PATH.read_text() != rendered:
|
||||
print(
|
||||
f"{SCHEMA_PATH.name} is out of sync with {PRICES_PATH.name}. "
|
||||
f"Run `python {Path(__file__).relative_to(REPO_ROOT)}` and commit the result."
|
||||
)
|
||||
return 1
|
||||
print(f"{SCHEMA_PATH.name} is in sync and {PRICES_PATH.name} validates against it")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
|
|
@ -0,0 +1,523 @@
|
|||
{
|
||||
"annotations": {
|
||||
"list": []
|
||||
},
|
||||
"editable": true,
|
||||
"fiscalYearStartMonth": 0,
|
||||
"graphTooltip": 0,
|
||||
"links": [],
|
||||
"panels": [
|
||||
{
|
||||
"type": "stat",
|
||||
"title": "Requests",
|
||||
"gridPos": {
|
||||
"h": 4,
|
||||
"w": 6,
|
||||
"x": 0,
|
||||
"y": 0
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"fieldConfig": {
|
||||
"defaults": {
|
||||
"unit": "short",
|
||||
"decimals": 0,
|
||||
"color": {
|
||||
"mode": "fixed",
|
||||
"fixedColor": "blue"
|
||||
}
|
||||
},
|
||||
"overrides": []
|
||||
},
|
||||
"options": {
|
||||
"reduceOptions": {
|
||||
"calcs": [
|
||||
"lastNotNull"
|
||||
],
|
||||
"fields": "",
|
||||
"values": false
|
||||
},
|
||||
"colorMode": "background",
|
||||
"graphMode": "none"
|
||||
},
|
||||
"targets": [
|
||||
{
|
||||
"refId": "A",
|
||||
"editorMode": "code",
|
||||
"instant": true,
|
||||
"expr": "sum(increase(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))"
|
||||
}
|
||||
],
|
||||
"id": 1
|
||||
},
|
||||
{
|
||||
"type": "stat",
|
||||
"title": "Spend",
|
||||
"description": "LiteLLM's computed cost for the selected window, from gen_ai.usage.cost",
|
||||
"gridPos": {
|
||||
"h": 4,
|
||||
"w": 6,
|
||||
"x": 6,
|
||||
"y": 0
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"fieldConfig": {
|
||||
"defaults": {
|
||||
"unit": "currencyUSD",
|
||||
"decimals": 4,
|
||||
"color": {
|
||||
"mode": "fixed",
|
||||
"fixedColor": "green"
|
||||
}
|
||||
},
|
||||
"overrides": []
|
||||
},
|
||||
"options": {
|
||||
"reduceOptions": {
|
||||
"calcs": [
|
||||
"lastNotNull"
|
||||
],
|
||||
"fields": "",
|
||||
"values": false
|
||||
},
|
||||
"colorMode": "background",
|
||||
"graphMode": "none"
|
||||
},
|
||||
"targets": [
|
||||
{
|
||||
"refId": "A",
|
||||
"editorMode": "code",
|
||||
"instant": true,
|
||||
"expr": "sum(increase(gen_ai_usage_cost_USD_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))"
|
||||
}
|
||||
],
|
||||
"id": 2
|
||||
},
|
||||
{
|
||||
"type": "stat",
|
||||
"title": "Tokens",
|
||||
"gridPos": {
|
||||
"h": 4,
|
||||
"w": 6,
|
||||
"x": 12,
|
||||
"y": 0
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"fieldConfig": {
|
||||
"defaults": {
|
||||
"unit": "short",
|
||||
"decimals": 0,
|
||||
"color": {
|
||||
"mode": "fixed",
|
||||
"fixedColor": "purple"
|
||||
}
|
||||
},
|
||||
"overrides": []
|
||||
},
|
||||
"options": {
|
||||
"reduceOptions": {
|
||||
"calcs": [
|
||||
"lastNotNull"
|
||||
],
|
||||
"fields": "",
|
||||
"values": false
|
||||
},
|
||||
"colorMode": "background",
|
||||
"graphMode": "none"
|
||||
},
|
||||
"targets": [
|
||||
{
|
||||
"refId": "A",
|
||||
"editorMode": "code",
|
||||
"instant": true,
|
||||
"expr": "sum(increase(gen_ai_client_token_usage_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))"
|
||||
}
|
||||
],
|
||||
"id": 3
|
||||
},
|
||||
{
|
||||
"type": "stat",
|
||||
"title": "p95 request duration",
|
||||
"gridPos": {
|
||||
"h": 4,
|
||||
"w": 6,
|
||||
"x": 18,
|
||||
"y": 0
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"fieldConfig": {
|
||||
"defaults": {
|
||||
"unit": "s",
|
||||
"decimals": 2,
|
||||
"color": {
|
||||
"mode": "fixed",
|
||||
"fixedColor": "orange"
|
||||
}
|
||||
},
|
||||
"overrides": []
|
||||
},
|
||||
"options": {
|
||||
"reduceOptions": {
|
||||
"calcs": [
|
||||
"lastNotNull"
|
||||
],
|
||||
"fields": "",
|
||||
"values": false
|
||||
},
|
||||
"colorMode": "background",
|
||||
"graphMode": "none"
|
||||
},
|
||||
"targets": [
|
||||
{
|
||||
"refId": "A",
|
||||
"editorMode": "code",
|
||||
"instant": true,
|
||||
"expr": "histogram_quantile(0.95, sum by (le) (rate(gen_ai_client_operation_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range])))"
|
||||
}
|
||||
],
|
||||
"id": 4
|
||||
},
|
||||
{
|
||||
"type": "timeseries",
|
||||
"title": "Request rate by model",
|
||||
"gridPos": {
|
||||
"h": 8,
|
||||
"w": 12,
|
||||
"x": 0,
|
||||
"y": 4
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"fieldConfig": {
|
||||
"defaults": {
|
||||
"unit": "reqpm",
|
||||
"custom": {
|
||||
"lineWidth": 2,
|
||||
"fillOpacity": 8,
|
||||
"showPoints": "never"
|
||||
}
|
||||
},
|
||||
"overrides": []
|
||||
},
|
||||
"options": {
|
||||
"legend": {
|
||||
"displayMode": "list",
|
||||
"placement": "bottom"
|
||||
},
|
||||
"tooltip": {
|
||||
"mode": "multi",
|
||||
"sort": "desc"
|
||||
}
|
||||
},
|
||||
"targets": [
|
||||
{
|
||||
"refId": "A",
|
||||
"editorMode": "code",
|
||||
"legendFormat": "{{gen_ai_request_model}}",
|
||||
"expr": "sum by (gen_ai_request_model) (rate(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 60"
|
||||
}
|
||||
],
|
||||
"id": 5
|
||||
},
|
||||
{
|
||||
"type": "timeseries",
|
||||
"title": "Spend rate by model",
|
||||
"description": "USD per hour, derived from the gen_ai.usage.cost histogram",
|
||||
"gridPos": {
|
||||
"h": 8,
|
||||
"w": 12,
|
||||
"x": 12,
|
||||
"y": 4
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"fieldConfig": {
|
||||
"defaults": {
|
||||
"unit": "currencyUSD",
|
||||
"custom": {
|
||||
"lineWidth": 2,
|
||||
"fillOpacity": 8,
|
||||
"showPoints": "never"
|
||||
}
|
||||
},
|
||||
"overrides": []
|
||||
},
|
||||
"options": {
|
||||
"legend": {
|
||||
"displayMode": "list",
|
||||
"placement": "bottom"
|
||||
},
|
||||
"tooltip": {
|
||||
"mode": "multi",
|
||||
"sort": "desc"
|
||||
}
|
||||
},
|
||||
"targets": [
|
||||
{
|
||||
"refId": "A",
|
||||
"editorMode": "code",
|
||||
"legendFormat": "{{gen_ai_request_model}}",
|
||||
"expr": "sum by (gen_ai_request_model) (rate(gen_ai_usage_cost_USD_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 3600"
|
||||
}
|
||||
],
|
||||
"id": 6
|
||||
},
|
||||
{
|
||||
"type": "timeseries",
|
||||
"title": "Tokens per minute by model and type",
|
||||
"description": "gen_ai.client.token.usage split by the gen_ai.token.type attribute",
|
||||
"gridPos": {
|
||||
"h": 8,
|
||||
"w": 12,
|
||||
"x": 0,
|
||||
"y": 12
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"fieldConfig": {
|
||||
"defaults": {
|
||||
"unit": "short",
|
||||
"custom": {
|
||||
"lineWidth": 2,
|
||||
"fillOpacity": 8,
|
||||
"showPoints": "never"
|
||||
}
|
||||
},
|
||||
"overrides": []
|
||||
},
|
||||
"options": {
|
||||
"legend": {
|
||||
"displayMode": "list",
|
||||
"placement": "bottom"
|
||||
},
|
||||
"tooltip": {
|
||||
"mode": "multi",
|
||||
"sort": "desc"
|
||||
}
|
||||
},
|
||||
"targets": [
|
||||
{
|
||||
"refId": "A",
|
||||
"editorMode": "code",
|
||||
"legendFormat": "{{gen_ai_request_model}} {{gen_ai_token_type}}",
|
||||
"expr": "sum by (gen_ai_request_model, gen_ai_token_type) (rate(gen_ai_client_token_usage_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 60"
|
||||
}
|
||||
],
|
||||
"id": 7
|
||||
},
|
||||
{
|
||||
"type": "timeseries",
|
||||
"title": "p95 request duration by model",
|
||||
"gridPos": {
|
||||
"h": 8,
|
||||
"w": 12,
|
||||
"x": 12,
|
||||
"y": 12
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"fieldConfig": {
|
||||
"defaults": {
|
||||
"unit": "s",
|
||||
"custom": {
|
||||
"lineWidth": 2,
|
||||
"fillOpacity": 0,
|
||||
"showPoints": "never"
|
||||
}
|
||||
},
|
||||
"overrides": []
|
||||
},
|
||||
"options": {
|
||||
"legend": {
|
||||
"displayMode": "list",
|
||||
"placement": "bottom"
|
||||
},
|
||||
"tooltip": {
|
||||
"mode": "multi",
|
||||
"sort": "desc"
|
||||
}
|
||||
},
|
||||
"targets": [
|
||||
{
|
||||
"refId": "A",
|
||||
"editorMode": "code",
|
||||
"legendFormat": "{{gen_ai_request_model}}",
|
||||
"expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_client_operation_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))"
|
||||
}
|
||||
],
|
||||
"id": 8
|
||||
},
|
||||
{
|
||||
"type": "timeseries",
|
||||
"title": "p95 time to first token (streaming)",
|
||||
"description": "gen_ai.server.time_to_first_token, recorded only for streaming requests",
|
||||
"gridPos": {
|
||||
"h": 8,
|
||||
"w": 12,
|
||||
"x": 0,
|
||||
"y": 20
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"fieldConfig": {
|
||||
"defaults": {
|
||||
"unit": "s",
|
||||
"custom": {
|
||||
"lineWidth": 2,
|
||||
"fillOpacity": 0,
|
||||
"showPoints": "never"
|
||||
}
|
||||
},
|
||||
"overrides": []
|
||||
},
|
||||
"options": {
|
||||
"legend": {
|
||||
"displayMode": "list",
|
||||
"placement": "bottom"
|
||||
},
|
||||
"tooltip": {
|
||||
"mode": "multi",
|
||||
"sort": "desc"
|
||||
}
|
||||
},
|
||||
"targets": [
|
||||
{
|
||||
"refId": "A",
|
||||
"editorMode": "code",
|
||||
"legendFormat": "{{gen_ai_request_model}}",
|
||||
"expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_server_time_to_first_token_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))"
|
||||
}
|
||||
],
|
||||
"id": 9
|
||||
},
|
||||
{
|
||||
"type": "timeseries",
|
||||
"title": "p95 provider generation time",
|
||||
"description": "gen_ai.client.response.duration, upstream generation time excluding LiteLLM overhead",
|
||||
"gridPos": {
|
||||
"h": 8,
|
||||
"w": 12,
|
||||
"x": 12,
|
||||
"y": 20
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"fieldConfig": {
|
||||
"defaults": {
|
||||
"unit": "s",
|
||||
"custom": {
|
||||
"lineWidth": 2,
|
||||
"fillOpacity": 0,
|
||||
"showPoints": "never"
|
||||
}
|
||||
},
|
||||
"overrides": []
|
||||
},
|
||||
"options": {
|
||||
"legend": {
|
||||
"displayMode": "list",
|
||||
"placement": "bottom"
|
||||
},
|
||||
"tooltip": {
|
||||
"mode": "multi",
|
||||
"sort": "desc"
|
||||
}
|
||||
},
|
||||
"targets": [
|
||||
{
|
||||
"refId": "A",
|
||||
"editorMode": "code",
|
||||
"legendFormat": "{{gen_ai_request_model}}",
|
||||
"expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_client_response_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))"
|
||||
}
|
||||
],
|
||||
"id": 10
|
||||
}
|
||||
],
|
||||
"preload": false,
|
||||
"refresh": "30s",
|
||||
"schemaVersion": 42,
|
||||
"tags": [
|
||||
"litellm",
|
||||
"genai",
|
||||
"opentelemetry"
|
||||
],
|
||||
"templating": {
|
||||
"list": [
|
||||
{
|
||||
"name": "datasource",
|
||||
"label": "Prometheus",
|
||||
"type": "datasource",
|
||||
"query": "prometheus",
|
||||
"current": {},
|
||||
"hide": 0
|
||||
},
|
||||
{
|
||||
"name": "service",
|
||||
"label": "Service",
|
||||
"type": "query",
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"query": "label_values(gen_ai_client_operation_duration_seconds_count, service_name)",
|
||||
"refresh": 2,
|
||||
"includeAll": true,
|
||||
"multi": true,
|
||||
"current": {
|
||||
"text": "All",
|
||||
"value": "$__all"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "model",
|
||||
"label": "Model",
|
||||
"type": "query",
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${datasource}"
|
||||
},
|
||||
"query": "label_values(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\"}, gen_ai_request_model)",
|
||||
"refresh": 2,
|
||||
"includeAll": true,
|
||||
"multi": true,
|
||||
"current": {
|
||||
"text": "All",
|
||||
"value": "$__all"
|
||||
}
|
||||
}
|
||||
]
|
||||
},
|
||||
"time": {
|
||||
"from": "now-1h",
|
||||
"to": "now"
|
||||
},
|
||||
"timepicker": {},
|
||||
"timezone": "browser",
|
||||
"title": "LiteLLM GenAI (OpenTelemetry)",
|
||||
"uid": "litellm-genai-otel",
|
||||
"weekStart": ""
|
||||
}
|
||||
|
|
@ -0,0 +1,35 @@
|
|||
# LiteLLM GenAI dashboard (OpenTelemetry metrics)
|
||||
|
||||
Dashboard for the `gen_ai.*` metrics the OpenTelemetry v2 integration emits, as opposed to the `litellm_*` Prometheus metrics the other dashboards in this folder chart.
|
||||
|
||||
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source. Panels: request count, spend, token count, p95 duration, request rate by model, spend rate per hour by model, tokens per minute split by input and output, p95 duration by model, p95 time to first token, and p95 provider generation time. Template variables for data source, service, and model.
|
||||
|
||||
## Pre-requisites
|
||||
|
||||
Metrics are off by default. In the proxy environment:
|
||||
|
||||
```shell
|
||||
LITELLM_OTEL_V2=true
|
||||
LITELLM_OTEL_INTEGRATION_ENABLE_METRICS=true
|
||||
OTEL_EXPORTER="otlp_http"
|
||||
OTEL_ENDPOINT="<your OTLP endpoint>"
|
||||
```
|
||||
|
||||
You also need the metric attribute filter, or the panels will plot flat lines at zero. LiteLLM's default attribute set includes per-request fields, so nearly every request lands in its own time series with a single sample, and `rate()` has nothing to compute over:
|
||||
|
||||
```yaml title="config.yaml"
|
||||
callback_settings:
|
||||
otel:
|
||||
attributes:
|
||||
include_list:
|
||||
- gen_ai.operation.name
|
||||
- gen_ai.system
|
||||
- gen_ai.request.model
|
||||
- gen_ai.framework
|
||||
```
|
||||
|
||||
See [Grafana Cloud](https://docs.litellm.ai/docs/observability/grafana_cloud) for the full setup, and [OpenTelemetry v2](https://docs.litellm.ai/docs/observability/opentelemetry_v2#metrics) for the metric reference.
|
||||
|
||||
## Note on Grafana's AI Observability integration
|
||||
|
||||
Grafana Cloud ships prebuilt GenAI dashboards that query these same metric names, so they look like a drop-in alternative to this one. They are not: twenty of their twenty-two panels filter on `telemetry_sdk_name="openlit"`, a label LiteLLM does not carry and cannot be configured to add, so those panels stay empty.
|
||||
|
|
@ -2,6 +2,10 @@
|
|||
|
||||
This folder contains the `json` for creating Grafana Dashboards
|
||||
|
||||
## [LiteLLM GenAI Dashboard (OpenTelemetry)](./dashboard_genai_otel)
|
||||
|
||||
Charts the `gen_ai.*` metrics from the OpenTelemetry v2 integration: spend, tokens, request rate, and latency percentiles by model. Separate from the dashboards below, which chart the `litellm_*` Prometheus metrics.
|
||||
|
||||
## [LiteLLM v2 Dashboard](./dashboard_v2)
|
||||
|
||||
<img width="1316" alt="grafana_1" src="https://github.com/user-attachments/assets/d0df802d-0cb9-4906-a679-941c547789ab">
|
||||
|
|
|
|||
46
db_scripts/backfill_daily_tool_spend.sql
Normal file
46
db_scripts/backfill_daily_tool_spend.sql
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
-- One-shot backfill of the LiteLLM_DailyToolSpend rollup from the per-request
|
||||
-- LiteLLM_SpendLogToolIndex x LiteLLM_SpendLogs tables.
|
||||
--
|
||||
-- This is an opt-in, manual operation. New deployments do not need it: the
|
||||
-- rollup is written at request time from the moment the release is deployed.
|
||||
-- Run it only if you want the Cost Optimization "Spend by tool" card to show
|
||||
-- history from before the deploy, and only once.
|
||||
--
|
||||
-- IMPORTANT caveats before running:
|
||||
--
|
||||
-- 1. Pre-deploy index rows may include tools that were merely DECLARED in a
|
||||
-- request body but never invoked (the release this ships with stops
|
||||
-- recording those). For agentic clients that declare many tools per
|
||||
-- request, backfilled history attributes each request's full spend to
|
||||
-- every declared tool, overstating per-tool spend. Post-deploy rows do not
|
||||
-- have this problem. If your traffic is mostly such clients, consider not
|
||||
-- backfilling.
|
||||
--
|
||||
-- 2. Coverage is bounded by spend-log retention: rows older than
|
||||
-- maximum_spend_logs_retention_period are already gone.
|
||||
--
|
||||
-- 3. Replace the cutover timestamp below with the time you deployed the
|
||||
-- release, so backfilled per-request rows cannot double-count on top of
|
||||
-- rollup rows the new writer already created. ON CONFLICT DO NOTHING is a
|
||||
-- second guard for (date, tool_name) buckets the writer already touched:
|
||||
-- such buckets keep the writer's numbers and skip the backfill's.
|
||||
--
|
||||
-- Usage:
|
||||
-- psql "$DATABASE_URL" -v cutover="'2026-07-25T00:00:00Z'" -f db_scripts/backfill_daily_tool_spend.sql
|
||||
|
||||
SET TIME ZONE 'UTC';
|
||||
|
||||
INSERT INTO "LiteLLM_DailyToolSpend" (date, tool_name, spend, total_tokens, request_count, created_at, updated_at)
|
||||
SELECT
|
||||
to_char(ti.start_time, 'YYYY-MM-DD') AS date,
|
||||
ti.tool_name,
|
||||
COALESCE(SUM(sl.spend), 0) AS spend,
|
||||
COALESCE(SUM(sl.total_tokens), 0) AS total_tokens,
|
||||
COUNT(*) AS request_count,
|
||||
now() AS created_at,
|
||||
now() AS updated_at
|
||||
FROM "LiteLLM_SpendLogToolIndex" ti
|
||||
JOIN "LiteLLM_SpendLogs" sl ON sl.request_id = ti.request_id
|
||||
WHERE ti.start_time < :cutover::timestamptz
|
||||
GROUP BY 1, 2
|
||||
ON CONFLICT (date, tool_name) DO NOTHING;
|
||||
|
|
@ -316,26 +316,34 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
where_clause: Dict[str, Any] = {"file_purpose": "batch", **owner_filter}
|
||||
|
||||
if after:
|
||||
where_clause["id"] = {"gt": after}
|
||||
cursor_row = (
|
||||
await self.prisma_client.db.litellm_managedobjecttable.find_first(
|
||||
where={**where_clause, "unified_object_id": after}
|
||||
)
|
||||
)
|
||||
if cursor_row is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Invalid 'after' cursor: no batch found with id '{after}'.",
|
||||
)
|
||||
|
||||
fetch_limit = limit or 20
|
||||
if target_model_names:
|
||||
# Oversample so post-fetch model-name filtering still has enough rows.
|
||||
fetch_limit = max(fetch_limit * 3, 100)
|
||||
page_size = limit or 20
|
||||
cursor_args: Dict[str, Any] = (
|
||||
{"cursor": {"unified_object_id": after}, "skip": 1} if after else {}
|
||||
)
|
||||
|
||||
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
|
||||
where=where_clause,
|
||||
take=fetch_limit,
|
||||
order={"created_at": "desc"},
|
||||
take=page_size + 1,
|
||||
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
|
||||
**cursor_args,
|
||||
)
|
||||
|
||||
batch_objects: List[LiteLLMBatch] = []
|
||||
for batch in batches:
|
||||
try:
|
||||
# Stop once we have enough after filtering
|
||||
if len(batch_objects) >= (limit or 20):
|
||||
break
|
||||
has_more = len(batches) > page_size
|
||||
|
||||
batch_objects: List[LiteLLMBatch] = []
|
||||
for batch in batches[:page_size]:
|
||||
try:
|
||||
batch_data = (
|
||||
json.loads(batch.file_object)
|
||||
if isinstance(batch.file_object, str)
|
||||
|
|
@ -351,9 +359,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
)
|
||||
continue
|
||||
|
||||
return build_list_page(
|
||||
batch_objects, has_more=len(batch_objects) == (limit or 20)
|
||||
)
|
||||
return build_list_page(batch_objects, has_more=has_more)
|
||||
|
||||
async def get_user_created_file_ids(
|
||||
self, user_api_key_dict: UserAPIKeyAuth, model_object_ids: List[str]
|
||||
|
|
|
|||
|
|
@ -11,7 +11,8 @@ Endpoints for /project operations
|
|||
#### PROJECT MANAGEMENT ####
|
||||
|
||||
import json
|
||||
from typing import List, Optional, Union
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
|
||||
|
|
@ -25,15 +26,24 @@ from litellm.proxy.management_helpers.utils import (
|
|||
)
|
||||
from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_models
|
||||
from prisma.actions import LiteLLM_TeamTableActions
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _team_table(prisma_client: PrismaClient) -> "LiteLLM_TeamTableActions[prisma_models.LiteLLM_TeamTable]":
|
||||
team_table: LiteLLM_TeamTableActions[prisma_models.LiteLLM_TeamTable] = prisma_client.db.litellm_teamtable
|
||||
return team_table
|
||||
|
||||
|
||||
async def _check_user_permission_for_project(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
team_id: Optional[str],
|
||||
team_id: str | None,
|
||||
prisma_client: PrismaClient,
|
||||
require_admin: bool = False,
|
||||
team_object: Optional[LiteLLM_TeamTable] = None,
|
||||
team_object: LiteLLM_TeamTable | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if user has permission to manage a project.
|
||||
|
|
@ -57,9 +67,7 @@ async def _check_user_permission_for_project(
|
|||
|
||||
team = team_object
|
||||
if team is None:
|
||||
team = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": team_id}
|
||||
)
|
||||
team = await _team_table(prisma_client).find_unique(where={"team_id": team_id})
|
||||
|
||||
if team and team.admins:
|
||||
return user_api_key_dict.user_id in team.admins
|
||||
|
|
@ -70,9 +78,9 @@ async def _check_user_permission_for_project(
|
|||
async def _validate_team_exists(
|
||||
team_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
):
|
||||
) -> "prisma_models.LiteLLM_TeamTable":
|
||||
"""Validate that a team exists. Returns the team row."""
|
||||
team = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
team = await _team_table(prisma_client).find_unique(
|
||||
where={"team_id": team_id},
|
||||
)
|
||||
|
||||
|
|
@ -89,7 +97,7 @@ async def _validate_team_exists(
|
|||
|
||||
def _check_team_project_limits(
|
||||
team_object: LiteLLM_TeamTable,
|
||||
data: Union[NewProjectRequest, UpdateProjectRequest],
|
||||
data: NewProjectRequest | UpdateProjectRequest,
|
||||
) -> None:
|
||||
"""
|
||||
Check that project limits respect its parent Team's limits.
|
||||
|
|
@ -108,16 +116,12 @@ def _check_team_project_limits(
|
|||
if data.max_budget is not None and data.max_budget < 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"max_budget cannot be negative. Received: {data.max_budget}"
|
||||
},
|
||||
detail={"error": f"max_budget cannot be negative. Received: {data.max_budget}"},
|
||||
)
|
||||
if data.soft_budget is not None and data.soft_budget < 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"
|
||||
},
|
||||
detail={"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"},
|
||||
)
|
||||
|
||||
# --- soft_budget < max_budget ---
|
||||
|
|
@ -131,7 +135,7 @@ def _check_team_project_limits(
|
|||
)
|
||||
|
||||
# --- Validate project models are a subset of team models ---
|
||||
project_models = getattr(data, "models", None)
|
||||
project_models = data.models
|
||||
team_models = team_object.models or []
|
||||
if project_models and len(team_models) > 0:
|
||||
# If team has 'all-proxy-models', skip validation as it allows all models
|
||||
|
|
@ -148,11 +152,7 @@ def _check_team_project_limits(
|
|||
# --- Validate project max_budget <= team max_budget ---
|
||||
# Team stores budget fields directly (max_budget, tpm_limit, rpm_limit)
|
||||
# unlike Project which uses a separate LiteLLM_BudgetTable relation
|
||||
if (
|
||||
data.max_budget is not None
|
||||
and team_object.max_budget is not None
|
||||
and data.max_budget > team_object.max_budget
|
||||
):
|
||||
if data.max_budget is not None and team_object.max_budget is not None and data.max_budget > team_object.max_budget:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
|
|
@ -161,11 +161,7 @@ def _check_team_project_limits(
|
|||
)
|
||||
|
||||
# --- Validate project tpm_limit <= team tpm_limit ---
|
||||
if (
|
||||
data.tpm_limit is not None
|
||||
and team_object.tpm_limit is not None
|
||||
and data.tpm_limit > team_object.tpm_limit
|
||||
):
|
||||
if data.tpm_limit is not None and team_object.tpm_limit is not None and data.tpm_limit > team_object.tpm_limit:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
|
|
@ -174,11 +170,7 @@ def _check_team_project_limits(
|
|||
)
|
||||
|
||||
# --- Validate project rpm_limit <= team rpm_limit ---
|
||||
if (
|
||||
data.rpm_limit is not None
|
||||
and team_object.rpm_limit is not None
|
||||
and data.rpm_limit > team_object.rpm_limit
|
||||
):
|
||||
if data.rpm_limit is not None and team_object.rpm_limit is not None and data.rpm_limit > team_object.rpm_limit:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
|
|
@ -189,19 +181,19 @@ def _check_team_project_limits(
|
|||
|
||||
async def _create_budget_for_project(
|
||||
data: NewProjectRequest,
|
||||
user_id: Optional[str],
|
||||
user_id: str | None,
|
||||
litellm_proxy_admin_name: str,
|
||||
prisma_client: PrismaClient,
|
||||
) -> str:
|
||||
"""Create a budget for the project and return budget_id."""
|
||||
budget_params = LiteLLM_BudgetTable.model_fields.keys()
|
||||
_json_data = data.json(exclude_none=True)
|
||||
_json_data: Mapping[str, object] = data.json(exclude_none=True)
|
||||
_budget_data = {k: v for k, v in _json_data.items() if k in budget_params}
|
||||
budget_row = LiteLLM_BudgetTable(**_budget_data)
|
||||
budget_row = LiteLLM_BudgetTable.model_validate(_budget_data)
|
||||
|
||||
new_budget = prisma_client.jsonify_object(budget_row.json(exclude_none=True))
|
||||
|
||||
_budget = await prisma_client.db.litellm_budgettable.create(
|
||||
_budget: prisma_models.LiteLLM_BudgetTable = await prisma_client.db.litellm_budgettable.create(
|
||||
data={
|
||||
**new_budget,
|
||||
"created_by": user_id or litellm_proxy_admin_name,
|
||||
|
|
@ -214,8 +206,8 @@ async def _create_budget_for_project(
|
|||
|
||||
async def _set_project_object_permission(
|
||||
data: NewProjectRequest,
|
||||
prisma_client: Optional[PrismaClient],
|
||||
) -> Optional[str]:
|
||||
prisma_client: PrismaClient | None,
|
||||
) -> str | None:
|
||||
"""
|
||||
Creates the LiteLLM_ObjectPermissionTable record for the project.
|
||||
Returns the object_permission_id if created, otherwise None.
|
||||
|
|
@ -224,7 +216,7 @@ async def _set_project_object_permission(
|
|||
return None
|
||||
|
||||
if data.object_permission is not None:
|
||||
created_object_permission = (
|
||||
created_object_permission: prisma_models.LiteLLM_ObjectPermissionTable = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=data.object_permission.model_dump(exclude_none=True),
|
||||
)
|
||||
|
|
@ -344,8 +336,7 @@ async def new_project(
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Only premium users can add tags to projects. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
"error": "Only premium users can add tags to projects. " + CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -353,8 +344,7 @@ async def new_project(
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Project management is an enterprise feature. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
"error": "Project management is an enterprise feature. " + CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -375,13 +365,11 @@ async def new_project(
|
|||
)
|
||||
|
||||
# Validate team exists and get team object with budget
|
||||
team_object = await _validate_team_exists(
|
||||
team_id=data.team_id, prisma_client=prisma_client
|
||||
)
|
||||
team_object = await _validate_team_exists(team_id=data.team_id, prisma_client=prisma_client)
|
||||
|
||||
# Validate project limits against team limits
|
||||
_check_team_project_limits(
|
||||
team_object=LiteLLM_TeamTable(**team_object.model_dump()),
|
||||
team_object=LiteLLM_TeamTable.model_validate(team_object.model_dump()),
|
||||
data=data,
|
||||
)
|
||||
|
||||
|
|
@ -391,7 +379,7 @@ async def new_project(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
team_id=data.team_id,
|
||||
prisma_client=prisma_client,
|
||||
team_object=LiteLLM_TeamTable(**team_object.model_dump()),
|
||||
team_object=LiteLLM_TeamTable.model_validate(team_object.model_dump()),
|
||||
)
|
||||
|
||||
if not has_permission:
|
||||
|
|
@ -449,17 +437,13 @@ async def new_project(
|
|||
value=getattr(data, field),
|
||||
)
|
||||
|
||||
new_project_row = prisma_client.jsonify_object(
|
||||
project_row.json(exclude_none=True)
|
||||
)
|
||||
new_project_row = prisma_client.jsonify_object(project_row.json(exclude_none=True))
|
||||
|
||||
# Remove budget fields (following organization_endpoints.py pattern)
|
||||
new_project_row = _remove_budget_fields_from_project_data(new_project_row)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"new_project_row: {json.dumps(new_project_row, indent=2)}"
|
||||
)
|
||||
response = await prisma_client.db.litellm_projecttable.create(
|
||||
verbose_proxy_logger.info(f"new_project_row: {json.dumps(new_project_row, indent=2)}")
|
||||
response: prisma_models.LiteLLM_ProjectTable = await prisma_client.db.litellm_projecttable.create(
|
||||
data={
|
||||
**new_project_row, # type: ignore
|
||||
},
|
||||
|
|
@ -469,9 +453,7 @@ async def new_project(
|
|||
return response
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.project_endpoints.new_project(): Exception occured - {}".format(
|
||||
str(e)
|
||||
)
|
||||
"litellm.proxy.management_endpoints.project_endpoints.new_project(): Exception occured - {}".format(str(e))
|
||||
)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
|
@ -539,8 +521,7 @@ async def update_project(
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Only premium users can add tags to projects. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
"error": "Only premium users can add tags to projects. " + CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -548,8 +529,7 @@ async def update_project(
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Project management is an enterprise feature. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
"error": "Project management is an enterprise feature. " + CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -576,9 +556,9 @@ async def update_project(
|
|||
)
|
||||
|
||||
# Fetch existing project
|
||||
existing_project = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
where={"project_id": data.project_id}
|
||||
)
|
||||
existing_project: (
|
||||
prisma_models.LiteLLM_ProjectTable | None
|
||||
) = await prisma_client.db.litellm_projecttable.find_unique(where={"project_id": data.project_id})
|
||||
|
||||
if existing_project is None:
|
||||
raise ProxyException(
|
||||
|
|
@ -595,9 +575,7 @@ async def update_project(
|
|||
target_team_id = data.team_id or existing_project.team_id
|
||||
target_team_obj = None
|
||||
if target_team_id is not None:
|
||||
target_team_obj = await _validate_team_exists(
|
||||
team_id=target_team_id, prisma_client=prisma_client
|
||||
)
|
||||
target_team_obj = await _validate_team_exists(team_id=target_team_id, prisma_client=prisma_client)
|
||||
|
||||
has_permission = await _check_user_permission_for_project(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -620,32 +598,26 @@ async def update_project(
|
|||
team_id=data.team_id,
|
||||
prisma_client=prisma_client,
|
||||
team_object=(
|
||||
LiteLLM_TeamTable(**target_team_obj.model_dump())
|
||||
if target_team_obj
|
||||
else None
|
||||
LiteLLM_TeamTable.model_validate(target_team_obj.model_dump()) if target_team_obj else None
|
||||
),
|
||||
)
|
||||
if not can_assign_to_target:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Cannot reassign project to a team you are not an admin of"
|
||||
},
|
||||
detail={"error": "Cannot reassign project to a team you are not an admin of"},
|
||||
)
|
||||
|
||||
# Validate project limits against team limits
|
||||
if target_team_obj is not None:
|
||||
_check_team_project_limits(
|
||||
team_object=LiteLLM_TeamTable(**target_team_obj.model_dump()),
|
||||
team_object=LiteLLM_TeamTable.model_validate(target_team_obj.model_dump()),
|
||||
data=data,
|
||||
)
|
||||
|
||||
# Prepare update data
|
||||
update_data = data.json(exclude_none=True, exclude={"project_id"})
|
||||
update_data = prisma_client.jsonify_object(update_data)
|
||||
update_data["updated_by"] = (
|
||||
user_api_key_dict.user_id or litellm_proxy_admin_name
|
||||
)
|
||||
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
|
||||
|
||||
# Handle budget updates
|
||||
budget_fields = LiteLLM_BudgetTable.model_fields.keys()
|
||||
|
|
@ -671,21 +643,17 @@ async def update_project(
|
|||
if existing_project.object_permission_id:
|
||||
# Update existing permission
|
||||
await prisma_client.db.litellm_objectpermissiontable.update(
|
||||
where={
|
||||
"object_permission_id": existing_project.object_permission_id
|
||||
},
|
||||
where={"object_permission_id": existing_project.object_permission_id},
|
||||
data=object_permission_data,
|
||||
)
|
||||
else:
|
||||
# Create new permission
|
||||
created_permission = (
|
||||
created_permission: prisma_models.LiteLLM_ObjectPermissionTable = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=object_permission_data,
|
||||
)
|
||||
)
|
||||
update_data["object_permission_id"] = (
|
||||
created_permission.object_permission_id
|
||||
)
|
||||
update_data["object_permission_id"] = created_permission.object_permission_id
|
||||
|
||||
# Handle metadata fields
|
||||
for field in LiteLLM_ManagementEndpoint_MetadataFields:
|
||||
|
|
@ -698,7 +666,7 @@ async def update_project(
|
|||
update_data = _remove_budget_fields_from_project_data(update_data)
|
||||
|
||||
# Update project
|
||||
updated_project = await prisma_client.db.litellm_projecttable.update(
|
||||
updated_project: prisma_models.LiteLLM_ProjectTable | None = await prisma_client.db.litellm_projecttable.update(
|
||||
where={"project_id": data.project_id},
|
||||
data=update_data,
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
|
|
@ -718,7 +686,7 @@ async def update_project(
|
|||
"/project/delete",
|
||||
tags=["project management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[LiteLLM_ProjectTable],
|
||||
response_model=list[LiteLLM_ProjectTable],
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def delete_project(
|
||||
|
|
@ -749,8 +717,7 @@ async def delete_project(
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Project management is an enterprise feature. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
"error": "Project management is an enterprise feature. " + CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -778,9 +745,7 @@ async def delete_project(
|
|||
|
||||
for project_id in data.project_ids:
|
||||
# Check if project exists
|
||||
existing_project = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
where={"project_id": project_id}
|
||||
)
|
||||
existing_project = await prisma_client.db.litellm_projecttable.find_unique(where={"project_id": project_id})
|
||||
|
||||
if existing_project is None:
|
||||
raise ProxyException(
|
||||
|
|
@ -791,11 +756,9 @@ async def delete_project(
|
|||
)
|
||||
|
||||
# Check if there are any keys associated with this project
|
||||
associated_keys = (
|
||||
await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={"project_id": project_id}
|
||||
)
|
||||
)
|
||||
associated_keys: Sequence[
|
||||
prisma_models.LiteLLM_VerificationToken
|
||||
] = await prisma_client.db.litellm_verificationtoken.find_many(where={"project_id": project_id})
|
||||
|
||||
if len(associated_keys) > 0:
|
||||
raise ProxyException(
|
||||
|
|
@ -806,9 +769,9 @@ async def delete_project(
|
|||
)
|
||||
|
||||
# Delete the project
|
||||
deleted_project = await prisma_client.db.litellm_projecttable.delete(
|
||||
where={"project_id": project_id}
|
||||
)
|
||||
deleted_project: (
|
||||
prisma_models.LiteLLM_ProjectTable | None
|
||||
) = await prisma_client.db.litellm_projecttable.delete(where={"project_id": project_id})
|
||||
|
||||
deleted_projects.append(deleted_project)
|
||||
|
||||
|
|
@ -854,7 +817,7 @@ async def project_info(
|
|||
)
|
||||
|
||||
# Fetch project
|
||||
project = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
project: prisma_models.LiteLLM_ProjectTable | None = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
where={"project_id": project_id},
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
)
|
||||
|
|
@ -872,17 +835,11 @@ async def project_info(
|
|||
is_team_member = False
|
||||
|
||||
if project.team_id and user_api_key_dict.user_id:
|
||||
team = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": project.team_id}
|
||||
)
|
||||
team = await _team_table(prisma_client).find_unique(where={"team_id": project.team_id})
|
||||
if team:
|
||||
caller_user_id = user_api_key_dict.user_id
|
||||
for m in team.members_with_roles or []:
|
||||
m_user_id = (
|
||||
m.get("user_id")
|
||||
if isinstance(m, dict)
|
||||
else getattr(m, "user_id", None)
|
||||
)
|
||||
m_user_id = m.get("user_id") if isinstance(m, dict) else getattr(m, "user_id", None)
|
||||
if m_user_id == caller_user_id:
|
||||
is_team_member = True
|
||||
break
|
||||
|
|
@ -896,9 +853,7 @@ async def project_info(
|
|||
return project
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.project_endpoints.project_info(): Exception occured - {}".format(
|
||||
str(e)
|
||||
)
|
||||
"litellm.proxy.management_endpoints.project_endpoints.project_info(): Exception occured - {}".format(str(e))
|
||||
)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
|
@ -907,7 +862,7 @@ async def project_info(
|
|||
"/project/list",
|
||||
tags=["project management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[LiteLLM_ProjectTable],
|
||||
response_model=list[LiteLLM_ProjectTable],
|
||||
)
|
||||
async def list_projects(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
|
|
@ -932,21 +887,19 @@ async def list_projects(
|
|||
|
||||
# If proxy admin, get all projects
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
projects = await prisma_client.db.litellm_projecttable.find_many(
|
||||
projects: Sequence[
|
||||
prisma_models.LiteLLM_ProjectTable
|
||||
] = await prisma_client.db.litellm_projecttable.find_many(
|
||||
include={"litellm_budget_table": True, "object_permission": True}
|
||||
)
|
||||
else:
|
||||
# Look up the user's team memberships via the reverse-index on
|
||||
# LiteLLM_UserTable.teams (maintained by team_member_add alongside
|
||||
# members_with_roles). This avoids a full scan of all team rows.
|
||||
user_record = await prisma_client.db.litellm_usertable.find_unique(
|
||||
user_record: prisma_models.LiteLLM_UserTable | None = await prisma_client.db.litellm_usertable.find_unique(
|
||||
where={"user_id": user_api_key_dict.user_id},
|
||||
)
|
||||
user_team_ids = (
|
||||
user_record.teams
|
||||
if user_record is not None and user_record.teams
|
||||
else []
|
||||
)
|
||||
user_team_ids: Sequence[str] = user_record.teams if user_record is not None and user_record.teams else []
|
||||
|
||||
projects = await prisma_client.db.litellm_projecttable.find_many(
|
||||
where={"team_id": {"in": user_team_ids}},
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.51"
|
||||
version = "0.1.52"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.1.51"
|
||||
version = "0.1.52"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -54,6 +54,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/messages",
|
||||
"/v1/skills",
|
||||
"/v1/a2a/",
|
||||
"/a2a/",
|
||||
# LiteLLM-native LLM surface
|
||||
"/v1/rerank",
|
||||
"/v2/rerank",
|
||||
|
|
|
|||
|
|
@ -5,5 +5,5 @@ dependencies:
|
|||
- name: redis
|
||||
repository: oci://registry-1.docker.io/bitnamicharts
|
||||
version: 18.19.1
|
||||
digest: sha256:8660fe6287f9941d08c0902f3f13731079b8cecd2a5da2fbc54e5b7aae4a6f62
|
||||
generated: "2024-03-10T02:28:52.275022+05:30"
|
||||
digest: sha256:38962e231f6596b93f82a8412bbe4cf5de696caecf5775dfbbd163383eb1c009
|
||||
generated: "2026-07-28T10:21:22.511401-07:00"
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ type: application
|
|||
# This is the chart version. This version number should be incremented each time you make changes
|
||||
# to the chart and its templates, including the app version.
|
||||
# Versions are expected to follow Semantic Versioning (https://semver.org/)
|
||||
version: 1.1.0
|
||||
version: 1.1.1
|
||||
|
||||
# This is the version number of the application being deployed. This version number should be
|
||||
# incremented each time you make changes to the application. Versions are not expected to
|
||||
|
|
@ -32,10 +32,10 @@ annotations:
|
|||
|
||||
dependencies:
|
||||
- name: "postgresql"
|
||||
version: ">=13.3.0"
|
||||
version: "14.3.1"
|
||||
repository: oci://registry-1.docker.io/bitnamicharts
|
||||
condition: db.deployStandalone
|
||||
- name: redis
|
||||
version: ">=18.0.0"
|
||||
version: "18.19.1"
|
||||
repository: oci://registry-1.docker.io/bitnamicharts
|
||||
condition: redis.enabled
|
||||
|
|
|
|||
|
|
@ -130,6 +130,16 @@ Set `billingMetrics.caSecretName` only when the collector is a private or test o
|
|||
| `db.deployStandalone` | Deploy a standalone, single instance deployment of Postgres, using the Bitnami postgresql chart. This is useful for getting started but doesn't provide HA or (by default) data backups. | `true` |
|
||||
| `postgresql.*` | If `db.deployStandalone` is `true`, configuration passed to the Bitnami postgresql chart. See the [Bitnami Documentation](https://github.com/bitnami/charts/tree/main/bitnami/postgresql) for full configuration details. See [values.yaml](./values.yaml) for the default configuration. | See [values.yaml](./values.yaml) |
|
||||
| `postgresql.auth.*` | If `db.deployStandalone` is `true`, care should be taken to ensure the default `password` and `postgres-password` values are **NOT** used. | `NoTaGrEaTpAsSwOrD` |
|
||||
| `postgresql.image.*` | If `db.deployStandalone` is `true`, the image for the bundled Postgres. Pinned to a `docker.io/bitnamilegacy` build because Bitnami retired the versioned tags under `docker.io/bitnami`. | `bitnamilegacy/postgresql:16.2.0-debian-12-r6` |
|
||||
| `redis.image.*` | If `redis.enabled` is `true`, the image for the bundled Redis. Pinned to a `docker.io/bitnamilegacy` build for the same reason. | `bitnamilegacy/redis:7.2.4-debian-12-r9` |
|
||||
|
||||
#### Bundled Postgres image
|
||||
|
||||
Bitnami removed the versioned tags from `docker.io/bitnami` and republished the archived builds under `docker.io/bitnamilegacy`, so the image defaults that ship inside the `postgresql` and `redis` subcharts no longer pull. The chart pins both to the `bitnamilegacy` copies of the exact builds those subchart versions were released with, which keeps the on-disk data directory layout unchanged for existing installs.
|
||||
|
||||
Keep `postgresql.image.tag` pinned. `docker.io/bitnami/postgresql` still publishes a floating `latest`, and pointing the bundled Postgres at a different major version starts the server against a data directory it cannot read (`database files are incompatible with server`). There is no in-place way back, so crossing a major version means dumping the database with the old image and restoring it into the new one. The chart refuses to render when the tag is empty or `latest`.
|
||||
|
||||
Those images no longer receive updates. For anything beyond getting started, run Postgres outside the chart and point at it with `db.useExisting`.
|
||||
|
||||
#### Example Postgres `db.useExisting` Secret
|
||||
|
||||
|
|
|
|||
|
|
@ -146,3 +146,18 @@ Get redis service port
|
|||
{{ .Values.redis.master.service.ports.redis }}
|
||||
{{- end -}}
|
||||
{{- end -}}
|
||||
|
||||
{{/*
|
||||
Reject an unpinned image tag for the bundled PostgreSQL.
|
||||
A floating tag lets a chart upgrade start a newer PostgreSQL major against the
|
||||
existing PersistentVolumeClaim. The server then refuses to start on a data
|
||||
directory written by another major version, and the only way back is a dump
|
||||
taken before the change, which by that point no longer exists.
|
||||
*/}}
|
||||
{{- define "litellm.validateBundledPostgresImageTag" -}}
|
||||
{{- $tag := .Values.postgresql.image.tag | default "" | toString -}}
|
||||
{{- $digest := .Values.postgresql.image.digest | default "" | toString -}}
|
||||
{{- if and (eq $digest "") (or (eq $tag "") (eq $tag "latest")) -}}
|
||||
{{- fail (printf "postgresql.image.tag must be pinned to an explicit version when db.deployStandalone is true (got %q). An unpinned tag can start a different PostgreSQL major against the existing data directory, which makes the database unreadable and is not recoverable in place. Crossing a major version requires a dump and restore." $tag) -}}
|
||||
{{- end -}}
|
||||
{{- end -}}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
{{- if .Values.db.deployStandalone -}}
|
||||
{{- include "litellm.validateBundledPostgresImageTag" . -}}
|
||||
apiVersion: v1
|
||||
kind: Secret
|
||||
metadata:
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ metadata:
|
|||
spec:
|
||||
containers:
|
||||
- name: test
|
||||
image: bitnami/kubectl:latest
|
||||
image: docker.io/bitnamilegacy/kubectl:1.29.2-debian-12-r3
|
||||
command: ['sh', '-c']
|
||||
args:
|
||||
- |
|
||||
|
|
|
|||
94
helm/litellm-helm/tests/bundled_db_images_tests.yaml
Normal file
94
helm/litellm-helm/tests/bundled_db_images_tests.yaml
Normal file
|
|
@ -0,0 +1,94 @@
|
|||
suite: test bundled database images
|
||||
templates:
|
||||
- charts/postgresql/templates/primary/statefulset.yaml
|
||||
- charts/redis/templates/master/application.yaml
|
||||
- charts/redis/templates/configmap.yaml
|
||||
- charts/redis/templates/health-configmap.yaml
|
||||
- charts/redis/templates/scripts-configmap.yaml
|
||||
- charts/redis/templates/secret.yaml
|
||||
- secret-dbcredentials.yaml
|
||||
- templates/tests/test-servicemonitor.yaml
|
||||
tests:
|
||||
- it: should pull the bundled postgres from a repository that still publishes the pinned tag
|
||||
template: charts/postgresql/templates/primary/statefulset.yaml
|
||||
set:
|
||||
db.deployStandalone: true
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].image
|
||||
value: docker.io/bitnamilegacy/postgresql:16.2.0-debian-12-r6
|
||||
|
||||
- it: should pull the bundled postgres metrics exporter from the same repository
|
||||
template: charts/postgresql/templates/primary/statefulset.yaml
|
||||
set:
|
||||
db.deployStandalone: true
|
||||
postgresql.metrics.enabled: true
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.containers[1].image
|
||||
value: docker.io/bitnamilegacy/postgres-exporter:0.15.0-debian-12-r14
|
||||
|
||||
- it: should run the bundled postgres init container from the same repository
|
||||
template: charts/postgresql/templates/primary/statefulset.yaml
|
||||
set:
|
||||
db.deployStandalone: true
|
||||
postgresql.volumePermissions.enabled: true
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.initContainers[0].image
|
||||
value: docker.io/bitnamilegacy/os-shell:12-debian-12-r16
|
||||
|
||||
- it: should pull the bundled redis from a repository that still publishes the pinned tag
|
||||
template: charts/redis/templates/master/application.yaml
|
||||
set:
|
||||
redis.enabled: true
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].image
|
||||
value: docker.io/bitnamilegacy/redis:7.2.4-debian-12-r9
|
||||
|
||||
- it: should reject a floating postgres tag that could cross a major version on an existing volume
|
||||
template: secret-dbcredentials.yaml
|
||||
set:
|
||||
db.deployStandalone: true
|
||||
postgresql.image.tag: latest
|
||||
asserts:
|
||||
- failedTemplate:
|
||||
errorMessage: 'postgresql.image.tag must be pinned to an explicit version when db.deployStandalone is true (got "latest"). An unpinned tag can start a different PostgreSQL major against the existing data directory, which makes the database unreadable and is not recoverable in place. Crossing a major version requires a dump and restore.'
|
||||
|
||||
- it: should reject an empty postgres tag
|
||||
template: secret-dbcredentials.yaml
|
||||
set:
|
||||
db.deployStandalone: true
|
||||
postgresql.image.tag: ""
|
||||
asserts:
|
||||
- failedTemplate:
|
||||
errorMessage: 'postgresql.image.tag must be pinned to an explicit version when db.deployStandalone is true (got ""). An unpinned tag can start a different PostgreSQL major against the existing data directory, which makes the database unreadable and is not recoverable in place. Crossing a major version requires a dump and restore.'
|
||||
|
||||
- it: should accept an empty postgres tag when the image is pinned by digest
|
||||
template: secret-dbcredentials.yaml
|
||||
set:
|
||||
db.deployStandalone: true
|
||||
postgresql.image.tag: ""
|
||||
postgresql.image.digest: sha256:0d0e2f1a5b3c4d6e7f8091a2b3c4d5e6f708192a3b4c5d6e7f8091a2b3c4d5e6
|
||||
asserts:
|
||||
- hasDocuments:
|
||||
count: 1
|
||||
|
||||
- it: should run the servicemonitor test pod from a pinned image
|
||||
template: templates/tests/test-servicemonitor.yaml
|
||||
set:
|
||||
serviceMonitor.enabled: true
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.containers[0].image
|
||||
value: docker.io/bitnamilegacy/kubectl:1.29.2-debian-12-r3
|
||||
|
||||
- it: should not constrain the postgres tag when the bundled database is not deployed
|
||||
template: secret-dbcredentials.yaml
|
||||
set:
|
||||
db.deployStandalone: false
|
||||
postgresql.image.tag: latest
|
||||
asserts:
|
||||
- hasDocuments:
|
||||
count: 0
|
||||
|
|
@ -328,8 +328,32 @@ lifecycle: {}
|
|||
|
||||
# Settings for Bitnami postgresql chart (if db.deployStandalone is true, ignored
|
||||
# otherwise)
|
||||
#
|
||||
# Bitnami retired the versioned tags under docker.io/bitnami and republished the
|
||||
# archived builds under docker.io/bitnamilegacy, so the subchart's own image
|
||||
# defaults no longer resolve. The repository below points at the same build the
|
||||
# subchart was released with, which keeps the on-disk data directory layout
|
||||
# identical for existing installs.
|
||||
#
|
||||
# Keep the tag pinned. docker.io/bitnami still publishes a floating `latest`,
|
||||
# and starting a newer PostgreSQL major against an existing data directory
|
||||
# leaves the server refusing to boot ("database files are incompatible with
|
||||
# server") with no way back other than a dump taken beforehand. Crossing a major
|
||||
# version is a dump-and-restore, not an image bump. The chart refuses to render
|
||||
# an unpinned tag for this reason
|
||||
postgresql:
|
||||
architecture: standalone
|
||||
image:
|
||||
repository: bitnamilegacy/postgresql
|
||||
tag: 16.2.0-debian-12-r6
|
||||
volumePermissions:
|
||||
image:
|
||||
repository: bitnamilegacy/os-shell
|
||||
tag: 12-debian-12-r16
|
||||
metrics:
|
||||
image:
|
||||
repository: bitnamilegacy/postgres-exporter
|
||||
tag: 0.15.0-debian-12-r14
|
||||
auth:
|
||||
username: litellm
|
||||
database: litellm
|
||||
|
|
@ -359,9 +383,36 @@ postgresql:
|
|||
# When `redis.sentinel.enabled` is set, the coordination block is rendered with
|
||||
# `sentinel_nodes` and `service_name` (from `redis.sentinel.masterSet`) instead
|
||||
# of host/port, because a plain Redis client cannot talk to the sentinel port
|
||||
#
|
||||
# The image repositories carry the same bitnamilegacy repoint as postgresql
|
||||
# above; the versioned tags the subchart ships with are gone from
|
||||
# docker.io/bitnami
|
||||
redis:
|
||||
enabled: false
|
||||
architecture: standalone
|
||||
image:
|
||||
repository: bitnamilegacy/redis
|
||||
tag: 7.2.4-debian-12-r9
|
||||
sentinel:
|
||||
image:
|
||||
repository: bitnamilegacy/redis-sentinel
|
||||
tag: 7.2.4-debian-12-r7
|
||||
metrics:
|
||||
image:
|
||||
repository: bitnamilegacy/redis-exporter
|
||||
tag: 1.58.0-debian-12-r4
|
||||
volumePermissions:
|
||||
image:
|
||||
repository: bitnamilegacy/os-shell
|
||||
tag: 12-debian-12-r16
|
||||
sysctl:
|
||||
image:
|
||||
repository: bitnamilegacy/os-shell
|
||||
tag: 12-debian-12-r16
|
||||
kubectl:
|
||||
image:
|
||||
repository: bitnamilegacy/kubectl
|
||||
tag: 1.29.2-debian-12-r3
|
||||
coordination:
|
||||
# Set to false to keep the bundled Redis for response caching only and leave
|
||||
# `general_settings.coordination_redis` out of the rendered config. A
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@
|
|||
"/v1/fine-tuning" "/fine-tuning" "/v1/responses" "/responses" "/v1/threads" "/threads"
|
||||
"/v1/assistants" "/assistants" "/v1/vector_stores" "/vector_stores" "/v1/indexes"
|
||||
"/v1/models" "/models" "/openai" "/engines"
|
||||
"/v1/messages" "/messages" "/v1/skills" "/v1/a2a"
|
||||
"/v1/messages" "/messages" "/v1/skills" "/v1/a2a" "/a2a"
|
||||
"/v1/rerank" "/v2/rerank" "/rerank" "/v1/ocr" "/ocr" "/v1/rag" "/rag"
|
||||
"/v1/video" "/v1/videos" "/video" "/videos" "/v1/search" "/search"
|
||||
"/v1/containers" "/containers" "/v1/evals" "/v1/memory" "/queue/chat"
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- CreateIndex
|
||||
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogToolIndex_start_time_idx" ON "LiteLLM_SpendLogToolIndex"("start_time");
|
||||
|
|
@ -0,0 +1,12 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_DailyToolSpend" (
|
||||
"date" TEXT NOT NULL,
|
||||
"tool_name" TEXT NOT NULL,
|
||||
"spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
|
||||
"total_tokens" BIGINT NOT NULL DEFAULT 0,
|
||||
"request_count" BIGINT NOT NULL DEFAULT 0,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"updated_at" TIMESTAMP(3) NOT NULL,
|
||||
|
||||
CONSTRAINT "LiteLLM_DailyToolSpend_pkey" PRIMARY KEY ("date","tool_name")
|
||||
);
|
||||
|
|
@ -1094,6 +1094,20 @@ model LiteLLM_SpendLogToolIndex {
|
|||
|
||||
@@id([request_id, tool_name])
|
||||
@@index([tool_name, start_time])
|
||||
@@index([start_time])
|
||||
}
|
||||
|
||||
// Daily tool spend rollup (one row per tool per day) – the Cost Optimization card reads this, never SpendLogs
|
||||
model LiteLLM_DailyToolSpend {
|
||||
date String
|
||||
tool_name String
|
||||
spend Float @default(0.0)
|
||||
total_tokens BigInt @default(0)
|
||||
request_count BigInt @default(0)
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
|
||||
@@id([date, tool_name])
|
||||
}
|
||||
|
||||
// Prompt table for storing prompt configurations
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.80"
|
||||
version = "0.4.81"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.80"
|
||||
version = "0.4.81"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
# Adding a provider / route to litellm-rust
|
||||
|
||||
Three layers, same for every route (see `ocr` and `realtime` as references):
|
||||
Everything for a route lives in `crates/core/src/<route>/`; `crates/core/src/messages` is the reference. A host (the axum gateway, the Python bridge) only calls the route's entrypoint.
|
||||
|
||||
1. **Transform contract (pure)** — `crates/core/src/<route>/transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) + types in `types.rs`. No network, env, or auth.
|
||||
2. **Provider config (pure)** — `crates/providers/src/<provider>/<route>/transformation.rs`: implement that trait as a `const <PROVIDER>_<ROUTE>_CONFIG`, mirroring the Python provider tree. Add parity unit tests.
|
||||
3. **HTTP / transport (the host)** — `crates/providers/src/<route>.rs` (e.g. `ocr.rs`, `realtime.rs`): the callable fn (`run_ocr`, `realtime`). It resolves the key, builds the auth header, builds URL + transforms via the config, then does the network call. This is the only layer allowed to do I/O.
|
||||
1. **Entrypoint** — `mod.rs`: `pub async fn <route>(request) -> CoreResult<Response>`, the Rust equivalent of `litellm.<route>()`, plus a `<route>_stream` variant when the route streams. It is the only thing a host touches.
|
||||
2. **Transform contract** — `transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) with types in `types.rs`.
|
||||
3. **Provider config** — `crates/core/src/providers/<provider>/<route>/transformation.rs`: implement that trait as a `const <PROVIDER>_<ROUTE>_CONFIG`, mirroring the Python provider tree. Add parity unit tests.
|
||||
4. **Prepare + handler** — `prepare.rs` resolves provider/model, credentials, auth headers, and URL, then transforms the request; `handler.rs` performs the provider call through the shared client in `client.rs` and transforms the response.
|
||||
|
||||
## Coding standards
|
||||
|
||||
|
|
@ -25,4 +26,4 @@ variants of it. The test for a good abstraction is that adding the next provider
|
|||
is a few declarative lines, not a new file of duplicated flow. Only diverge from
|
||||
the base when behavior is genuinely different, and say so explicitly in the PR.
|
||||
|
||||
**Calling:** the host invokes the route fn — the Python bridge calls `run_ocr`; the `ai-gateway` server calls `realtime`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`.
|
||||
**Calling:** hosts invoke the core entrypoint — the Python bridge and the `ai-gateway` route service both call `litellm_core::messages::messages`. Never add a provider handler to `ai-gateway`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`.
|
||||
|
|
|
|||
|
|
@ -4,14 +4,30 @@ litellm-rust has exactly THREE crates. A crate is a LAYER, not a route. Routes (
|
|||
|
||||
## Crates
|
||||
|
||||
| Crate | Role | Pure / I/O |
|
||||
|-------|------|------------|
|
||||
| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure |
|
||||
| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O |
|
||||
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding |
|
||||
| Crate | Role |
|
||||
|-------|------|
|
||||
| litellm-core | The LiteLLM SDK in Rust. One public entrypoint per top-level call (`messages::messages()`), owning types, transforms, provider resolution, auth, and the provider HTTP call. Call it, get a typed response. |
|
||||
| litellm-ai-gateway | The axum server (behind the `server` feature) plus the WebSocket hosts. Translates HTTP/WS to core entrypoints; owns no provider logic and no handlers. |
|
||||
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
|
||||
|
||||
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
|
||||
|
||||
## Where a route lives
|
||||
|
||||
A top-level LiteLLM call is a module under `crates/core/src/<route>/`, shaped like `messages`:
|
||||
|
||||
```
|
||||
core/src/messages/
|
||||
mod.rs # pub async fn messages(..) -> CoreResult<..> (+ messages_stream for SSE)
|
||||
types.rs # request/response types, MessagesRequest
|
||||
transformation.rs # the provider template trait
|
||||
prepare.rs # provider resolution, auth headers, URL
|
||||
handler.rs # the provider call
|
||||
client.rs # the shared reqwest client
|
||||
```
|
||||
|
||||
Handlers never live in `ai-gateway`. `ocr`, `audio_transcription`, and `realtime` are still hosted there from before this rule; they move to `core` as they are touched.
|
||||
|
||||
Adding a crate: default to a MODULE. New crate ONLY on a real trigger — separate artifact (binary/cdylib), proc-macro, shared foundation, or publishable standalone. A new provider or route is none of these.
|
||||
|
||||
Adding a crate fails crates/core/tests/workspace_crate_allowlist.rs until you update its allowlist and this file — intentional.
|
||||
|
|
|
|||
|
|
@ -23,21 +23,34 @@ the base when behavior is genuinely different, and say so explicitly in the PR.
|
|||
|
||||
## Crates (exactly three — see AGENTS.md)
|
||||
|
||||
`litellm-core` describes work; `litellm-ai-gateway` executes it; `litellm-python-bridge`
|
||||
exposes it to the Python SDK. A crate is a **layer**, not a route — add modules, not crates.
|
||||
`litellm-core` **is** the LiteLLM SDK in Rust: it makes the LLM call.
|
||||
`litellm-ai-gateway` is an HTTP/WebSocket server in front of it, and
|
||||
`litellm-python-bridge` exposes it to the Python SDK. A crate is a **layer**, not
|
||||
a route — add modules, not crates.
|
||||
|
||||
## Core Boundary
|
||||
|
||||
`litellm-core` is the pure translation layer; the `litellm-ai-gateway` host executes work.
|
||||
`litellm-core` owns the whole call. The Rust equivalent of `litellm.messages()`
|
||||
is `litellm_core::messages::messages(request).await`: you call it, it does the
|
||||
provider call, and you get a typed non-streaming response back.
|
||||
|
||||
Route-level Rust structure mirrors LiteLLM's Python responsibilities:
|
||||
- `core/src/<route>/` owns the route contract, shared types, and provider
|
||||
template traits. For OCR, this means `core/src/ocr`.
|
||||
- `core/src/<route>/` owns the route end to end: the public entrypoint fn named
|
||||
after the route in `mod.rs`, the request/response types (`types.rs`), the
|
||||
provider template trait (`transformation.rs`), the provider/auth/URL
|
||||
resolution (`prepare.rs`), the HTTP client (`client.rs`), and the handler that
|
||||
performs the call (`handler.rs`). `core/src/messages` is the reference.
|
||||
- `core/src/providers/<provider>/<route>/transformation.rs` owns the
|
||||
provider-specific transform. For Mistral OCR, this means
|
||||
`core/src/providers/mistral/ocr/transformation.rs`.
|
||||
- Network execution lives in the host crate `ai-gateway` (`ai-gateway/src/io/`),
|
||||
never inside `core`.
|
||||
provider-specific transform. For Anthropic Messages, this means
|
||||
`core/src/providers/anthropic/messages/transformation.rs`.
|
||||
- Handlers live in `core`, never in a host. `ai-gateway` must not contain a
|
||||
route handler that talks to a provider; its axum route reads the HTTP request,
|
||||
picks a deployment, and calls the `core` entrypoint. `python-bridge` marshals
|
||||
Python objects and calls the same entrypoint.
|
||||
|
||||
Streaming keeps the same shape: the route entrypoint has a `<route>_stream`
|
||||
variant in `core` that returns the upstream response so a host can splice it to
|
||||
its own caller; the host still owns no provider logic.
|
||||
|
||||
Call-hook and lifecycle instrumentation, including phase timing, usage
|
||||
accumulation, and callback payload construction, always lives in `core`.
|
||||
|
|
@ -45,21 +58,31 @@ Hosts feed observed events into core and dispatch the completed payloads through
|
|||
their I/O logger; hosts must not own callback orchestration.
|
||||
|
||||
Allowed in `core`:
|
||||
- Pure request transforms
|
||||
- Pure response transforms
|
||||
- Pure stream chunk normalization
|
||||
- The public entrypoint for a top-level LiteLLM call
|
||||
- Request/response transforms and stream chunk normalization
|
||||
- Provider resolution, auth header construction, and URL building
|
||||
- The provider HTTP call itself, through a shared reused client with connect and
|
||||
request timeouts
|
||||
- Shared data types and validation errors
|
||||
- Deterministic token/cost helper logic
|
||||
|
||||
Not allowed in `core`:
|
||||
- Network calls
|
||||
- Environment variable or secret reads
|
||||
- Serving HTTP: axum routes, extractors, and transport concerns stay in the host
|
||||
- Filesystem access
|
||||
- Database or cache access
|
||||
- Provider SDK signing or auth flows
|
||||
- Database access
|
||||
- Config file reading and rollout state
|
||||
- Logging callbacks, spend writes, or custom callbacks
|
||||
- Global mutable runtime state
|
||||
|
||||
Env reads in `core` are limited to credential fallback inside a route's
|
||||
`prepare.rs` (the `env_lookup` closure), mirroring what the Python SDK does when
|
||||
no key is passed. Everything else config-shaped is resolved by the host and
|
||||
passed in.
|
||||
|
||||
Routes still hosted in `ai-gateway` (`ocr`, `audio_transcription`, `realtime`)
|
||||
predate this rule and are being moved into `core` route modules; do not add new
|
||||
ones there, and prefer moving one when you touch it.
|
||||
|
||||
Python owns rollout state and fallback while Rust is being introduced. Rust
|
||||
paths must be off by default until parity tests prove equivalence with Python.
|
||||
A new provider/route may instead be implemented rust-only with no Python
|
||||
|
|
@ -93,10 +116,10 @@ the first PR:
|
|||
- Preserve Python output shape intentionally. If a field is always serialized as
|
||||
`null` for Python parity, leave a short comment explaining that parity choice.
|
||||
|
||||
## Host I/O Rules
|
||||
## Network I/O Rules
|
||||
|
||||
These rules apply when adding future crates or modules that execute network I/O,
|
||||
such as `ai-gateway`, router hosts, or standalone servers:
|
||||
These rules apply to every module that executes network I/O, whether it is a
|
||||
`core` route handler or a host such as `ai-gateway`:
|
||||
|
||||
- Set connect and full-request timeouts. No unbounded waits.
|
||||
- Reuse HTTP clients; do not construct clients per request.
|
||||
|
|
|
|||
|
|
@ -2,18 +2,31 @@
|
|||
|
||||
This workspace contains the staged Rust implementation for LiteLLM.
|
||||
|
||||
Rust starts as a pure transform core used by the existing Python host. Python
|
||||
continues to own auth, configuration, network I/O, retries, routing, logging,
|
||||
`litellm-core` is the LiteLLM SDK in Rust: one entrypoint per top-level call
|
||||
that makes the LLM call and hands back a typed response, the same shape as
|
||||
`litellm.messages()` in Python.
|
||||
|
||||
```rust
|
||||
let response = litellm_core::messages::messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body,
|
||||
api_key: Some(key),
|
||||
..
|
||||
})
|
||||
.await?;
|
||||
```
|
||||
|
||||
Python continues to own configuration, retries, routing policy, logging,
|
||||
callbacks, spend tracking, and customer plugins until each Rust path has parity
|
||||
coverage and production evidence.
|
||||
|
||||
## Crates
|
||||
|
||||
| Crate | Role | Pure / I/O |
|
||||
|-------|------|------------|
|
||||
| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure |
|
||||
| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O |
|
||||
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding |
|
||||
| Crate | Role |
|
||||
|-------|------|
|
||||
| litellm-core | The SDK. Per-route entrypoints (`messages::messages()`), types, provider transforms (modules under `providers/`), provider resolution, auth, the provider HTTP call, and the router. |
|
||||
| litellm-ai-gateway | The axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
|
||||
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
|
||||
|
||||
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
|
||||
|
||||
|
|
@ -21,16 +34,16 @@ Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-
|
|||
|
||||
```text
|
||||
crates/
|
||||
core/ Route contracts, shared pure types, errors, and templates.
|
||||
src/ocr/
|
||||
providers/ Provider-specific pure transforms.
|
||||
src/mistral/ocr/transformation.rs
|
||||
core/ The SDK: route modules + provider transforms.
|
||||
src/messages/ mod.rs (entrypoint), types, transformation, prepare, handler, client
|
||||
src/providers/anthropic/messages/transformation.rs
|
||||
ai-gateway/ Axum server + WebSocket hosts; calls core entrypoints.
|
||||
python-bridge/ PyO3 bridge for Python LiteLLM.
|
||||
```
|
||||
|
||||
The folder shape should follow the Python provider tree:
|
||||
`providers/src/<provider>/<route>/transformation.rs`. The bridge should expose
|
||||
one function per top-level route, starting with `ocr(payload)`.
|
||||
The folder shape follows the Python provider tree:
|
||||
`core/src/providers/<provider>/<route>/transformation.rs`. The bridge exposes one
|
||||
function per top-level route, mirroring the core entrypoints.
|
||||
|
||||
## Checks
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
# Provider coding standards (litellm-rust)
|
||||
|
||||
Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MISTRAL_OCR_CONFIG`) is the reference; `messages` (`ANTHROPIC_MESSAGES_CONFIG`) is the next port.
|
||||
Rules for adding or changing an LLM provider/route in `litellm-rust`. `messages` (`core/src/messages`, `ANTHROPIC_MESSAGES_CONFIG`) is the reference: a route is a `core` module with a public entrypoint that makes the call and returns a typed response.
|
||||
|
||||
## Provider resolution
|
||||
|
||||
|
|
@ -16,10 +16,10 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST
|
|||
|
||||
## Boundaries
|
||||
|
||||
7. Layers never cross: `core` = pure transforms/types (no network, env, secrets, auth, logging, global mutable state); `ai-gateway` = all I/O, auth headers, HTTP/SSE, lifecycle hooks; `python-bridge` = thin PyO3 adapter.
|
||||
7. Layers never cross: `core` = the call itself (entrypoint, types, transforms, provider resolution, auth headers, provider HTTP, lifecycle hooks); `ai-gateway` = serving HTTP/WS (routing, extractors, auth of *our* callers, streaming to the client); `python-bridge` = thin PyO3 adapter. Hosts call the core entrypoint; they never build a provider request.
|
||||
8. Generic/route files contain zero provider-specific branches. A provider is one module under `core/src/providers/<provider>/<route>/`; a route is a module, never a new crate.
|
||||
9. Route entry point stays thin: `<route>()` -> `prepare_*` -> `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing. Handlers validate and delegate; no business logic in them.
|
||||
10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Env reads happen only at the host/config layer, with the `DEFAULT_*` fallback defined in `constants.rs`.
|
||||
9. Route entry point stays thin: `core::<route>::<route>()` -> `prepare_*` -> handler (or `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing). Axum handlers validate and delegate to a service that calls the entrypoint; no business logic in them.
|
||||
10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Config-shaped env reads happen at the host/config layer with the `DEFAULT_*` fallback defined in `constants.rs`; the only env read in `core` is the credential fallback in a route's `prepare.rs`.
|
||||
|
||||
## Types and errors
|
||||
|
||||
|
|
@ -33,7 +33,7 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST
|
|||
|
||||
16. Never log request/response bodies, base64 payloads, document contents, or secrets. Truncate and bound any upstream body before it crosses a host boundary.
|
||||
17. Treat empty/whitespace credentials, URLs, and config values as absent at the host resolution layer.
|
||||
18. Host I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS.
|
||||
18. Network I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS.
|
||||
|
||||
## Tests and rollout
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,9 @@
|
|||
# ai-gateway — folder architecture
|
||||
|
||||
The Axum server that fronts the Rust gateway. It owns transport + config + auth
|
||||
only; deployment selection lives in `core::router`, transforms in `core`/`providers`.
|
||||
only; deployment selection lives in `core::router`, and the LLM call itself
|
||||
(transforms, auth headers, provider HTTP) lives behind a `core` route entrypoint
|
||||
such as `litellm_core::messages::messages`. No provider handler lives here.
|
||||
|
||||
```
|
||||
src/
|
||||
|
|
@ -32,6 +34,11 @@ src/
|
|||
args; it runs during extraction. Never re-implement the check per route.
|
||||
- **Handlers are thin.** A handler validates and delegates to its `service`. No
|
||||
business logic, no provider calls, no transforms in handlers.
|
||||
- **Services call `core`, they don't reimplement it.** A `service` picks the
|
||||
deployment and calls the `core` route entrypoint. Provider resolution, auth
|
||||
headers, URL building, and the HTTP call are `core`'s job; a service that
|
||||
builds a provider request itself is a bug (`routes/messages/service.rs` is
|
||||
the reference).
|
||||
- **State is shared and cheap to clone.** Long-lived handles live behind `Arc` in
|
||||
`state.rs`; read env/config only in `main.rs` when building state.
|
||||
|
||||
|
|
|
|||
|
|
@ -8,11 +8,11 @@ dials OpenAI upstream, and splices the two sockets frame-by-frame.
|
|||
|
||||
`litellm-rust` is exactly three crates (a crate is a **layer**, not a route):
|
||||
|
||||
| Crate | Role | Pure / I/O |
|
||||
|-------|------|------------|
|
||||
| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under `providers/`), and the router. Builds requests/responses; no network. | Pure |
|
||||
| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under `io/`) plus the Axum server binary (behind the `server` feature). | I/O |
|
||||
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding |
|
||||
| Crate | Role |
|
||||
|-------|------|
|
||||
| litellm-core | The LiteLLM SDK in Rust — per-route entrypoints (`messages::messages()`) that resolve the provider, transform, and make the call; plus types, provider transforms, and the router. |
|
||||
| litellm-ai-gateway | The Axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
|
||||
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
|
||||
|
||||
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
|
||||
|
||||
|
|
|
|||
|
|
@ -29,18 +29,6 @@ pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500;
|
|||
#[cfg(feature = "server")]
|
||||
pub(crate) const DEFAULT_PROVIDER: &str = "openai";
|
||||
|
||||
/// Full-request timeout ceiling for Anthropic Messages provider calls, in
|
||||
/// seconds. Mirrors the Python Anthropic Messages default. The per-request
|
||||
/// timeout from `litellm_params` still overrides this on the request builder.
|
||||
pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
/// Connect timeout for Anthropic Messages provider calls, in seconds.
|
||||
pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10;
|
||||
|
||||
/// Max characters of an upstream error body echoed across the host boundary
|
||||
/// before truncation, so provider bodies are bounded and data-minimized.
|
||||
pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256;
|
||||
|
||||
pub(crate) const DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS: u64 = 10;
|
||||
pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300;
|
||||
|
||||
|
|
@ -48,10 +36,6 @@ pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300;
|
|||
#[cfg(feature = "server")]
|
||||
pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages";
|
||||
|
||||
/// Provider name used by the Anthropic Messages route when a deployment's
|
||||
/// provider model does not carry an explicit provider prefix.
|
||||
pub(crate) const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
|
||||
|
||||
/// Request headers owned by the gateway and never forwarded upstream.
|
||||
#[cfg(feature = "server")]
|
||||
pub(crate) const MESSAGES_HEADERS_NOT_FORWARDED: &[&str] =
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
pub use crate::messages::{MessagesRequest, messages};
|
||||
|
|
@ -1,5 +1,4 @@
|
|||
pub mod audio_transcription;
|
||||
pub mod messages;
|
||||
pub mod ocr;
|
||||
pub mod realtime;
|
||||
pub mod realtime_pool;
|
||||
|
|
|
|||
|
|
@ -4,7 +4,9 @@
|
|||
//! without pulling in the HTTP server:
|
||||
//!
|
||||
//! - Call-type modules such as [`ocr`]: provider transforms, lifecycle hooks,
|
||||
//! and provider I/O. Always available — no feature required.
|
||||
//! and provider I/O. Always available — no feature required. These predate the
|
||||
//! rule that a route's entrypoint and handler live in `litellm-core` (see
|
||||
//! `litellm_core::messages`) and move there as they are touched.
|
||||
//! - [`io`]: compatibility exports and realtime WebSocket splice helpers.
|
||||
//! - The server modules ([`auth`], [`routes`], [`state`]) and anything pulling
|
||||
//! `axum` are gated behind the `server` feature, which the `litellm-ai-gateway`
|
||||
|
|
@ -14,7 +16,6 @@
|
|||
pub mod audio_transcription;
|
||||
mod client;
|
||||
pub mod io;
|
||||
pub mod messages;
|
||||
pub mod ocr;
|
||||
|
||||
/// GIL-activity tracking. Pure (atomics only); shared by the `server` routes and
|
||||
|
|
|
|||
|
|
@ -1,49 +0,0 @@
|
|||
use litellm_core::CoreResult;
|
||||
use serde_json::Value;
|
||||
|
||||
mod client;
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
mod types;
|
||||
|
||||
pub use types::MessagesRequest;
|
||||
|
||||
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
|
||||
use prepare::prepare_messages_call;
|
||||
|
||||
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<Value> {
|
||||
match execute_messages(request, false).await? {
|
||||
MessagesResponse::Json(body) => Ok(body),
|
||||
MessagesResponse::Stream(response) => {
|
||||
drop(response);
|
||||
Err(litellm_core::CoreError::InvalidResponse(
|
||||
"non-streaming messages execution returned a stream".to_string(),
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) enum MessagesResponse {
|
||||
Json(Value),
|
||||
Stream(reqwest::Response),
|
||||
}
|
||||
|
||||
pub(crate) async fn execute_messages(
|
||||
request: MessagesRequest<'_>,
|
||||
stream: bool,
|
||||
) -> CoreResult<MessagesResponse> {
|
||||
let prepared = prepare_messages_call(request)?;
|
||||
if stream {
|
||||
execute_messages_provider_stream(prepared)
|
||||
.await
|
||||
.map(MessagesResponse::Stream)
|
||||
} else {
|
||||
execute_messages_provider_call(prepared)
|
||||
.await
|
||||
.map(MessagesResponse::Json)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
@ -1,24 +0,0 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_core::messages::transformation::AnthropicMessagesProviderConfig;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub struct MessagesRequest<'a> {
|
||||
pub model: &'a str,
|
||||
pub body: Value,
|
||||
pub api_key: Option<&'a str>,
|
||||
pub api_base: Option<&'a str>,
|
||||
pub custom_llm_provider: Option<&'a str>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub(crate) struct ProviderMessagesRequest {
|
||||
pub(crate) provider: String,
|
||||
pub(crate) model: String,
|
||||
pub(crate) config: &'static dyn AnthropicMessagesProviderConfig,
|
||||
pub(crate) url: String,
|
||||
pub(crate) body: Value,
|
||||
pub(crate) upstream_headers: Vec<(String, String)>,
|
||||
pub(crate) timeout: Option<Duration>,
|
||||
}
|
||||
|
|
@ -19,7 +19,10 @@ async fn handle(...) -> impl IntoResponse { ... }
|
|||
When a route has business logic worth testing without axum, put it in a sibling
|
||||
`service` (a file, or a folder if the route grows). The route file stays the
|
||||
**axum surface** (router + handler + any socket/SSE adapter); `service` is plain
|
||||
Rust with **no axum types**. `realtime/` is the example:
|
||||
Rust with **no axum types**, and its job is to pick the deployment and call the
|
||||
`core` route entrypoint (see `messages/service.rs` calling
|
||||
`litellm_core::messages::messages`). Never build a provider request, resolve a
|
||||
key, or perform the provider call here. `realtime/` is the older example:
|
||||
```
|
||||
realtime/
|
||||
mod.rs # axum surface: router() + handler + the WS<->events adapter
|
||||
|
|
@ -33,6 +36,8 @@ genuinely gets hard to read.
|
|||
`crate::auth::RequireMasterKey` to its arguments; it runs during extraction.
|
||||
Never re-implement the check per route.
|
||||
- **Handlers contain no business logic; `service` contains no axum types.**
|
||||
- **No provider handlers in this crate.** Transforms, auth headers, and the
|
||||
provider HTTP call live in `core/src/<route>/`.
|
||||
- A route owns its paths in its own `router()`; `mod.rs` only merges.
|
||||
- Cross-cutting concerns (logging, CORS, timeouts) → Tower layers in `mod.rs`,
|
||||
not duplicated in handlers.
|
||||
|
|
|
|||
|
|
@ -1,12 +1,12 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_core::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use litellm_core::messages::types::MessagesRequest;
|
||||
use litellm_core::messages::{messages, messages_stream};
|
||||
use litellm_core::router::Router;
|
||||
use litellm_core::{CoreError, CoreResult};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use crate::messages::{MessagesRequest, execute_messages};
|
||||
|
||||
pub(crate) enum MessagesResponse {
|
||||
Json(Value),
|
||||
Stream(reqwest::Response),
|
||||
|
|
@ -52,13 +52,14 @@ pub async fn run(
|
|||
extra_headers,
|
||||
timeout: None,
|
||||
};
|
||||
let stream = request.body.get("stream").and_then(Value::as_bool) == Some(true);
|
||||
execute_messages(request, stream)
|
||||
.await
|
||||
.map(|response| match response {
|
||||
crate::messages::MessagesResponse::Json(body) => MessagesResponse::Json(body),
|
||||
crate::messages::MessagesResponse::Stream(upstream) => {
|
||||
MessagesResponse::Stream(upstream)
|
||||
}
|
||||
if request.body.get("stream").and_then(Value::as_bool) == Some(true) {
|
||||
return messages_stream(request).await.map(MessagesResponse::Stream);
|
||||
}
|
||||
|
||||
let response = messages(request).await?;
|
||||
serde_json::to_value(response)
|
||||
.map(MessagesResponse::Json)
|
||||
.map_err(|err| {
|
||||
CoreError::InvalidResponse(format!("failed to serialize messages response: {err}"))
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,7 @@
|
|||
litellm-core is the PURE translation layer — types, route contracts (traits), provider transforms (modules under `providers/`), and the router. No network, no I/O, no env reads.
|
||||
litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route (`messages::messages()`, the Rust equivalent of `litellm.messages()`): you call it and get a typed non-streaming response back.
|
||||
|
||||
Routes (ocr, realtime) and providers (mistral, openai) are modules, not crates.
|
||||
A route module owns everything the call needs: types, the provider template trait, provider transforms (under `providers/`), provider/auth/URL resolution, and the handler that performs the HTTP call. Handlers belong here, never in a host crate.
|
||||
|
||||
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback dispatch. Env reads are limited to credential fallback in a route's `prepare.rs`.
|
||||
|
||||
Routes (messages, ocr, realtime) and providers (anthropic, mistral, openai) are modules, not crates.
|
||||
|
|
|
|||
|
|
@ -4,20 +4,28 @@ Rules for `litellm-rust/crates/core`.
|
|||
|
||||
## Responsibility
|
||||
|
||||
`core` owns shared data types, typed errors, and deterministic helper contracts.
|
||||
It must stay pure and host-independent.
|
||||
`core` is the LiteLLM SDK in Rust: it makes the LLM call. Every top-level
|
||||
LiteLLM call has a public entrypoint here, named after the route
|
||||
(`messages::messages()` is the Rust equivalent of `litellm.messages()`), and
|
||||
calling it returns a typed non-streaming response.
|
||||
|
||||
Allowed:
|
||||
- The public entrypoint for a route, plus its `<route>_stream` variant when the
|
||||
route supports streaming.
|
||||
- Provider resolution, auth header construction, URL building, and the provider
|
||||
HTTP call (shared reused client, connect + request timeouts).
|
||||
- Shared request/response structs.
|
||||
- Typed errors with stable, non-sensitive messages.
|
||||
- Deterministic validation helpers.
|
||||
- Serialization helpers that intentionally mirror Python output shape.
|
||||
- Route templates that match Python base config responsibilities, such as
|
||||
`ocr::transformation::OcrProviderConfig`.
|
||||
`messages::transformation::AnthropicMessagesProviderConfig`.
|
||||
|
||||
Not allowed:
|
||||
- Network, filesystem, database, cache, or environment access.
|
||||
- Secret reads or auth/header construction.
|
||||
- Serving HTTP: axum routers, extractors, and other transport concerns.
|
||||
- Filesystem, database, or cache access.
|
||||
- Config file reading or rollout state; the host resolves those and passes them
|
||||
in. Env reads are limited to credential fallback in a route's `prepare.rs`.
|
||||
- Logging callbacks, tracing spans, spend writes, or customer callbacks.
|
||||
- Provider-specific branching that belongs in `providers`.
|
||||
- Panics for user/provider-controlled input.
|
||||
|
|
@ -33,10 +41,21 @@ typed field on a struct, not a raw string threaded through the API.
|
|||
|
||||
## Structure
|
||||
|
||||
Use route names directly under `src/`: `ocr`, future `messages`,
|
||||
Use route names directly under `src/`: `messages`, `ocr`, future
|
||||
`chat_completions`, `embeddings`, and similar top-level LiteLLM calls. Do not
|
||||
invent broad names like `engine` for route contracts.
|
||||
|
||||
`src/messages` is the reference shape for a route module:
|
||||
|
||||
```
|
||||
mod.rs pub async fn messages(..) (+ messages_stream)
|
||||
types.rs request/response types
|
||||
transformation.rs the provider template trait
|
||||
prepare.rs provider resolution, auth headers, URL
|
||||
handler.rs the provider call
|
||||
client.rs the shared reqwest client
|
||||
```
|
||||
|
||||
## Parity Rules
|
||||
|
||||
- Every shared type used by a provider transform needs unit tests for
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ repository.workspace = true
|
|||
|
||||
[dependencies]
|
||||
rand.workspace = true
|
||||
reqwest.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
|
|
@ -30,5 +31,4 @@ bedrock-auth = [
|
|||
]
|
||||
|
||||
[dev-dependencies]
|
||||
reqwest.workspace = true
|
||||
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }
|
||||
|
|
|
|||
|
|
@ -1,3 +1,19 @@
|
|||
pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com";
|
||||
pub const OPENAI_RESPONSES_DEFAULT_API_BASE: &str = "https://api.openai.com/v1";
|
||||
pub const OPENAI_RESPONSES_PATH: &str = "/responses";
|
||||
|
||||
/// Full-request timeout ceiling for Anthropic Messages provider calls, in
|
||||
/// seconds. Mirrors the Python Anthropic Messages default. The per-request
|
||||
/// timeout from the caller still overrides this on the request builder.
|
||||
pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
/// Connect timeout for Anthropic Messages provider calls, in seconds.
|
||||
pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10;
|
||||
|
||||
/// Max characters of an upstream error body echoed across the call boundary
|
||||
/// before truncation, so provider bodies are bounded and data-minimized.
|
||||
pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256;
|
||||
|
||||
/// Provider name used for Anthropic Messages when a deployment's provider model
|
||||
/// does not carry an explicit provider prefix.
|
||||
pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
|
||||
|
|
|
|||
|
|
@ -1,11 +1,11 @@
|
|||
use litellm_core::CoreResult;
|
||||
use litellm_core::error::{CoreError, json_type_name};
|
||||
use litellm_core::messages::transformation::AnthropicMessagesProviderConfig;
|
||||
use litellm_core::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG;
|
||||
use litellm_core::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS;
|
||||
use crate::error::{CoreError, CoreResult, json_type_name};
|
||||
use crate::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG;
|
||||
use crate::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG;
|
||||
|
||||
use super::transformation::AnthropicMessagesProviderConfig;
|
||||
|
||||
pub(super) fn truncate_error_body(body: &str) -> String {
|
||||
if body.chars().count() <= MESSAGES_ERROR_BODY_MAX_CHARS {
|
||||
|
|
@ -1,15 +1,13 @@
|
|||
use litellm_core::CoreResult;
|
||||
use litellm_core::error::CoreError;
|
||||
use serde_json::Value;
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
|
||||
use super::client::http_client;
|
||||
use super::common_utils::truncate_error_body;
|
||||
use super::types::ProviderMessagesRequest;
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use super::types::{AnthropicMessagesResponse, ProviderMessagesRequest};
|
||||
|
||||
pub(super) async fn execute_messages_provider_call(
|
||||
request: ProviderMessagesRequest,
|
||||
) -> CoreResult<Value> {
|
||||
) -> CoreResult<AnthropicMessagesResponse> {
|
||||
let mut request_builder = http_client().post(&request.url).json(&request.body);
|
||||
for (key, value) in &request.upstream_headers {
|
||||
request_builder = request_builder.header(key, value);
|
||||
|
|
@ -39,12 +37,7 @@ pub(super) async fn execute_messages_provider_call(
|
|||
let response = serde_json::from_str(&text).map_err(|err| {
|
||||
CoreError::InvalidResponse(format!("invalid messages response JSON: {err}"))
|
||||
})?;
|
||||
let transformed = request
|
||||
.config
|
||||
.transform_response(&request.model, response)?;
|
||||
serde_json::to_value(transformed).map_err(|err| {
|
||||
CoreError::InvalidResponse(format!("failed to serialize messages response: {err}"))
|
||||
})
|
||||
request.config.transform_response(&request.model, response)
|
||||
}
|
||||
|
||||
pub(super) async fn execute_messages_provider_stream(
|
||||
|
|
@ -1,2 +1,32 @@
|
|||
//! The Anthropic Messages call, the Rust equivalent of Python's
|
||||
//! `litellm.messages()`.
|
||||
//!
|
||||
//! [`messages`] is the top-level entrypoint: give it a model, a body, and
|
||||
//! credentials, and it resolves the provider, transforms the request, calls the
|
||||
//! provider, and returns a typed non-streaming response. [`messages_stream`]
|
||||
//! is the streaming variant; it hands the raw upstream response back so a host
|
||||
//! can splice the event stream to its own caller.
|
||||
|
||||
mod client;
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod transformation;
|
||||
pub mod types;
|
||||
|
||||
use crate::error::CoreResult;
|
||||
|
||||
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
|
||||
use prepare::prepare_messages_call;
|
||||
use types::{AnthropicMessagesResponse, MessagesRequest};
|
||||
|
||||
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<AnthropicMessagesResponse> {
|
||||
execute_messages_provider_call(prepare_messages_call(request)?).await
|
||||
}
|
||||
|
||||
pub async fn messages_stream(request: MessagesRequest<'_>) -> CoreResult<reqwest::Response> {
|
||||
execute_messages_provider_stream(prepare_messages_call(request)?).await
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -1,9 +1,8 @@
|
|||
use litellm_core::CoreError;
|
||||
use litellm_core::CoreResult;
|
||||
use litellm_core::messages::transformation::MessagesAuthStrategy;
|
||||
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
|
||||
use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers};
|
||||
use super::transformation::MessagesAuthStrategy;
|
||||
use super::types::{MessagesRequest, ProviderMessagesRequest};
|
||||
|
||||
pub(super) fn prepare_messages_call(
|
||||
|
|
@ -1,14 +1,16 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_core::error::CoreError;
|
||||
use serde_json::{Map, Value, json};
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::{TcpListener, TcpStream};
|
||||
|
||||
use crate::error::CoreError;
|
||||
|
||||
use super::common_utils::{
|
||||
has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body,
|
||||
};
|
||||
use super::{MessagesRequest, messages};
|
||||
use super::messages;
|
||||
use super::types::MessagesRequest;
|
||||
|
||||
async fn read_http_request(socket: &mut TcpStream) -> String {
|
||||
let mut request = Vec::new();
|
||||
|
|
@ -152,8 +154,8 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through()
|
|||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
||||
assert_eq!(response["content"][0]["text"], "hi");
|
||||
assert_eq!(response["stop_reason"], "end_turn");
|
||||
assert_eq!(response.content[0]["text"], "hi");
|
||||
assert_eq!(response.stop_reason.as_deref(), Some("end_turn"));
|
||||
|
||||
let request = server.await.expect("server task completes");
|
||||
let (head, body) = request.split_once("\r\n\r\n").expect("has body");
|
||||
|
|
@ -208,8 +210,8 @@ async fn messages_round_trip_builds_native_anthropic_request() {
|
|||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
||||
assert_eq!(response["content"][0]["text"], "hi");
|
||||
assert_eq!(response["stop_reason"], "end_turn");
|
||||
assert_eq!(response.content[0]["text"], "hi");
|
||||
assert_eq!(response.stop_reason.as_deref(), Some("end_turn"));
|
||||
|
||||
let request = server.await.expect("server task completes");
|
||||
let (head, _) = request.split_once("\r\n\r\n").expect("has body");
|
||||
|
|
@ -1,6 +1,30 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::transformation::AnthropicMessagesProviderConfig;
|
||||
|
||||
pub struct MessagesRequest<'a> {
|
||||
pub model: &'a str,
|
||||
pub body: Value,
|
||||
pub api_key: Option<&'a str>,
|
||||
pub api_base: Option<&'a str>,
|
||||
pub custom_llm_provider: Option<&'a str>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub(super) struct ProviderMessagesRequest {
|
||||
pub(super) provider: String,
|
||||
pub(super) model: String,
|
||||
pub(super) config: &'static dyn AnthropicMessagesProviderConfig,
|
||||
pub(super) url: String,
|
||||
pub(super) body: Value,
|
||||
pub(super) upstream_headers: Vec<(String, String)>,
|
||||
pub(super) timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum SystemPrompt {
|
||||
|
|
|
|||
|
|
@ -1,3 +1,3 @@
|
|||
litellm-python-bridge is the PyO3 cdylib that exposes Rust to the litellm Python SDK — a thin adapter (Python objects → Rust calls → Python results) over litellm-ai-gateway.
|
||||
litellm-python-bridge is the PyO3 cdylib that exposes Rust to the litellm Python SDK — a thin adapter (Python objects → Rust calls → Python results) over the litellm-core route entrypoints (e.g. `litellm_core::messages::messages`).
|
||||
|
||||
Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call into litellm-ai-gateway.
|
||||
Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call the core entrypoint.
|
||||
|
|
|
|||
|
|
@ -11,11 +11,11 @@ Python-compatible dictionaries.
|
|||
## Bridge Shape
|
||||
|
||||
- Prefer one stable method per top-level LiteLLM route, for example
|
||||
`ocr(payload)`.
|
||||
`messages(...)`, calling the matching `litellm-core` entrypoint.
|
||||
- Do not add one exported PyO3 function per provider helper unless there is a
|
||||
measured reason.
|
||||
- Provider dispatch belongs in Rust route modules such as
|
||||
`litellm_providers::ocr`, not in this PyO3 crate.
|
||||
- Provider dispatch belongs in the `litellm-core` route module (e.g.
|
||||
`litellm_core::messages`), not in this PyO3 crate.
|
||||
- Python owns rollout state and fallback. Rust should return errors; Python
|
||||
decides whether to raise or fall back. For a rust-only provider/route (no
|
||||
Python reference), the Python side is a thin dispatch that calls Rust and
|
||||
|
|
|
|||
|
|
@ -4,10 +4,11 @@ use std::time::Duration;
|
|||
use litellm_ai_gateway::io::audio_transcription::{
|
||||
AudioTranscriptionRequest, audio_transcription as run_audio_transcription,
|
||||
};
|
||||
use litellm_ai_gateway::io::messages::{MessagesRequest, messages as run_messages};
|
||||
use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
|
||||
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
|
||||
use litellm_core::error::CoreError;
|
||||
use litellm_core::messages::messages as run_messages;
|
||||
use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest};
|
||||
use pyo3::exceptions::{PyRuntimeError, PyValueError};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyAny, PyDict};
|
||||
|
|
@ -35,6 +36,15 @@ fn json_to_py(py: Python<'_>, value: Value) -> PyResult<Py<PyAny>> {
|
|||
Ok(json.call_method1("loads", (encoded,))?.unbind())
|
||||
}
|
||||
|
||||
fn messages_response_to_py(
|
||||
py: Python<'_>,
|
||||
response: AnthropicMessagesResponse,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let value =
|
||||
serde_json::to_value(response).map_err(|err| PyValueError::new_err(err.to_string()))?;
|
||||
json_to_py(py, value)
|
||||
}
|
||||
|
||||
fn core_error_to_pyerr(err: CoreError) -> PyErr {
|
||||
match err {
|
||||
CoreError::Auth(message) => PyValueError::new_err(message),
|
||||
|
|
@ -382,7 +392,7 @@ fn messages(
|
|||
});
|
||||
|
||||
match result {
|
||||
Ok(value) => json_to_py(py, value),
|
||||
Ok(response) => messages_response_to_py(py, response),
|
||||
Err(err) => Err(core_error_to_pyerr(err)),
|
||||
}
|
||||
}
|
||||
|
|
@ -404,7 +414,7 @@ fn amessages(
|
|||
marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?;
|
||||
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
let value = run_messages(MessagesRequest {
|
||||
let response = run_messages(MessagesRequest {
|
||||
model: &model,
|
||||
body,
|
||||
api_key: api_key.as_deref(),
|
||||
|
|
@ -416,7 +426,7 @@ fn amessages(
|
|||
.await
|
||||
.map_err(core_error_to_pyerr)?;
|
||||
|
||||
Python::attach(|py| json_to_py(py, value))
|
||||
Python::attach(|py| messages_response_to_py(py, response))
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import json
|
||||
from typing import Any, Iterator, List, Literal, Optional, Tuple
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Iterable, Iterator, List, Literal, Optional, Tuple
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -24,20 +25,20 @@ async def calculate_batch_cost_and_usage(
|
|||
deployment-specific pricing (e.g. input_cost_per_token_batches)
|
||||
is used instead of the global cost map.
|
||||
"""
|
||||
batch_cost = _batch_cost_calculator(
|
||||
if (
|
||||
custom_llm_provider == "vertex_ai"
|
||||
and model_name
|
||||
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
|
||||
):
|
||||
batch_cost, batch_usage = calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
|
||||
return batch_cost, batch_usage, [model_name]
|
||||
|
||||
return _aggregate_batch_cost_usage_models(
|
||||
entries=file_content_dictionary,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
file_content_dictionary=file_content_dictionary,
|
||||
model_name=model_name,
|
||||
model_info=model_info,
|
||||
)
|
||||
batch_usage = _get_batch_job_total_usage_from_file_content(
|
||||
file_content_dictionary=file_content_dictionary,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_name=model_name,
|
||||
)
|
||||
batch_models = _get_batch_models_from_file_content(file_content_dictionary, model_name, custom_llm_provider)
|
||||
|
||||
return batch_cost, batch_usage, batch_models
|
||||
|
||||
|
||||
async def _handle_completed_batch(
|
||||
|
|
@ -46,7 +47,9 @@ async def _handle_completed_batch(
|
|||
model_name: Optional[str] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> Tuple[float, Usage, List[str]]:
|
||||
"""Helper function to process a completed batch and handle logging
|
||||
"""Fetch a completed batch's output file and aggregate its cost, usage, and
|
||||
models in a single pass over the JSONL lines, so the parsed file content is
|
||||
never materialized in memory.
|
||||
|
||||
Args:
|
||||
batch: The batch object
|
||||
|
|
@ -54,75 +57,109 @@ async def _handle_completed_batch(
|
|||
model_name: Optional model name
|
||||
litellm_params: Optional litellm parameters containing credentials (api_key, api_base, etc.)
|
||||
"""
|
||||
# Get batch results
|
||||
file_content_dictionary = await _get_batch_output_file_content_as_dictionary(
|
||||
batch, custom_llm_provider, litellm_params=litellm_params
|
||||
)
|
||||
file_content = await _fetch_batch_output_file_content(batch, custom_llm_provider, litellm_params=litellm_params)
|
||||
|
||||
# Calculate costs and usage
|
||||
batch_cost = _batch_cost_calculator(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
file_content_dictionary=file_content_dictionary,
|
||||
model_name=model_name,
|
||||
)
|
||||
batch_usage = _get_batch_job_total_usage_from_file_content(
|
||||
file_content_dictionary=file_content_dictionary,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_name=model_name,
|
||||
)
|
||||
|
||||
batch_models = _get_batch_models_from_file_content(file_content_dictionary, model_name, custom_llm_provider)
|
||||
|
||||
return batch_cost, batch_usage, batch_models
|
||||
|
||||
|
||||
def _get_batch_models_from_file_content(
|
||||
file_content_dictionary: List[dict],
|
||||
model_name: Optional[str] = None,
|
||||
custom_llm_provider: str = "openai",
|
||||
) -> List[str]:
|
||||
"""
|
||||
Get the models from the file content
|
||||
"""
|
||||
if model_name:
|
||||
return [model_name]
|
||||
batch_models = []
|
||||
for _item in file_content_dictionary:
|
||||
if _batch_response_was_successful(_item, custom_llm_provider):
|
||||
_response_body = _get_response_from_batch_job_output_file(_item, custom_llm_provider)
|
||||
_model = _response_body.get("model")
|
||||
if _model:
|
||||
batch_models.append(_model)
|
||||
return batch_models
|
||||
|
||||
|
||||
def _batch_cost_calculator(
|
||||
file_content_dictionary: List[dict],
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
|
||||
model_name: Optional[str] = None,
|
||||
model_info: Optional[ModelInfo] = None,
|
||||
) -> float:
|
||||
"""
|
||||
Calculate the cost of a batch based on the output file id
|
||||
"""
|
||||
if (
|
||||
custom_llm_provider == "vertex_ai"
|
||||
and model_name
|
||||
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
|
||||
):
|
||||
batch_cost, _ = calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
|
||||
verbose_logger.debug("vertex_ai_total_cost=%s", batch_cost)
|
||||
return batch_cost
|
||||
batch_cost, batch_usage = calculate_vertex_ai_batch_cost_and_usage(
|
||||
_get_file_content_as_dictionary(file_content), model_name
|
||||
)
|
||||
return batch_cost, batch_usage, [model_name]
|
||||
|
||||
# For other providers, use the existing logic
|
||||
total_cost = _get_batch_job_cost_from_file_content(
|
||||
file_content_dictionary=file_content_dictionary,
|
||||
return _aggregate_batch_cost_usage_models(
|
||||
entries=_iter_batch_input_entries(file_content),
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_name=model_name,
|
||||
model_info=model_info,
|
||||
)
|
||||
verbose_logger.debug("total_cost=%s", total_cost)
|
||||
return total_cost
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BatchOutputLineStats:
|
||||
cost: float
|
||||
prompt_tokens: int
|
||||
completion_tokens: int
|
||||
total_tokens: int
|
||||
cache_read_tokens: int
|
||||
cache_creation_tokens: int
|
||||
model: Optional[str]
|
||||
|
||||
|
||||
def _iter_successful_output_line_stats(
|
||||
entries: Iterable[dict],
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
|
||||
model_name: Optional[str],
|
||||
model_info: Optional[ModelInfo],
|
||||
) -> Iterator[_BatchOutputLineStats]:
|
||||
from litellm.cost_calculator import batch_cost_calculator
|
||||
|
||||
for entry in entries:
|
||||
if not _batch_response_was_successful(entry, custom_llm_provider):
|
||||
continue
|
||||
response_body = _get_response_from_batch_job_output_file(entry, custom_llm_provider)
|
||||
usage = _get_batch_job_usage_from_response_body(response_body, custom_llm_provider)
|
||||
prompt_details = _parse_prompt_tokens_details(usage)
|
||||
raw_model = response_body.get("model")
|
||||
response_model = raw_model if isinstance(raw_model, str) and raw_model else None
|
||||
if model_info is not None or custom_llm_provider in ("anthropic", "bedrock"):
|
||||
if custom_llm_provider == "bedrock" and model_name:
|
||||
cost_model = model_name
|
||||
else:
|
||||
cost_model = response_model or model_name or ""
|
||||
prompt_cost, completion_cost = batch_cost_calculator(
|
||||
usage=usage,
|
||||
model=cost_model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_info=model_info,
|
||||
)
|
||||
line_cost = prompt_cost + completion_cost
|
||||
else:
|
||||
line_cost = litellm.completion_cost(
|
||||
completion_response=response_body,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
call_type=CallTypes.aretrieve_batch.value,
|
||||
)
|
||||
yield _BatchOutputLineStats(
|
||||
cost=line_cost,
|
||||
prompt_tokens=usage.prompt_tokens,
|
||||
completion_tokens=usage.completion_tokens,
|
||||
total_tokens=usage.total_tokens,
|
||||
cache_read_tokens=prompt_details["cache_hit_tokens"],
|
||||
cache_creation_tokens=prompt_details["cache_creation_tokens"],
|
||||
model=response_model,
|
||||
)
|
||||
|
||||
|
||||
def _aggregate_batch_cost_usage_models(
|
||||
entries: Iterable[dict],
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
|
||||
model_name: Optional[str] = None,
|
||||
model_info: Optional[ModelInfo] = None,
|
||||
) -> Tuple[float, Usage, List[str]]:
|
||||
"""Aggregate cost, usage, and models from batch output entries in a single
|
||||
pass, holding one small stats record per line instead of the parsed file."""
|
||||
line_stats = tuple(_iter_successful_output_line_stats(entries, custom_llm_provider, model_name, model_info))
|
||||
|
||||
cache_token_params = {
|
||||
key: tokens
|
||||
for key, tokens in (
|
||||
("cache_read_input_tokens", sum(stats.cache_read_tokens for stats in line_stats)),
|
||||
("cache_creation_input_tokens", sum(stats.cache_creation_tokens for stats in line_stats)),
|
||||
)
|
||||
if tokens > 0
|
||||
}
|
||||
batch_usage = Usage(
|
||||
total_tokens=sum(stats.total_tokens for stats in line_stats),
|
||||
prompt_tokens=sum(stats.prompt_tokens for stats in line_stats),
|
||||
completion_tokens=sum(stats.completion_tokens for stats in line_stats),
|
||||
**cache_token_params,
|
||||
)
|
||||
batch_models = [model_name] if model_name else [stats.model for stats in line_stats if stats.model]
|
||||
total_cost = sum((stats.cost for stats in line_stats), 0.0)
|
||||
verbose_logger.debug("batch output aggregate: cost=%s usage=%s models=%s", total_cost, batch_usage, batch_models)
|
||||
return total_cost, batch_usage, batch_models
|
||||
|
||||
|
||||
def calculate_vertex_ai_batch_cost_and_usage(
|
||||
|
|
@ -193,13 +230,13 @@ def calculate_vertex_ai_batch_cost_and_usage(
|
|||
)
|
||||
|
||||
|
||||
async def _get_batch_output_file_content_as_dictionary(
|
||||
async def _fetch_batch_output_file_content(
|
||||
batch: Batch,
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> List[dict]:
|
||||
) -> bytes:
|
||||
"""
|
||||
Get the batch output file content as a list of dictionaries
|
||||
Fetch the batch output file and return its raw JSONL bytes
|
||||
|
||||
Args:
|
||||
batch: The batch object
|
||||
|
|
@ -212,9 +249,6 @@ async def _get_batch_output_file_content_as_dictionary(
|
|||
_is_base64_encoded_unified_file_id,
|
||||
)
|
||||
|
||||
if custom_llm_provider == "vertex_ai":
|
||||
raise ValueError("Vertex AI does not support file content retrieval")
|
||||
|
||||
if batch.output_file_id is None:
|
||||
raise ValueError("Output file id is None cannot retrieve file content")
|
||||
|
||||
|
|
@ -240,7 +274,7 @@ async def _get_batch_output_file_content_as_dictionary(
|
|||
file_content_kwargs.update(credentials)
|
||||
|
||||
_file_content = await afile_content(**file_content_kwargs) # type: ignore[reportArgumentType]
|
||||
return _get_file_content_as_dictionary(_file_content.content)
|
||||
return _file_content.content
|
||||
|
||||
|
||||
def _extract_file_access_credentials(litellm_params: Optional[dict]) -> dict:
|
||||
|
|
@ -270,6 +304,8 @@ def _extract_file_access_credentials(litellm_params: Optional[dict]) -> dict:
|
|||
"vertex_project",
|
||||
"vertex_location",
|
||||
"vertex_credentials",
|
||||
"gcs_bucket_name",
|
||||
"bucket_name",
|
||||
"timeout",
|
||||
"max_retries",
|
||||
]
|
||||
|
|
@ -284,17 +320,7 @@ def _get_file_content_as_dictionary(file_content: bytes) -> List[dict]:
|
|||
"""
|
||||
Get the file content as a list of dictionaries from JSON Lines format
|
||||
"""
|
||||
try:
|
||||
_file_content_str = file_content.decode("utf-8")
|
||||
# Split by newlines and parse each line as a separate JSON object
|
||||
json_objects = []
|
||||
for line in _file_content_str.strip().split("\n"):
|
||||
if line: # Skip empty lines
|
||||
json_objects.append(json.loads(line))
|
||||
verbose_logger.debug("json_objects=%s", json.dumps(json_objects, indent=4))
|
||||
return json_objects
|
||||
except Exception as e:
|
||||
raise e
|
||||
return list(_iter_batch_input_entries(file_content))
|
||||
|
||||
|
||||
def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
|
||||
|
|
@ -361,101 +387,6 @@ def _count_entry_tokens(
|
|||
return 0
|
||||
|
||||
|
||||
def _get_batch_job_cost_from_file_content(
|
||||
file_content_dictionary: List[dict],
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
|
||||
model_name: Optional[str] = None,
|
||||
model_info: Optional[ModelInfo] = None,
|
||||
) -> float:
|
||||
"""
|
||||
Get the cost of a batch job from the file content
|
||||
"""
|
||||
from litellm.cost_calculator import batch_cost_calculator
|
||||
|
||||
try:
|
||||
total_cost: float = 0.0
|
||||
# parse the file content as json
|
||||
verbose_logger.debug("file_content_dictionary=%s", json.dumps(file_content_dictionary, indent=4))
|
||||
for _item in file_content_dictionary:
|
||||
if _batch_response_was_successful(_item, custom_llm_provider):
|
||||
_response_body = _get_response_from_batch_job_output_file(_item, custom_llm_provider)
|
||||
if model_info is not None or custom_llm_provider in ("anthropic", "bedrock"):
|
||||
usage = _get_batch_job_usage_from_response_body(_response_body, custom_llm_provider)
|
||||
# Bedrock batch output lines report a short internal model id
|
||||
# (e.g. "claude-sonnet-4-6") that is not in the cost map; use the
|
||||
# deployment model name for pricing when available.
|
||||
if custom_llm_provider == "bedrock" and model_name:
|
||||
model = model_name
|
||||
else:
|
||||
model = _response_body.get("model") or model_name or ""
|
||||
prompt_cost, completion_cost = batch_cost_calculator(
|
||||
usage=usage,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_info=model_info,
|
||||
)
|
||||
total_cost += prompt_cost + completion_cost
|
||||
else:
|
||||
total_cost += litellm.completion_cost(
|
||||
completion_response=_response_body,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
call_type=CallTypes.aretrieve_batch.value,
|
||||
)
|
||||
verbose_logger.debug("total_cost=%s", total_cost)
|
||||
return total_cost
|
||||
except Exception as e:
|
||||
verbose_logger.error("error in _get_batch_job_cost_from_file_content", e)
|
||||
raise e
|
||||
|
||||
|
||||
def _get_batch_job_total_usage_from_file_content(
|
||||
file_content_dictionary: List[dict],
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
|
||||
model_name: Optional[str] = None,
|
||||
) -> Usage:
|
||||
"""
|
||||
Get the tokens of a batch job from the file content
|
||||
"""
|
||||
if (
|
||||
custom_llm_provider == "vertex_ai"
|
||||
and model_name
|
||||
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
|
||||
):
|
||||
_, batch_usage = calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
|
||||
return batch_usage
|
||||
|
||||
# For other providers, use the existing logic
|
||||
total_tokens: int = 0
|
||||
prompt_tokens: int = 0
|
||||
completion_tokens: int = 0
|
||||
cache_read_tokens: int = 0
|
||||
cache_creation_tokens: int = 0
|
||||
for _item in file_content_dictionary:
|
||||
if _batch_response_was_successful(_item, custom_llm_provider):
|
||||
_response_body = _get_response_from_batch_job_output_file(_item, custom_llm_provider)
|
||||
usage: Usage = _get_batch_job_usage_from_response_body(_response_body, custom_llm_provider)
|
||||
total_tokens += usage.total_tokens
|
||||
prompt_tokens += usage.prompt_tokens
|
||||
completion_tokens += usage.completion_tokens
|
||||
prompt_details = _parse_prompt_tokens_details(usage)
|
||||
cache_read_tokens += prompt_details["cache_hit_tokens"]
|
||||
cache_creation_tokens += prompt_details["cache_creation_tokens"]
|
||||
cache_token_params = {
|
||||
key: tokens
|
||||
for key, tokens in (
|
||||
("cache_read_input_tokens", cache_read_tokens),
|
||||
("cache_creation_input_tokens", cache_creation_tokens),
|
||||
)
|
||||
if tokens > 0
|
||||
}
|
||||
return Usage(
|
||||
total_tokens=total_tokens,
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
**cache_token_params,
|
||||
)
|
||||
|
||||
|
||||
def _count_prompt_or_input_tokens(model: str, value: Any) -> int:
|
||||
"""Token-count a ``prompt`` / ``input`` field that the OpenAI batch
|
||||
schema allows in four shapes:
|
||||
|
|
|
|||
|
|
@ -209,7 +209,15 @@ class ResponsesToCompletionBridgeHandler:
|
|||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
elif isinstance(result, ModelResponse):
|
||||
return result
|
||||
if not stream:
|
||||
return result
|
||||
return self._completed_response_as_stream(
|
||||
response=result,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
logging_obj=logging_obj,
|
||||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
elif not stream:
|
||||
responses_api_response = self._collect_response_from_stream(result)
|
||||
return self.transformation_handler.transform_response(
|
||||
|
|
@ -299,7 +307,15 @@ class ResponsesToCompletionBridgeHandler:
|
|||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
elif isinstance(result, ModelResponse):
|
||||
return result
|
||||
if not stream:
|
||||
return result
|
||||
return self._completed_response_as_stream(
|
||||
response=result,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
logging_obj=logging_obj,
|
||||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
elif not stream:
|
||||
responses_api_response = await self._collect_response_from_stream_async(result)
|
||||
return self.transformation_handler.transform_response(
|
||||
|
|
@ -331,6 +347,25 @@ class ResponsesToCompletionBridgeHandler:
|
|||
)
|
||||
return self._apply_post_stream_processing(streamwrapper, model, custom_llm_provider)
|
||||
|
||||
def _completed_response_as_stream(
|
||||
self,
|
||||
response: "ModelResponse",
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
json_mode: bool | None,
|
||||
) -> "CustomStreamWrapper":
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
|
||||
streamwrapper = CustomStreamWrapper(
|
||||
completion_stream=MockResponseIterator(model_response=response, json_mode=json_mode),
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
return self._apply_post_stream_processing(streamwrapper, model, custom_llm_provider)
|
||||
|
||||
@staticmethod
|
||||
def _apply_post_stream_processing(
|
||||
stream: "CustomStreamWrapper",
|
||||
|
|
|
|||
|
|
@ -1077,6 +1077,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
||||
def __init__(self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False):
|
||||
super().__init__(streaming_response, sync_stream, json_mode)
|
||||
self._chat_completion_id: str | None = None
|
||||
|
||||
def _handle_string_chunk(
|
||||
self, str_line: Union[str, "BaseModel"]
|
||||
|
|
@ -1384,4 +1385,13 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
ModelResponseStream: OpenAI-formatted streaming chunk
|
||||
"""
|
||||
verbose_logger.debug(f"Chat provider: transform_streaming_response called with chunk: {chunk}")
|
||||
return OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(chunk)
|
||||
return self._with_stream_scoped_id(
|
||||
OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(chunk)
|
||||
)
|
||||
|
||||
def _with_stream_scoped_id(self, chunk: "ModelResponseStream") -> "ModelResponseStream":
|
||||
if self._chat_completion_id is None:
|
||||
self._chat_completion_id = chunk.id
|
||||
else:
|
||||
chunk.id = self._chat_completion_id
|
||||
return chunk
|
||||
|
|
|
|||
|
|
@ -1418,6 +1418,8 @@ LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE = int(
|
|||
os.getenv("LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE", 1000)
|
||||
)
|
||||
LITELLM_PROXY_ADMIN_NAME = "default_user_id"
|
||||
LITELLM_PROXY_BUDGET_NAME = "litellm-proxy-budget"
|
||||
GLOBAL_PROXY_SPEND_CACHE_KEY = f"{LITELLM_PROXY_ADMIN_NAME}:spend"
|
||||
|
||||
########################### CLI SSO AUTHENTICATION CONSTANTS ###########################
|
||||
LITELLM_CLI_SOURCE_IDENTIFIER = "litellm-cli"
|
||||
|
|
@ -1455,6 +1457,7 @@ SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(os.getenv("SPEND_LOG_CLEA
|
|||
SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS = float(
|
||||
os.getenv("SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.5)
|
||||
)
|
||||
TOOL_SPEND_TOP_TOOLS = 100
|
||||
SPEND_LOG_PARTITION_INTERVAL = os.getenv("SPEND_LOG_PARTITION_INTERVAL", "day")
|
||||
SPEND_LOG_PARTITION_PRECREATE_AHEAD = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7))
|
||||
SPEND_LOG_QUEUE_SIZE_THRESHOLD = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100))
|
||||
|
|
|
|||
|
|
@ -69,6 +69,8 @@ from litellm.exceptions import (
|
|||
# proxy's metadata sanitizer.
|
||||
_PRE_CALL_EXECUTED_TOKEN = secrets.token_hex(16)
|
||||
|
||||
_GUARDRAIL_BLOCK_STATUS_CODES = frozenset({400, 403, 422})
|
||||
|
||||
_guardrail_self_recorded: contextvars.ContextVar[bool] = contextvars.ContextVar(
|
||||
"litellm_guardrail_self_recorded", default=False
|
||||
)
|
||||
|
|
@ -1055,8 +1057,15 @@ class CustomGuardrail(CustomLogger):
|
|||
- GuardrailRaisedException (generic guardrail API, tool permission)
|
||||
- BlockedPiiEntityError (Presidio PII detection)
|
||||
- SensitiveDataRouteException (sensitive-data reroute to on-premise model)
|
||||
- HTTPException with status 400 (content policy violation)
|
||||
- HTTPException with a block-signalling status (400, 403, 422)
|
||||
- ModifyResponseException (passthrough mode violation)
|
||||
|
||||
Only the statuses guardrails use in-tree to signal a deliberate rejection
|
||||
count as an intervention: 400 (content policy), 403 (e.g. akto) and 422
|
||||
(e.g. llm_as_a_judge). Other 4xx codes are commonly propagated from an
|
||||
upstream guardrail provider response (401 bad key, 408 timeout, 429 rate
|
||||
limit, or a raw upstream status), which are technical failures, not
|
||||
blocks, so they stay guardrail_failed_to_respond.
|
||||
"""
|
||||
if isinstance(e, ModifyResponseException):
|
||||
return True
|
||||
|
|
@ -1069,7 +1078,11 @@ class CustomGuardrail(CustomLogger):
|
|||
),
|
||||
):
|
||||
return True
|
||||
if HTTPException is not None and isinstance(e, HTTPException) and e.status_code == 400:
|
||||
if (
|
||||
HTTPException is not None
|
||||
and isinstance(e, HTTPException)
|
||||
and e.status_code in _GUARDRAIL_BLOCK_STATUS_CODES
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
|
|
|||
|
|
@ -133,6 +133,15 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
"dotted_order": metadata.get("dotted_order", None),
|
||||
}
|
||||
|
||||
def _redact_metadata(self, metadata: dict) -> dict:
|
||||
# helper is shallow; also scrub nested requester_metadata since
|
||||
# LangSmith forwards the whole dict into the run
|
||||
redacted = redact_user_api_key_info(metadata=dict(metadata))
|
||||
nested = redacted.get("requester_metadata")
|
||||
if isinstance(nested, dict):
|
||||
redacted["requester_metadata"] = redact_user_api_key_info(metadata=nested)
|
||||
return redacted
|
||||
|
||||
def _build_extra_metadata(self, metadata: Dict):
|
||||
extra_metadata = dict(metadata)
|
||||
requester_metadata = extra_metadata.get("requester_metadata")
|
||||
|
|
@ -141,13 +150,7 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
if key in requester_metadata and key not in extra_metadata:
|
||||
extra_metadata[key] = requester_metadata[key]
|
||||
|
||||
# helper is shallow; also scrub nested requester_metadata since
|
||||
# LangSmith forwards the whole dict into `extra`
|
||||
extra_metadata = redact_user_api_key_info(metadata=extra_metadata)
|
||||
nested = extra_metadata.get("requester_metadata")
|
||||
if isinstance(nested, dict):
|
||||
extra_metadata["requester_metadata"] = redact_user_api_key_info(metadata=nested)
|
||||
return extra_metadata
|
||||
return self._redact_metadata(extra_metadata)
|
||||
|
||||
def _build_outputs_with_usage(self, payload: StandardLoggingPayload) -> Dict[str, Any]:
|
||||
response = payload["response"]
|
||||
|
|
@ -200,12 +203,13 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
|
||||
metadata = payload["metadata"]
|
||||
extra_metadata = self._build_extra_metadata(dict(metadata))
|
||||
inputs = {**payload, "metadata": self._redact_metadata(dict(metadata))}
|
||||
outputs = self._build_outputs_with_usage(payload)
|
||||
|
||||
data = {
|
||||
"name": fields["run_name"],
|
||||
"run_type": "llm",
|
||||
"inputs": payload,
|
||||
"inputs": inputs,
|
||||
"outputs": outputs,
|
||||
"session_name": fields["project_name"],
|
||||
"start_time": payload["startTime"],
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from litellm.integrations.opentelemetry_utils.gen_ai_semconv import (
|
|||
OTELSemconvCategory,
|
||||
parse_semconv_opt_in,
|
||||
)
|
||||
from litellm.integrations.otel.model.semconv import Metric
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.litellm_core_utils.secret_redaction import redact_string
|
||||
from litellm.secret_managers.main import get_secret_bool, str_to_bool
|
||||
|
|
@ -597,32 +598,32 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
meter = meter_provider.get_meter(__name__)
|
||||
|
||||
self._operation_duration_histogram = meter.create_histogram(
|
||||
name="gen_ai.client.operation.duration", # Replace with semconv constant in otel 1.38
|
||||
name=Metric.OPERATION_DURATION,
|
||||
description="GenAI operation duration",
|
||||
unit="s",
|
||||
)
|
||||
self._token_usage_histogram = meter.create_histogram(
|
||||
name="gen_ai.client.token.usage", # Replace with semconv constant in otel 1.38
|
||||
name=Metric.TOKEN_USAGE,
|
||||
description="GenAI token usage",
|
||||
unit="{token}",
|
||||
)
|
||||
self._cost_histogram = meter.create_histogram(
|
||||
name="gen_ai.client.token.cost",
|
||||
name=Metric.TOKEN_COST,
|
||||
description="GenAI request cost",
|
||||
unit="USD",
|
||||
)
|
||||
self._time_to_first_token_histogram = meter.create_histogram(
|
||||
name="gen_ai.client.response.time_to_first_token",
|
||||
name=Metric.TIME_TO_FIRST_TOKEN,
|
||||
description="Time to first token for streaming requests",
|
||||
unit="s",
|
||||
)
|
||||
self._time_per_output_token_histogram = meter.create_histogram(
|
||||
name="gen_ai.client.response.time_per_output_token",
|
||||
name=Metric.TIME_PER_OUTPUT_TOKEN,
|
||||
description="Average time per output token (generation time / completion tokens)",
|
||||
unit="s",
|
||||
)
|
||||
self._response_duration_histogram = meter.create_histogram(
|
||||
name="gen_ai.client.response.duration",
|
||||
name=Metric.RESPONSE_DURATION,
|
||||
description="Total LLM API generation time (excludes LiteLLM overhead)",
|
||||
unit="s",
|
||||
)
|
||||
|
|
@ -2980,10 +2981,12 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
def _get_metric_reader(self):
|
||||
"""
|
||||
Get the appropriate metric reader based on the configuration.
|
||||
|
||||
Histograms keep the SDK's default cumulative temporality: Prometheus-backed
|
||||
OTLP receivers reject delta histograms and drop the whole batch, while
|
||||
backends that prefer delta still accept cumulative.
|
||||
"""
|
||||
from opentelemetry.sdk.metrics import Histogram
|
||||
from opentelemetry.sdk.metrics.export import (
|
||||
AggregationTemporality,
|
||||
ConsoleMetricExporter,
|
||||
PeriodicExportingMetricReader,
|
||||
)
|
||||
|
|
@ -3014,7 +3017,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
exporter = OTLPMetricExporter(
|
||||
endpoint=normalized_endpoint,
|
||||
headers=_split_otel_headers,
|
||||
preferred_temporality={Histogram: AggregationTemporality.DELTA},
|
||||
)
|
||||
return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
|
||||
|
||||
|
|
@ -3032,7 +3034,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
exporter = OTLPMetricExporter(
|
||||
endpoint=normalized_endpoint,
|
||||
headers=_split_otel_headers,
|
||||
preferred_temporality={Histogram: AggregationTemporality.DELTA},
|
||||
)
|
||||
return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
|
||||
|
||||
|
|
|
|||
|
|
@ -222,7 +222,29 @@ lives in [`plumbing/`](./plumbing):
|
|||
otherwise the operator's globally configured `MeterProvider` is reused so its
|
||||
readers/exporters receive them alongside the server metrics, and one is built
|
||||
and registered as the global only when none is set (mirroring how V2 owns trace
|
||||
export).
|
||||
export). A **failed** call records `gen_ai.client.operation.duration` too,
|
||||
carrying the semconv `error.type` (the mapped provider exception's class name),
|
||||
so the histogram covers the whole traffic and failure-rate panels are buildable;
|
||||
the other five instruments describe a completed generation and are skipped
|
||||
rather than filled with a fabricated zero. `error.type` is stamped after the
|
||||
cardinality filter, so an `otel.attributes` list cannot merge the failure series
|
||||
back into the success series. A proxy-gate rejection (auth / rate limit) records
|
||||
nothing, for the same reason it gets no span: no upstream call happened.
|
||||
Both paths cap their attributes at `METRIC_ATTRIBUTE_CEILING` before the
|
||||
operator's own `otel.attributes` filter runs, so the filter can narrow the set
|
||||
but never widen it. The ceiling is what keeps series count bounded by the
|
||||
deployment's own key/team/user/deployment count instead of by its traffic: a
|
||||
label value that moves per request mints a time series per request, which both
|
||||
bills per request on a hosted backend and leaves a histogram that cannot be
|
||||
aggregated. So client-supplied and per-request metadata (`requester_metadata`,
|
||||
`spend_logs_metadata`, `user_api_key_end_user_id`, `requester_ip_address`) is
|
||||
metric-ineligible and stays on the span, where cardinality is free, and the
|
||||
`hidden_params` label carries only `model_id`, the deployment identity a
|
||||
per-deployment panel joins on. `api_base` is excluded despite naming the same
|
||||
deployment, because it is a documented per-call parameter and so is caller-chosen
|
||||
in SDK use. Because the shared validator accepts every span attribute name, a
|
||||
filter that names a metric-ineligible one logs a warning once when the filter
|
||||
resolves rather than silently emitting nothing for it.
|
||||
- [`events.py`](./plumbing/events.py) — GenAI client events. Gated on
|
||||
`enable_events` (`LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS`), a failed LLM call
|
||||
records the semconv `gen_ai.client.operation.exception` log event at severity
|
||||
|
|
@ -267,7 +289,10 @@ lives in [`plumbing/`](./plumbing):
|
|||
|
||||
- **A new attribute vocabulary for a backend**: add a mapper in `mappers/`
|
||||
(a class with a `map(data) -> AttributeMap` method, typically built from
|
||||
`key -> extractor` tables) and register it in `mappers/__init__._MAPPER_BY_NAME`.
|
||||
`key -> extractor` tables) and register it in `mappers/__init__._PLAIN_MAPPERS`.
|
||||
If it spells declared tool definitions out per index, register it in
|
||||
`_TOOL_DEFINITION_MAPPERS` instead and take the shared attribute budget in its
|
||||
constructor, so the family stays bounded span-wide rather than per vocabulary.
|
||||
- **A new integration**: add a preset in `presets/` that returns an
|
||||
`OpenTelemetryV2Config`, and register it in `presets/__init__.PRESET_BY_CALLBACK`.
|
||||
If it supports dynamic credentials, add a header builder to
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from litellm.integrations.otel.model.baggage import promoted_baggage
|
|||
from litellm.integrations.otel.model.config import OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.plumbing.context import (
|
||||
is_recordable_span,
|
||||
mcp_message_transport_span,
|
||||
request_root_span,
|
||||
resolve_mcp_span_context,
|
||||
resolve_parent_context,
|
||||
|
|
@ -280,13 +281,29 @@ class OpenTelemetryV2(CustomLogger):
|
|||
self._record_metrics(kwargs, response_obj, start_time, end_time)
|
||||
|
||||
def _record_metrics(self, kwargs, response_obj, start_time, end_time) -> None:
|
||||
"""Record the GenAI metrics for a successful LLM call. Best-effort: a
|
||||
recording failure (e.g. a malformed payload) must never break the span
|
||||
close or the request itself."""
|
||||
"""Record the GenAI metrics for a successful LLM call."""
|
||||
self._guarded_record(lambda recorder: recorder.record(kwargs, response_obj, start_time, end_time))
|
||||
|
||||
def _record_failure_metrics(self, kwargs, start_time, end_time) -> None:
|
||||
"""Record the GenAI metrics for a failed LLM call, so the duration
|
||||
histogram covers the whole traffic rather than only what survived.
|
||||
|
||||
A synthetic proxy-gate log (auth / rate-limit rejection) is skipped for the
|
||||
same reason it gets no span: no upstream call happened, so its duration is
|
||||
not a GenAI operation's duration and would pull the histogram down."""
|
||||
if LLMCallEvent.from_dict(kwargs).is_no_upstream_call:
|
||||
return
|
||||
self._guarded_record(lambda recorder: recorder.record_failure(kwargs, start_time, end_time))
|
||||
|
||||
def _guarded_record(self, record: "Callable[[GenAIMetricRecorder], None]") -> None:
|
||||
"""Run one metric recording. Best-effort: a recording failure (e.g. a
|
||||
malformed payload) must never break the span close or the request itself. A
|
||||
misconfigured attribute filter is operator-fixable, so it is surfaced once
|
||||
at ERROR instead of being swallowed."""
|
||||
if self._metrics_recorder is None:
|
||||
return
|
||||
try:
|
||||
self._metrics_recorder.record(kwargs, response_obj, start_time, end_time)
|
||||
record(self._metrics_recorder)
|
||||
except ValueError as exc:
|
||||
if not self._metric_filter_error_logged:
|
||||
verbose_logger.error(
|
||||
|
|
@ -303,6 +320,7 @@ class OpenTelemetryV2(CustomLogger):
|
|||
if self._emit_mcp_list_tools(kwargs, start_time, end_time):
|
||||
return
|
||||
self._close_llm_call(kwargs, start_time, end_time)
|
||||
self._record_failure_metrics(kwargs, start_time, end_time)
|
||||
|
||||
def _seed_identity_baggage(self, identity: RequestIdentity, model: str | None, context: Context) -> Context:
|
||||
"""Seed authenticated request-identity Baggage onto ``context`` so the Baggage
|
||||
|
|
@ -641,8 +659,14 @@ class OpenTelemetryV2(CustomLogger):
|
|||
endpoint, auth failure), so the failed request carries the same error keys
|
||||
a failed LLM call does. v1's ``OpenTelemetry`` implemented this same hook;
|
||||
v2 lost it when it stopped subclassing ``OpenTelemetry``, which is the
|
||||
LIT-4179 regression for pre-call failures."""
|
||||
span = request_root_span() or user_api_key_dict.parent_otel_span
|
||||
LIT-4179 regression for pre-call failures.
|
||||
|
||||
An MCP message is handled on the session's task, where the request-root
|
||||
anchor is whatever request opened the session, so prefer the transport the
|
||||
gateway published for this specific message. Without that, a failed tool
|
||||
call aimed its error at the ``initialize`` request's finished span and the
|
||||
SDK dropped it, leaving the POST that actually failed unmarked."""
|
||||
span = mcp_message_transport_span() or request_root_span() or user_api_key_dict.parent_otel_span
|
||||
if span is None or not is_recordable_span(span):
|
||||
return None
|
||||
stamp_error(span, _span_error_from_exception(original_exception, traceback_str=traceback_str))
|
||||
|
|
|
|||
|
|
@ -18,13 +18,19 @@ from litellm.integrations.otel.mappers.langfuse import LangfuseMapper
|
|||
from litellm.integrations.otel.mappers.langtrace import LangtraceMapper
|
||||
from litellm.integrations.otel.mappers.legacy import LegacyMapper
|
||||
from litellm.integrations.otel.mappers.openinference import OpenInferenceMapper
|
||||
from litellm.integrations.otel.mappers.utils import tool_attr_budget
|
||||
from litellm.integrations.otel.mappers.weave import WeaveMapper
|
||||
|
||||
# Registry keyed by ``config.mapper_names`` entries.
|
||||
_MAPPER_BY_NAME: dict[str, Callable[[], AttributeMapper]] = {
|
||||
# Registries keyed by ``config.mapper_names`` entries, split by whether the
|
||||
# vocabulary spells declared tool definitions out per index. Those share one
|
||||
# span-wide attribute ceiling, so resolution has to know how many of them are
|
||||
# active before it can build them.
|
||||
_TOOL_DEFINITION_MAPPERS: dict[str, Callable[[int], AttributeMapper]] = {
|
||||
"genai": GenAIMapper,
|
||||
"legacy": LegacyMapper,
|
||||
"openinference": OpenInferenceMapper,
|
||||
}
|
||||
_PLAIN_MAPPERS: dict[str, Callable[[], AttributeMapper]] = {
|
||||
"langfuse": LangfuseMapper,
|
||||
"weave": WeaveMapper,
|
||||
"langtrace": LangtraceMapper,
|
||||
|
|
@ -33,13 +39,19 @@ _MAPPER_BY_NAME: dict[str, Callable[[], AttributeMapper]] = {
|
|||
|
||||
def resolve_mappers(names: Iterable[str]) -> list[AttributeMapper]:
|
||||
"""Resolve mapper names to instances. Unknown names raise ``ValueError``."""
|
||||
out: list[AttributeMapper] = []
|
||||
for name in names:
|
||||
factory = _MAPPER_BY_NAME.get(name)
|
||||
if factory is None:
|
||||
raise ValueError(f"unknown mapper name {name!r}; known: {sorted(_MAPPER_BY_NAME)}")
|
||||
out.append(factory())
|
||||
return out
|
||||
ordered = tuple(names)
|
||||
for name in ordered:
|
||||
if name not in _TOOL_DEFINITION_MAPPERS and name not in _PLAIN_MAPPERS:
|
||||
known = sorted((*_TOOL_DEFINITION_MAPPERS, *_PLAIN_MAPPERS))
|
||||
raise ValueError(f"unknown mapper name {name!r}; known: {known}")
|
||||
# Distinct vocabularies each write the tool family under their own keys, so
|
||||
# the ceiling is split by how many of them are configured. Repeating a name
|
||||
# rewrites the same keys, so only distinct ones count.
|
||||
budget = tool_attr_budget(len({*ordered} & _TOOL_DEFINITION_MAPPERS.keys()))
|
||||
return [
|
||||
_TOOL_DEFINITION_MAPPERS[name](budget) if name in _TOOL_DEFINITION_MAPPERS else _PLAIN_MAPPERS[name]()
|
||||
for name in ordered
|
||||
]
|
||||
|
||||
|
||||
__all__ = [
|
||||
|
|
|
|||
|
|
@ -11,10 +11,11 @@ from typing import Callable
|
|||
|
||||
from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, SpanData
|
||||
from litellm.integrations.otel.mappers.utils import (
|
||||
MAX_TOOL_DEFINITION_ATTRS_PER_SPAN,
|
||||
collect,
|
||||
drop_none,
|
||||
output_messages,
|
||||
serialize_messages,
|
||||
tool_definition_attrs,
|
||||
)
|
||||
from litellm.integrations.otel.model.payloads import (
|
||||
GuardrailSpanData,
|
||||
|
|
@ -135,6 +136,9 @@ class GenAIMapper:
|
|||
LiteLLM.SERVICE_CALL_TYPE: lambda d: d.call_type,
|
||||
}
|
||||
|
||||
def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None:
|
||||
self._tool_attr_budget = tool_attr_budget
|
||||
|
||||
def map(self, data: SpanData) -> AttributeMap:
|
||||
match data:
|
||||
case LLMCallSpanData():
|
||||
|
|
@ -150,18 +154,18 @@ class GenAIMapper:
|
|||
case _:
|
||||
return {}
|
||||
|
||||
@classmethod
|
||||
def _llm_call(cls, data: LLMCallSpanData) -> AttributeMap:
|
||||
attrs = collect(cls._LLM_CALL_ATTRS, data)
|
||||
attrs.update(
|
||||
drop_none(
|
||||
{
|
||||
f"gen_ai.tool.{idx}.{suffix}": extract(tool)
|
||||
for idx, tool in enumerate(data.tools)
|
||||
for suffix, extract in cls._TOOL_ATTRS.items()
|
||||
}
|
||||
def _llm_call(self, data: LLMCallSpanData) -> AttributeMap:
|
||||
attrs = collect(self._LLM_CALL_ATTRS, data)
|
||||
if data.tools:
|
||||
attrs[LiteLLM.TOOLS_DECLARED] = len(data.tools)
|
||||
attrs.update(
|
||||
tool_definition_attrs(
|
||||
lambda idx, suffix: f"gen_ai.tool.{idx}.{suffix}",
|
||||
data.tools,
|
||||
self._TOOL_ATTRS,
|
||||
self._tool_attr_budget,
|
||||
)
|
||||
)
|
||||
)
|
||||
return attrs
|
||||
|
||||
@classmethod
|
||||
|
|
|
|||
|
|
@ -12,7 +12,11 @@ Like ``GenAIMapper``, each span kind declares its schema as a flat
|
|||
from typing import Callable, Final
|
||||
|
||||
from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, SpanData
|
||||
from litellm.integrations.otel.mappers.utils import collect, drop_none
|
||||
from litellm.integrations.otel.mappers.utils import (
|
||||
MAX_TOOL_DEFINITION_ATTRS_PER_SPAN,
|
||||
collect,
|
||||
tool_definition_attrs,
|
||||
)
|
||||
from litellm.integrations.otel.model.payloads import (
|
||||
LLMCallSpanData,
|
||||
ServiceSpanData,
|
||||
|
|
@ -63,6 +67,9 @@ class LegacyMapper:
|
|||
_LEGACY_ERROR: lambda d: d.error.message if d.error is not None and d.error.message else None,
|
||||
}
|
||||
|
||||
def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None:
|
||||
self._tool_attr_budget = tool_attr_budget
|
||||
|
||||
def map(self, data: SpanData) -> AttributeMap:
|
||||
match data:
|
||||
case LLMCallSpanData():
|
||||
|
|
@ -72,16 +79,14 @@ class LegacyMapper:
|
|||
case _:
|
||||
return {}
|
||||
|
||||
@classmethod
|
||||
def _llm_call(cls, data: LLMCallSpanData) -> AttributeMap:
|
||||
attrs = collect(cls._LLM_CALL_ATTRS, data)
|
||||
def _llm_call(self, data: LLMCallSpanData) -> AttributeMap:
|
||||
attrs = collect(self._LLM_CALL_ATTRS, data)
|
||||
attrs.update(
|
||||
drop_none(
|
||||
{
|
||||
f"llm.request.functions.{idx}.{suffix}": extract(tool)
|
||||
for idx, tool in enumerate(data.tools)
|
||||
for suffix, extract in cls._TOOL_ATTRS.items()
|
||||
}
|
||||
tool_definition_attrs(
|
||||
lambda idx, suffix: f"llm.request.functions.{idx}.{suffix}",
|
||||
data.tools,
|
||||
self._TOOL_ATTRS,
|
||||
self._tool_attr_budget,
|
||||
)
|
||||
)
|
||||
return attrs
|
||||
|
|
|
|||
|
|
@ -13,9 +13,11 @@ from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, Span
|
|||
from litellm.integrations.otel.mappers.utils import (
|
||||
collect,
|
||||
drop_none,
|
||||
MAX_TOOL_DEFINITION_ATTRS_PER_SPAN,
|
||||
json_if,
|
||||
message_content,
|
||||
output_messages,
|
||||
tool_definition_attrs,
|
||||
)
|
||||
from litellm.integrations.otel.model.payloads import (
|
||||
LLMCallSpanData,
|
||||
|
|
@ -70,6 +72,9 @@ class OpenInferenceMapper:
|
|||
),
|
||||
}
|
||||
|
||||
def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None:
|
||||
self._tool_attr_budget = tool_attr_budget
|
||||
|
||||
def map(self, data: SpanData) -> AttributeMap:
|
||||
match data:
|
||||
case LLMCallSpanData():
|
||||
|
|
@ -77,14 +82,13 @@ class OpenInferenceMapper:
|
|||
case _:
|
||||
return {}
|
||||
|
||||
@classmethod
|
||||
def _llm_call(cls, data: LLMCallSpanData) -> AttributeMap:
|
||||
def _llm_call(self, data: LLMCallSpanData) -> AttributeMap:
|
||||
return {
|
||||
**collect(cls._LLM_CALL_ATTRS, data),
|
||||
**collect(cls._BLOB_ATTRS, data),
|
||||
**cls._messages("llm.input_messages", "input.value", data.messages_in),
|
||||
**cls._messages("llm.output_messages", "output.value", output_messages(data)),
|
||||
**cls._tools(data),
|
||||
**collect(self._LLM_CALL_ATTRS, data),
|
||||
**collect(self._BLOB_ATTRS, data),
|
||||
**self._messages("llm.input_messages", "input.value", data.messages_in),
|
||||
**self._messages("llm.output_messages", "output.value", output_messages(data)),
|
||||
**self._tools(data),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -108,12 +112,10 @@ class OpenInferenceMapper:
|
|||
attrs[value_key] = json.dumps([{"role": role, "content": content} for role, content in parsed])
|
||||
return attrs
|
||||
|
||||
@classmethod
|
||||
def _tools(cls, data: LLMCallSpanData) -> AttributeMap:
|
||||
return drop_none(
|
||||
{
|
||||
f"llm.tools.{idx}.{suffix}": extract(tool)
|
||||
for idx, tool in enumerate(data.tools)
|
||||
for suffix, extract in cls._TOOL_ATTRS.items()
|
||||
}
|
||||
def _tools(self, data: LLMCallSpanData) -> AttributeMap:
|
||||
return tool_definition_attrs(
|
||||
lambda idx, suffix: f"llm.tools.{idx}.{suffix}",
|
||||
data.tools,
|
||||
self._TOOL_ATTRS,
|
||||
self._tool_attr_budget,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -6,10 +6,34 @@ they live in one place.
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Callable, Mapping, Sequence
|
||||
from typing import Callable, Final, Mapping, Sequence
|
||||
|
||||
from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue
|
||||
from litellm.integrations.otel.model.payloads import LLMCallSpanData
|
||||
from litellm.integrations.otel.model.payloads import LLMCallSpanData, ToolDefinition
|
||||
|
||||
DEFAULT_SPAN_ATTRIBUTE_LIMIT: Final = 128
|
||||
"""The OTel SDK's default per-span attribute count limit."""
|
||||
|
||||
MAX_TOOL_DEFINITION_ATTRS_PER_SPAN: Final = DEFAULT_SPAN_ATTRIBUTE_LIMIT // 4
|
||||
"""Span-wide ceiling on attributes spent spelling out declared tool definitions.
|
||||
|
||||
Tool definitions are an unbounded attribute family: one entry per declared
|
||||
tool, per field, per active vocabulary. Agentic clients declare hundreds, which
|
||||
overruns the span attribute limit. That limit evicts oldest-first, so an
|
||||
uncapped family silently destroys the core ``gen_ai.*`` attributes written
|
||||
before it.
|
||||
|
||||
The ceiling is span-wide rather than per-mapper because several vocabularies
|
||||
can be active at once and each spells the same tools out under its own keys, so
|
||||
a per-mapper allowance multiplies by the number of vocabularies and reaches the
|
||||
limit again. Reserving a quarter of the span for tool detail leaves the rest to
|
||||
core telemetry no matter how many vocabularies are configured.
|
||||
"""
|
||||
|
||||
|
||||
def tool_attr_budget(vocabularies: int) -> int:
|
||||
"""Split the span-wide tool-definition ceiling across active vocabularies."""
|
||||
return MAX_TOOL_DEFINITION_ATTRS_PER_SPAN // max(vocabularies, 1)
|
||||
|
||||
|
||||
def drop_none(values: Mapping[str, AttrValue | None]) -> AttributeMap:
|
||||
|
|
@ -17,6 +41,29 @@ def drop_none(values: Mapping[str, AttrValue | None]) -> AttributeMap:
|
|||
return {k: v for k, v in values.items() if v is not None}
|
||||
|
||||
|
||||
def tool_definition_attrs(
|
||||
key_for: Callable[[int, str], str],
|
||||
tools: Sequence[ToolDefinition],
|
||||
extractors: Mapping[str, Callable[[ToolDefinition], AttrValue | None]],
|
||||
attr_budget: int,
|
||||
) -> AttributeMap:
|
||||
"""Per-index attributes for as many tools as ``attr_budget`` affords.
|
||||
|
||||
``key_for`` builds a vocabulary's key from the tool's index and the field
|
||||
name, so each mapper keeps its own naming while sharing the budget. One tool
|
||||
always keeps its detail, so the family stays legible even when many
|
||||
vocabularies split the ceiling.
|
||||
"""
|
||||
max_tools = max(attr_budget // max(len(extractors), 1), 1)
|
||||
return drop_none(
|
||||
{
|
||||
key_for(idx, suffix): extract(tool)
|
||||
for idx, tool in enumerate(tools[:max_tools])
|
||||
for suffix, extract in extractors.items()
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def collect(table: Mapping[str, Callable], source: object) -> AttributeMap:
|
||||
"""Apply an extractor table to ``source``, dropping ``None`` results."""
|
||||
return drop_none({key: extract(source) for key, extract in table.items()})
|
||||
|
|
|
|||
|
|
@ -233,6 +233,7 @@ class LiteLLM:
|
|||
# ``litellm_params.model``), distinct from the user-facing ``gen_ai.request.model``.
|
||||
PROVIDER_MODEL: Final = "litellm.provider.model"
|
||||
REQUEST_STREAMING: Final = "litellm.request.streaming"
|
||||
TOOLS_DECLARED: Final = "litellm.request.tools.declared"
|
||||
GUARDRAIL_NAME: Final = "litellm.guardrail.name"
|
||||
GUARDRAIL_MODE: Final = "litellm.guardrail.mode"
|
||||
GUARDRAIL_STATUS: Final = "litellm.guardrail.status"
|
||||
|
|
@ -257,13 +258,27 @@ class LiteLLM:
|
|||
|
||||
|
||||
class Metric:
|
||||
"""GenAI metric instrument names."""
|
||||
"""GenAI metric instrument names.
|
||||
|
||||
Every name here that a convention or a backend defines uses that name, so a
|
||||
consumer charting GenAI telemetry finds litellm's series where it looks for
|
||||
them. ``TOKEN_USAGE``, ``OPERATION_DURATION``, ``TIME_TO_FIRST_TOKEN`` and
|
||||
``TIME_PER_OUTPUT_TOKEN`` are semconv instruments, defined in the GenAI
|
||||
conventions; the ``gen_ai.client.response.*`` spellings litellm used for the
|
||||
latter two are not conventions at all, so nothing downstream could chart
|
||||
them. Cost has no semconv instrument, so it takes ``gen_ai.usage.cost``, the
|
||||
name backends already query for spend.
|
||||
|
||||
``RESPONSE_DURATION`` keeps its vendor spelling deliberately: the closest
|
||||
convention, ``gen_ai.server.request.duration``, would collide in meaning with
|
||||
``OPERATION_DURATION``, which litellm already emits for the whole operation.
|
||||
"""
|
||||
|
||||
TOKEN_USAGE: Final = "gen_ai.client.token.usage"
|
||||
OPERATION_DURATION: Final = "gen_ai.client.operation.duration"
|
||||
TOKEN_COST: Final = "gen_ai.client.token.cost"
|
||||
TIME_TO_FIRST_TOKEN: Final = "gen_ai.client.response.time_to_first_token"
|
||||
TIME_PER_OUTPUT_TOKEN: Final = "gen_ai.client.response.time_per_output_token"
|
||||
TOKEN_COST: Final = "gen_ai.usage.cost"
|
||||
TIME_TO_FIRST_TOKEN: Final = "gen_ai.server.time_to_first_token"
|
||||
TIME_PER_OUTPUT_TOKEN: Final = "gen_ai.server.time_per_output_token"
|
||||
RESPONSE_DURATION: Final = "gen_ai.client.response.duration"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
"""Shared, OpenTelemetry-free helpers for the otel integration.
|
||||
|
||||
Generic value coercion (for reading heterogeneous logging-payload dicts), time
|
||||
conversion, and header parsing — pulled out of the individual modules so they
|
||||
live in one place. Deliberately free of any ``opentelemetry`` import so the
|
||||
OTel-free sources of truth (payloads, semconv, spans, config) can use it too.
|
||||
Generic value coercion (for reading heterogeneous logging-payload dicts) and
|
||||
time conversion — pulled out of the individual modules so they live in one
|
||||
place. Deliberately free of any ``opentelemetry`` import so the OTel-free
|
||||
sources of truth (payloads, semconv, spans, config) can use it too. OTLP header
|
||||
parsing lives in :mod:`litellm.integrations.otel.plumbing.providers` instead,
|
||||
because it delegates to the OTel SDK's own W3C Baggage parser.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
|
|
@ -89,15 +91,3 @@ def to_seconds(value: datetime | float | int | str | None) -> float | None:
|
|||
except ValueError:
|
||||
continue
|
||||
return None
|
||||
|
||||
|
||||
def parse_headers(raw: str | None) -> dict[str, str]:
|
||||
"""Parse an OTLP ``"k=v,k=v"`` header string into a dict."""
|
||||
headers: dict[str, str] = {}
|
||||
if not raw:
|
||||
return headers
|
||||
for pair in raw.split(","):
|
||||
if "=" in pair:
|
||||
key, _, value = pair.partition("=")
|
||||
headers[key.strip()] = value.strip()
|
||||
return headers
|
||||
|
|
|
|||
|
|
@ -79,47 +79,67 @@ def reset_mcp_message_trace_carrier(token: "Token[Mapping[str, str] | None]") ->
|
|||
_mcp_message_trace_carrier.reset(token)
|
||||
|
||||
|
||||
# The transport span of the HTTP request carrying the CURRENT MCP message, as a
|
||||
# plain ``SpanContext`` so it can cross a task boundary.
|
||||
# The transport span of the HTTP request carrying the CURRENT MCP message.
|
||||
#
|
||||
# ``_request_root_span`` above cannot be used for MCP: a *stateful* streamable-HTTP
|
||||
# session runs every message on the single task spawned by that session's
|
||||
# ``initialize`` POST, so the ContextVar the ASGI request task writes at auth time
|
||||
# is frozen at ``initialize`` there and never sees the later ``tools/call`` POSTs.
|
||||
# Reading it from the message handler would parent every tool call in the session
|
||||
# to the first request's (already ended) server span. The gateway instead resolves
|
||||
# the current message's transport span on the request task and hands it over the
|
||||
# same way it hands over per-request auth, and the handler publishes it here for
|
||||
# the span emitter to pick up.
|
||||
_mcp_message_transport_span_context: "ContextVar[SpanContext | None]" = ContextVar(
|
||||
"litellm_otel_mcp_message_transport_span_context", default=None
|
||||
# Reading it from the message handler parents every tool call in the session to the
|
||||
# first request's server span and aims that call's ``error.*`` at it — a span that
|
||||
# ended long ago, so the SDK drops the write and the failure reaches no request at
|
||||
# all. The gateway instead resolves the current message's transport span on the
|
||||
# request task and hands it over the same way it hands over per-request auth, and
|
||||
# the handler publishes it here for the span emitter and the failure hook.
|
||||
_mcp_message_transport_span: "ContextVar[Span | None]" = ContextVar(
|
||||
"litellm_otel_mcp_message_transport_span", default=None
|
||||
)
|
||||
|
||||
|
||||
def set_mcp_message_transport_span_context(
|
||||
span_context: "SpanContext | None",
|
||||
) -> "Token[SpanContext | None]":
|
||||
def set_mcp_message_transport_span(span: object) -> "Token[Span | None]":
|
||||
"""Publish the transport span of the request carrying the current MCP message.
|
||||
|
||||
Also re-anchors the request root, so everything else the message emits or stamps
|
||||
— the identity attributes seeded onto the server span, a guardrail span, a
|
||||
proxy-level failure — lands on this request instead of on the one that opened
|
||||
the session. The MCP SDK dispatches each message on its own task, so the anchor
|
||||
is scoped to this message; the handler re-publishes it for the next one either
|
||||
way. Only a transport still open for writes is anchored: replacing the anchor
|
||||
with a request that already answered would just move the dropped writes from one
|
||||
finished span to another.
|
||||
|
||||
Takes ``object`` because the gateway reads it back out of the ASGI scope, whose
|
||||
values are untyped; anything that is not a usable span is stored as ``None``
|
||||
rather than trusted.
|
||||
|
||||
Returns the reset token; the caller must reset it once the message is handled
|
||||
so the transport never leaks to the next message on the same session task.
|
||||
"""
|
||||
return _mcp_message_transport_span_context.set(span_context)
|
||||
transport = span if isinstance(span, Span) and is_recordable_span(span) else None
|
||||
if transport is not None and transport.is_recording():
|
||||
set_request_root_span(transport)
|
||||
return _mcp_message_transport_span.set(transport)
|
||||
|
||||
|
||||
def reset_mcp_message_transport_span_context(token: "Token[SpanContext | None]") -> None:
|
||||
_mcp_message_transport_span_context.reset(token)
|
||||
def reset_mcp_message_transport_span(token: "Token[Span | None]") -> None:
|
||||
_mcp_message_transport_span.reset(token)
|
||||
|
||||
|
||||
def request_root_span_context() -> "SpanContext | None":
|
||||
"""The anchored request root span's context, safe to hand to another task.
|
||||
def mcp_message_transport_span() -> "Span | None":
|
||||
"""The published transport span, only while it is still open for writes.
|
||||
|
||||
A ``SpanContext`` is an immutable value, unlike the live ``Span``, so passing it
|
||||
across the MCP session-task boundary cannot keep a finished span alive or invite
|
||||
writes to it from the wrong request.
|
||||
Recording — not merely valid — is the bar here because this span is the target
|
||||
of ``error.*`` stamping from another task, and the publisher's validity check
|
||||
cannot speak for a span that has since ended. A finished span keeps a valid
|
||||
context forever, so it would otherwise be handed back for a write the SDK then
|
||||
refuses. The POST carrying a ``tools/call`` stays open until the result is
|
||||
written, so it is recording for the life of the call; a notification POST can
|
||||
answer first, and this returns ``None`` for it rather than writing into the void.
|
||||
"""
|
||||
span = request_root_span()
|
||||
return span.get_span_context() if span is not None else None
|
||||
span = _mcp_message_transport_span.get()
|
||||
if span is None or not span.is_recording():
|
||||
return None
|
||||
return span
|
||||
|
||||
|
||||
def _mcp_transport_span_context() -> "SpanContext | None":
|
||||
|
|
@ -127,12 +147,16 @@ def _mcp_transport_span_context() -> "SpanContext | None":
|
|||
|
||||
Prefers the transport the gateway published for this specific message; falls
|
||||
back to the ambient request anchor for paths that emit an MCP span on the
|
||||
request task itself (the REST MCP endpoints, the SDK).
|
||||
request task itself (the REST MCP endpoints, the SDK). Parenting and linking
|
||||
only need the immutable context, and unlike ``mcp_message_transport_span`` they
|
||||
stay correct against a transport that has already finished, so this does not
|
||||
require the span to still be recording.
|
||||
"""
|
||||
published = _mcp_message_transport_span_context.get()
|
||||
if published is not None and published.is_valid:
|
||||
return published
|
||||
return request_root_span_context()
|
||||
published = _mcp_message_transport_span.get()
|
||||
if published is not None:
|
||||
return published.get_span_context()
|
||||
span = request_root_span()
|
||||
return span.get_span_context() if span is not None else None
|
||||
|
||||
|
||||
def set_request_baggage(values: Mapping[str, str], context: Context | None = None) -> Context:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
"""GenAI client metrics: the six ``gen_ai.client.*`` histograms plus the
|
||||
recorder that builds attributes, applies the shared cardinality filter, and
|
||||
records a request's metrics in the success path.
|
||||
records a request's metrics on both the success and the failure path.
|
||||
|
||||
The instrument names/units/descriptions and the recording + timing math mirror
|
||||
the v1 :mod:`litellm.integrations.opentelemetry` integration so both engines emit
|
||||
|
|
@ -10,11 +10,12 @@ identical metrics. The attribute cardinality filter is reused from v1 by import
|
|||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any, FrozenSet, Mapping, Optional
|
||||
from typing import Any, Final, FrozenSet, Mapping, Optional, TypeAlias
|
||||
|
||||
from opentelemetry.metrics import Histogram, Meter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.opentelemetry import (
|
||||
METRIC_METADATA_KEYS,
|
||||
TOKEN_TYPE_ATTRIBUTE,
|
||||
|
|
@ -22,7 +23,7 @@ from litellm.integrations.opentelemetry import (
|
|||
_resolve_metric_attribute_filter,
|
||||
)
|
||||
from litellm.integrations.otel.model.metadata import time_to_first_chunk_seconds
|
||||
from litellm.integrations.otel.model.semconv import Metric, resolve_operation
|
||||
from litellm.integrations.otel.model.semconv import Error, Metric, resolve_operation
|
||||
from litellm.integrations.otel.model.utils import to_seconds
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
|
|
@ -72,8 +73,82 @@ def create_genai_metrics(meter: Meter) -> GenAIMetrics:
|
|||
)
|
||||
|
||||
|
||||
# A metric datapoint's attributes. Values are the strings the recorder builds, except
|
||||
# the request model, which is whatever the caller passed and may be absent.
|
||||
MetricAttributes: TypeAlias = Mapping[str, "str | None"]
|
||||
|
||||
ERROR_TYPE_FALLBACK: Final = "_OTHER"
|
||||
|
||||
# Every attribute a metric datapoint may carry, on either path. A label value that
|
||||
# is unique per request is a new time series that will never be written to again, so
|
||||
# this set is what keeps the series count bounded by the deployment's own
|
||||
# key/team/user/deployment count rather than by its traffic. Each entry is a fixed
|
||||
# enum or an operator-provisioned identifier.
|
||||
#
|
||||
# Deliberately excluded is everything the *client* supplies or that moves per
|
||||
# request: ``metadata.requester_metadata`` and ``metadata.spend_logs_metadata`` (both
|
||||
# free-form from the request body), ``metadata.user_api_key_end_user_id`` (the body's
|
||||
# ``user`` field), and ``metadata.requester_ip_address``. Those stay on the span,
|
||||
# where cardinality is free and where they already are.
|
||||
# ``metadata.user_api_key_user_email`` is left out too: it is bounded, but it is PII
|
||||
# duplicating the user id already here.
|
||||
#
|
||||
# This is a CEILING, applied before the operator's own include/exclude filter, so an
|
||||
# operator can narrow it but never widen it back to an unbounded attribute.
|
||||
METRIC_ATTRIBUTE_CEILING: Final[frozenset[str]] = frozenset(
|
||||
(
|
||||
"gen_ai.operation.name",
|
||||
"gen_ai.system",
|
||||
"gen_ai.request.model",
|
||||
"gen_ai.framework",
|
||||
"metadata.user_api_key_hash",
|
||||
"metadata.user_api_key_alias",
|
||||
"metadata.user_api_key_team_id",
|
||||
"metadata.user_api_key_team_alias",
|
||||
"metadata.user_api_key_org_id",
|
||||
"metadata.user_api_key_user_id",
|
||||
"hidden_params",
|
||||
)
|
||||
)
|
||||
|
||||
# The only ``hidden_params`` field that becomes part of the ``hidden_params`` label.
|
||||
# The object as a whole is per-request by construction -- ``response_cost``,
|
||||
# ``litellm_overhead_time_ms``, ``cache_key``, ``usage_object`` and the provider's
|
||||
# ``additional_headers`` rate-limit counters all move on every call -- so dumping it
|
||||
# whole made one series per request out of every instrument.
|
||||
#
|
||||
# ``model_id`` is the router's own deployment id, so it is bounded by the deployment
|
||||
# list and is what a per-deployment panel joins on. ``api_base`` is deliberately NOT
|
||||
# here even though it names the same thing: it is a documented per-call parameter, so
|
||||
# in SDK use it is chosen by the caller rather than provisioned by the operator, and a
|
||||
# caller varying it would put the per-request cardinality straight back.
|
||||
BOUNDED_HIDDEN_PARAM_KEYS: Final[tuple[str, ...]] = ("model_id",)
|
||||
|
||||
|
||||
def resolve_error_type(kwargs: Mapping[str, Any]) -> str:
|
||||
"""The ``error.type`` value for a failed request.
|
||||
|
||||
Bounded by construction: the mapped provider exception's class name (the same
|
||||
``error_information.error_class`` the failure span stamps), else the provider
|
||||
status code, else the raw exception's class name, else ``_OTHER`` — the value
|
||||
the convention reserves for a failure the instrumentation cannot classify. The
|
||||
exception *message* is unbounded and never becomes a label; it stays on the
|
||||
span and its exception event, where high cardinality is free.
|
||||
"""
|
||||
std_log = kwargs.get("standard_logging_object")
|
||||
info = getattr(std_log, "error_information", None) or (std_log or {}).get("error_information") or {}
|
||||
error_class = info.get("error_class") or info.get("error_code")
|
||||
if error_class:
|
||||
return str(error_class)
|
||||
exception = kwargs.get("exception")
|
||||
if exception is not None:
|
||||
return type(exception).__name__
|
||||
return ERROR_TYPE_FALLBACK
|
||||
|
||||
|
||||
class GenAIMetricRecorder:
|
||||
"""Records the six GenAI histograms for one successful LLM call.
|
||||
"""Records the six GenAI histograms for one successful LLM call, and the
|
||||
duration histogram alone for one failed LLM call (see :meth:`record_failure`).
|
||||
|
||||
The cardinality filter is resolved lazily on the first record: the proxy
|
||||
populates ``callback_settings.otel.attributes`` after the logger is built, so
|
||||
|
|
@ -96,7 +171,7 @@ class GenAIMetricRecorder:
|
|||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> None:
|
||||
common_attrs = self._filter_attributes(self._common_attributes(kwargs))
|
||||
common_attrs = self._filter_attributes(self._bounded_attributes(kwargs))
|
||||
duration_s = (end_time - start_time).total_seconds()
|
||||
|
||||
self._metrics.operation_duration.record(duration_s, attributes=common_attrs)
|
||||
|
|
@ -110,6 +185,38 @@ class GenAIMetricRecorder:
|
|||
self._record_time_per_output_token(kwargs, response_obj, end_time, duration_s, common_attrs)
|
||||
self._record_response_duration(kwargs, end_time, common_attrs)
|
||||
|
||||
def record_failure(
|
||||
self,
|
||||
kwargs: Mapping[str, Any],
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> None:
|
||||
"""Record the one metric a failed request can honestly report: the
|
||||
operation's duration, tagged with ``error.type``.
|
||||
|
||||
The other five instruments all describe a completed generation and have
|
||||
nothing to measure here. litellm hands the failure callback no
|
||||
``response_obj`` at all, so there is no usage to split into input/output
|
||||
tokens and no completion-token count to divide generation time by; it also
|
||||
zeroes ``response_cost`` on failure. Recording them anyway would put a
|
||||
fabricated zero into series that dashboards average.
|
||||
|
||||
The attribute set is :data:`METRIC_ATTRIBUTE_CEILING`, the same cap the
|
||||
success path uses. A failure needs no provider spend, so a caller who can put
|
||||
a unique value into a client-supplied attribute could mint one histogram
|
||||
series per request for free; the cap is what makes that impossible on either
|
||||
path.
|
||||
|
||||
``error.type`` is stamped after both filters, exactly like
|
||||
``gen_ai.token.type``, so an operator's include/exclude list cannot strip
|
||||
the discriminator and silently merge failures back into the success series.
|
||||
"""
|
||||
attributes = {
|
||||
**self._filter_attributes(self._bounded_attributes(kwargs)),
|
||||
Error.TYPE: resolve_error_type(kwargs),
|
||||
}
|
||||
self._metrics.operation_duration.record((end_time - start_time).total_seconds(), attributes=attributes)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Attribute building + cardinality filter
|
||||
# ------------------------------------------------------------------ #
|
||||
|
|
@ -136,11 +243,25 @@ class GenAIMetricRecorder:
|
|||
common_attrs[f"metadata.{key}"] = str(value)
|
||||
|
||||
hidden_params = getattr(std_log, "hidden_params", None) or (std_log or {}).get("hidden_params", {})
|
||||
if hidden_params:
|
||||
common_attrs["hidden_params"] = safe_dumps(hidden_params)
|
||||
bounded_hidden_params = {
|
||||
key: hidden_params[key]
|
||||
for key in BOUNDED_HIDDEN_PARAM_KEYS
|
||||
if isinstance(hidden_params, Mapping) and hidden_params.get(key) is not None
|
||||
}
|
||||
if bounded_hidden_params:
|
||||
common_attrs["hidden_params"] = safe_dumps(bounded_hidden_params)
|
||||
|
||||
return common_attrs
|
||||
|
||||
def _bounded_attributes(self, kwargs: Mapping[str, Any]) -> MetricAttributes:
|
||||
"""The datapoint attributes, capped at :data:`METRIC_ATTRIBUTE_CEILING`.
|
||||
|
||||
The cap runs BEFORE the operator's include/exclude filter so the filter can
|
||||
only narrow it. An operator who names an excluded attribute in an include
|
||||
list gets nothing for it rather than reintroducing an unbounded label.
|
||||
"""
|
||||
return {k: v for k, v in self._common_attributes(kwargs).items() if k in METRIC_ATTRIBUTE_CEILING}
|
||||
|
||||
def _ensure_filter(self) -> None:
|
||||
if self._filter_resolved:
|
||||
return
|
||||
|
|
@ -157,8 +278,29 @@ class GenAIMetricRecorder:
|
|||
# without reconstructing the recorder.
|
||||
self._include, self._exclude = _resolve_metric_attribute_filter(attributes)
|
||||
self._filter_resolved = True
|
||||
self._warn_about_metric_ineligible_names()
|
||||
|
||||
def _filter_attributes(self, attrs: dict) -> dict:
|
||||
def _warn_about_metric_ineligible_names(self) -> None:
|
||||
"""Say so when the operator's filter names an attribute the ceiling removes.
|
||||
|
||||
The shared validator accepts every span attribute name, so a name that is
|
||||
legal on a span but metric-ineligible would otherwise be a silent no-op: an
|
||||
``include_list`` naming it emits nothing for it and an ``exclude_list`` naming
|
||||
it looks like it worked. Logged once, when the filter resolves, rather than
|
||||
per request.
|
||||
"""
|
||||
named = (self._include or frozenset()) | (self._exclude or frozenset())
|
||||
ineligible = sorted(named - METRIC_ATTRIBUTE_CEILING - {TOKEN_TYPE_ATTRIBUTE})
|
||||
if ineligible:
|
||||
verbose_logger.warning(
|
||||
"OTel metrics: %s cannot be a metric attribute and is being ignored; it varies "
|
||||
"per request or is client-supplied, so it would make one time series per request. "
|
||||
"It is still on the span. Metric attributes are limited to: %s",
|
||||
", ".join(ineligible),
|
||||
", ".join(sorted(METRIC_ATTRIBUTE_CEILING)),
|
||||
)
|
||||
|
||||
def _filter_attributes(self, attrs: MetricAttributes) -> MetricAttributes:
|
||||
self._ensure_filter()
|
||||
if self._include is not None:
|
||||
return {k: v for k, v in attrs.items() if k in self._include}
|
||||
|
|
|
|||
|
|
@ -29,15 +29,13 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
|
|||
InMemorySpanExporter,
|
||||
)
|
||||
from opentelemetry.trace import Span, SpanKind, Tracer
|
||||
from opentelemetry.util.re import parse_env_headers
|
||||
|
||||
from litellm._version import version as litellm_version
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.semconv import LiteLLM
|
||||
from litellm.integrations.otel.model.spans import LiteLLMSpanKind
|
||||
|
||||
# Re-exported so ``providers.parse_headers`` remains a stable entry point.
|
||||
from litellm.integrations.otel.model.utils import parse_headers as parse_headers
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.metrics import Meter
|
||||
from opentelemetry.sdk.metrics.export import MetricReader
|
||||
|
|
@ -119,6 +117,23 @@ def _otlp_traces_endpoint(endpoint: str | None) -> str | None:
|
|||
return endpoint + "/v1/traces"
|
||||
|
||||
|
||||
def parse_headers(raw: str | None) -> dict[str, str]:
|
||||
"""Parse an OTLP ``"k=v,k=v"`` header string into a dict.
|
||||
|
||||
``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded per the OTLP spec, so
|
||||
values are percent-decoded: a vendor that documents
|
||||
``Authorization=Basic%20<token>`` (Grafana Cloud does, because a bare space
|
||||
is not representable there) has to reach the exporter as ``Basic <token>``,
|
||||
not with a literal ``%20`` that the backend rejects as malformed. The SDK's
|
||||
own parser is used so litellm decodes exactly what the OTLP exporters do
|
||||
when they read the env var themselves; ``liberal`` keeps values that are not
|
||||
percent-encoded (``Authorization=Bearer <token>``) working unchanged.
|
||||
"""
|
||||
if not raw:
|
||||
return {}
|
||||
return dict(parse_env_headers(raw, liberal=True))
|
||||
|
||||
|
||||
def _exporter_from_spec(spec: ExporterSpec) -> SpanExporter:
|
||||
kind = (spec.kind or "console").lower()
|
||||
factory = _EXPORTER_FACTORIES.get(kind)
|
||||
|
|
@ -191,6 +206,13 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
|
|||
``console`` (and any unrecognized kind) exports to the console; ``otlp_http``
|
||||
and ``otlp_grpc`` export over OTLP with the configured endpoint/headers. The
|
||||
reader exports on a 5s period, matching v1.
|
||||
|
||||
Histograms keep the SDK's default cumulative temporality. Prometheus-backed
|
||||
OTLP receivers (Grafana Cloud / Mimir, and the Prometheus OTLP endpoint)
|
||||
reject delta histograms outright with ``invalid temporality and type
|
||||
combination``, which drops the whole metric batch, while backends that
|
||||
prefer delta still accept cumulative. The enterprise billing exporter
|
||||
already relies on the same default.
|
||||
"""
|
||||
from opentelemetry.sdk.metrics.export import (
|
||||
ConsoleMetricExporter,
|
||||
|
|
@ -202,18 +224,12 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
|
|||
from opentelemetry.exporter.otlp.proto.http.metric_exporter import (
|
||||
OTLPMetricExporter as HTTPMetricExporter,
|
||||
)
|
||||
from opentelemetry.sdk.metrics import Histogram
|
||||
from opentelemetry.sdk.metrics.export import AggregationTemporality
|
||||
|
||||
exporter: Any = HTTPMetricExporter(
|
||||
endpoint=_otlp_metrics_endpoint(config.endpoint),
|
||||
headers=parse_headers(config.headers),
|
||||
preferred_temporality={Histogram: AggregationTemporality.DELTA},
|
||||
)
|
||||
elif kind in ("otlp_grpc", "grpc"):
|
||||
from opentelemetry.sdk.metrics import Histogram
|
||||
from opentelemetry.sdk.metrics.export import AggregationTemporality
|
||||
|
||||
try:
|
||||
from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import (
|
||||
OTLPMetricExporter as GRPCMetricExporter,
|
||||
|
|
@ -227,7 +243,6 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
|
|||
exporter = GRPCMetricExporter(
|
||||
endpoint=config.endpoint,
|
||||
headers=parse_headers(config.headers),
|
||||
preferred_temporality={Histogram: AggregationTemporality.DELTA},
|
||||
)
|
||||
else:
|
||||
exporter = ConsoleMetricExporter()
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from typing import (
|
|||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Sequence,
|
||||
Tuple,
|
||||
|
|
@ -68,6 +69,15 @@ else:
|
|||
|
||||
_DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT = 5.0
|
||||
|
||||
# Tiers a caller may name in a request, across the providers that accept the
|
||||
# parameter: OpenAI ("auto", "default", "flex", "priority", "scale"), Bedrock and
|
||||
# Groq (subsets of those), Anthropic ("auto", "standard_only") and Vertex, which
|
||||
# maps "default" to "standard". Used to bound the caller-controlled fallback in
|
||||
# ``get_service_tier_from_standard_logging_payload``.
|
||||
KNOWN_REQUEST_SERVICE_TIERS = frozenset(
|
||||
{"auto", "batch", "default", "flex", "priority", "scale", "standard", "standard_only"}
|
||||
)
|
||||
|
||||
|
||||
def _get_budget_metrics_per_request_timeout() -> float:
|
||||
raw = os.getenv("PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT")
|
||||
|
|
@ -1244,6 +1254,7 @@ class PrometheusLogger(CustomLogger):
|
|||
client_ip=standard_logging_payload["metadata"].get("requester_ip_address"),
|
||||
user_agent=standard_logging_payload["metadata"].get("user_agent"),
|
||||
stream=(str(standard_logging_payload.get("stream")) if litellm.prometheus_emit_stream_label else None),
|
||||
service_tier=get_service_tier_from_standard_logging_payload(standard_logging_payload),
|
||||
)
|
||||
|
||||
if user_api_key is not None and isinstance(user_api_key, str) and user_api_key.startswith("sk-"):
|
||||
|
|
@ -1449,6 +1460,8 @@ class PrometheusLogger(CustomLogger):
|
|||
prompt_details = usage_object.get("prompt_tokens_details") or {}
|
||||
completion_details = usage_object.get("completion_tokens_details") or {}
|
||||
|
||||
cache_creation_detail_tokens = PrometheusLogger._resolve_cache_write_tokens(prompt_details)
|
||||
|
||||
detail_metrics: List[Tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]] = [
|
||||
(
|
||||
self.litellm_input_cached_tokens_metric,
|
||||
|
|
@ -1458,7 +1471,7 @@ class PrometheusLogger(CustomLogger):
|
|||
(
|
||||
self.litellm_input_cache_creation_tokens_metric,
|
||||
"litellm_input_cache_creation_tokens_metric",
|
||||
(prompt_details.get("cache_creation_tokens") if isinstance(prompt_details, dict) else None),
|
||||
cache_creation_detail_tokens,
|
||||
),
|
||||
(
|
||||
self.litellm_input_audio_tokens_metric,
|
||||
|
|
@ -1597,27 +1610,12 @@ class PrometheusLogger(CustomLogger):
|
|||
)
|
||||
|
||||
# Provider prompt caching metrics are independent of LiteLLM cache_hit.
|
||||
provider_cache_read_tokens = 0
|
||||
provider_cache_creation_tokens = 0
|
||||
usage_obj = (standard_logging_payload.get("metadata", {}) or {}).get("usage_object")
|
||||
if isinstance(usage_obj, dict):
|
||||
# Prefer explicit provider cache fields when available.
|
||||
_read = usage_obj.get("cache_read_input_tokens")
|
||||
_write = usage_obj.get("cache_creation_input_tokens")
|
||||
|
||||
if isinstance(_read, int):
|
||||
provider_cache_read_tokens = _read
|
||||
if isinstance(_write, int):
|
||||
provider_cache_creation_tokens = _write
|
||||
|
||||
# Fallback to prompt_tokens_details.cached_tokens (common normalization point).
|
||||
# Only fallback when the explicit field is genuinely absent (None).
|
||||
if _read is None:
|
||||
prompt_details = usage_obj.get("prompt_tokens_details")
|
||||
if isinstance(prompt_details, dict):
|
||||
cached_tokens = prompt_details.get("cached_tokens")
|
||||
if isinstance(cached_tokens, int):
|
||||
provider_cache_read_tokens = cached_tokens
|
||||
(
|
||||
provider_cache_read_tokens,
|
||||
provider_cache_creation_tokens,
|
||||
) = PrometheusLogger._resolve_provider_cache_tokens(usage_obj)
|
||||
|
||||
if provider_cache_read_tokens > 0:
|
||||
PrometheusLogger._inc_labeled_counter(
|
||||
|
|
@ -1639,6 +1637,40 @@ class PrometheusLogger(CustomLogger):
|
|||
amount=float(provider_cache_creation_tokens),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_provider_cache_tokens(usage_obj: Mapping[str, object]) -> tuple[int, int]:
|
||||
# Prefer explicit provider cache fields when available.
|
||||
_read = usage_obj.get("cache_read_input_tokens")
|
||||
_write = usage_obj.get("cache_creation_input_tokens")
|
||||
|
||||
provider_cache_read_tokens = _read if isinstance(_read, int) else 0
|
||||
provider_cache_creation_tokens = _write if isinstance(_write, int) else 0
|
||||
|
||||
# Fallback to prompt_tokens_details (common normalization point).
|
||||
# Only fallback when the explicit field is genuinely absent (None).
|
||||
prompt_details = usage_obj.get("prompt_tokens_details")
|
||||
if _read is None and isinstance(prompt_details, dict):
|
||||
cached_tokens = prompt_details.get("cached_tokens")
|
||||
if isinstance(cached_tokens, int):
|
||||
provider_cache_read_tokens = cached_tokens
|
||||
|
||||
if _write is None:
|
||||
write_tokens = PrometheusLogger._resolve_cache_write_tokens(prompt_details)
|
||||
if write_tokens is not None:
|
||||
provider_cache_creation_tokens = write_tokens
|
||||
|
||||
return provider_cache_read_tokens, provider_cache_creation_tokens
|
||||
|
||||
@staticmethod
|
||||
def _resolve_cache_write_tokens(prompt_details: object) -> int | None:
|
||||
if not isinstance(prompt_details, dict):
|
||||
return None
|
||||
for key in ("cache_write_tokens", "cache_creation_tokens"):
|
||||
value = prompt_details.get(key)
|
||||
if isinstance(value, int) and not isinstance(value, bool):
|
||||
return value
|
||||
return None
|
||||
|
||||
def _increment_mcp_tool_call_metrics(
|
||||
self,
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
|
|
@ -4076,6 +4108,44 @@ def get_custom_labels_from_metadata(metadata: dict) -> Dict[str, str]:
|
|||
return result
|
||||
|
||||
|
||||
def get_service_tier_from_standard_logging_payload(
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
) -> str | None:
|
||||
"""
|
||||
Resolve the service tier a request ran on, for the ``service_tier`` label.
|
||||
|
||||
The tier the provider actually served wins over the tier the caller asked for,
|
||||
so latency and spend stay segmentable when the request said ``auto`` and the
|
||||
provider picked the concrete tier. Providers report the served tier either at
|
||||
the top level of the response (OpenAI, Bedrock, Groq) or on the usage object
|
||||
(Anthropic).
|
||||
|
||||
Streaming responses carry no served tier, so the requested tier is the
|
||||
fallback. That value is caller-controlled and survives param mapping even
|
||||
where the provider then ignores it (Bedrock and Groq accept the request and
|
||||
drop an unrecognized tier), so it is only labelled when it names a known
|
||||
tier; otherwise one caller could mint a Prometheus series per string. Values
|
||||
the provider itself reports are not caller-controlled and stay unrestricted,
|
||||
so a tier a provider adds later is still labelled correctly.
|
||||
"""
|
||||
response = standard_logging_payload.get("response")
|
||||
usage_object = standard_logging_payload.get("metadata", {}).get("usage_object")
|
||||
|
||||
served_candidates: tuple[object, ...] = (
|
||||
response.get("service_tier") if isinstance(response, dict) else None,
|
||||
usage_object.get("service_tier") if isinstance(usage_object, dict) else None,
|
||||
)
|
||||
served_tier = next((tier for tier in served_candidates if isinstance(tier, str) and tier), None)
|
||||
if served_tier is not None:
|
||||
return served_tier
|
||||
|
||||
model_parameters = standard_logging_payload.get("model_parameters")
|
||||
requested_tier = model_parameters.get("service_tier") if isinstance(model_parameters, dict) else None
|
||||
if isinstance(requested_tier, str) and requested_tier in KNOWN_REQUEST_SERVICE_TIERS:
|
||||
return requested_tier
|
||||
return None
|
||||
|
||||
|
||||
def _get_combined_custom_metadata_from_standard_logging_payload(
|
||||
standard_logging_payload: Optional[dict],
|
||||
) -> Dict[str, Any]:
|
||||
|
|
|
|||
|
|
@ -1878,7 +1878,14 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
elif standard_logging_object is not None:
|
||||
self.model_call_details["standard_logging_object"] = standard_logging_object
|
||||
else:
|
||||
self.model_call_details["response_cost"] = None
|
||||
# Streaming reaches here before its cost is known, so the cost
|
||||
# is seeded to None, but only when nothing has already
|
||||
# established one. A stream that assembles into a response
|
||||
# object recomputes the cost right after this; a pass-through
|
||||
# stream cannot (its body is opaque) and carries the cost its
|
||||
# upstream reported in the response headers, which an
|
||||
# unconditional reset would discard.
|
||||
self.model_call_details.setdefault("response_cost", None)
|
||||
|
||||
result = self._transform_usage_objects(result=result)
|
||||
|
||||
|
|
@ -4673,6 +4680,7 @@ class StandardLoggingPayloadSetup:
|
|||
applied_guardrails=applied_guardrails,
|
||||
mcp_tool_call_metadata=mcp_tool_call_metadata,
|
||||
vector_store_request_metadata=vector_store_request_metadata,
|
||||
routing_decision=None,
|
||||
usage_object=usage_object,
|
||||
requester_custom_headers=None,
|
||||
cold_storage_object_key=None,
|
||||
|
|
@ -5512,6 +5520,7 @@ def get_standard_logging_metadata(
|
|||
applied_guardrails=None,
|
||||
mcp_tool_call_metadata=None,
|
||||
vector_store_request_metadata=None,
|
||||
routing_decision=None,
|
||||
usage_object=None,
|
||||
requester_custom_headers=None,
|
||||
user_api_key_request_route=None,
|
||||
|
|
@ -5547,18 +5556,6 @@ def scrub_sensitive_keys_in_metadata(litellm_params: Optional[dict]):
|
|||
litellm_params["_langfuse_masking_function"] = masking_fn
|
||||
litellm_params["metadata"] = metadata
|
||||
|
||||
## check user_api_key_metadata for sensitive logging keys
|
||||
cleaned_user_api_key_metadata = {}
|
||||
if "user_api_key_metadata" in metadata and isinstance(metadata["user_api_key_metadata"], dict):
|
||||
for k, v in metadata["user_api_key_metadata"].items():
|
||||
if k == "logging": # prevent logging user logging keys
|
||||
cleaned_user_api_key_metadata[k] = "scrubbed_by_litellm_for_sensitive_keys"
|
||||
else:
|
||||
cleaned_user_api_key_metadata[k] = v
|
||||
|
||||
metadata["user_api_key_metadata"] = cleaned_user_api_key_metadata
|
||||
litellm_params["metadata"] = metadata
|
||||
|
||||
return litellm_params
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import mimetypes
|
|||
import re
|
||||
import xml.etree.ElementTree as ET
|
||||
from enum import Enum
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Dict, List, Optional, Set, Tuple, TypedDict, Union, cast, overload
|
||||
|
||||
from jinja2.sandbox import ImmutableSandboxedEnvironment
|
||||
|
|
@ -1266,7 +1267,7 @@ def _get_dummy_thought_signature() -> str:
|
|||
def convert_to_gemini_tool_call_invoke(
|
||||
message: ChatCompletionAssistantMessage,
|
||||
model: Optional[str] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
forward_function_call_id: bool = False,
|
||||
) -> List[VertexPartType]:
|
||||
"""
|
||||
OpenAI tool invokes:
|
||||
|
|
@ -1316,16 +1317,12 @@ def convert_to_gemini_tool_call_invoke(
|
|||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
forward_tool_call_id = bool(
|
||||
model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider)
|
||||
)
|
||||
|
||||
if tool_calls is not None:
|
||||
for idx, tool in enumerate(tool_calls):
|
||||
if "function" in tool:
|
||||
gemini_function_call: Optional[VertexFunctionCall] = _gemini_tool_call_invoke_helper(
|
||||
function_call_params=tool["function"],
|
||||
tool_call_id=(tool.get("id") if forward_tool_call_id else None),
|
||||
tool_call_id=(tool.get("id") if forward_function_call_id else None),
|
||||
)
|
||||
if gemini_function_call is not None:
|
||||
part_dict: VertexPartType = {"function_call": gemini_function_call}
|
||||
|
|
@ -1377,8 +1374,7 @@ def convert_to_gemini_tool_call_invoke(
|
|||
def convert_to_gemini_tool_call_result(
|
||||
message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage],
|
||||
last_message_with_tool_calls: Optional[dict],
|
||||
model: Optional[str] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
forward_function_call_id: bool = False,
|
||||
) -> Union[VertexPartType, List[VertexPartType]]:
|
||||
"""
|
||||
OpenAI message with a tool result looks like:
|
||||
|
|
@ -1500,14 +1496,8 @@ def convert_to_gemini_tool_call_result(
|
|||
name = tool.get("function", {}).get("name", "")
|
||||
|
||||
# Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix).
|
||||
# Only Google AI Studio Gemini 3+ accepts `id` on function_response parts.
|
||||
# Vertex AI and older Gemini models reject the field with HTTP 400.
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
gemini_call_id: Optional[str] = None
|
||||
if model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider):
|
||||
if forward_function_call_id:
|
||||
raw_tool_call_id = message.get("tool_call_id")
|
||||
if raw_tool_call_id and isinstance(raw_tool_call_id, str):
|
||||
stripped_id = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0]
|
||||
|
|
@ -5350,7 +5340,9 @@ def prompt_factory(
|
|||
def get_attribute_or_key(tool_or_function, attribute, default=None):
|
||||
if hasattr(tool_or_function, attribute):
|
||||
return getattr(tool_or_function, attribute)
|
||||
return tool_or_function.get(attribute, default)
|
||||
if isinstance(tool_or_function, Mapping):
|
||||
return tool_or_function.get(attribute, default)
|
||||
return default
|
||||
|
||||
|
||||
class NormalizedToolCall(TypedDict):
|
||||
|
|
@ -5379,14 +5371,18 @@ def _parse_tool_call_arguments(raw: Any, tool_name: Optional[str], context: str)
|
|||
return parsed if isinstance(parsed, dict) else {}
|
||||
|
||||
|
||||
def _tool_calls_from_chat_completion_response(response: Any) -> list[NormalizedToolCall]:
|
||||
def _tool_calls_from_chat_completion_response(
|
||||
response: Any, include_all_choices: bool = False
|
||||
) -> list[NormalizedToolCall]:
|
||||
choices = get_attribute_or_key(response, "choices", None)
|
||||
if not (isinstance(choices, list) and choices):
|
||||
return []
|
||||
message = get_attribute_or_key(choices[0], "message", None)
|
||||
tool_calls = get_attribute_or_key(message, "tool_calls", None) if message else None
|
||||
if not isinstance(tool_calls, list):
|
||||
return []
|
||||
tool_calls: list[Any] = []
|
||||
for choice in choices if include_all_choices else choices[:1]:
|
||||
message = get_attribute_or_key(choice, "message", None)
|
||||
choice_tool_calls = get_attribute_or_key(message, "tool_calls", None) if message else None
|
||||
if isinstance(choice_tool_calls, list):
|
||||
tool_calls.extend(choice_tool_calls)
|
||||
result: list[NormalizedToolCall] = []
|
||||
for tc in tool_calls:
|
||||
fn = get_attribute_or_key(tc, "function", None)
|
||||
|
|
@ -5449,7 +5445,7 @@ def _tool_calls_from_anthropic_messages_response(response: Any) -> list[Normaliz
|
|||
return result
|
||||
|
||||
|
||||
def get_tool_calls_from_response(response: Any) -> list[NormalizedToolCall]:
|
||||
def get_tool_calls_from_response(response: Any, include_all_choices: bool = False) -> list[NormalizedToolCall]:
|
||||
"""
|
||||
Extract tool/function calls from a response object into a normalized
|
||||
``{"id", "name", "arguments"}`` shape, regardless of which API surface
|
||||
|
|
@ -5457,11 +5453,20 @@ def get_tool_calls_from_response(response: Any) -> list[NormalizedToolCall]:
|
|||
the Responses API (``output`` items of type ``function_call``), or the
|
||||
Anthropic Messages API (``content`` blocks of type ``tool_use``).
|
||||
|
||||
``include_all_choices`` decides the chat-completions scope: the default
|
||||
reads only ``choices[0]``, which is what consumers that act on THE reply
|
||||
(e.g. guardrails rebuilding the primary assistant message) want; usage
|
||||
accounting passes True because every choice of an ``n>1`` request costs
|
||||
money and its tool calls really ran. The other surfaces have a single
|
||||
output, so the flag has no effect on them.
|
||||
|
||||
Callers that only care about a specific tool should filter the result by
|
||||
``name`` themselves -- this returns every tool call found.
|
||||
"""
|
||||
chat_tool_calls = _tool_calls_from_chat_completion_response(response, include_all_choices=include_all_choices)
|
||||
if chat_tool_calls:
|
||||
return chat_tool_calls
|
||||
for extractor in (
|
||||
_tool_calls_from_chat_completion_response,
|
||||
_tool_calls_from_responses_api_response,
|
||||
_tool_calls_from_anthropic_messages_response,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -393,24 +393,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
if compaction_event is not None:
|
||||
return compaction_event
|
||||
|
||||
if self.sent_content_block_start is False:
|
||||
self.sent_content_block_start = True
|
||||
self.sent_content_block_finish = False
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
}
|
||||
)
|
||||
return self.chunk_queue.popleft()
|
||||
|
||||
for chunk in self.completion_stream:
|
||||
if chunk == "None" or chunk is None:
|
||||
raise Exception
|
||||
|
||||
should_start_new_block = self._should_start_new_content_block(chunk)
|
||||
if should_start_new_block:
|
||||
is_opening_first_block = self.sent_content_block_start is False
|
||||
if is_opening_first_block and self._is_blank_delta(chunk):
|
||||
continue
|
||||
if is_opening_first_block:
|
||||
self.sent_content_block_start = True
|
||||
self.sent_content_block_finish = False
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": self.current_content_block_start,
|
||||
}
|
||||
)
|
||||
elif should_start_new_block:
|
||||
self._increment_content_block_index()
|
||||
|
||||
# applied_edits only needs to flow to the final message_delta
|
||||
|
|
@ -447,7 +448,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
# ``not self.queued_usage_chunk``.
|
||||
continue
|
||||
|
||||
if should_start_new_block and not self.sent_content_block_finish:
|
||||
if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish:
|
||||
# Queue the sequence: content_block_stop -> content_block_start
|
||||
# -> (optionally) the trigger chunk's delta.
|
||||
#
|
||||
|
|
@ -615,25 +616,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
if compaction_event is not None:
|
||||
return compaction_event
|
||||
|
||||
if self.sent_content_block_start is False:
|
||||
self.sent_content_block_start = True
|
||||
self.sent_content_block_finish = False
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
}
|
||||
)
|
||||
return self.chunk_queue.popleft()
|
||||
|
||||
async for chunk in self.completion_stream:
|
||||
if chunk == "None" or chunk is None:
|
||||
raise Exception
|
||||
|
||||
# Check if we need to start a new content block
|
||||
should_start_new_block = self._should_start_new_content_block(chunk)
|
||||
if should_start_new_block:
|
||||
is_opening_first_block = self.sent_content_block_start is False
|
||||
if is_opening_first_block and self._is_blank_delta(chunk):
|
||||
continue
|
||||
if is_opening_first_block:
|
||||
self.sent_content_block_start = True
|
||||
self.sent_content_block_finish = False
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": self.current_content_block_start,
|
||||
}
|
||||
)
|
||||
elif should_start_new_block:
|
||||
self._increment_content_block_index()
|
||||
|
||||
# applied_edits only needs to flow to the final message_delta
|
||||
|
|
@ -664,7 +665,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
# Check if this processed chunk has a stop_reason - hold it for next chunk
|
||||
|
||||
if not self.queued_usage_chunk:
|
||||
if should_start_new_block and not self.sent_content_block_finish:
|
||||
if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish:
|
||||
# Queue the sequence: content_block_stop -> content_block_start
|
||||
# -> (optionally) the trigger chunk's delta.
|
||||
#
|
||||
|
|
@ -875,6 +876,22 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
return False
|
||||
return bool(delta.get(_delta_payload_field(delta_type)))
|
||||
|
||||
@staticmethod
|
||||
def _is_blank_delta(chunk: "ModelResponseStream") -> bool:
|
||||
choice = chunk.choices[0]
|
||||
if choice.finish_reason is not None:
|
||||
return False
|
||||
delta = choice.delta
|
||||
if getattr(delta, "tool_calls", None):
|
||||
return False
|
||||
if getattr(delta, "content", None):
|
||||
return False
|
||||
if getattr(delta, "reasoning_content", None):
|
||||
return False
|
||||
if getattr(delta, "thinking_blocks", None):
|
||||
return False
|
||||
return True
|
||||
|
||||
def _should_start_new_content_block(self, chunk: "ModelResponseStream") -> bool:
|
||||
"""
|
||||
Determine if we should start a new content block based on the processed chunk.
|
||||
|
|
|
|||
|
|
@ -331,6 +331,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
"thinking",
|
||||
"output_format",
|
||||
"output_config",
|
||||
"stop_sequences",
|
||||
]
|
||||
|
||||
def _is_web_search_tool(self, tool: Dict[str, Any]) -> bool:
|
||||
|
|
@ -615,7 +616,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
thinking_type = thinking.get("type", "disabled")
|
||||
|
||||
if thinking_type == "disabled":
|
||||
return None
|
||||
return "none"
|
||||
elif thinking_type == "enabled":
|
||||
return reasoning_effort_from_thinking_budget(thinking.get("budget_tokens", 0))
|
||||
elif thinking_type == "adaptive":
|
||||
|
|
@ -683,25 +684,37 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
thinking
|
||||
)
|
||||
if reasoning_effort:
|
||||
summary = thinking.get("summary") if isinstance(thinking, dict) else None
|
||||
auto_summary = is_reasoning_auto_summary_enabled()
|
||||
if summary:
|
||||
return {
|
||||
"reasoning_effort": {
|
||||
"effort": reasoning_effort,
|
||||
"summary": summary,
|
||||
}
|
||||
}
|
||||
elif auto_summary:
|
||||
return {
|
||||
"reasoning_effort": {
|
||||
"effort": reasoning_effort,
|
||||
"summary": "detailed",
|
||||
}
|
||||
}
|
||||
return {"reasoning_effort": reasoning_effort}
|
||||
return {
|
||||
"reasoning_effort": LiteLLMAnthropicMessagesAdapter._apply_reasoning_summary_wrapping(
|
||||
reasoning_effort, thinking
|
||||
)
|
||||
}
|
||||
return {}
|
||||
|
||||
@staticmethod
|
||||
def _apply_reasoning_summary_wrapping(
|
||||
reasoning_effort: str,
|
||||
thinking: Dict[str, Any],
|
||||
) -> Any:
|
||||
"""
|
||||
Apply the reasoning_effort/summary wrapping rules shared by every
|
||||
thinking->reasoning_effort translation path.
|
||||
|
||||
Disabled thinking always stays a plain string - there's no reasoning
|
||||
trace to summarize, and non-Claude providers (e.g. Fireworks) expect
|
||||
reasoning_effort as a plain string, not a summary dict.
|
||||
"""
|
||||
thinking_type = thinking.get("type") if isinstance(thinking, dict) else None
|
||||
if thinking_type == "disabled":
|
||||
return reasoning_effort
|
||||
|
||||
summary = thinking.get("summary") if isinstance(thinking, dict) else None
|
||||
if summary:
|
||||
return {"effort": reasoning_effort, "summary": summary}
|
||||
if is_reasoning_auto_summary_enabled():
|
||||
return {"effort": reasoning_effort, "summary": "detailed"}
|
||||
return reasoning_effort
|
||||
|
||||
def translate_anthropic_tool_choice_to_openai(
|
||||
self, tool_choice: AnthropicMessagesToolChoice
|
||||
) -> ChatCompletionToolChoiceValues:
|
||||
|
|
@ -919,6 +932,18 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
tool_choice=cast(AnthropicMessagesToolChoice, tool_choice)
|
||||
)
|
||||
|
||||
def _translate_stop_sequences_to_openai(
|
||||
self,
|
||||
anthropic_message_request: AnthropicMessagesRequest,
|
||||
new_kwargs: ChatCompletionRequest,
|
||||
) -> None:
|
||||
if "stop_sequences" not in anthropic_message_request:
|
||||
return
|
||||
stop_sequences = anthropic_message_request["stop_sequences"]
|
||||
if not stop_sequences:
|
||||
return
|
||||
new_kwargs["stop"] = stop_sequences
|
||||
|
||||
def _translate_tools_to_openai(
|
||||
self,
|
||||
anthropic_message_request: AnthropicMessagesRequest,
|
||||
|
|
@ -976,32 +1001,17 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
if not reasoning_effort:
|
||||
return
|
||||
|
||||
thinking_type = thinking.get("type") if isinstance(thinking, dict) else None
|
||||
|
||||
# For adaptive thinking, override with output_config.effort if available
|
||||
if isinstance(thinking, dict) and thinking.get("type") == "adaptive":
|
||||
if thinking_type == "adaptive":
|
||||
output_config = anthropic_message_request.get("output_config")
|
||||
if isinstance(output_config, dict) and output_config.get("effort"):
|
||||
reasoning_effort = output_config["effort"]
|
||||
|
||||
summary = thinking.get("summary") if isinstance(thinking, dict) else None
|
||||
auto_summary = is_reasoning_auto_summary_enabled()
|
||||
if summary:
|
||||
new_kwargs["reasoning_effort"] = cast(
|
||||
Any,
|
||||
{
|
||||
"effort": reasoning_effort,
|
||||
"summary": summary,
|
||||
},
|
||||
)
|
||||
elif auto_summary:
|
||||
new_kwargs["reasoning_effort"] = cast(
|
||||
Any,
|
||||
{
|
||||
"effort": reasoning_effort,
|
||||
"summary": "detailed",
|
||||
},
|
||||
)
|
||||
else:
|
||||
new_kwargs["reasoning_effort"] = reasoning_effort
|
||||
new_kwargs["reasoning_effort"] = self._apply_reasoning_summary_wrapping(
|
||||
reasoning_effort, cast(Dict[str, Any], thinking)
|
||||
)
|
||||
|
||||
def _translate_output_format_to_openai(
|
||||
self,
|
||||
|
|
@ -1098,6 +1108,11 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
anthropic_message_request=anthropic_message_request,
|
||||
new_kwargs=new_kwargs,
|
||||
)
|
||||
## CONVERT STOP_SEQUENCES
|
||||
self._translate_stop_sequences_to_openai(
|
||||
anthropic_message_request=anthropic_message_request,
|
||||
new_kwargs=new_kwargs,
|
||||
)
|
||||
## CONVERT OUTPUT_FORMAT to RESPONSE_FORMAT
|
||||
self._translate_output_format_to_openai(
|
||||
anthropic_message_request=anthropic_message_request,
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ import litellm
|
|||
import litellm.constants as _c
|
||||
from litellm.litellm_core_utils.url_utils import validate_url
|
||||
from litellm.llms.anthropic.common_utils import strip_advisor_blocks_from_messages
|
||||
from litellm.router_utils.cooldown_handlers import mark_advisor_orchestration_failure
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
)
|
||||
|
|
@ -124,30 +125,36 @@ class AdvisorOrchestrationHandler(MessagesInterceptor):
|
|||
|
||||
iteration += 1
|
||||
if iteration > max_uses:
|
||||
raise AdvisorMaxIterationsError(
|
||||
max_iterations_error = AdvisorMaxIterationsError(
|
||||
f"Advisor orchestration loop exceeded max_uses={max_uses}. "
|
||||
"Increase max_uses in the advisor tool definition or cap the request."
|
||||
)
|
||||
mark_advisor_orchestration_failure(max_iterations_error)
|
||||
raise max_iterations_error
|
||||
|
||||
# --- Build advisor context ---
|
||||
advisor_messages = _build_advisor_context(current_messages, executor_response, advisor_use_block)
|
||||
|
||||
# --- Advisor sub-call (always non-streaming, no tools) ---
|
||||
advisor_response: AnthropicMessagesResponse = await _call_messages_handler(
|
||||
model=advisor_model,
|
||||
messages=advisor_messages,
|
||||
tools=None,
|
||||
stream=False,
|
||||
max_tokens=max_tokens,
|
||||
custom_llm_provider=None, # let litellm resolve from model name
|
||||
metadata={
|
||||
**metadata_base,
|
||||
"advisor_sub_call": True,
|
||||
"parent_request_id": parent_request_id,
|
||||
},
|
||||
api_key=advisor_api_key,
|
||||
api_base=advisor_api_base,
|
||||
)
|
||||
try:
|
||||
advisor_response: AnthropicMessagesResponse = await _call_messages_handler(
|
||||
model=advisor_model,
|
||||
messages=advisor_messages,
|
||||
tools=None,
|
||||
stream=False,
|
||||
max_tokens=max_tokens,
|
||||
custom_llm_provider=None, # let litellm resolve from model name
|
||||
metadata={
|
||||
**metadata_base,
|
||||
"advisor_sub_call": True,
|
||||
"parent_request_id": parent_request_id,
|
||||
},
|
||||
api_key=advisor_api_key,
|
||||
api_base=advisor_api_base,
|
||||
)
|
||||
except Exception as advisor_sub_call_exception:
|
||||
mark_advisor_orchestration_failure(advisor_sub_call_exception)
|
||||
raise
|
||||
|
||||
advisor_text = _extract_response_text(advisor_response)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,8 @@
|
|||
Azure Batches API Handler
|
||||
"""
|
||||
|
||||
from typing import Any, Coroutine, Optional, Union, cast
|
||||
from collections.abc import Coroutine
|
||||
from typing import cast
|
||||
|
||||
import httpx
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
|
|
@ -33,32 +34,30 @@ class AzureBatchesAPI(BaseAzureLLM):
|
|||
async def acreate_batch(
|
||||
self,
|
||||
create_batch_data: CreateBatchRequest,
|
||||
azure_client: Union[AsyncAzureOpenAI, AsyncOpenAI],
|
||||
azure_client: AsyncAzureOpenAI | AsyncOpenAI,
|
||||
) -> LiteLLMBatch:
|
||||
response = await azure_client.batches.create(**create_batch_data) # type: ignore[arg-type]
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
def create_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
create_batch_data: CreateBatchRequest,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
max_retries: Optional[int],
|
||||
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]:
|
||||
azure_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = (
|
||||
self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
api_version: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
) -> LiteLLMBatch | Coroutine[object, object, LiteLLMBatch]:
|
||||
azure_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
if azure_client is None:
|
||||
raise ValueError(
|
||||
|
|
@ -73,38 +72,36 @@ class AzureBatchesAPI(BaseAzureLLM):
|
|||
return self.acreate_batch( # type: ignore
|
||||
create_batch_data=create_batch_data, azure_client=azure_client
|
||||
)
|
||||
response = cast(Union[AzureOpenAI, OpenAI], azure_client).batches.create(**create_batch_data) # type: ignore[arg-type]
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
response = cast(AzureOpenAI | OpenAI, azure_client).batches.create(**create_batch_data) # type: ignore[arg-type]
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
async def aretrieve_batch(
|
||||
self,
|
||||
retrieve_batch_data: RetrieveBatchRequest,
|
||||
client: Union[AsyncAzureOpenAI, AsyncOpenAI],
|
||||
client: AsyncAzureOpenAI | AsyncOpenAI,
|
||||
) -> LiteLLMBatch:
|
||||
response = await client.batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type]
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
def retrieve_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
retrieve_batch_data: RetrieveBatchRequest,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
max_retries: Optional[int],
|
||||
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
api_version: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
):
|
||||
azure_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = (
|
||||
self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
azure_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
if azure_client is None:
|
||||
raise ValueError(
|
||||
|
|
@ -119,38 +116,36 @@ class AzureBatchesAPI(BaseAzureLLM):
|
|||
return self.aretrieve_batch( # type: ignore
|
||||
retrieve_batch_data=retrieve_batch_data, client=azure_client
|
||||
)
|
||||
response = cast(Union[AzureOpenAI, OpenAI], azure_client).batches.retrieve(**retrieve_batch_data)
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
response = cast(AzureOpenAI | OpenAI, azure_client).batches.retrieve(**retrieve_batch_data)
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
async def acancel_batch(
|
||||
self,
|
||||
cancel_batch_data: CancelBatchRequest,
|
||||
client: Union[AsyncAzureOpenAI, AsyncOpenAI],
|
||||
client: AsyncAzureOpenAI | AsyncOpenAI,
|
||||
) -> LiteLLMBatch:
|
||||
response = await client.batches.cancel(**cancel_batch_data)
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
def cancel_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
cancel_batch_data: CancelBatchRequest,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
max_retries: Optional[int],
|
||||
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
api_version: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
):
|
||||
azure_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = (
|
||||
self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
azure_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
if azure_client is None:
|
||||
raise ValueError(
|
||||
|
|
@ -172,13 +167,13 @@ class AzureBatchesAPI(BaseAzureLLM):
|
|||
"Azure client is not an instance of AzureOpenAI or OpenAI. Make sure you passed a sync client."
|
||||
)
|
||||
response = azure_client.batches.cancel(**cancel_batch_data)
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
async def alist_batches(
|
||||
self,
|
||||
client: Union[AsyncAzureOpenAI, AsyncOpenAI],
|
||||
after: Optional[str] = None,
|
||||
limit: Optional[int] = None,
|
||||
client: AsyncAzureOpenAI | AsyncOpenAI,
|
||||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
):
|
||||
response = await client.batches.list(after=after, limit=limit) # type: ignore
|
||||
return response
|
||||
|
|
@ -186,25 +181,23 @@ class AzureBatchesAPI(BaseAzureLLM):
|
|||
def list_batches(
|
||||
self,
|
||||
_is_async: bool,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
max_retries: Optional[int],
|
||||
after: Optional[str] = None,
|
||||
limit: Optional[int] = None,
|
||||
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
api_version: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
):
|
||||
azure_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]] = (
|
||||
self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
azure_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
litellm_params=litellm_params or {},
|
||||
)
|
||||
if azure_client is None:
|
||||
raise ValueError(
|
||||
|
|
|
|||
|
|
@ -4,8 +4,6 @@ Azure AI Anthropic CountTokens API transformation logic.
|
|||
Extends the base Anthropic CountTokens transformation with Azure authentication.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION
|
||||
from litellm.llms.anthropic.count_tokens.transformation import (
|
||||
AnthropicCountTokensConfig,
|
||||
|
|
@ -25,8 +23,8 @@ class AzureAIAnthropicCountTokensConfig(AnthropicCountTokensConfig):
|
|||
def get_required_headers(
|
||||
self,
|
||||
api_key: str,
|
||||
litellm_params: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, str]:
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
Get the required headers for the Azure AI Anthropic CountTokens API.
|
||||
|
||||
|
|
@ -53,7 +51,7 @@ class AzureAIAnthropicCountTokensConfig(AnthropicCountTokensConfig):
|
|||
if "api_key" not in litellm_params:
|
||||
litellm_params["api_key"] = api_key
|
||||
|
||||
litellm_params_obj = GenericLiteLLMParams(**litellm_params)
|
||||
litellm_params_obj = GenericLiteLLMParams.model_validate(litellm_params)
|
||||
|
||||
# Get Azure auth headers (api-key or Authorization)
|
||||
azure_headers = BaseAzureLLM._base_validate_azure_environment(headers={}, litellm_params=litellm_params_obj)
|
||||
|
|
|
|||
|
|
@ -78,8 +78,10 @@ class BaseAWSLLM:
|
|||
# Storage is in-process memory only: default ``DualCache()`` has no Redis backend unless attached
|
||||
# elsewhere. Entry TTL: static access-key + secret + region use ``_get_default_ttl_for_boto3_credentials``
|
||||
# (~59 minutes); ambient env (``_auth_with_env_vars`` returns ``ttl=None``) uses ``InMemoryCache``'s
|
||||
# ``default_ttl`` (600 seconds / 10 minutes). AssumeRole, web identity, profiles, and explicit
|
||||
# session-token tuples are not cached — see ``get_credentials`` and ``_get_or_set_cached_credentials``.
|
||||
# ``default_ttl`` (600 seconds / 10 minutes); web identity STS credentials use
|
||||
# ``_get_default_ttl_for_boto3_credentials`` (~59 minutes), keyed on all aws_* credential args
|
||||
# plus ssl_verify. AssumeRole, profiles, and explicit session-token tuples are not cached — see
|
||||
# ``get_credentials`` and ``_get_or_set_cached_credentials``.
|
||||
_shared_iam_cache: ClassVar[DualCache] = DualCache()
|
||||
|
||||
def __init__(self) -> None:
|
||||
|
|
@ -136,11 +138,12 @@ class BaseAWSLLM:
|
|||
which ``InMemoryCache.set_cache`` resolves to ``default_ttl`` (600 seconds / 10 minutes by
|
||||
default).
|
||||
|
||||
Used only for static access-key credentials and ambient credentials from
|
||||
Used for static access-key credentials, ambient credentials from
|
||||
``_auth_with_env_vars`` (including when skipping AssumeRole because the runtime identity
|
||||
already matches ``aws_role_name``).
|
||||
already matches ``aws_role_name``), and web identity STS credentials (plain
|
||||
non-refreshable ``Credentials`` cached ~59 min, inside the 3600s STS session).
|
||||
|
||||
AssumeRole, web identity exchange, profiles, and explicit session-token tuples are not
|
||||
AssumeRole, profiles, and explicit session-token tuples are not
|
||||
cached here — shared ``Credentials`` / refresh state must not span logical sessions.
|
||||
"""
|
||||
cache_key = self.get_cache_key(credential_args)
|
||||
|
|
@ -266,23 +269,26 @@ class BaseAWSLLM:
|
|||
# Credentials - boto3.Credentials
|
||||
# cache ttl - Optional[int]. If None, the credentials are not cached. Some auth flows have no expiry time.
|
||||
#
|
||||
# iam_cache: static keys and ambient env only (including skip-AssumeRole path).
|
||||
# Do not cache AssumeRole / web identity / profile / explicit session-token paths here.
|
||||
# iam_cache: static keys, ambient env (including skip-AssumeRole path), and web identity.
|
||||
# Do not cache AssumeRole / profile / explicit session-token paths here.
|
||||
#########################################################
|
||||
if self._is_auth_with_web_identity_token(
|
||||
aws_web_identity_token,
|
||||
aws_role_name,
|
||||
aws_session_name,
|
||||
):
|
||||
credentials, _cache_ttl = self._auth_with_web_identity_token(
|
||||
aws_web_identity_token=cast(str, aws_web_identity_token),
|
||||
aws_role_name=cast(str, aws_role_name),
|
||||
aws_session_name=cast(str, aws_session_name),
|
||||
aws_region_name=aws_region_name,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
aws_external_id=aws_external_id,
|
||||
return self._get_or_set_cached_credentials(
|
||||
args,
|
||||
lambda: self._auth_with_web_identity_token(
|
||||
aws_web_identity_token=cast(str, aws_web_identity_token),
|
||||
aws_role_name=cast(str, aws_role_name),
|
||||
aws_session_name=cast(str, aws_session_name),
|
||||
aws_region_name=aws_region_name,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
aws_external_id=aws_external_id,
|
||||
ssl_verify=ssl_verify,
|
||||
),
|
||||
)
|
||||
return credentials
|
||||
elif self._is_auth_with_aws_role(aws_role_name):
|
||||
# Same role (IRSA/ECS/EC2): ambient creds via _get_or_set_cached_credentials like the
|
||||
# default env branch; never pre-read cache (must run _is_already_running_as_role first).
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ from .invoke_handler import (
|
|||
AmazonAnthropicClaudeStreamDecoder,
|
||||
AmazonDeepSeekR1StreamDecoder,
|
||||
AWSEventStreamDecoder,
|
||||
BedrockLLM,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,19 +1,10 @@
|
|||
"""
|
||||
TODO: DELETE FILE. Bedrock LLM is no longer used. Goto `litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py`
|
||||
"""
|
||||
|
||||
import copy
|
||||
import time
|
||||
import types
|
||||
from functools import partial
|
||||
from typing import (
|
||||
AsyncIterator,
|
||||
Callable,
|
||||
Iterator,
|
||||
Optional,
|
||||
Tuple,
|
||||
cast,
|
||||
get_args,
|
||||
)
|
||||
|
||||
import httpx # type: ignore
|
||||
|
|
@ -25,16 +16,6 @@ from litellm.caching.caching import InMemoryCache
|
|||
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
|
||||
from litellm.litellm_core_utils.core_helpers import map_finish_reason
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
cohere_message_pt,
|
||||
construct_tool_use_system_prompt,
|
||||
contains_tag,
|
||||
custom_prompt,
|
||||
extract_between_tags,
|
||||
parse_xml_params,
|
||||
prompt_factory,
|
||||
)
|
||||
from litellm.llms.anthropic.chat.handler import (
|
||||
ModelResponseIterator as AnthropicModelResponseIterator,
|
||||
)
|
||||
|
|
@ -64,12 +45,9 @@ from litellm.types.utils import (
|
|||
StreamingChoices,
|
||||
Usage,
|
||||
)
|
||||
from litellm.utils import CustomStreamWrapper, get_secret
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
from ..common_utils import (
|
||||
BedrockError,
|
||||
ModelResponseIterator,
|
||||
build_bedrock_stream_error,
|
||||
get_bedrock_response_stream_shape,
|
||||
get_bedrock_tool_name,
|
||||
|
|
@ -77,9 +55,6 @@ from ..common_utils import (
|
|||
|
||||
bedrock_tool_name_mappings: InMemoryCache = InMemoryCache(max_size_in_memory=50, default_ttl=600)
|
||||
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
|
||||
from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import (
|
||||
AmazonBedrockOpenAIConfig,
|
||||
)
|
||||
|
||||
converse_config = AmazonConverseConfig()
|
||||
|
||||
|
|
@ -351,932 +326,6 @@ def make_sync_call(
|
|||
raise BedrockError(status_code=500, message=str(e))
|
||||
|
||||
|
||||
class BedrockLLM(BaseAWSLLM):
|
||||
"""
|
||||
Example call
|
||||
|
||||
```
|
||||
curl --location --request POST 'https://bedrock-runtime.{aws_region_name}.amazonaws.com/model/{bedrock_model_name}/invoke' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--header 'Accept: application/json' \
|
||||
--user "$AWS_ACCESS_KEY_ID":"$AWS_SECRET_ACCESS_KEY" \
|
||||
--aws-sigv4 "aws:amz:us-east-1:bedrock" \
|
||||
--data-raw '{
|
||||
"prompt": "Hi",
|
||||
"temperature": 0,
|
||||
"p": 0.9,
|
||||
"max_tokens": 4096
|
||||
}'
|
||||
```
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
@staticmethod
|
||||
def is_claude_messages_api_model(model: str) -> bool:
|
||||
"""
|
||||
Check if the model uses the Claude Messages API (Claude 3+).
|
||||
|
||||
Handles:
|
||||
- Regional prefixes: eu.anthropic.claude-*, us.anthropic.claude-*
|
||||
- Claude 3 models: claude-3-haiku, claude-3-sonnet, claude-3-opus, claude-3-5-*, claude-3-7-*
|
||||
- Claude 4 models: claude-opus-4, claude-sonnet-4, claude-haiku-4
|
||||
"""
|
||||
# Normalize model string to lowercase for matching
|
||||
model_lower = model.lower()
|
||||
|
||||
# Claude 3+ indicators (all use Messages API)
|
||||
messages_api_indicators = [
|
||||
"claude-3", # Claude 3.x models
|
||||
"claude-opus-4", # Claude Opus 4
|
||||
"claude-sonnet-4", # Claude Sonnet 4
|
||||
"claude-haiku-4", # Claude Haiku 4
|
||||
]
|
||||
|
||||
return any(indicator in model_lower for indicator in messages_api_indicators)
|
||||
|
||||
def convert_messages_to_prompt(self, model, messages, provider, custom_prompt_dict) -> Tuple[str, Optional[list]]:
|
||||
# handle anthropic prompts and amazon titan prompts
|
||||
prompt = ""
|
||||
chat_history: Optional[list] = None
|
||||
## CUSTOM PROMPT
|
||||
if model in custom_prompt_dict:
|
||||
# check if the model has a registered custom prompt
|
||||
model_prompt_details = custom_prompt_dict[model]
|
||||
prompt = custom_prompt(
|
||||
role_dict=model_prompt_details["roles"],
|
||||
initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""),
|
||||
final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
|
||||
messages=messages,
|
||||
)
|
||||
return prompt, None
|
||||
## ELSE
|
||||
if provider == "anthropic" or provider == "amazon":
|
||||
prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
|
||||
elif provider == "mistral":
|
||||
prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
|
||||
elif provider == "meta" or provider == "llama":
|
||||
prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
|
||||
elif provider == "openai":
|
||||
# OpenAI uses messages directly, no prompt conversion needed
|
||||
# Return empty prompt as it won't be used
|
||||
prompt = ""
|
||||
elif provider == "cohere":
|
||||
prompt, chat_history = cohere_message_pt(messages=messages)
|
||||
else:
|
||||
prompt = ""
|
||||
for message in messages:
|
||||
if "role" in message:
|
||||
if message["role"] == "user":
|
||||
prompt += f"{message['content']}"
|
||||
else:
|
||||
prompt += f"{message['content']}"
|
||||
else:
|
||||
prompt += f"{message['content']}"
|
||||
return prompt, chat_history # type: ignore
|
||||
|
||||
def process_response(
|
||||
self,
|
||||
model: str,
|
||||
response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
stream: Optional[bool],
|
||||
logging_obj: Logging,
|
||||
optional_params: dict,
|
||||
api_key: str,
|
||||
data: Union[dict, str],
|
||||
messages: List,
|
||||
print_verbose,
|
||||
encoding,
|
||||
) -> Union[ModelResponse, CustomStreamWrapper]:
|
||||
provider = self.get_bedrock_invoke_provider(model)
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
original_response=response.text,
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
print_verbose(f"raw model_response: {response.text}")
|
||||
|
||||
## RESPONSE OBJECT
|
||||
try:
|
||||
completion_response = response.json()
|
||||
except Exception:
|
||||
raise BedrockError(message=response.text, status_code=422)
|
||||
|
||||
outputText: Optional[str] = None
|
||||
try:
|
||||
if provider == "cohere":
|
||||
if "text" in completion_response:
|
||||
outputText = completion_response["text"] # type: ignore
|
||||
elif "generations" in completion_response:
|
||||
outputText = completion_response["generations"][0]["text"]
|
||||
model_response.choices[0].finish_reason = map_finish_reason(
|
||||
completion_response["generations"][0]["finish_reason"]
|
||||
)
|
||||
elif provider == "anthropic":
|
||||
if self.is_claude_messages_api_model(model):
|
||||
json_schemas: dict = {}
|
||||
_is_function_call = False
|
||||
## Handle Tool Calling
|
||||
if "tools" in optional_params:
|
||||
_is_function_call = True
|
||||
for tool in optional_params["tools"]:
|
||||
json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None)
|
||||
outputText = completion_response.get("content")[0].get("text", None)
|
||||
if outputText is not None and contains_tag("invoke", outputText): # OUTPUT PARSE FUNCTION CALL
|
||||
function_name = extract_between_tags("tool_name", outputText)[0]
|
||||
function_arguments_str = extract_between_tags("invoke", outputText)[0].strip()
|
||||
function_arguments_str = f"<invoke>{function_arguments_str}</invoke>"
|
||||
function_arguments = parse_xml_params(
|
||||
function_arguments_str,
|
||||
json_schema=json_schemas.get(
|
||||
function_name, None
|
||||
), # check if we have a json schema for this function name)
|
||||
)
|
||||
_message = litellm.Message(
|
||||
tool_calls=[
|
||||
{
|
||||
"id": f"call_{uuid.uuid4()}",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": function_name,
|
||||
"arguments": json.dumps(function_arguments),
|
||||
},
|
||||
}
|
||||
],
|
||||
content=None,
|
||||
)
|
||||
model_response.choices[0].message = _message # type: ignore
|
||||
model_response._hidden_params["original_response"] = (
|
||||
outputText # allow user to access raw anthropic tool calling response
|
||||
)
|
||||
if _is_function_call is True and stream is not None and stream is True:
|
||||
print_verbose("INSIDE BEDROCK STREAMING TOOL CALLING CONDITION BLOCK")
|
||||
# return an iterator
|
||||
streaming_model_response = ModelResponseStream()
|
||||
streaming_model_response.choices[0].finish_reason = getattr(
|
||||
model_response.choices[0], "finish_reason", "stop"
|
||||
)
|
||||
# streaming_model_response.choices = [litellm.utils.StreamingChoices()]
|
||||
streaming_choice = litellm.utils.StreamingChoices()
|
||||
streaming_choice.index = model_response.choices[0].index
|
||||
_tool_calls = []
|
||||
print_verbose(f"type of model_response.choices[0]: {type(model_response.choices[0])}")
|
||||
print_verbose(f"type of streaming_choice: {type(streaming_choice)}")
|
||||
if isinstance(model_response.choices[0], litellm.Choices):
|
||||
if getattr(
|
||||
model_response.choices[0].message, "tool_calls", None
|
||||
) is not None and isinstance(model_response.choices[0].message.tool_calls, list):
|
||||
for tool_call in model_response.choices[0].message.tool_calls:
|
||||
_tool_call = {**tool_call.dict(), "index": 0}
|
||||
_tool_calls.append(_tool_call)
|
||||
delta_obj = Delta(
|
||||
content=getattr(model_response.choices[0].message, "content", None),
|
||||
role=model_response.choices[0].message.role,
|
||||
tool_calls=_tool_calls,
|
||||
)
|
||||
streaming_choice.delta = delta_obj
|
||||
streaming_model_response.choices = [streaming_choice]
|
||||
completion_stream = ModelResponseIterator(model_response=streaming_model_response)
|
||||
print_verbose(
|
||||
"Returns anthropic CustomStreamWrapper with 'cached_response' streaming object"
|
||||
)
|
||||
return litellm.CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="cached_response",
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
model_response.choices[0].finish_reason = map_finish_reason(
|
||||
completion_response.get("stop_reason", "")
|
||||
)
|
||||
_usage = litellm.Usage(
|
||||
prompt_tokens=completion_response["usage"]["input_tokens"],
|
||||
completion_tokens=completion_response["usage"]["output_tokens"],
|
||||
total_tokens=completion_response["usage"]["input_tokens"]
|
||||
+ completion_response["usage"]["output_tokens"],
|
||||
)
|
||||
setattr(model_response, "usage", _usage)
|
||||
else:
|
||||
outputText = completion_response["completion"]
|
||||
|
||||
model_response.choices[0].finish_reason = completion_response["stop_reason"]
|
||||
elif provider == "ai21":
|
||||
outputText = completion_response.get("completions")[0].get("data").get("text")
|
||||
elif provider == "meta" or provider == "llama":
|
||||
outputText = completion_response["generation"]
|
||||
elif provider == "openai":
|
||||
# OpenAI imported models use OpenAI Chat Completions format
|
||||
if "choices" in completion_response and len(completion_response["choices"]) > 0:
|
||||
choice = completion_response["choices"][0]
|
||||
if "message" in choice:
|
||||
outputText = choice["message"].get("content")
|
||||
elif "text" in choice: # fallback for completion format
|
||||
outputText = choice["text"]
|
||||
|
||||
# Set finish reason
|
||||
if "finish_reason" in choice:
|
||||
model_response.choices[0].finish_reason = map_finish_reason(choice["finish_reason"])
|
||||
|
||||
# Set usage if available
|
||||
if "usage" in completion_response:
|
||||
usage = completion_response["usage"]
|
||||
_usage = litellm.Usage(
|
||||
prompt_tokens=usage.get("prompt_tokens", 0),
|
||||
completion_tokens=usage.get("completion_tokens", 0),
|
||||
total_tokens=usage.get("total_tokens", 0),
|
||||
)
|
||||
setattr(model_response, "usage", _usage)
|
||||
elif provider == "mistral":
|
||||
outputText = completion_response["outputs"][0]["text"]
|
||||
model_response.choices[0].finish_reason = completion_response["outputs"][0]["stop_reason"]
|
||||
else: # amazon titan
|
||||
outputText = completion_response.get("results")[0].get("outputText")
|
||||
except Exception as e:
|
||||
raise BedrockError(
|
||||
message="Error processing={}, Received error={}".format(response.text, str(e)),
|
||||
status_code=422,
|
||||
)
|
||||
|
||||
try:
|
||||
if (
|
||||
outputText is not None
|
||||
and len(outputText) > 0
|
||||
and hasattr(model_response.choices[0], "message")
|
||||
and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore
|
||||
is None
|
||||
):
|
||||
model_response.choices[0].message.content = outputText # type: ignore
|
||||
elif (
|
||||
hasattr(model_response.choices[0], "message")
|
||||
and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore
|
||||
is not None
|
||||
):
|
||||
pass
|
||||
else:
|
||||
raise Exception()
|
||||
except Exception as e:
|
||||
raise BedrockError(
|
||||
message="Error parsing received text={}.\nError-{}".format(outputText, str(e)),
|
||||
status_code=response.status_code,
|
||||
)
|
||||
|
||||
if stream and provider == "ai21":
|
||||
streaming_model_response = ModelResponseStream()
|
||||
streaming_model_response.choices[0].finish_reason = model_response.choices[ # type: ignore
|
||||
0
|
||||
].finish_reason
|
||||
# streaming_model_response.choices = [litellm.utils.StreamingChoices()]
|
||||
streaming_choice = litellm.utils.StreamingChoices()
|
||||
streaming_choice.index = model_response.choices[0].index
|
||||
delta_obj = litellm.utils.Delta(
|
||||
content=getattr(model_response.choices[0].message, "content", None), # type: ignore
|
||||
role=model_response.choices[0].message.role, # type: ignore
|
||||
)
|
||||
streaming_choice.delta = delta_obj
|
||||
streaming_model_response.choices = [streaming_choice]
|
||||
mri = ModelResponseIterator(model_response=streaming_model_response)
|
||||
return CustomStreamWrapper(
|
||||
completion_stream=mri,
|
||||
model=model,
|
||||
custom_llm_provider="cached_response",
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
## CALCULATING USAGE - bedrock returns usage in the headers
|
||||
# Skip if usage was already set (e.g., from JSON response for OpenAI provider)
|
||||
if not hasattr(model_response, "usage") or getattr(model_response, "usage", None) is None:
|
||||
bedrock_input_tokens = response.headers.get("x-amzn-bedrock-input-token-count", None)
|
||||
bedrock_output_tokens = response.headers.get("x-amzn-bedrock-output-token-count", None)
|
||||
|
||||
prompt_tokens = int(bedrock_input_tokens or litellm.token_counter(messages=messages))
|
||||
|
||||
completion_tokens = int(
|
||||
bedrock_output_tokens
|
||||
or litellm.token_counter(
|
||||
text=model_response.choices[0].message.content, # type: ignore
|
||||
count_response_tokens=True,
|
||||
)
|
||||
)
|
||||
|
||||
model_response.created = int(time.time())
|
||||
model_response.model = model
|
||||
usage = Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
)
|
||||
setattr(model_response, "usage", usage)
|
||||
else:
|
||||
# Ensure created and model are set even if usage was already set
|
||||
model_response.created = int(time.time())
|
||||
model_response.model = model
|
||||
|
||||
return model_response
|
||||
|
||||
def completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
api_base: Optional[str],
|
||||
custom_prompt_dict: dict,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
encoding,
|
||||
logging_obj: Logging,
|
||||
optional_params: dict,
|
||||
acompletion: bool,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
litellm_params=None,
|
||||
logger_fn=None,
|
||||
extra_headers: Optional[dict] = None,
|
||||
client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None,
|
||||
) -> Union[ModelResponse, CustomStreamWrapper]:
|
||||
try:
|
||||
from botocore.credentials import Credentials
|
||||
except ImportError:
|
||||
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
|
||||
|
||||
## SETUP ##
|
||||
stream = optional_params.pop("stream", None)
|
||||
stream_chunk_size = optional_params.pop("stream_chunk_size", None)
|
||||
|
||||
provider = self.get_bedrock_invoke_provider(model)
|
||||
modelId = self.get_bedrock_model_id(
|
||||
model=model,
|
||||
provider=provider,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
## CREDENTIALS ##
|
||||
# pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them
|
||||
aws_secret_access_key = optional_params.pop("aws_secret_access_key", None)
|
||||
aws_access_key_id = optional_params.pop("aws_access_key_id", None)
|
||||
aws_session_token = optional_params.pop("aws_session_token", None)
|
||||
aws_region_name = optional_params.pop("aws_region_name", None)
|
||||
aws_role_name = optional_params.pop("aws_role_name", None)
|
||||
aws_session_name = optional_params.pop("aws_session_name", None)
|
||||
aws_profile_name = optional_params.pop("aws_profile_name", None)
|
||||
aws_bedrock_runtime_endpoint = optional_params.pop(
|
||||
"aws_bedrock_runtime_endpoint", None
|
||||
) # https://bedrock-runtime.{region_name}.amazonaws.com
|
||||
aws_web_identity_token = optional_params.pop("aws_web_identity_token", None)
|
||||
aws_sts_endpoint = optional_params.pop("aws_sts_endpoint", None)
|
||||
ssl_verify = optional_params.pop("ssl_verify", None)
|
||||
|
||||
### SET REGION NAME ###
|
||||
if aws_region_name is None:
|
||||
# check env #
|
||||
litellm_aws_region_name = get_secret("AWS_REGION_NAME", None)
|
||||
|
||||
if litellm_aws_region_name is not None and isinstance(litellm_aws_region_name, str):
|
||||
aws_region_name = litellm_aws_region_name
|
||||
|
||||
standard_aws_region_name = get_secret("AWS_REGION", None)
|
||||
if standard_aws_region_name is not None and isinstance(standard_aws_region_name, str):
|
||||
aws_region_name = standard_aws_region_name
|
||||
|
||||
if aws_region_name is None:
|
||||
aws_region_name = "us-west-2"
|
||||
|
||||
credentials: Credentials = self.get_credentials(
|
||||
aws_access_key_id=aws_access_key_id,
|
||||
aws_secret_access_key=aws_secret_access_key,
|
||||
aws_session_token=aws_session_token,
|
||||
aws_region_name=aws_region_name,
|
||||
aws_session_name=aws_session_name,
|
||||
aws_profile_name=aws_profile_name,
|
||||
aws_role_name=aws_role_name,
|
||||
aws_web_identity_token=aws_web_identity_token,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
### SET RUNTIME ENDPOINT ###
|
||||
endpoint_url, proxy_endpoint_url = self.get_runtime_endpoint(
|
||||
api_base=api_base,
|
||||
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
|
||||
aws_region_name=aws_region_name,
|
||||
)
|
||||
|
||||
if (stream is not None and stream is True) and provider != "ai21":
|
||||
endpoint_url = f"{endpoint_url}/model/{modelId}/invoke-with-response-stream"
|
||||
proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke-with-response-stream"
|
||||
else:
|
||||
endpoint_url = f"{endpoint_url}/model/{modelId}/invoke"
|
||||
proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke"
|
||||
|
||||
if acompletion and provider == "anthropic" and self.is_claude_messages_api_model(model):
|
||||
if isinstance(client, HTTPHandler):
|
||||
client = None
|
||||
return self._async_anthropic_messages_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
endpoint_url=endpoint_url,
|
||||
proxy_endpoint_url=proxy_endpoint_url,
|
||||
credentials=credentials,
|
||||
aws_region_name=aws_region_name,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
encoding=encoding,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
) # type: ignore[return-value]
|
||||
|
||||
prompt, chat_history = self.convert_messages_to_prompt(model, messages, provider, custom_prompt_dict)
|
||||
inference_params = copy.deepcopy(optional_params)
|
||||
json_schemas: dict = {}
|
||||
if provider == "cohere":
|
||||
if model.startswith("cohere.command-r"):
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonCohereChatConfig().get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
_data = {"message": prompt, **inference_params}
|
||||
if chat_history is not None:
|
||||
_data["chat_history"] = chat_history
|
||||
data = json.dumps(_data)
|
||||
else:
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonCohereConfig.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
if stream is True:
|
||||
inference_params["stream"] = True # cohere requires stream = True in inference params
|
||||
data = json.dumps({"prompt": prompt, **inference_params})
|
||||
elif provider == "anthropic":
|
||||
if self.is_claude_messages_api_model(model):
|
||||
# Separate system prompt from rest of message
|
||||
system_prompt_idx: list[int] = []
|
||||
system_messages: list[str] = []
|
||||
for idx, message in enumerate(messages):
|
||||
if message["role"] == "system":
|
||||
system_messages.append(message["content"])
|
||||
system_prompt_idx.append(idx)
|
||||
if len(system_prompt_idx) > 0:
|
||||
inference_params["system"] = "\n".join(system_messages)
|
||||
messages = [i for j, i in enumerate(messages) if j not in system_prompt_idx]
|
||||
# Format rest of message according to anthropic guidelines
|
||||
messages = prompt_factory(model=model, messages=messages, custom_llm_provider="anthropic_xml") # type: ignore
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonAnthropicClaudeConfig.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
## Handle Tool Calling
|
||||
if "tools" in inference_params:
|
||||
_is_function_call = True
|
||||
for tool in inference_params["tools"]:
|
||||
json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None)
|
||||
tool_calling_system_prompt = construct_tool_use_system_prompt(tools=inference_params["tools"])
|
||||
inference_params["system"] = (
|
||||
inference_params.get("system", "\n") + tool_calling_system_prompt
|
||||
) # add the anthropic tool calling prompt to the system prompt
|
||||
inference_params.pop("tools")
|
||||
data = json.dumps({"messages": messages, **inference_params})
|
||||
else:
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonAnthropicConfig.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
data = json.dumps({"prompt": prompt, **inference_params})
|
||||
elif provider == "ai21":
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonAI21Config.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
|
||||
data = json.dumps({"prompt": prompt, **inference_params})
|
||||
elif provider == "mistral":
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonMistralConfig.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
|
||||
data = json.dumps({"prompt": prompt, **inference_params})
|
||||
elif provider == "amazon": # amazon titan
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonTitanConfig.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
|
||||
data = json.dumps(
|
||||
{
|
||||
"inputText": prompt,
|
||||
"textGenerationConfig": inference_params,
|
||||
}
|
||||
)
|
||||
elif provider == "meta" or provider == "llama":
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonLlamaConfig.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
data = json.dumps({"prompt": prompt, **inference_params})
|
||||
elif provider == "openai":
|
||||
## OpenAI imported models use OpenAI Chat Completions format (messages-based)
|
||||
# Use AmazonBedrockOpenAIConfig for proper OpenAI transformation
|
||||
openai_config = AmazonBedrockOpenAIConfig()
|
||||
supported_params = openai_config.get_supported_openai_params(model=model)
|
||||
|
||||
# Filter to only supported OpenAI params
|
||||
filtered_params = {k: v for k, v in inference_params.items() if k in supported_params}
|
||||
|
||||
# OpenAI uses messages format, not prompt
|
||||
data = json.dumps({"messages": messages, **filtered_params})
|
||||
else:
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": inference_params,
|
||||
},
|
||||
)
|
||||
raise BedrockError(
|
||||
status_code=404,
|
||||
message="Bedrock Invoke HTTPX: Unknown provider={}, model={}. Try calling via converse route - `bedrock/converse/<model>`.".format(
|
||||
provider, model
|
||||
),
|
||||
)
|
||||
|
||||
## COMPLETION CALL
|
||||
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if extra_headers is not None:
|
||||
headers = {"Content-Type": "application/json", **extra_headers}
|
||||
prepped = self.get_request_headers(
|
||||
credentials=credentials,
|
||||
aws_region_name=aws_region_name,
|
||||
extra_headers=extra_headers,
|
||||
endpoint_url=endpoint_url,
|
||||
data=data,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
|
||||
### ROUTING (ASYNC, STREAMING, SYNC)
|
||||
if acompletion:
|
||||
if isinstance(client, HTTPHandler):
|
||||
client = None
|
||||
if stream is True and provider != "ai21":
|
||||
return self.async_streaming(
|
||||
model=model,
|
||||
messages=messages,
|
||||
data=data,
|
||||
api_base=proxy_endpoint_url,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
encoding=encoding,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
stream=True,
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
headers=prepped.headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
) # type: ignore
|
||||
### ASYNC COMPLETION
|
||||
return self.async_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
data=data,
|
||||
api_base=proxy_endpoint_url,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
encoding=encoding,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
stream=stream, # type: ignore
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
headers=prepped.headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
) # type: ignore
|
||||
|
||||
if client is None or isinstance(client, AsyncHTTPHandler):
|
||||
_params = {}
|
||||
if timeout is not None:
|
||||
if isinstance(timeout, float) or isinstance(timeout, int):
|
||||
timeout = httpx.Timeout(timeout)
|
||||
_params["timeout"] = timeout
|
||||
self.client = _get_httpx_client(_params) # type: ignore
|
||||
else:
|
||||
self.client = client
|
||||
if (stream is not None and stream is True) and provider != "ai21":
|
||||
response = self.client.post(
|
||||
url=proxy_endpoint_url,
|
||||
headers=prepped.headers, # type: ignore
|
||||
data=data,
|
||||
stream=stream,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise BedrockError(status_code=response.status_code, message=str(response.read()))
|
||||
|
||||
decoder = AWSEventStreamDecoder(model=model)
|
||||
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
|
||||
streaming_response = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
original_response=streaming_response,
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
return streaming_response
|
||||
|
||||
try:
|
||||
response = self.client.post(
|
||||
url=proxy_endpoint_url,
|
||||
headers=dict(prepped.headers),
|
||||
data=data,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code = err.response.status_code
|
||||
raise BedrockError(status_code=error_code, message=err.response.text)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
|
||||
return self.process_response(
|
||||
model=model,
|
||||
response=response,
|
||||
model_response=model_response,
|
||||
stream=stream,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
api_key="",
|
||||
data=data,
|
||||
messages=messages,
|
||||
print_verbose=print_verbose,
|
||||
encoding=encoding,
|
||||
)
|
||||
|
||||
async def _async_anthropic_messages_completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
endpoint_url: str,
|
||||
proxy_endpoint_url: str,
|
||||
credentials,
|
||||
aws_region_name: str,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
encoding,
|
||||
logging_obj: Logging,
|
||||
optional_params: dict,
|
||||
stream,
|
||||
litellm_params=None,
|
||||
logger_fn=None,
|
||||
extra_headers: Optional[dict] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
stream_chunk_size: Optional[int] = None,
|
||||
) -> Union[ModelResponse, CustomStreamWrapper]:
|
||||
transformed_request = await litellm.AmazonAnthropicClaudeConfig().async_transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params or {},
|
||||
headers=extra_headers or {},
|
||||
)
|
||||
data = json.dumps(transformed_request)
|
||||
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if extra_headers is not None:
|
||||
headers = {"Content-Type": "application/json", **extra_headers}
|
||||
prepped = self.get_request_headers(
|
||||
credentials=credentials,
|
||||
aws_region_name=aws_region_name,
|
||||
extra_headers=extra_headers,
|
||||
endpoint_url=endpoint_url,
|
||||
data=data,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
|
||||
if stream is True:
|
||||
return await self.async_streaming(
|
||||
model=model,
|
||||
messages=messages,
|
||||
data=data,
|
||||
api_base=proxy_endpoint_url,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
encoding=encoding,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
stream=True,
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
headers=prepped.headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
)
|
||||
return await self.async_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
data=data,
|
||||
api_base=proxy_endpoint_url,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
encoding=encoding,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
stream=stream, # type: ignore
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
headers=prepped.headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
)
|
||||
|
||||
async def async_completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
api_base: str,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
data: str,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
encoding,
|
||||
logging_obj: Logging,
|
||||
stream,
|
||||
optional_params: dict,
|
||||
litellm_params=None,
|
||||
logger_fn=None,
|
||||
headers={},
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
) -> Union[ModelResponse, CustomStreamWrapper]:
|
||||
if client is None:
|
||||
_params = {}
|
||||
if timeout is not None:
|
||||
if isinstance(timeout, float) or isinstance(timeout, int):
|
||||
timeout = httpx.Timeout(timeout)
|
||||
_params["timeout"] = timeout
|
||||
client = get_async_httpx_client(params=_params, llm_provider=litellm.LlmProviders.BEDROCK) # type: ignore
|
||||
else:
|
||||
client = client # type: ignore
|
||||
|
||||
try:
|
||||
response = await client.post(
|
||||
api_base,
|
||||
headers=headers,
|
||||
data=data,
|
||||
timeout=timeout,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code = err.response.status_code
|
||||
raise BedrockError(status_code=error_code, message=err.response.text)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
|
||||
return self.process_response(
|
||||
model=model,
|
||||
response=response,
|
||||
model_response=model_response,
|
||||
stream=stream if isinstance(stream, bool) else False,
|
||||
logging_obj=logging_obj,
|
||||
api_key="",
|
||||
data=data,
|
||||
messages=messages,
|
||||
print_verbose=print_verbose,
|
||||
optional_params=optional_params,
|
||||
encoding=encoding,
|
||||
)
|
||||
|
||||
@track_llm_api_timing() # for streaming, we need to instrument the function calling the wrapper
|
||||
async def async_streaming(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
api_base: str,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
data: str,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
encoding,
|
||||
logging_obj: Logging,
|
||||
stream,
|
||||
optional_params: dict,
|
||||
litellm_params=None,
|
||||
logger_fn=None,
|
||||
headers={},
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
stream_chunk_size: Optional[int] = None,
|
||||
) -> CustomStreamWrapper:
|
||||
# The call is not made here; instead, we prepare the necessary objects for the stream.
|
||||
|
||||
streaming_response = CustomStreamWrapper(
|
||||
completion_stream=None,
|
||||
make_call=partial(
|
||||
make_call,
|
||||
client=client,
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
data=data, # type: ignore
|
||||
model=model,
|
||||
messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
fake_stream=True if "ai21" in api_base else False,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
),
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
return streaming_response
|
||||
|
||||
@staticmethod
|
||||
def _get_provider_from_model_path(
|
||||
model_path: str,
|
||||
) -> Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL]:
|
||||
"""
|
||||
Helper function to get the provider from a model path with format: provider/model-name
|
||||
|
||||
Args:
|
||||
model_path (str): The model path (e.g., 'llama/arn:aws:bedrock:us-east-1:086734376398:imported-model/r4c4kewx2s0n' or 'anthropic/model-name')
|
||||
|
||||
Returns:
|
||||
Optional[str]: The provider name, or None if no valid provider found
|
||||
"""
|
||||
parts = model_path.split("/")
|
||||
if len(parts) >= 1:
|
||||
provider = parts[0]
|
||||
if provider in get_args(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL):
|
||||
return cast(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL, provider)
|
||||
return None
|
||||
|
||||
|
||||
class AWSEventStreamDecoder:
|
||||
def __init__(self, model: str, json_mode: Optional[bool] = False) -> None:
|
||||
from botocore.parsers import EventStreamJSONParser
|
||||
|
|
|
|||
|
|
@ -1109,8 +1109,10 @@ def get_bedrock_chat_config(model: str):
|
|||
Returns:
|
||||
The appropriate Bedrock config class instance
|
||||
"""
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
bedrock_route = BedrockModelInfo.get_bedrock_route(model)
|
||||
bedrock_invoke_provider = litellm.BedrockLLM.get_bedrock_invoke_provider(model=model)
|
||||
bedrock_invoke_provider = BaseAWSLLM.get_bedrock_invoke_provider(model=model)
|
||||
base_model = BedrockModelInfo.get_base_model(model)
|
||||
|
||||
# Handle explicit routes first
|
||||
|
|
|
|||
|
|
@ -143,13 +143,26 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
client: Union[ClientSession, Callable[[], ClientSession]],
|
||||
ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None,
|
||||
owns_session: bool = True,
|
||||
session_factory: Callable[[], ClientSession] | None = None,
|
||||
):
|
||||
self.client = client
|
||||
self._ssl_verify = ssl_verify # Store for per-request SSL override
|
||||
super().__init__(client=client, owns_session=owns_session)
|
||||
# Store the client factory for recreating sessions when needed
|
||||
if callable(client):
|
||||
self._client_factory = client
|
||||
default_factory: Callable[[], ClientSession] = client if callable(client) else ClientSession
|
||||
self._client_factory: Callable[[], ClientSession] = session_factory or default_factory
|
||||
|
||||
def _rebuild_session(self) -> ClientSession:
|
||||
"""
|
||||
Build a replacement session from the configured factory.
|
||||
|
||||
The replacement is reachable only from this transport, so the transport
|
||||
owns it from here on even when it was originally handed a session it did
|
||||
not own (the proxy's shared session).
|
||||
"""
|
||||
session = self._client_factory()
|
||||
self._owns_session = True
|
||||
return session
|
||||
|
||||
def _get_valid_client_session(self) -> ClientSession:
|
||||
"""
|
||||
|
|
@ -158,24 +171,16 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
This handles the case where the session was created in a different
|
||||
event loop that may have been closed (common in CI/CD environments).
|
||||
"""
|
||||
from aiohttp.client import ClientSession
|
||||
|
||||
# If we don't have a client or it's not a ClientSession, create one
|
||||
if not isinstance(self.client, ClientSession):
|
||||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
self.client = self._rebuild_session()
|
||||
# Don't return yet - check if the newly created session is valid
|
||||
|
||||
# Check if the session itself is closed
|
||||
if self.client.closed:
|
||||
verbose_logger.debug("Session is closed, creating new session")
|
||||
# Create a new session
|
||||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
self.client = self._rebuild_session()
|
||||
return self.client
|
||||
|
||||
# Check if the existing session is still valid for the current event loop
|
||||
|
|
@ -188,7 +193,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
# Close old session to prevent leaks
|
||||
old_session = self.client
|
||||
try:
|
||||
if not old_session.closed:
|
||||
if self._owns_session and not old_session.closed:
|
||||
try:
|
||||
asyncio.create_task(old_session.close())
|
||||
except RuntimeError:
|
||||
|
|
@ -198,17 +203,11 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
verbose_logger.debug(f"Error closing old session: {e}")
|
||||
|
||||
# Create a new session in the current event loop
|
||||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
self.client = self._rebuild_session()
|
||||
|
||||
except (RuntimeError, AttributeError):
|
||||
# If we can't check the loop or session is invalid, recreate it
|
||||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
self.client = self._rebuild_session()
|
||||
|
||||
return self.client
|
||||
|
||||
|
|
@ -303,10 +302,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
if "Session is closed" in str(e):
|
||||
verbose_logger.debug(f"Session closed during request, retrying with new session: {e}")
|
||||
# Force creation of a new session
|
||||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
self.client = self._rebuild_session()
|
||||
client_session = self.client
|
||||
|
||||
# Retry the request with the new session
|
||||
|
|
|
|||
|
|
@ -1013,17 +1013,6 @@ class AsyncHTTPHandler:
|
|||
|
||||
verbose_logger.debug("Creating AiohttpTransport...")
|
||||
|
||||
# Use shared session if provided and valid
|
||||
if shared_session is not None and not shared_session.closed:
|
||||
verbose_logger.debug(f"SHARED SESSION: Reusing existing ClientSession (ID: {id(shared_session)})")
|
||||
return LiteLLMAiohttpTransport(
|
||||
client=shared_session,
|
||||
ssl_verify=ssl_for_transport,
|
||||
owns_session=False,
|
||||
)
|
||||
|
||||
# Create new session only if none provided or existing one is invalid
|
||||
verbose_logger.debug("NEW SESSION: Creating new ClientSession (no shared session provided)")
|
||||
transport_connector_kwargs = {
|
||||
"keepalive_timeout": AIOHTTP_KEEPALIVE_TIMEOUT,
|
||||
"ttl_dns_cache": AIOHTTP_TTL_DNS_CACHE,
|
||||
|
|
@ -1041,11 +1030,26 @@ class AsyncHTTPHandler:
|
|||
if socket_factory is not None:
|
||||
transport_connector_kwargs["socket_factory"] = socket_factory
|
||||
|
||||
return LiteLLMAiohttpTransport(
|
||||
client=lambda: ClientSession(
|
||||
def session_factory() -> ClientSession:
|
||||
return ClientSession(
|
||||
connector=TCPConnector(**transport_connector_kwargs),
|
||||
trust_env=trust_env,
|
||||
),
|
||||
)
|
||||
|
||||
# Use shared session if provided and valid
|
||||
if shared_session is not None and not shared_session.closed:
|
||||
verbose_logger.debug(f"SHARED SESSION: Reusing existing ClientSession (ID: {id(shared_session)})")
|
||||
return LiteLLMAiohttpTransport(
|
||||
client=shared_session,
|
||||
ssl_verify=ssl_for_transport,
|
||||
owns_session=False,
|
||||
session_factory=session_factory,
|
||||
)
|
||||
|
||||
# Create new session only if none provided or existing one is invalid
|
||||
verbose_logger.debug("NEW SESSION: Creating new ClientSession (no shared session provided)")
|
||||
return LiteLLMAiohttpTransport(
|
||||
client=session_factory,
|
||||
ssl_verify=ssl_for_transport,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
OpenAI Evals API configuration and transformations
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
from collections.abc import Mapping
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -31,6 +31,10 @@ from litellm.types.router import GenericLiteLLMParams
|
|||
from litellm.types.utils import LlmProviders
|
||||
|
||||
|
||||
def _parsed_response_json(raw_response: httpx.Response) -> Mapping[str, object]:
|
||||
return raw_response.json()
|
||||
|
||||
|
||||
class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
||||
"""OpenAI-specific Evals API configuration"""
|
||||
|
||||
|
|
@ -38,7 +42,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
def custom_llm_provider(self) -> LlmProviders:
|
||||
return LlmProviders.OPENAI
|
||||
|
||||
def validate_environment(self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]) -> dict:
|
||||
def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict:
|
||||
"""Add OpenAI-specific headers"""
|
||||
import litellm
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
|
@ -61,9 +65,9 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_base: str | None,
|
||||
endpoint: str,
|
||||
eval_id: Optional[str] = None,
|
||||
eval_id: str | None = None,
|
||||
) -> str:
|
||||
"""Get complete URL for OpenAI Evals API"""
|
||||
if api_base is None:
|
||||
|
|
@ -79,7 +83,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
create_request: CreateEvalRequest,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Dict:
|
||||
) -> dict:
|
||||
"""Transform create eval request for OpenAI"""
|
||||
verbose_logger.debug("Transforming create eval request: %s", create_request)
|
||||
|
||||
|
|
@ -94,17 +98,17 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Eval:
|
||||
"""Transform OpenAI response to Eval object"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming create eval response: %s", response_json)
|
||||
|
||||
return Eval(**response_json)
|
||||
return Eval.model_validate(response_json)
|
||||
|
||||
def transform_list_evals_request(
|
||||
self,
|
||||
list_params: ListEvalsParams,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform list evals request for OpenAI"""
|
||||
api_base = "https://api.openai.com"
|
||||
if litellm_params and litellm_params.api_base:
|
||||
|
|
@ -113,7 +117,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
url = self.get_complete_url(api_base=api_base, endpoint="evals")
|
||||
|
||||
# Build query parameters
|
||||
query_params: Dict[str, Any] = {}
|
||||
query_params: dict[str, object] = {}
|
||||
if "limit" in list_params and list_params["limit"]:
|
||||
query_params["limit"] = list_params["limit"]
|
||||
if "after" in list_params and list_params["after"]:
|
||||
|
|
@ -138,10 +142,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ListEvalsResponse:
|
||||
"""Transform OpenAI response to ListEvalsResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming list evals response: %s", response_json)
|
||||
|
||||
return ListEvalsResponse(**response_json)
|
||||
return ListEvalsResponse.model_validate(response_json)
|
||||
|
||||
def transform_get_eval_request(
|
||||
self,
|
||||
|
|
@ -149,7 +153,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform get eval request for OpenAI"""
|
||||
url = self.get_complete_url(api_base=api_base, endpoint="evals", eval_id=eval_id)
|
||||
|
||||
|
|
@ -163,10 +167,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Eval:
|
||||
"""Transform OpenAI response to Eval object"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming get eval response: %s", response_json)
|
||||
|
||||
return Eval(**response_json)
|
||||
return Eval.model_validate(response_json)
|
||||
|
||||
def transform_update_eval_request(
|
||||
self,
|
||||
|
|
@ -175,7 +179,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict, Dict]:
|
||||
) -> tuple[str, dict, dict]:
|
||||
"""Transform update eval request for OpenAI"""
|
||||
url = self.get_complete_url(api_base=api_base, endpoint="evals", eval_id=eval_id)
|
||||
|
||||
|
|
@ -192,10 +196,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Eval:
|
||||
"""Transform OpenAI response to Eval object"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming update eval response: %s", response_json)
|
||||
|
||||
return Eval(**response_json)
|
||||
return Eval.model_validate(response_json)
|
||||
|
||||
def transform_delete_eval_request(
|
||||
self,
|
||||
|
|
@ -203,7 +207,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform delete eval request for OpenAI"""
|
||||
url = self.get_complete_url(api_base=api_base, endpoint="evals", eval_id=eval_id)
|
||||
|
||||
|
|
@ -217,10 +221,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> DeleteEvalResponse:
|
||||
"""Transform OpenAI response to DeleteEvalResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming delete eval response: %s", response_json)
|
||||
|
||||
return DeleteEvalResponse(**response_json)
|
||||
return DeleteEvalResponse.model_validate(response_json)
|
||||
|
||||
def transform_cancel_eval_request(
|
||||
self,
|
||||
|
|
@ -228,12 +232,12 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict, Dict]:
|
||||
) -> tuple[str, dict, dict]:
|
||||
"""Transform cancel eval request for OpenAI"""
|
||||
url = f"{self.get_complete_url(api_base=api_base, endpoint='evals', eval_id=eval_id)}/cancel"
|
||||
|
||||
# Empty body for cancel request
|
||||
request_body: Dict[str, Any] = {}
|
||||
request_body: dict[str, object] = {}
|
||||
|
||||
verbose_logger.debug("Cancel eval request - URL: %s", url)
|
||||
|
||||
|
|
@ -245,10 +249,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> CancelEvalResponse:
|
||||
"""Transform OpenAI response to CancelEvalResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming cancel eval response: %s", response_json)
|
||||
|
||||
return CancelEvalResponse(**response_json)
|
||||
return CancelEvalResponse.model_validate(response_json)
|
||||
|
||||
# Run API Transformations
|
||||
def transform_create_run_request(
|
||||
|
|
@ -257,7 +261,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
create_request: CreateRunRequest,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform create run request for OpenAI"""
|
||||
api_base = "https://api.openai.com"
|
||||
if litellm_params and litellm_params.api_base:
|
||||
|
|
@ -279,10 +283,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Run:
|
||||
"""Transform OpenAI response to Run object"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming create run response: %s", response_json)
|
||||
|
||||
return Run(**response_json)
|
||||
return Run.model_validate(response_json)
|
||||
|
||||
def transform_list_runs_request(
|
||||
self,
|
||||
|
|
@ -290,7 +294,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
list_params: ListRunsParams,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform list runs request for OpenAI"""
|
||||
api_base = "https://api.openai.com"
|
||||
if litellm_params and litellm_params.api_base:
|
||||
|
|
@ -300,7 +304,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
url = f"{api_base}/v1/evals/{encoded_eval_id}/runs"
|
||||
|
||||
# Build query parameters
|
||||
query_params: Dict[str, Any] = {}
|
||||
query_params: dict[str, object] = {}
|
||||
if "limit" in list_params and list_params["limit"]:
|
||||
query_params["limit"] = list_params["limit"]
|
||||
if "after" in list_params and list_params["after"]:
|
||||
|
|
@ -323,10 +327,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ListRunsResponse:
|
||||
"""Transform OpenAI response to ListRunsResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming list runs response: %s", response_json)
|
||||
|
||||
return ListRunsResponse(**response_json)
|
||||
return ListRunsResponse.model_validate(response_json)
|
||||
|
||||
def transform_get_run_request(
|
||||
self,
|
||||
|
|
@ -335,7 +339,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform get run request for OpenAI"""
|
||||
encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id")
|
||||
encoded_run_id = encode_url_path_segment(run_id, field_name="run_id")
|
||||
|
|
@ -351,10 +355,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Run:
|
||||
"""Transform OpenAI response to Run object"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming get run response: %s", response_json)
|
||||
|
||||
return Run(**response_json)
|
||||
return Run.model_validate(response_json)
|
||||
|
||||
def transform_cancel_run_request(
|
||||
self,
|
||||
|
|
@ -363,14 +367,14 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict, Dict]:
|
||||
) -> tuple[str, dict, dict]:
|
||||
"""Transform cancel run request for OpenAI"""
|
||||
encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id")
|
||||
encoded_run_id = encode_url_path_segment(run_id, field_name="run_id")
|
||||
url = f"{api_base}/v1/evals/{encoded_eval_id}/runs/{encoded_run_id}/cancel"
|
||||
|
||||
# Empty body for cancel request
|
||||
request_body: Dict[str, Any] = {}
|
||||
request_body: dict[str, object] = {}
|
||||
|
||||
verbose_logger.debug("Cancel run request - URL: %s", url)
|
||||
|
||||
|
|
@ -382,10 +386,10 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> CancelRunResponse:
|
||||
"""Transform OpenAI response to CancelRunResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming cancel run response: %s", response_json)
|
||||
|
||||
return CancelRunResponse(**response_json)
|
||||
return CancelRunResponse.model_validate(response_json)
|
||||
|
||||
def transform_delete_run_request(
|
||||
self,
|
||||
|
|
@ -394,14 +398,14 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict, Dict]:
|
||||
) -> tuple[str, dict, dict]:
|
||||
"""Transform delete run request for OpenAI"""
|
||||
encoded_eval_id = encode_url_path_segment(eval_id, field_name="eval_id")
|
||||
encoded_run_id = encode_url_path_segment(run_id, field_name="run_id")
|
||||
url = f"{api_base}/v1/evals/{encoded_eval_id}/runs/{encoded_run_id}"
|
||||
|
||||
# Empty body for delete request
|
||||
request_body: Dict[str, Any] = {}
|
||||
request_body: dict[str, object] = {}
|
||||
|
||||
verbose_logger.debug("Delete run request - URL: %s", url)
|
||||
|
||||
|
|
@ -413,7 +417,7 @@ class OpenAIEvalsConfig(BaseEvalsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> RunDeleteResponse:
|
||||
"""Transform OpenAI response to RunDeleteResponse"""
|
||||
response_json = raw_response.json()
|
||||
response_json = _parsed_response_json(raw_response)
|
||||
verbose_logger.debug("Transforming delete run response: %s", response_json)
|
||||
|
||||
return RunDeleteResponse(**response_json)
|
||||
return RunDeleteResponse.model_validate(response_json)
|
||||
|
|
|
|||
|
|
@ -1626,7 +1626,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
openai_client: AsyncOpenAI,
|
||||
) -> OpenAIFileObject:
|
||||
response = await openai_client.files.create(**create_file_data) # type: ignore[arg-type]
|
||||
return OpenAIFileObject(**response.model_dump())
|
||||
return OpenAIFileObject.model_validate(response.model_dump())
|
||||
|
||||
def create_file(
|
||||
self,
|
||||
|
|
@ -1662,7 +1662,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
create_file_data=create_file_data, openai_client=openai_client
|
||||
)
|
||||
response = cast(OpenAI, openai_client).files.create(**create_file_data) # type: ignore[arg-type]
|
||||
return OpenAIFileObject(**response.model_dump())
|
||||
return OpenAIFileObject.model_validate(response.model_dump())
|
||||
|
||||
async def afile_content(
|
||||
self,
|
||||
|
|
@ -1986,7 +1986,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
openai_client: AsyncOpenAI,
|
||||
) -> LiteLLMBatch:
|
||||
response = await openai_client.batches.create(**create_batch_data) # type: ignore[arg-type]
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
def create_batch(
|
||||
self,
|
||||
|
|
@ -2023,7 +2023,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
)
|
||||
response = cast(OpenAI, openai_client).batches.create(**create_batch_data) # type: ignore[arg-type]
|
||||
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
async def aretrieve_batch(
|
||||
self,
|
||||
|
|
@ -2032,7 +2032,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
) -> LiteLLMBatch:
|
||||
verbose_logger.debug("retrieving batch, args= %s", retrieve_batch_data)
|
||||
response = await openai_client.batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type]
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
def retrieve_batch(
|
||||
self,
|
||||
|
|
@ -2068,7 +2068,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
retrieve_batch_data=retrieve_batch_data, openai_client=openai_client
|
||||
)
|
||||
response = cast(OpenAI, openai_client).batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type]
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
async def acancel_batch(
|
||||
self,
|
||||
|
|
@ -2077,7 +2077,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
) -> LiteLLMBatch:
|
||||
verbose_logger.debug("async cancelling batch, args= %s", cancel_batch_data)
|
||||
response = await openai_client.batches.cancel(**cancel_batch_data)
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
def cancel_batch(
|
||||
self,
|
||||
|
|
@ -2117,7 +2117,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
if not isinstance(openai_client, OpenAI):
|
||||
raise ValueError("OpenAI client is not an instance of OpenAI. Make sure you passed a sync OpenAI client.")
|
||||
response = openai_client.batches.cancel(**cancel_batch_data)
|
||||
return LiteLLMBatch(**response.model_dump())
|
||||
return LiteLLMBatch.model_validate(response.model_dump())
|
||||
|
||||
async def alist_batches(
|
||||
self,
|
||||
|
|
@ -2477,9 +2477,9 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
response_obj: Optional[OpenAIMessage] = None
|
||||
if getattr(thread_message, "status", None) is None:
|
||||
thread_message.status = "completed"
|
||||
response_obj = OpenAIMessage(**thread_message.dict())
|
||||
response_obj = OpenAIMessage.model_validate(thread_message.dict())
|
||||
else:
|
||||
response_obj = OpenAIMessage(**thread_message.dict())
|
||||
response_obj = OpenAIMessage.model_validate(thread_message.dict())
|
||||
return response_obj
|
||||
|
||||
# fmt: off
|
||||
|
|
@ -2556,9 +2556,9 @@ class OpenAIAssistantsAPI(BaseLLM):
|
|||
response_obj: Optional[OpenAIMessage] = None
|
||||
if getattr(thread_message, "status", None) is None:
|
||||
thread_message.status = "completed"
|
||||
response_obj = OpenAIMessage(**thread_message.dict())
|
||||
response_obj = OpenAIMessage.model_validate(thread_message.dict())
|
||||
else:
|
||||
response_obj = OpenAIMessage(**thread_message.dict())
|
||||
response_obj = OpenAIMessage.model_validate(thread_message.dict())
|
||||
return response_obj
|
||||
|
||||
async def async_get_messages(
|
||||
|
|
|
|||
|
|
@ -280,7 +280,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
raw_response_headers = dict(raw_response.headers)
|
||||
processed_headers = process_response_headers(raw_response_headers)
|
||||
try:
|
||||
response = ResponsesAPIResponse(**raw_response_json)
|
||||
response = ResponsesAPIResponse.model_validate(raw_response_json)
|
||||
except Exception:
|
||||
verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct")
|
||||
response = ResponsesAPIResponse.model_construct(**raw_response_json)
|
||||
|
|
@ -506,7 +506,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code)
|
||||
raw_response_headers = dict(raw_response.headers)
|
||||
processed_headers = process_response_headers(raw_response_headers)
|
||||
response = ResponsesAPIResponse(**raw_response_json)
|
||||
response = ResponsesAPIResponse.model_validate(raw_response_json)
|
||||
response._hidden_params["additional_headers"] = processed_headers
|
||||
response._hidden_params["headers"] = raw_response_headers
|
||||
|
||||
|
|
@ -588,7 +588,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
raw_response_headers = dict(raw_response.headers)
|
||||
processed_headers = process_response_headers(raw_response_headers)
|
||||
|
||||
response = ResponsesAPIResponse(**raw_response_json)
|
||||
response = ResponsesAPIResponse.model_validate(raw_response_json)
|
||||
response._hidden_params["additional_headers"] = processed_headers
|
||||
response._hidden_params["headers"] = raw_response_headers
|
||||
|
||||
|
|
@ -647,7 +647,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
processed_headers = process_response_headers(raw_response_headers)
|
||||
|
||||
try:
|
||||
response = ResponsesAPIResponse(**raw_response_json)
|
||||
response = ResponsesAPIResponse.model_validate(raw_response_json)
|
||||
except Exception:
|
||||
verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct")
|
||||
response = ResponsesAPIResponse.model_construct(**raw_response_json)
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
|
|||
"""
|
||||
|
||||
import re
|
||||
from typing import List, Optional, Tuple, Literal
|
||||
from typing import List, Optional, Sequence, Tuple, Literal
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.llms.vertex_ai import CachedContentRequestBody
|
||||
|
|
@ -152,6 +152,20 @@ def separate_cached_messages(
|
|||
return cached_messages, non_cached_messages
|
||||
|
||||
|
||||
def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageValues]) -> bool:
|
||||
"""
|
||||
The cachedContents API rejects contents ending on a model turn, which is how it
|
||||
classifies both assistant messages and tool results, with HTTP 400
|
||||
"Requests ending with a model turn are not supported". System messages are
|
||||
extracted into system_instruction before contents are built, so the terminal
|
||||
turn is the last non-system message.
|
||||
"""
|
||||
non_system_messages = tuple(message for message in cached_messages if message.get("role") != "system")
|
||||
if not non_system_messages:
|
||||
return bool(cached_messages)
|
||||
return non_system_messages[-1].get("role") not in ("assistant", "tool", "function")
|
||||
|
||||
|
||||
def transform_openai_messages_to_gemini_context_caching(
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ from litellm.types.llms.vertex_ai import (
|
|||
from ..common_utils import VertexAIError, get_vertex_base_url
|
||||
from ..vertex_llm_base import VertexBase
|
||||
from .transformation import (
|
||||
cached_messages_end_on_supported_turn,
|
||||
separate_cached_messages,
|
||||
transform_openai_messages_to_gemini_context_caching,
|
||||
)
|
||||
|
|
@ -308,6 +309,14 @@ class ContextCachingEndpoints(VertexBase):
|
|||
if len(cached_messages) == 0:
|
||||
return messages, optional_params, None
|
||||
|
||||
if not cached_messages_end_on_supported_turn(cached_messages):
|
||||
verbose_logger.debug(
|
||||
"Vertex AI context caching: cached message block ends on a model turn once "
|
||||
"system messages are extracted, which the cachedContents API rejects. "
|
||||
"Skipping context caching."
|
||||
)
|
||||
return messages, optional_params, None
|
||||
|
||||
# Gemini requires a minimum of 1024 tokens for context caching.
|
||||
# Skip caching if the cached content is too small to avoid API errors.
|
||||
if not is_prompt_caching_valid_prompt(
|
||||
|
|
@ -459,6 +468,14 @@ class ContextCachingEndpoints(VertexBase):
|
|||
if len(cached_messages) == 0:
|
||||
return messages, optional_params, None
|
||||
|
||||
if not cached_messages_end_on_supported_turn(cached_messages):
|
||||
verbose_logger.debug(
|
||||
"Vertex AI context caching: cached message block ends on a model turn once "
|
||||
"system messages are extracted, which the cachedContents API rejects. "
|
||||
"Skipping context caching."
|
||||
)
|
||||
return messages, optional_params, None
|
||||
|
||||
# Gemini requires a minimum of 1024 tokens for context caching.
|
||||
# Skip caching if the cached content is too small to avoid API errors.
|
||||
if not is_prompt_caching_valid_prompt(
|
||||
|
|
|
|||
|
|
@ -1,7 +1,9 @@
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from urllib.parse import unquote
|
||||
from typing import Any, Coroutine, Optional, Tuple, Union
|
||||
from typing import Any, Coroutine, Mapping, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -10,6 +12,7 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import (
|
|||
GCSBucketBase,
|
||||
GCSLoggingConfig,
|
||||
)
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
VERTEX_AI_MANAGED_GCS_PREFIX,
|
||||
should_allow_legacy_cloud_file_ids,
|
||||
|
|
@ -39,6 +42,35 @@ class VertexAIFilesHandler(GCSBucketBase):
|
|||
llm_provider=LlmProviders.VERTEX_AI,
|
||||
)
|
||||
|
||||
def _resolve_read_gcs_config(
|
||||
self,
|
||||
litellm_params: Mapping[str, object] | None,
|
||||
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""
|
||||
Resolve the GCS bucket and service-account credentials for the read/content path.
|
||||
|
||||
Sources them from the deployment's ``litellm_params`` (``gcs_bucket_name`` /
|
||||
``bucket_name`` and ``vertex_credentials``), mirroring the write path in
|
||||
``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the global
|
||||
``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` env vars. This lets Vertex batch
|
||||
run entirely at the model-group level, so output written to a per-model bucket is
|
||||
readable without setting the global env vars.
|
||||
"""
|
||||
params: Mapping[str, object] = litellm_params or {}
|
||||
bucket_candidate = params.get("gcs_bucket_name") or params.get("bucket_name")
|
||||
configured_bucket_name = bucket_candidate if isinstance(bucket_candidate, str) else os.getenv("GCS_BUCKET_NAME")
|
||||
|
||||
credentials = params.get("vertex_credentials") or vertex_credentials
|
||||
if isinstance(credentials, dict):
|
||||
path_service_account: str | None = json.dumps(credentials)
|
||||
elif isinstance(credentials, str):
|
||||
path_service_account = credentials
|
||||
else:
|
||||
path_service_account = os.getenv("GCS_PATH_SERVICE_ACCOUNT")
|
||||
|
||||
return configured_bucket_name, path_service_account
|
||||
|
||||
def _extract_bucket_and_object_from_file_id(
|
||||
self,
|
||||
file_id: str,
|
||||
|
|
@ -91,7 +123,17 @@ class VertexAIFilesHandler(GCSBucketBase):
|
|||
if not file_id:
|
||||
raise ValueError("file_id is required in file_content_request")
|
||||
|
||||
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(kwargs={})
|
||||
configured_bucket_name, path_service_account = self._resolve_read_gcs_config(
|
||||
litellm_params=litellm_params,
|
||||
vertex_credentials=vertex_credentials,
|
||||
)
|
||||
dynamic_params = StandardCallbackDynamicParams(
|
||||
gcs_bucket_name=configured_bucket_name,
|
||||
gcs_path_service_account=path_service_account,
|
||||
)
|
||||
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(
|
||||
kwargs={"standard_callback_dynamic_params": dynamic_params}
|
||||
)
|
||||
bucket_name, object_path = self._extract_bucket_and_object_from_file_id(
|
||||
file_id=file_id,
|
||||
configured_bucket_name=gcs_logging_config["bucket_name"],
|
||||
|
|
|
|||
|
|
@ -661,6 +661,10 @@ def _gemini_convert_messages_with_history(
|
|||
vertex_project = litellm_params.get("vertex_project") or litellm_params.get("vertex_ai_project")
|
||||
vertex_credentials = litellm_params.get("vertex_credentials") or litellm_params.get("vertex_ai_credentials")
|
||||
|
||||
from .vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
|
||||
forward_function_call_id = VertexGeminiConfig._forward_gemini_function_call_id(model or "")
|
||||
|
||||
try:
|
||||
while msg_i < len(messages):
|
||||
user_content: List[PartType] = []
|
||||
|
|
@ -910,7 +914,7 @@ def _gemini_convert_messages_with_history(
|
|||
gemini_tool_call_parts = convert_to_gemini_tool_call_invoke(
|
||||
assistant_msg,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
forward_function_call_id=forward_function_call_id,
|
||||
)
|
||||
## check if gemini_tool_call already exists in assistant_content
|
||||
for gemini_tool_call_part in gemini_tool_call_parts:
|
||||
|
|
@ -973,8 +977,7 @@ def _gemini_convert_messages_with_history(
|
|||
_part = convert_to_gemini_tool_call_result(
|
||||
messages[msg_i], # type: ignore
|
||||
last_message_with_tool_calls, # type: ignore
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
forward_function_call_id=forward_function_call_id,
|
||||
)
|
||||
msg_i += 1
|
||||
# Handle both single part and list of parts (for Computer Use with images)
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue