mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge branch 'litellm_internal_staging' into feature/improve-gigachat-provider
This commit is contained in:
commit
8599009b28
167 changed files with 9599 additions and 2287 deletions
2
.github/pull_request_template.md
vendored
2
.github/pull_request_template.md
vendored
|
|
@ -4,7 +4,7 @@
|
|||
|
||||
## Linear ticket
|
||||
|
||||
<!-- if you are an internal contributor, add the Linear ticket e.g. "Resolves LIT-1234" to magically link the Linear ticket to the GitHub PR -->
|
||||
<!-- if you are an internal contributor (e.g., your username is postfixed with -berri or -berriai), add "Resolves " followed by the Linear ticket e.g. "Resolves LIT-1234" to magically link the Linear ticket to the GitHub PR -->
|
||||
|
||||
## Pre-Submission checklist
|
||||
|
||||
|
|
|
|||
2
.github/workflows/check-ui-api-types.yml
vendored
2
.github/workflows/check-ui-api-types.yml
vendored
|
|
@ -54,7 +54,7 @@ jobs:
|
|||
run: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
|
||||
with:
|
||||
node-version: "20"
|
||||
cache: "npm"
|
||||
|
|
|
|||
6
.github/workflows/codeql.yml
vendored
6
.github/workflows/codeql.yml
vendored
|
|
@ -43,14 +43,14 @@ jobs:
|
|||
persist-credentials: false
|
||||
|
||||
- name: Initialize CodeQL
|
||||
uses: github/codeql-action/init@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3
|
||||
uses: github/codeql-action/init@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1
|
||||
with:
|
||||
languages: ${{ matrix.language }}
|
||||
build-mode: ${{ matrix.build-mode }}
|
||||
config-file: ./.github/codeql/codeql-config.yml
|
||||
|
||||
- name: Perform CodeQL Analysis
|
||||
uses: github/codeql-action/analyze@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3
|
||||
uses: github/codeql-action/analyze@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1
|
||||
with:
|
||||
category: "/language:${{ matrix.language }}"
|
||||
output: sarif-results
|
||||
|
|
@ -77,7 +77,7 @@ jobs:
|
|||
output: sarif-results/python.sarif
|
||||
|
||||
- name: Upload SARIF
|
||||
uses: github/codeql-action/upload-sarif@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3
|
||||
uses: github/codeql-action/upload-sarif@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1
|
||||
with:
|
||||
sarif_file: sarif-results
|
||||
category: "/language:${{ matrix.language }}"
|
||||
|
|
|
|||
131
.github/workflows/review_gate.yml
vendored
131
.github/workflows/review_gate.yml
vendored
|
|
@ -1,131 +0,0 @@
|
|||
name: Agent Shin — review gate
|
||||
|
||||
# Keeps the `ready for review` label in sync with whether an external PR
|
||||
# currently clears BOTH the LLM rubric AND Greptile's confidence score.
|
||||
#
|
||||
# pass -> add `ready for review` + a "passed / all clear" comment
|
||||
# regress -> remove the label + a "what's missing" comment (PR stays open)
|
||||
# fail, <24h old -> a one-time "what's missing" notice (grace window)
|
||||
# fail, >24h old -> close + a comment (reopen via `@agent-shin reconsider`)
|
||||
#
|
||||
# DRY-RUN BY DEFAULT. Every side effect (label add/remove, comment, close) is
|
||||
# gated behind `--close`, which is only added when the repo variable
|
||||
# `AGENT_SHIN_ENABLED == "true"`. Until then runs only write the verdict to the
|
||||
# workflow step summary.
|
||||
#
|
||||
# Manual single PR: gh workflow run "Agent Shin — review gate" -f pr_number=NNN
|
||||
# Manual dry-run: gh workflow run "Agent Shin — review gate" -f close=false
|
||||
#
|
||||
# We use `pull_request_target` so the workflow can read repo secrets and run
|
||||
# against fork PRs. Fork code is never checked out — only PR metadata is read
|
||||
# via `gh api`.
|
||||
|
||||
on:
|
||||
pull_request_target:
|
||||
types: [opened, reopened, synchronize, ready_for_review]
|
||||
schedule:
|
||||
# Daily at 09:30 UTC — re-reconciles labels as Greptile re-reviews land.
|
||||
- cron: "30 9 * * *"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
pr_number:
|
||||
description: "Single PR to reconcile (omit to sweep all open PRs)."
|
||||
required: false
|
||||
close:
|
||||
description: "If AGENT_SHIN_ENABLED=true, actually act (false = dry run)."
|
||||
required: false
|
||||
default: "false"
|
||||
type: choice
|
||||
options:
|
||||
- "true"
|
||||
- "false"
|
||||
grace_days:
|
||||
description: "Hours/24 a failing, un-tagged PR may stay open before close."
|
||||
required: false
|
||||
default: "1"
|
||||
min_greptile_score:
|
||||
description: "Greptile score below which a PR counts as not passing (1-5)."
|
||||
required: false
|
||||
default: "4"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
issues: write
|
||||
pull-requests: write
|
||||
|
||||
jobs:
|
||||
review-gate:
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout triage script
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
sparse-checkout: .github/scripts
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Install LLM client
|
||||
run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt
|
||||
|
||||
- name: Run review gate
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
# Mirror the triage workflow: only expose the LLM key when the bot is
|
||||
# enabled or a collaborator triggers it manually, so an external user
|
||||
# can't force paid LLM calls by churning a fork PR while the bot is
|
||||
# still in dry-run.
|
||||
OPENAI_API_KEY: ${{ (vars.AGENT_SHIN_ENABLED == 'true' || github.event_name == 'workflow_dispatch') && secrets.OPENAI_API_KEY || '' }}
|
||||
OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }}
|
||||
TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }}
|
||||
AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }}
|
||||
CLOSE_FLAG: ${{ github.event.inputs.close || 'false' }}
|
||||
GRACE_DAYS: ${{ github.event.inputs.grace_days || '1' }}
|
||||
MIN_GREPTILE_SCORE: ${{ github.event.inputs.min_greptile_score || '4' }}
|
||||
EVENT_PR: ${{ github.event.pull_request.number }}
|
||||
INPUT_PR: ${{ github.event.inputs.pr_number }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
COMMON=(--review-gate --grace-days "${GRACE_DAYS}" --min-greptile-score "${MIN_GREPTILE_SCORE}")
|
||||
|
||||
# Fail-safe gating, identical philosophy to the Greptile closer:
|
||||
# - AGENT_SHIN_ENABLED must be the EXACT string "true" to act at all.
|
||||
# - A manual dispatch can still preview with close=false.
|
||||
# - Automatic triggers (PR events, schedule) act once enabled — that
|
||||
# is the whole point of the gate (re-tag / un-tag automatically).
|
||||
DO_CLOSE="false"
|
||||
if [ "${AGENT_SHIN_ENABLED:-false}" != "true" ]; then
|
||||
echo "::notice::AGENT_SHIN_ENABLED is not 'true' -> dry-run (no labels/comments/closes)."
|
||||
elif [ "${GITHUB_EVENT_NAME:-}" = "workflow_dispatch" ] && [ "${CLOSE_FLAG:-false}" = "true" ]; then
|
||||
DO_CLOSE="true"
|
||||
echo "::notice::Manual run -> acting for real."
|
||||
elif [ "${GITHUB_EVENT_NAME:-}" != "workflow_dispatch" ]; then
|
||||
DO_CLOSE="true"
|
||||
echo "::notice::Enabled automatic trigger (${GITHUB_EVENT_NAME:-}) -> acting for real."
|
||||
else
|
||||
echo "::notice::Manual dispatch with close=false -> dry-run."
|
||||
fi
|
||||
if [ "${DO_CLOSE}" = "true" ]; then
|
||||
COMMON+=(--close)
|
||||
fi
|
||||
|
||||
# Single PR (PR event or explicit input) vs. sweep over all open PRs.
|
||||
TARGET_PR="${EVENT_PR:-${INPUT_PR:-}}"
|
||||
if [ -n "${TARGET_PR}" ]; then
|
||||
python3 .github/scripts/triage_with_llm.py --repo "${{ github.repository }}" --pr "${TARGET_PR}" "${COMMON[@]}"
|
||||
else
|
||||
echo "::notice::Sweeping all open PRs."
|
||||
# Match GH_LIST_ALL_LIMIT in agent_shin_shared.py: gh lists newest-first,
|
||||
# so any cap below the real backlog silently drops the *oldest* PRs —
|
||||
# exactly the stale ones this daily sweep is meant to reconcile.
|
||||
mapfile -t NUMBERS < <(gh pr list --repo "${{ github.repository }}" --state open --limit 100000 --json number --jq '.[].number')
|
||||
for n in "${NUMBERS[@]}"; do
|
||||
echo "::group::PR #${n}"
|
||||
python3 .github/scripts/triage_with_llm.py --repo "${{ github.repository }}" --pr "${n}" "${COMMON[@]}" || echo "::warning::review gate errored on #${n}"
|
||||
echo "::endgroup::"
|
||||
done
|
||||
fi
|
||||
4
.github/workflows/test-litellm-ui-build.yml
vendored
4
.github/workflows/test-litellm-ui-build.yml
vendored
|
|
@ -25,7 +25,7 @@ jobs:
|
|||
persist-credentials: false
|
||||
|
||||
- name: Setup Node.js
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
|
||||
with:
|
||||
node-version: "20"
|
||||
cache: "npm"
|
||||
|
|
@ -77,7 +77,7 @@ jobs:
|
|||
|
||||
- name: Setup Node.js
|
||||
if: steps.changed.outputs.has_files == 'true'
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
|
||||
with:
|
||||
node-version: "20"
|
||||
cache: "npm"
|
||||
|
|
|
|||
10
.github/workflows/test-unit-proxy-endpoints.yml
vendored
10
.github/workflows/test-unit-proxy-endpoints.yml
vendored
|
|
@ -11,8 +11,6 @@ on:
|
|||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -20,6 +18,10 @@ concurrency:
|
|||
|
||||
jobs:
|
||||
proxy-endpoints:
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: >-
|
||||
|
|
@ -52,6 +54,10 @@ jobs:
|
|||
# is independent and its coverage artifact is uploaded separately.
|
||||
# See: https://www.notion.so/36c43b8acdab81ee845fd5365128a2fc
|
||||
proxy-server:
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: tests/test_litellm/proxy/proxy_server
|
||||
|
|
|
|||
7
.github/workflows/test_server_root_path.yml
vendored
7
.github/workflows/test_server_root_path.yml
vendored
|
|
@ -32,17 +32,16 @@ jobs:
|
|||
df -h /
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3.12
|
||||
uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3.12.0
|
||||
|
||||
- name: Build Docker image
|
||||
uses: docker/build-push-action@0adf9959216b96bec444f325f1e493d4aa344497 #v6.14
|
||||
uses: docker/build-push-action@0adf9959216b96bec444f325f1e493d4aa344497 # v6.14.0
|
||||
with:
|
||||
context: .
|
||||
file: ./docker/Dockerfile.non_root
|
||||
tags: litellm-test:${{ github.sha }}
|
||||
load: true
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
push: false
|
||||
|
||||
- name: Start LiteLLM container with SERVER_ROOT_PATH
|
||||
run: |
|
||||
|
|
|
|||
110
.github/workflows/triage_pr_with_llm.yml
vendored
110
.github/workflows/triage_pr_with_llm.yml
vendored
|
|
@ -1,110 +0,0 @@
|
|||
name: Agent Shin — PR triage
|
||||
|
||||
# LLM-as-judge triage for external pull requests.
|
||||
#
|
||||
# DRY-RUN BY DEFAULT. Closures and public comments are gated on the repo
|
||||
# variable `AGENT_SHIN_ENABLED` being set to the string `"true"`. Until then,
|
||||
# every run only writes its verdict to the workflow step summary so the team
|
||||
# can QA the judge's decisions before flipping it on.
|
||||
#
|
||||
# To enable for real:
|
||||
# 1. Add a repo secret `OPENAI_API_KEY` (or compatible).
|
||||
# 2. Set repo variable `AGENT_SHIN_ENABLED` to `true`
|
||||
# (Settings > Secrets and variables > Actions > Variables).
|
||||
#
|
||||
# We use `pull_request_target` so the workflow has access to repo secrets
|
||||
# and runs against PRs from forks. We never check out fork code — only read
|
||||
# PR metadata via `gh api`, so this is safe.
|
||||
|
||||
on:
|
||||
pull_request_target:
|
||||
types: [opened, reopened]
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
pr_number:
|
||||
description: "PR number to triage manually."
|
||||
required: true
|
||||
close:
|
||||
description: "If true and AGENT_SHIN_ENABLED=true, actually close on fail."
|
||||
required: false
|
||||
default: "false"
|
||||
type: choice
|
||||
options:
|
||||
- "true"
|
||||
- "false"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
issues: write
|
||||
pull-requests: write
|
||||
|
||||
jobs:
|
||||
triage:
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout triage script
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
sparse-checkout: .github/scripts
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Install LLM client
|
||||
run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt
|
||||
|
||||
- name: Run Agent Shin
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
# Only expose the LLM key when the bot is enabled or a collaborator
|
||||
# triggers it manually, so an external user can't force paid LLM
|
||||
# calls by churning a fork PR while the bot is still in dry-run.
|
||||
# The Python script calls the LLM whenever this var is set
|
||||
# (regardless of `--close`); stripping `--close` doesn't suppress
|
||||
# the API call, only the destructive side effects.
|
||||
OPENAI_API_KEY: ${{ (vars.AGENT_SHIN_ENABLED == 'true' || github.event_name == 'workflow_dispatch') && secrets.OPENAI_API_KEY || '' }}
|
||||
OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }}
|
||||
TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }}
|
||||
AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }}
|
||||
DISPATCH_CLOSE: ${{ github.event.inputs.close }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number || github.event.inputs.pr_number }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
ARGS=(--repo "${{ github.repository }}" --pr "${PR_NUMBER}")
|
||||
# Fail-safe gating: only the EXACT string "true" enables the
|
||||
# destructive --close path. The workflow_dispatch input is a
|
||||
# `choice` dropdown of "true"/"false" so the UI is constrained,
|
||||
# but the API (`gh workflow run -f close=...`) accepts any
|
||||
# string, and a `!= "false"` check would treat "True", "yes",
|
||||
# "1", "TRUE", typos, and accidental whitespace as enabling
|
||||
# closure. Mirror the Greptile closer's `= "true"` pattern.
|
||||
if [ "${AGENT_SHIN_ENABLED:-false}" = "true" ] && [ "${DISPATCH_CLOSE:-false}" = "true" ]; then
|
||||
ARGS+=(--close)
|
||||
echo "::notice::Agent Shin is ENABLED and running in close-on-fail mode."
|
||||
elif [ "${AGENT_SHIN_ENABLED:-false}" = "true" ]; then
|
||||
echo "::notice::Agent Shin is ENABLED but this trigger is dry-run (workflow_dispatch close != 'true' or scheduled event)."
|
||||
else
|
||||
echo "::notice::Agent Shin is in DRY-RUN mode (AGENT_SHIN_ENABLED is not 'true'). No comments will be posted; no PRs will be closed."
|
||||
fi
|
||||
# On the scheduled/automatic pull_request_target trigger we default to
|
||||
# dry-run regardless, so the team can review verdicts in the step
|
||||
# summary before any contributor sees a comment. Only the manual
|
||||
# workflow_dispatch path (with close=true) closes PRs.
|
||||
if [ "${GITHUB_EVENT_NAME:-}" = "pull_request_target" ]; then
|
||||
# strip any --close added above (filter out, don't substitute
|
||||
# to empty string — that would leave a stray "" positional arg
|
||||
# that argparse rejects)
|
||||
FILTERED=()
|
||||
for arg in "${ARGS[@]}"; do
|
||||
if [ "${arg}" != "--close" ]; then
|
||||
FILTERED+=("${arg}")
|
||||
fi
|
||||
done
|
||||
ARGS=("${FILTERED[@]}")
|
||||
echo "::notice::pull_request_target trigger -> forcing dry-run."
|
||||
fi
|
||||
python3 .github/scripts/triage_with_llm.py "${ARGS[@]}"
|
||||
13
.github/workflows/zizmor.yml
vendored
13
.github/workflows/zizmor.yml
vendored
|
|
@ -2,9 +2,9 @@ name: GitHub Actions Security Analysis
|
|||
|
||||
on:
|
||||
push:
|
||||
branches: [main]
|
||||
branches: [main, litellm_internal_staging]
|
||||
pull_request:
|
||||
branches: [main]
|
||||
branches: [main, litellm_internal_staging]
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -18,9 +18,7 @@ jobs:
|
|||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
permissions:
|
||||
security-events: write
|
||||
contents: read
|
||||
actions: read
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
|
|
@ -28,4 +26,9 @@ jobs:
|
|||
persist-credentials: false
|
||||
|
||||
- name: Run zizmor
|
||||
uses: zizmorcore/zizmor-action@71321a20a9ded102f6e9ce5718a2fcec2c4f70d8 # v0.5.2
|
||||
uses: zizmorcore/zizmor-action@5f14fd08f7cf1cb1609c1e344975f152c7ee938d # v0.5.6
|
||||
with:
|
||||
version: "1.24.1"
|
||||
min-severity: medium
|
||||
advanced-security: false
|
||||
annotations: true
|
||||
|
|
|
|||
|
|
@ -345,6 +345,7 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
|
|||
| [OVHCloud AI Endpoints (`ovhcloud`)](https://docs.litellm.ai/docs/providers/ovhcloud) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Perplexity AI (`perplexity`)](https://docs.litellm.ai/docs/providers/perplexity) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Petals (`petals`)](https://docs.litellm.ai/docs/providers/petals) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Pinstripes (`pinstripes`)](https://docs.litellm.ai/docs/providers/pinstripes) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Predibase (`predibase`)](https://docs.litellm.ai/docs/providers/predibase) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Recraft (`recraft`)](https://docs.litellm.ai/docs/providers/recraft) | | | | | ✅ | | | | | |
|
||||
| [Replicate (`replicate`)](https://docs.litellm.ai/docs/providers/replicate) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
|
|
|
|||
|
|
@ -1,31 +1,31 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"baseline": 24954,
|
||||
"baseline": 24989,
|
||||
"slack": 2500
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"baseline": 1863,
|
||||
"baseline": 1934,
|
||||
"slack": 180
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"baseline": 220,
|
||||
"slack": 3
|
||||
"slack": 22
|
||||
},
|
||||
"reportAttributeAccessIssue": {
|
||||
"baseline": 335,
|
||||
"slack": 3
|
||||
"baseline": 346,
|
||||
"slack": 35
|
||||
},
|
||||
"reportCallIssue": {
|
||||
"baseline": 77,
|
||||
"baseline": 87,
|
||||
"slack": 10
|
||||
},
|
||||
"reportConstantRedefinition": {
|
||||
"baseline": 39,
|
||||
"slack": 3
|
||||
"slack": 4
|
||||
},
|
||||
"reportDeprecated": {
|
||||
"baseline": 217,
|
||||
"slack": 10
|
||||
"slack": 22
|
||||
},
|
||||
"reportDuplicateImport": {
|
||||
"baseline": 28,
|
||||
|
|
@ -41,11 +41,11 @@
|
|||
},
|
||||
"reportGeneralTypeIssues": {
|
||||
"baseline": 151,
|
||||
"slack": 3
|
||||
"slack": 15
|
||||
},
|
||||
"reportIncompatibleMethodOverride": {
|
||||
"baseline": 52,
|
||||
"slack": 10
|
||||
"slack": 5
|
||||
},
|
||||
"reportIncompatibleVariableOverride": {
|
||||
"baseline": 8,
|
||||
|
|
@ -73,7 +73,7 @@
|
|||
},
|
||||
"reportMissingParameterType": {
|
||||
"baseline": 3933,
|
||||
"slack": 10
|
||||
"slack": 390
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"baseline": 10612,
|
||||
|
|
@ -97,7 +97,7 @@
|
|||
},
|
||||
"reportOptionalMemberAccess": {
|
||||
"baseline": 724,
|
||||
"slack": 10
|
||||
"slack": 72
|
||||
},
|
||||
"reportOptionalOperand": {
|
||||
"baseline": 3,
|
||||
|
|
@ -120,8 +120,8 @@
|
|||
"slack": 3
|
||||
},
|
||||
"reportReturnType": {
|
||||
"baseline": 118,
|
||||
"slack": 10
|
||||
"baseline": 126,
|
||||
"slack": 13
|
||||
},
|
||||
"reportTypedDictNotRequiredAccess": {
|
||||
"baseline": 20,
|
||||
|
|
@ -136,19 +136,19 @@
|
|||
"slack": 3000
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"baseline": 76,
|
||||
"baseline": 75,
|
||||
"slack": 10
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"baseline": 27322,
|
||||
"baseline": 27037,
|
||||
"slack": 2500
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"baseline": 13636,
|
||||
"baseline": 13612,
|
||||
"slack": 1000
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"baseline": 21776,
|
||||
"baseline": 21445,
|
||||
"slack": 2000
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
|
|
@ -156,7 +156,7 @@
|
|||
"slack": 10
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"baseline": 680,
|
||||
"baseline": 683,
|
||||
"slack": 10
|
||||
},
|
||||
"reportUnnecessaryContains": {
|
||||
|
|
@ -164,12 +164,12 @@
|
|||
"slack": 3
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"baseline": 807,
|
||||
"slack": 10
|
||||
"baseline": 808,
|
||||
"slack": 80
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"baseline": 110,
|
||||
"slack": 3
|
||||
"slack": 11
|
||||
},
|
||||
"reportUntypedFunctionDecorator": {
|
||||
"baseline": 22,
|
||||
|
|
@ -185,10 +185,10 @@
|
|||
},
|
||||
"reportUnusedImport": {
|
||||
"baseline": 670,
|
||||
"slack": 10
|
||||
"slack": 50
|
||||
},
|
||||
"reportUnusedVariable": {
|
||||
"baseline": 865,
|
||||
"slack": 10
|
||||
"slack": 50
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
# What is this?
|
||||
## This hook is used to check for LiteLLM managed files in the request body, and replace them with model-specific file id
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -1472,8 +1472,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
error_message += f" (showing {MAX_BATCHES_IN_ERROR} most recent): {', '.join(batch_statuses)}. "
|
||||
|
||||
error_message += (
|
||||
f"To delete this file before complete cost tracking, please delete or cancel the referencing batch(es) first. "
|
||||
f"Alternatively, wait for all batches to complete and for cost to be computed (batch_processed=true)."
|
||||
"To delete this file before complete cost tracking, please delete or cancel the referencing batch(es) first. "
|
||||
"Alternatively, wait for all batches to complete and for cost to be computed (batch_processed=true)."
|
||||
)
|
||||
|
||||
# Record blocked deletion metric
|
||||
|
|
@ -1550,9 +1550,22 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
|
||||
if specific_model_file_id_mapping:
|
||||
exception_dict = {}
|
||||
for model_id, file_id in specific_model_file_id_mapping.items():
|
||||
for model_id, provider_file_id in specific_model_file_id_mapping.items():
|
||||
try:
|
||||
return await llm_router.afile_content(model=model_id, file_id=file_id, **data) # type: ignore
|
||||
# Cloud-storage providers (e.g. Bedrock S3) validate file ids
|
||||
# against the deployment's configured bucket, which they only
|
||||
# trust from this immutable server-side snapshot, never from
|
||||
# request params.
|
||||
credentials = llm_router.get_deployment_credentials_with_provider(
|
||||
model_id=model_id
|
||||
)
|
||||
if credentials is not None:
|
||||
data["_litellm_internal_model_credentials"] = cast(
|
||||
Dict, MappingProxyType(dict(credentials))
|
||||
)
|
||||
else:
|
||||
data.pop("_litellm_internal_model_credentials", None)
|
||||
return await llm_router.afile_content(model=model_id, file_id=provider_file_id, **data) # type: ignore
|
||||
except Exception as e:
|
||||
exception_dict[model_id] = str(e)
|
||||
raise Exception(
|
||||
|
|
|
|||
|
|
@ -100,6 +100,8 @@ class Cache:
|
|||
gcs_path: Optional[str] = None,
|
||||
redis_semantic_cache_embedding_model: str = "text-embedding-ada-002",
|
||||
redis_semantic_cache_index_name: Optional[str] = None,
|
||||
valkey_semantic_cache_embedding_model: str = "text-embedding-ada-002",
|
||||
valkey_semantic_cache_index_name: str | None = None,
|
||||
redis_flush_size: Optional[int] = None,
|
||||
redis_startup_nodes: Optional[List] = None,
|
||||
disk_cache_dir: Optional[str] = None,
|
||||
|
|
@ -208,6 +210,21 @@ class Cache:
|
|||
index_name=redis_semantic_cache_index_name,
|
||||
**kwargs,
|
||||
)
|
||||
elif type == LiteLLMCacheType.VALKEY_SEMANTIC:
|
||||
# Imported here, not at module top, so the optional redis dependency
|
||||
# is only required when this backend is actually selected.
|
||||
from .valkey_semantic_cache import ValkeySemanticCache
|
||||
|
||||
self.cache = ValkeySemanticCache(
|
||||
host=host,
|
||||
port=port,
|
||||
password=password,
|
||||
similarity_threshold=similarity_threshold,
|
||||
embedding_model=valkey_semantic_cache_embedding_model,
|
||||
index_name=valkey_semantic_cache_index_name,
|
||||
startup_nodes=redis_startup_nodes,
|
||||
**kwargs,
|
||||
)
|
||||
elif type == LiteLLMCacheType.QDRANT_SEMANTIC:
|
||||
self.cache = QdrantSemanticCache(
|
||||
qdrant_api_base=qdrant_api_base,
|
||||
|
|
@ -267,12 +284,50 @@ class Cache:
|
|||
if (
|
||||
self.type == LiteLLMCacheType.REDIS
|
||||
or self.type == LiteLLMCacheType.REDIS_SEMANTIC
|
||||
or self.type == LiteLLMCacheType.VALKEY_SEMANTIC
|
||||
) and default_in_redis_ttl is not None:
|
||||
self.ttl = default_in_redis_ttl
|
||||
|
||||
if self.namespace is not None and isinstance(self.cache, RedisCache):
|
||||
self.cache.namespace = self.namespace
|
||||
|
||||
# Params whose values carry prompt content. Excluded from semantic-cache
|
||||
# scope keys so differently worded prompts share a bucket and match via
|
||||
# vector similarity rather than being split into per-wording buckets.
|
||||
_SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS: frozenset = frozenset(
|
||||
{"messages", "prompt", "input"}
|
||||
)
|
||||
|
||||
# Server-set identity (from proxy auth) used to isolate semantic-cache
|
||||
# buckets per tenant. Required once the prompt is out of the scope key, so a
|
||||
# similar prompt from another key/team/org stays in a separate bucket.
|
||||
_SEMANTIC_CACHE_TENANT_SCOPE_FIELDS: tuple[str, ...] = (
|
||||
"user_api_key",
|
||||
"user_api_key_team_id",
|
||||
"user_api_key_org_id",
|
||||
)
|
||||
|
||||
def _is_semantic_cache(self) -> bool:
|
||||
return self.type in (
|
||||
LiteLLMCacheType.REDIS_SEMANTIC,
|
||||
LiteLLMCacheType.QDRANT_SEMANTIC,
|
||||
LiteLLMCacheType.VALKEY_SEMANTIC,
|
||||
)
|
||||
|
||||
def _get_semantic_cache_tenant_scope(self, kwargs: dict) -> str:
|
||||
metadata: dict = kwargs.get("metadata") or {}
|
||||
litellm_params: dict = kwargs.get("litellm_params") or {}
|
||||
metadata_in_litellm_params: dict = litellm_params.get("metadata") or {}
|
||||
|
||||
scope = ""
|
||||
for field in self._SEMANTIC_CACHE_TENANT_SCOPE_FIELDS:
|
||||
value = metadata.get(field)
|
||||
if value is None:
|
||||
value = metadata_in_litellm_params.get(field)
|
||||
if value is not None:
|
||||
scope += f"{field}: {value}"
|
||||
return scope
|
||||
|
||||
def get_cache_key(self, **kwargs) -> str:
|
||||
"""
|
||||
Get the cache key for the given arguments.
|
||||
|
|
@ -293,7 +348,15 @@ class Cache:
|
|||
|
||||
combined_kwargs = ModelParamHelper._get_all_llm_api_params()
|
||||
litellm_param_kwargs = all_litellm_params
|
||||
is_semantic_cache = self._is_semantic_cache()
|
||||
scope_excluded_params = (
|
||||
self._SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS
|
||||
if is_semantic_cache
|
||||
else frozenset()
|
||||
)
|
||||
for param in kwargs:
|
||||
if param in scope_excluded_params:
|
||||
continue
|
||||
if param in combined_kwargs:
|
||||
param_value: Optional[str] = self._get_param_value(param, kwargs)
|
||||
if param_value is not None:
|
||||
|
|
@ -309,6 +372,9 @@ class Cache:
|
|||
param_value = kwargs[param]
|
||||
cache_key += f"{str(param)}: {str(param_value)}"
|
||||
|
||||
if is_semantic_cache:
|
||||
cache_key += self._get_semantic_cache_tenant_scope(kwargs)
|
||||
|
||||
hashed_cache_key = Cache._get_hashed_cache_key(cache_key)
|
||||
hashed_cache_key = self._add_namespace_to_cache_key(hashed_cache_key, **kwargs)
|
||||
verbose_logger.debug(
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Supports syncing responses to Google Cloud Storage Buckets using HTTP requests.
|
|||
import json
|
||||
import asyncio
|
||||
from typing import Optional
|
||||
from urllib.parse import quote
|
||||
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase
|
||||
|
|
@ -48,7 +49,7 @@ class GCSCache(BaseCache):
|
|||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}"
|
||||
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={quote(object_name, safe='')}"
|
||||
data = json.dumps(value)
|
||||
self.sync_client.post(url=url, data=data, headers=headers)
|
||||
except Exception as e:
|
||||
|
|
@ -59,7 +60,7 @@ class GCSCache(BaseCache):
|
|||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}"
|
||||
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={quote(object_name, safe='')}"
|
||||
data = json.dumps(value)
|
||||
await self.async_client.post(url=url, data=data, headers=headers)
|
||||
except Exception as e:
|
||||
|
|
@ -72,7 +73,7 @@ class GCSCache(BaseCache):
|
|||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{object_name}?alt=media"
|
||||
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{quote(object_name, safe='')}?alt=media"
|
||||
response = self.sync_client.get(url=url, headers=headers)
|
||||
if response.status_code == 200:
|
||||
cached_response = json.loads(response.text)
|
||||
|
|
@ -91,7 +92,7 @@ class GCSCache(BaseCache):
|
|||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{object_name}?alt=media"
|
||||
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{quote(object_name, safe='')}?alt=media"
|
||||
response = await self.async_client.get(url=url, headers=headers)
|
||||
if response.status_code == 200:
|
||||
return json.loads(response.text)
|
||||
|
|
|
|||
|
|
@ -903,6 +903,43 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
raise e
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_set_max(
|
||||
self,
|
||||
key: str,
|
||||
value: float,
|
||||
ttl: int | None = None,
|
||||
) -> float | None:
|
||||
"""Atomically set ``key`` to ``value`` only when ``value`` is greater
|
||||
than the stored value (or the key is unset), refreshing the TTL.
|
||||
|
||||
Monotonic by construction: it never lowers the stored value, so a repair
|
||||
that writes an authoritative-but-slightly-stale total cannot clobber a
|
||||
concurrent increment that has already pushed the counter higher. The
|
||||
GET/compare/SET runs in a single Lua call, so it is also atomic across
|
||||
racing callers and pods. Returns the resulting value.
|
||||
"""
|
||||
_redis_client = self.init_async_client()
|
||||
_used_ttl = self.get_ttl(ttl=ttl)
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
lua = (
|
||||
"local cur = redis.call('GET', KEYS[1]) "
|
||||
"if cur == false or tonumber(cur) < tonumber(ARGV[1]) then "
|
||||
"redis.call('SET', KEYS[1], ARGV[1]) "
|
||||
"if tonumber(ARGV[2]) > 0 then redis.call('EXPIRE', KEYS[1], ARGV[2]) end "
|
||||
"return ARGV[1] end "
|
||||
"return cur"
|
||||
)
|
||||
result = cast(
|
||||
"str | bytes | int | float | None",
|
||||
await _redis_client.eval(lua, 1, key, str(value), str(int(_used_ttl or 0))),
|
||||
)
|
||||
if result is None:
|
||||
return None
|
||||
if isinstance(result, bytes):
|
||||
result = result.decode()
|
||||
return float(result)
|
||||
|
||||
async def flush_cache_buffer(self):
|
||||
print_verbose(
|
||||
f"flushing to redis....reached size of buffer {len(self.redis_batch_writing_buffer)}"
|
||||
|
|
|
|||
353
litellm/caching/valkey_semantic_cache.py
Normal file
353
litellm/caching/valkey_semantic_cache.py
Normal file
|
|
@ -0,0 +1,353 @@
|
|||
"""
|
||||
Valkey Semantic Cache implementation for LiteLLM
|
||||
|
||||
Backs semantic caching with Valkey (for example AWS ElastiCache for Valkey)
|
||||
running the valkey-search module.
|
||||
|
||||
RedisVL cannot drive valkey-search: it gates on a RediSearch module version
|
||||
that valkey-search does not report, and its SemanticCache index uses a TEXT
|
||||
field that valkey-search does not implement. This backend therefore talks to
|
||||
valkey-search directly over redis-py, building a vector index from the field
|
||||
types valkey-search does support (TAG for cache-key isolation and VECTOR for
|
||||
the prompt embedding) and running KNN queries for retrieval. Prompt extraction,
|
||||
embedding generation, and cached-response parsing are reused from
|
||||
RedisSemanticCache since those are backend agnostic.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import os
|
||||
import struct
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from redis import Redis
|
||||
from redis.asyncio import Redis as AsyncRedis
|
||||
from redis.commands.search.field import TagField, VectorField
|
||||
from redis.commands.search.indexDefinition import IndexDefinition, IndexType
|
||||
from redis.commands.search.query import Query
|
||||
|
||||
from litellm._logging import print_verbose
|
||||
from litellm._uuid import uuid
|
||||
|
||||
from .redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ValkeyCacheHit:
|
||||
response: str
|
||||
distance: float
|
||||
|
||||
|
||||
class ValkeySemanticCache(RedisSemanticCache):
|
||||
"""Valkey-backed semantic cache for LLM responses."""
|
||||
|
||||
DEFAULT_VALKEY_INDEX_NAME: str = "litellm_semantic_cache_index"
|
||||
EMBEDDING_FIELD_NAME: str = "embedding"
|
||||
PROMPT_FIELD_NAME: str = "prompt"
|
||||
RESPONSE_FIELD_NAME: str = "response"
|
||||
DISTANCE_FIELD_NAME: str = "vector_distance"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
host: str | None = None,
|
||||
port: str | None = None,
|
||||
password: str | None = None,
|
||||
redis_url: str | None = None,
|
||||
similarity_threshold: float | None = None,
|
||||
embedding_model: str = "text-embedding-ada-002",
|
||||
index_name: str | None = None,
|
||||
ssl: bool = False,
|
||||
startup_nodes: list | None = None,
|
||||
sync_client: Redis | None = None,
|
||||
async_client: AsyncRedis | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
if similarity_threshold is None:
|
||||
raise ValueError("similarity_threshold must be provided, passed None")
|
||||
|
||||
if startup_nodes:
|
||||
raise ValueError(
|
||||
"valkey-semantic does not support cluster-mode-enabled (multi-shard) "
|
||||
"endpoints. The async cluster client cannot route the FT.* search "
|
||||
"commands reliably. Point it at a cluster-mode-disabled endpoint "
|
||||
"instead (a primary with replicas is fine; only horizontal sharding "
|
||||
"is unsupported), or pass a single redis_url. On AWS, vector search "
|
||||
"needs ElastiCache for Valkey 8.2+ on a node-based cluster."
|
||||
)
|
||||
|
||||
self.similarity_threshold = similarity_threshold
|
||||
self.embedding_model = embedding_model
|
||||
self.index_name = index_name or self.DEFAULT_VALKEY_INDEX_NAME
|
||||
self.key_prefix = f"{self.index_name}:"
|
||||
self._index_dim: int | None = None
|
||||
|
||||
resolved_url = None
|
||||
if sync_client is None or async_client is None:
|
||||
resolved_url = redis_url or self._build_valkey_url(
|
||||
host, port, password, ssl
|
||||
)
|
||||
self.sync_client = (
|
||||
sync_client if sync_client is not None else Redis.from_url(resolved_url) # type: ignore[arg-type]
|
||||
)
|
||||
self.async_client = (
|
||||
async_client
|
||||
if async_client is not None
|
||||
else AsyncRedis.from_url(resolved_url) # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
print_verbose(f"Valkey semantic-cache initializing index - {self.index_name}")
|
||||
|
||||
@staticmethod
|
||||
def _build_valkey_url(
|
||||
host: str | None, port: str | None, password: str | None, ssl: bool = False
|
||||
) -> str:
|
||||
host = host or os.environ.get("VALKEY_HOST") or os.environ.get("REDIS_HOST")
|
||||
port = port or os.environ.get("VALKEY_PORT") or os.environ.get("REDIS_PORT")
|
||||
password = (
|
||||
password
|
||||
or os.environ.get("VALKEY_PASSWORD")
|
||||
or os.environ.get("REDIS_PASSWORD")
|
||||
)
|
||||
|
||||
if not host or not port:
|
||||
raise ValueError(
|
||||
"Missing required Valkey configuration. Provide host and port "
|
||||
"(or VALKEY_HOST/VALKEY_PORT), or pass redis_url."
|
||||
)
|
||||
|
||||
credentials = f":{password}@" if password else ""
|
||||
scheme = "rediss" if ssl else "redis"
|
||||
return f"{scheme}://{credentials}{host}:{port}"
|
||||
|
||||
@classmethod
|
||||
def _scope_tag(cls, key: str) -> str:
|
||||
# valkey-search TAG fields tokenize on punctuation and do not honour
|
||||
# backslash escaping, so an arbitrary cache key cannot be matched
|
||||
# verbatim. Hashing to hex yields a token that is always exact-match
|
||||
# safe and still uniquely isolates a caller's scope.
|
||||
return hashlib.sha256(str(key).encode("utf-8")).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _embedding_to_bytes(embedding: list[float]) -> bytes:
|
||||
return struct.pack(f"<{len(embedding)}f", *embedding)
|
||||
|
||||
def _index_schema(self, dim: int) -> tuple[TagField, VectorField]:
|
||||
return (
|
||||
TagField(self.CACHE_KEY_FIELD_NAME),
|
||||
VectorField(
|
||||
self.EMBEDDING_FIELD_NAME,
|
||||
"HNSW",
|
||||
{"TYPE": "FLOAT32", "DIM": dim, "DISTANCE_METRIC": "COSINE"},
|
||||
),
|
||||
)
|
||||
|
||||
def _index_definition(self) -> IndexDefinition:
|
||||
return IndexDefinition(prefix=[self.key_prefix], index_type=IndexType.HASH)
|
||||
|
||||
@staticmethod
|
||||
def _is_index_exists_error(exc: Exception) -> bool:
|
||||
return "already exists" in str(exc).lower()
|
||||
|
||||
@staticmethod
|
||||
def _extract_index_dim(info: dict) -> int | None:
|
||||
# FT.INFO nests the vector field's "dimensions" one level inside its
|
||||
# "index" block, so flatten each field descriptor a single level and
|
||||
# scan for the dimensions marker.
|
||||
for field in info.get("attributes") or []:
|
||||
if not isinstance(field, (list, tuple)):
|
||||
continue
|
||||
flat = [
|
||||
sub
|
||||
for item in field
|
||||
for sub in (item if isinstance(item, (list, tuple)) else [item])
|
||||
]
|
||||
for i, marker in enumerate(flat):
|
||||
if marker in (b"dimensions", "dimensions") and i + 1 < len(flat):
|
||||
return int(flat[i + 1])
|
||||
return None
|
||||
|
||||
def _assert_dim_matches(self, info: dict, dim: int) -> None:
|
||||
existing_dim = self._extract_index_dim(info)
|
||||
if existing_dim is not None and existing_dim != dim:
|
||||
raise ValueError(
|
||||
f"Valkey semantic-cache index '{self.index_name}' already exists with "
|
||||
f"embedding dimension {existing_dim}, but the configured embedding "
|
||||
f"model produced dimension {dim}. Use a different "
|
||||
f"valkey_semantic_cache_index_name or drop the existing index."
|
||||
)
|
||||
|
||||
def _ensure_index_sync(self, dim: int) -> None:
|
||||
if self._index_dim == dim:
|
||||
return
|
||||
try:
|
||||
self.sync_client.ft(self.index_name).create_index(
|
||||
self._index_schema(dim), definition=self._index_definition()
|
||||
)
|
||||
except Exception as exc:
|
||||
if not self._is_index_exists_error(exc):
|
||||
raise
|
||||
self._assert_dim_matches(self.sync_client.ft(self.index_name).info(), dim)
|
||||
self._index_dim = dim
|
||||
|
||||
async def _ensure_index_async(self, dim: int) -> None:
|
||||
if self._index_dim == dim:
|
||||
return
|
||||
try:
|
||||
await self.async_client.ft(self.index_name).create_index(
|
||||
self._index_schema(dim), definition=self._index_definition()
|
||||
)
|
||||
except Exception as exc:
|
||||
if not self._is_index_exists_error(exc):
|
||||
raise
|
||||
info = await self.async_client.ft(self.index_name).info()
|
||||
self._assert_dim_matches(info, dim)
|
||||
self._index_dim = dim
|
||||
|
||||
def _doc_key(self, key: str) -> str:
|
||||
return f"{self.key_prefix}{self._scope_tag(key)}:{uuid.uuid4()}"
|
||||
|
||||
def _doc_mapping(
|
||||
self, key: str, prompt: str, value_str: str, embedding: list[float]
|
||||
) -> dict:
|
||||
return {
|
||||
self.CACHE_KEY_FIELD_NAME: self._scope_tag(key),
|
||||
self.PROMPT_FIELD_NAME: prompt,
|
||||
self.RESPONSE_FIELD_NAME: value_str,
|
||||
self.EMBEDDING_FIELD_NAME: self._embedding_to_bytes(embedding),
|
||||
}
|
||||
|
||||
def _knn_query(self, key: str) -> Query:
|
||||
scope = self._scope_tag(key)
|
||||
query_string = (
|
||||
f"(@{self.CACHE_KEY_FIELD_NAME}:{{{scope}}})"
|
||||
f"=>[KNN 1 @{self.EMBEDDING_FIELD_NAME} $vec AS {self.DISTANCE_FIELD_NAME}]"
|
||||
)
|
||||
return (
|
||||
Query(query_string)
|
||||
.return_fields(self.RESPONSE_FIELD_NAME, self.DISTANCE_FIELD_NAME)
|
||||
.dialect(2)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _first_hit(cls, search_result: Any) -> _ValkeyCacheHit | None:
|
||||
docs = getattr(search_result, "docs", [])
|
||||
if not docs:
|
||||
return None
|
||||
doc = docs[0]
|
||||
return _ValkeyCacheHit(
|
||||
response=str(getattr(doc, cls.RESPONSE_FIELD_NAME)),
|
||||
distance=float(getattr(doc, cls.DISTANCE_FIELD_NAME)),
|
||||
)
|
||||
|
||||
def _resolve_hit(self, hit: _ValkeyCacheHit | None, key: str, **kwargs: Any) -> Any:
|
||||
if hit is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
similarity = 1 - hit.distance
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity
|
||||
|
||||
if similarity < self.similarity_threshold:
|
||||
return None
|
||||
return self._get_cache_logic(cached_response=hit.response)
|
||||
|
||||
def set_cache(self, key: str, value: Any, **kwargs: Any) -> None:
|
||||
print_verbose(f"Valkey semantic-cache set_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
embedding = self._get_embedding(prompt)
|
||||
self._ensure_index_sync(len(embedding))
|
||||
|
||||
doc_key = self._doc_key(key)
|
||||
self.sync_client.hset(
|
||||
doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding)
|
||||
)
|
||||
ttl = self._get_ttl(**kwargs)
|
||||
if ttl is not None:
|
||||
self.sync_client.expire(doc_key, ttl)
|
||||
except Exception as e:
|
||||
print_verbose(f"Error in Valkey semantic-cache set_cache: {str(e)}")
|
||||
|
||||
def get_cache(self, key: str, **kwargs: Any) -> Any:
|
||||
print_verbose(f"Valkey semantic-cache get_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
embedding = self._get_embedding(prompt)
|
||||
self._ensure_index_sync(len(embedding))
|
||||
|
||||
search_result = self.sync_client.ft(self.index_name).search(
|
||||
self._knn_query(key),
|
||||
query_params={"vec": self._embedding_to_bytes(embedding)},
|
||||
)
|
||||
return self._resolve_hit(self._first_hit(search_result), key, **kwargs)
|
||||
except Exception as e:
|
||||
print_verbose(f"Error in Valkey semantic-cache get_cache: {str(e)}")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
|
||||
async def async_set_cache(self, key: str, value: Any, **kwargs: Any) -> None:
|
||||
print_verbose(f"Async Valkey semantic-cache set_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
embedding = await self._get_async_embedding(prompt, **kwargs)
|
||||
await self._ensure_index_async(len(embedding))
|
||||
|
||||
doc_key = self._doc_key(key)
|
||||
await self.async_client.hset(
|
||||
doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding)
|
||||
)
|
||||
ttl = self._get_ttl(**kwargs)
|
||||
if ttl is not None:
|
||||
await self.async_client.expire(doc_key, ttl)
|
||||
except Exception as e:
|
||||
print_verbose(f"Error in async Valkey semantic-cache set_cache: {str(e)}")
|
||||
|
||||
async def async_get_cache(self, key: str, **kwargs: Any) -> Any:
|
||||
print_verbose(f"Async Valkey semantic-cache get_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
embedding = await self._get_async_embedding(prompt, **kwargs)
|
||||
await self._ensure_index_async(len(embedding))
|
||||
|
||||
search_result = await self.async_client.ft(self.index_name).search(
|
||||
self._knn_query(key),
|
||||
query_params={"vec": self._embedding_to_bytes(embedding)},
|
||||
)
|
||||
return self._resolve_hit(self._first_hit(search_result), key, **kwargs)
|
||||
except Exception as e:
|
||||
print_verbose(f"Error in async Valkey semantic-cache get_cache: {str(e)}")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
|
||||
async def async_set_cache_pipeline(
|
||||
self, cache_list: list[tuple[str, Any]], **kwargs: Any
|
||||
) -> None:
|
||||
try:
|
||||
await asyncio.gather(
|
||||
*[
|
||||
self.async_set_cache(key, value, **kwargs)
|
||||
for key, value in cache_list
|
||||
]
|
||||
)
|
||||
except Exception as e:
|
||||
print_verbose(
|
||||
f"Error in Valkey semantic-cache async_set_cache_pipeline: {str(e)}"
|
||||
)
|
||||
|
||||
async def _index_info(self) -> dict:
|
||||
return await self.async_client.ft(self.index_name).info()
|
||||
|
|
@ -802,6 +802,7 @@ openai_compatible_endpoints: List = [
|
|||
"https://api.inference.wandb.ai/v1",
|
||||
"https://api.clarifai.com/v2/ext/openai/v1",
|
||||
"https://api.libertai.io/v1",
|
||||
"https://pinstripes.io/v1",
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -865,6 +866,7 @@ openai_compatible_providers: List = [
|
|||
"clarifai",
|
||||
"docker_model_runner",
|
||||
"ragflow",
|
||||
"pinstripes", # Pinstripes - JSON-configured provider
|
||||
]
|
||||
openai_text_completion_compatible_providers: List = (
|
||||
[ # providers that support `/v1/completions`
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
from collections import OrderedDict
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Iterator, Mapping, cast
|
||||
from typing import TYPE_CHECKING, Any, Callable, Iterator, Mapping, Sequence, cast
|
||||
|
||||
from opentelemetry.context import attach, get_current
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
|
|
@ -546,6 +546,58 @@ class OpenTelemetryV2(CustomLogger):
|
|||
return span
|
||||
|
||||
|
||||
def select_global_otel_v2_logger(
|
||||
in_memory_loggers: Sequence[object],
|
||||
registered: "OpenTelemetryV2 | None" = None,
|
||||
) -> "OpenTelemetryV2":
|
||||
"""The single ``OpenTelemetryV2`` whose provider should become the OTel global.
|
||||
|
||||
The callback factory designates one logger as canonical the moment it builds
|
||||
the first one (``_init_otel_logger_on_litellm_proxy`` sets
|
||||
``proxy_server.open_telemetry_logger``), and every other v2 entry point —
|
||||
guardrail, identity seeding, phase spans — already routes through that same
|
||||
``registered`` owner. Reuse it here too so the global provider has one source
|
||||
of truth instead of a second, independently-derived guess; this is the logger
|
||||
a preset (arize, langfuse, …) folds the ``OTEL_*`` base exporter and its own
|
||||
exporter into, so the FastAPI server span and the gen-ai spans share one
|
||||
provider and one trace.
|
||||
|
||||
Fall back to ``in_memory_loggers`` for the SDK path, where no proxy global is
|
||||
set (selecting from there, not ``service_callback``, which a preset logger does
|
||||
not always reach), and build a generic logger from ``OTEL_*`` only when none was
|
||||
configured at all. Each fallback still avoids the second generic logger that
|
||||
orphaned the gen-ai spans onto a different backend than the server span.
|
||||
"""
|
||||
if registered is not None:
|
||||
return registered
|
||||
existing = next(
|
||||
(cb for cb in in_memory_loggers if isinstance(cb, OpenTelemetryV2)), None
|
||||
)
|
||||
return existing if existing is not None else OpenTelemetryV2()
|
||||
|
||||
|
||||
def publish_global_otel_v2_provider(
|
||||
in_memory_loggers: Sequence[object],
|
||||
set_global_provider: Callable[[TracerProvider], None],
|
||||
registered: "OpenTelemetryV2 | None" = None,
|
||||
) -> "OpenTelemetryV2":
|
||||
"""Select the single v2 logger and publish its provider as the OTel global.
|
||||
|
||||
The proxy calls this once at startup, after callbacks are initialized, so the
|
||||
preset logger already exists; it passes ``registered`` (the canonical owner the
|
||||
factory designated as ``proxy_server.open_telemetry_logger``) so the global
|
||||
provider reuses the same logger the rest of the v2 code emits through (see
|
||||
:func:`select_global_otel_v2_logger`). Both ``registered`` and
|
||||
``set_global_provider`` (the proxy passes
|
||||
``opentelemetry.trace.set_tracer_provider``) are injected so the publish step is
|
||||
unit-testable without reading or mutating real global OTel state. Returns the
|
||||
logger whose provider was published.
|
||||
"""
|
||||
logger = select_global_otel_v2_logger(in_memory_loggers, registered=registered)
|
||||
set_global_provider(logger._tracer_provider)
|
||||
return logger
|
||||
|
||||
|
||||
def _registered_v2_logger() -> "OpenTelemetryV2 | None":
|
||||
try:
|
||||
from litellm.proxy import proxy_server
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Typed configuration for the OpenTelemetry instrumentation."""
|
||||
|
||||
from enum import Enum
|
||||
from typing import Any, List
|
||||
|
||||
from pydantic import AliasChoices, BaseModel, Field, field_validator, model_validator
|
||||
|
|
@ -23,6 +24,23 @@ class CaptureMessageContent(str):
|
|||
SPAN_AND_EVENT = "span_and_event"
|
||||
|
||||
|
||||
class ExporterOwner(str, Enum):
|
||||
"""The preset that contributed an exporter. Values match the callback names
|
||||
in ``presets.PRESET_BY_CALLBACK`` so per-request dynamic-credential routing
|
||||
can match an exporter's owner against the credential source's callback name.
|
||||
A ``str`` enum so the value compares equal to the bare callback-name string."""
|
||||
|
||||
# Arize AX (the hosted platform) and Arize Phoenix (the open-source / Phoenix
|
||||
# Cloud tracer) are distinct backends with separate config and auth, so they
|
||||
# are separate owners. The member value stays the public callback name.
|
||||
ARIZE_AX = "arize"
|
||||
ARIZE_PHOENIX = "arize_phoenix"
|
||||
LANGFUSE_OTEL = "langfuse_otel"
|
||||
WEAVE_OTEL = "weave_otel"
|
||||
LEVO = "levo"
|
||||
AGENTOPS = "agentops"
|
||||
|
||||
|
||||
class _OTelV2Flag(BaseSettings):
|
||||
model_config = SettingsConfigDict(extra="ignore")
|
||||
|
||||
|
|
@ -49,6 +67,15 @@ class ExporterSpec(BaseModel):
|
|||
)
|
||||
endpoint: str | None = None
|
||||
headers: str | None = None
|
||||
owner: ExporterOwner | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"The preset that contributed this exporter. Per-request dynamic OTLP "
|
||||
"credentials are applied only to the exporter whose owner matches the "
|
||||
"credential source, so one tenant's vendor key never lands on a "
|
||||
"different backend's exporter."
|
||||
),
|
||||
)
|
||||
options: dict[str, str] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
|
|
|
|||
|
|
@ -88,13 +88,23 @@ class TenantTracerCache:
|
|||
return get_tracer(provider, self._tracer_name)
|
||||
|
||||
def _config_with_headers(self, headers: Mapping[str, str]) -> OpenTelemetryV2Config:
|
||||
"""Clone the config, replacing OTLP exporter headers with ``headers``."""
|
||||
"""Clone the config, stamping ``headers`` onto the credential's own exporter.
|
||||
|
||||
``headers`` are the per-request credentials of ``self._callback_name`` (the
|
||||
integration that built this cache), so they apply only to the exporter that
|
||||
integration contributed (``spec.owner``). A request that carries one
|
||||
tenant's Arize key must never rewrite the headers of a co-configured
|
||||
Langfuse or self-hosted collector exporter, which would leak that key to a
|
||||
different backend.
|
||||
"""
|
||||
header_str = ",".join(f"{key}={value}" for key, value in headers.items())
|
||||
header_update: dict[str, str] = {"headers": header_str}
|
||||
exporters = [
|
||||
(
|
||||
spec
|
||||
if spec.kind.lower() in _NON_OTLP_KINDS
|
||||
else spec.model_copy(update={"headers": header_str})
|
||||
spec.model_copy(update=header_update)
|
||||
if spec.owner == self._callback_name
|
||||
and spec.kind.lower() not in _NON_OTLP_KINDS
|
||||
else spec
|
||||
)
|
||||
for spec in self._config.exporters
|
||||
]
|
||||
|
|
|
|||
|
|
@ -16,7 +16,11 @@ from pydantic import Field
|
|||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.plumbing.providers import register_exporter_factory
|
||||
|
||||
_AGENTOPS_ENDPOINT = "https://otlp.agentops.cloud/v1/traces"
|
||||
|
|
@ -59,6 +63,7 @@ def agentops_preset(
|
|||
options=(
|
||||
{"api_key": settings.api_key} if settings.api_key else None
|
||||
),
|
||||
owner=ExporterOwner.AGENTOPS,
|
||||
),
|
||||
],
|
||||
"resource_attributes": {
|
||||
|
|
|
|||
|
|
@ -4,7 +4,11 @@ from pydantic import Field
|
|||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
from litellm.integrations.arize.arize import ArizeLogger as _V1ArizeLogger
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.presets.utils import ensure_mappers
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
|
|
@ -34,6 +38,7 @@ def arize_preset(
|
|||
kind=arize_cfg.protocol or "otlp_grpc",
|
||||
endpoint=arize_cfg.endpoint or "https://otlp.arize.com/v1",
|
||||
headers=headers,
|
||||
owner=ExporterOwner.ARIZE_AX,
|
||||
),
|
||||
],
|
||||
"mapper_names": ensure_mappers(base.mapper_names, "openinference"),
|
||||
|
|
|
|||
|
|
@ -3,7 +3,11 @@
|
|||
from litellm.integrations.langfuse.langfuse_otel import (
|
||||
LangfuseOtelLogger as _V1Langfuse,
|
||||
)
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.presets.utils import ensure_mappers
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
|
|
@ -23,6 +27,7 @@ def langfuse_preset(
|
|||
kind=kind,
|
||||
endpoint=cfg.endpoint,
|
||||
headers=cfg.headers,
|
||||
owner=ExporterOwner.LANGFUSE_OTEL,
|
||||
),
|
||||
],
|
||||
"mapper_names": ensure_mappers(base.mapper_names, "langfuse"),
|
||||
|
|
|
|||
|
|
@ -1,7 +1,11 @@
|
|||
"""Levo preset — OTLP/HTTP to a Levo collector with org+workspace headers."""
|
||||
|
||||
from litellm.integrations.levo.levo import LevoLogger as _V1Levo
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
|
||||
|
||||
def levo_preset(
|
||||
|
|
@ -18,6 +22,7 @@ def levo_preset(
|
|||
kind="otlp_http",
|
||||
endpoint=cfg.endpoint,
|
||||
headers=cfg.otlp_auth_headers,
|
||||
owner=ExporterOwner.LEVO,
|
||||
),
|
||||
],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,7 +6,11 @@ from pydantic_settings import BaseSettings, SettingsConfigDict
|
|||
from litellm.integrations.arize.arize_phoenix import (
|
||||
ArizePhoenixLogger as _V1Phoenix,
|
||||
)
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.presets.utils import ensure_mappers
|
||||
|
||||
|
||||
|
|
@ -37,6 +41,7 @@ def phoenix_preset(
|
|||
kind=cfg.protocol if hasattr(cfg, "protocol") else "otlp_http",
|
||||
endpoint=cfg.endpoint,
|
||||
headers=headers,
|
||||
owner=ExporterOwner.ARIZE_PHOENIX,
|
||||
),
|
||||
],
|
||||
"mapper_names": ensure_mappers(base.mapper_names, "openinference"),
|
||||
|
|
|
|||
|
|
@ -1,6 +1,10 @@
|
|||
"""Weave (W&B) preset."""
|
||||
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.presets.utils import ensure_mappers
|
||||
from litellm.integrations.weave.weave_otel import (
|
||||
_get_weave_authorization_header,
|
||||
|
|
@ -23,6 +27,7 @@ def weave_preset(
|
|||
kind=weave_cfg.protocol or "otlp_http",
|
||||
endpoint=weave_cfg.endpoint,
|
||||
headers=weave_cfg.otlp_auth_headers,
|
||||
owner=ExporterOwner.WEAVE_OTEL,
|
||||
),
|
||||
],
|
||||
# Weave consumes OpenInference + a small Weave-specific overlay.
|
||||
|
|
|
|||
|
|
@ -15,8 +15,23 @@ BEDROCK_MANAGED_S3_PREFIXES = (
|
|||
BEDROCK_MANAGED_S3_UPLOAD_PREFIX,
|
||||
BEDROCK_MANAGED_S3_OUTPUT_PREFIX,
|
||||
)
|
||||
MANAGED_CLOUD_STORAGE_SCHEMES = ("s3://", "gs://")
|
||||
_MAPPING_PROXY_TYPE: type = type(MappingProxyType({}))
|
||||
|
||||
|
||||
def is_managed_cloud_storage_uri(file_id: str) -> bool:
|
||||
"""
|
||||
True if file_id is a raw cloud-storage object URI (e.g. ``s3://bucket/key``).
|
||||
|
||||
These are internal provider artifacts. On the multi-tenant proxy they must be
|
||||
retrieved through their managed unified file id so owner/team access is enforced;
|
||||
a raw URI supplied by a caller bypasses that check.
|
||||
"""
|
||||
return isinstance(file_id, str) and file_id.startswith(
|
||||
MANAGED_CLOUD_STORAGE_SCHEMES
|
||||
)
|
||||
|
||||
|
||||
_SAFE_OBJECT_COMPONENT_PATTERN = re.compile(r"[^A-Za-z0-9._-]+")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -388,6 +388,9 @@ def get_llm_provider(
|
|||
elif endpoint == "https://api.inference.wandb.ai/v1":
|
||||
custom_llm_provider = "wandb"
|
||||
dynamic_api_key = get_secret_str("WANDB_API_KEY")
|
||||
elif endpoint == "https://pinstripes.io/v1":
|
||||
custom_llm_provider = "pinstripes"
|
||||
dynamic_api_key = get_secret_str("PINSTRIPES_API_KEY")
|
||||
elif endpoint == "https://gigachat.devices.sberbank.ru/api/v1":
|
||||
custom_llm_provider = "gigachat"
|
||||
dynamic_api_key = get_secret_str("GIGACHAT_API_KEY")
|
||||
|
|
@ -646,7 +649,7 @@ def _get_openai_compatible_provider_info(
|
|||
api_base,
|
||||
dynamic_api_key,
|
||||
) = litellm.BedrockMantleChatConfig()._get_openai_compatible_provider_info(
|
||||
api_base, api_key, litellm_params=litellm_params
|
||||
api_base, api_key, litellm_params=litellm_params, model=model
|
||||
)
|
||||
elif custom_llm_provider == "nvidia_nim":
|
||||
# nvidia_nim is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.endpoints.anyscale.com/v1
|
||||
|
|
|
|||
|
|
@ -44,8 +44,8 @@ class HealthCheckHelpers:
|
|||
model_params["litellm_logging_obj"] = litellm_logging_obj
|
||||
model_params["fallbacks"] = fallback_models
|
||||
model_params["max_tokens"] = model_params.get(
|
||||
"max_tokens", 10
|
||||
) # gpt-5-nano throws errors for max_tokens=1
|
||||
"max_tokens", 16
|
||||
) # GPT-5 models require max_output_tokens >= 16
|
||||
await acompletion(**model_params)
|
||||
return {}
|
||||
|
||||
|
|
|
|||
|
|
@ -2980,7 +2980,12 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
self.model_call_details["end_time"] = end_time
|
||||
self.model_call_details.setdefault("original_response", None)
|
||||
self.model_call_details["response_cost"] = 0
|
||||
# A stream interrupted mid-flight still billed the provider for the
|
||||
# chunks already delivered; the router stashes that recovered usage as
|
||||
# ``combined_usage_object`` and pre-computes its cost, so preserve it
|
||||
# here instead of zeroing the spend on an otherwise-failed request.
|
||||
if self.model_call_details.get("combined_usage_object") is None:
|
||||
self.model_call_details["response_cost"] = 0
|
||||
|
||||
if hasattr(exception, "headers") and isinstance(exception.headers, dict):
|
||||
self.model_call_details.setdefault("litellm_params", {})
|
||||
|
|
|
|||
|
|
@ -2290,6 +2290,7 @@ class CustomStreamWrapper:
|
|||
litellm.request_timeout
|
||||
)
|
||||
if self.logging_obj is not None:
|
||||
self._record_partial_usage_for_failure()
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=self.logging_obj.failure_handler,
|
||||
|
|
@ -2303,6 +2304,7 @@ class CustomStreamWrapper:
|
|||
except Exception as e:
|
||||
traceback_exception = traceback.format_exc()
|
||||
if self.logging_obj is not None:
|
||||
self._record_partial_usage_for_failure()
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=self.logging_obj.failure_handler,
|
||||
|
|
@ -2314,6 +2316,33 @@ class CustomStreamWrapper:
|
|||
)
|
||||
self._handle_stream_fallback_error(e)
|
||||
|
||||
def _record_partial_usage_for_failure(self) -> None:
|
||||
"""
|
||||
A stream that breaks mid-flight still billed the provider for the chunks
|
||||
already delivered. Recover that partial usage from the chunks seen so
|
||||
far and stash it, with its cost, on the logging object so the failure
|
||||
handler records the real partial spend instead of zero. A request that
|
||||
later recovers via a router fallback overwrites this with the combined
|
||||
success log on the same request id, so this never double counts.
|
||||
"""
|
||||
if self.logging_obj is None or not self.chunks:
|
||||
return
|
||||
try:
|
||||
partial_response = litellm.stream_chunk_builder(chunks=self.chunks)
|
||||
usage = cast(Optional[Usage], getattr(partial_response, "usage", None))
|
||||
if usage is None:
|
||||
return
|
||||
self.logging_obj.model_call_details["combined_usage_object"] = usage
|
||||
self.logging_obj.model_call_details["response_cost"] = (
|
||||
self.logging_obj._response_cost_calculator(result=partial_response)
|
||||
or 0.0
|
||||
)
|
||||
except Exception as recover_error:
|
||||
verbose_logger.debug(
|
||||
"could not recover partial usage for interrupted stream: %s",
|
||||
recover_error,
|
||||
)
|
||||
|
||||
def _handle_stream_fallback_error(self, e: Exception) -> "NoReturn":
|
||||
"""
|
||||
Common error handling for both __next__ and __anext__.
|
||||
|
|
|
|||
|
|
@ -859,7 +859,17 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
"""
|
||||
new_tools: List[ChatCompletionToolParam] = []
|
||||
tool_name_mapping: Dict[str, str] = {}
|
||||
mapped_tool_params = ["name", "input_schema", "description", "cache_control"]
|
||||
# "type" is the Anthropic tool type (e.g. "custom"); it must not be
|
||||
# merged into the OpenAI function `parameters` schema below, or it
|
||||
# overwrites the real parameters.type ("object") and the provider
|
||||
# rejects the request. See #30557.
|
||||
mapped_tool_params = [
|
||||
"name",
|
||||
"input_schema",
|
||||
"description",
|
||||
"cache_control",
|
||||
"type",
|
||||
]
|
||||
|
||||
for idx, tool in enumerate(tools):
|
||||
# Check if this is an Anthropic-native tool that should be kept as-is
|
||||
|
|
|
|||
|
|
@ -1,8 +1,6 @@
|
|||
import asyncio
|
||||
import base64
|
||||
import os
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Coroutine, Mapping, Optional, Tuple, Union, cast
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Coroutine, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -17,7 +15,6 @@ from litellm.types.llms.openai import (
|
|||
FileContentRequest,
|
||||
HttpxBinaryResponseContent,
|
||||
)
|
||||
from litellm.types.utils import SpecialEnums
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
|
||||
|
|
@ -37,40 +34,9 @@ class BedrockFilesHandler(BaseAWSLLM):
|
|||
)
|
||||
|
||||
def _extract_s3_uri_from_file_id(self, file_id: str) -> str:
|
||||
"""
|
||||
Extract S3 URI from encoded file ID.
|
||||
from .transformation import extract_s3_uri_from_file_id
|
||||
|
||||
The file ID can be in two formats:
|
||||
1. Base64-encoded unified file ID containing: llm_output_file_id,s3://bucket/path
|
||||
2. Direct S3 URI: s3://bucket/litellm-managed-prefix/path
|
||||
|
||||
Args:
|
||||
file_id: Encoded file ID or direct S3 URI
|
||||
|
||||
Returns:
|
||||
S3 URI (e.g., "s3://bucket-name/path/to/file")
|
||||
"""
|
||||
# First, try to decode if it's a base64-encoded unified file ID
|
||||
try:
|
||||
# Add padding if needed
|
||||
padded = file_id + "=" * (-len(file_id) % 4)
|
||||
decoded = base64.urlsafe_b64decode(padded).decode()
|
||||
|
||||
# Check if it's a unified file ID format
|
||||
if decoded.startswith(SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value):
|
||||
# Extract llm_output_file_id from the decoded string
|
||||
if "llm_output_file_id," in decoded:
|
||||
s3_uri = decoded.split("llm_output_file_id,")[1].split(";")[0]
|
||||
return s3_uri
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# If not base64 encoded or doesn't contain llm_output_file_id, accept only
|
||||
# explicit S3 URIs. Bucket and key validation happens before any S3 call.
|
||||
if file_id.startswith("s3://"):
|
||||
return file_id
|
||||
|
||||
raise ValueError("file_id must be a managed LiteLLM S3 file id")
|
||||
return extract_s3_uri_from_file_id(file_id)
|
||||
|
||||
def _parse_s3_uri(
|
||||
self,
|
||||
|
|
@ -95,26 +61,12 @@ class BedrockFilesHandler(BaseAWSLLM):
|
|||
allow_legacy_cloud_file_ids=allow_legacy_cloud_file_ids,
|
||||
)
|
||||
|
||||
def _get_configured_s3_bucket_name(self, litellm_params: dict) -> str:
|
||||
trusted_model_credentials = litellm_params.get(
|
||||
"_litellm_internal_model_credentials"
|
||||
)
|
||||
bucket_name = None
|
||||
if isinstance(trusted_model_credentials, type(MappingProxyType({}))):
|
||||
trusted_model_credentials_mapping = cast(
|
||||
Mapping[str, Any], trusted_model_credentials
|
||||
)
|
||||
candidate_bucket_name = trusted_model_credentials_mapping.get(
|
||||
"s3_bucket_name"
|
||||
)
|
||||
if isinstance(candidate_bucket_name, str):
|
||||
bucket_name = candidate_bucket_name
|
||||
bucket_name = bucket_name or os.getenv("AWS_S3_BUCKET_NAME")
|
||||
if not bucket_name:
|
||||
raise ValueError(
|
||||
"S3 bucket_name is required. Set 's3_bucket_name' in proxy config or AWS_S3_BUCKET_NAME for Bedrock file content retrieval."
|
||||
)
|
||||
return bucket_name
|
||||
def _get_configured_s3_bucket_name(
|
||||
self, litellm_params: Mapping[str, object]
|
||||
) -> str:
|
||||
from .transformation import get_configured_s3_bucket_name
|
||||
|
||||
return get_configured_s3_bucket_name(litellm_params)
|
||||
|
||||
async def afile_content(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1,23 +1,37 @@
|
|||
import base64
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
from collections.abc import Mapping, MutableMapping
|
||||
from types import MappingProxyType
|
||||
from typing import (
|
||||
Any,
|
||||
Dict,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
Union,
|
||||
)
|
||||
from urllib.parse import unquote
|
||||
|
||||
import httpx
|
||||
from httpx import Headers, Response
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.files.utils import FilesAPIUtils
|
||||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
BEDROCK_MANAGED_S3_BATCH_PREFIX,
|
||||
BEDROCK_MANAGED_S3_PREFIXES,
|
||||
BEDROCK_MANAGED_S3_UPLOAD_PREFIX,
|
||||
build_managed_cloud_object_name,
|
||||
encode_s3_object_key_for_url,
|
||||
sanitize_cloud_object_component,
|
||||
should_allow_legacy_cloud_file_ids,
|
||||
split_configured_cloud_bucket_name,
|
||||
validate_managed_cloud_file_id,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
|
@ -28,18 +42,98 @@ from litellm.llms.base_llm.files.transformation import (
|
|||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
CreateFileRequest,
|
||||
FileContentRequest,
|
||||
FileTypes,
|
||||
HttpxBinaryResponseContent,
|
||||
OpenAICreateFileRequestOptionalParams,
|
||||
OpenAIFileObject,
|
||||
PathLike,
|
||||
)
|
||||
from litellm.types.utils import ExtractedFileData, LlmProviders
|
||||
from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums
|
||||
from litellm.utils import get_llm_provider
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
from ..common_utils import BedrockError
|
||||
|
||||
# litellm_params key used to hand the SigV4-signed GET headers from
|
||||
# `transform_file_content_request` to `validate_environment` (the only hook
|
||||
# the shared file-content HTTP handler exposes for setting request headers).
|
||||
# Same pattern as the `upload_url` handoff in `transform_create_file_request`.
|
||||
S3_SIGNED_GET_HEADERS_PARAM = "_s3_signed_get_headers"
|
||||
|
||||
|
||||
class _BedrockS3RequestParams(BaseModel):
|
||||
"""Typed view of the credential/region params the S3 GetObject path reads."""
|
||||
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
aws_access_key_id: str | None = None
|
||||
aws_secret_access_key: str | None = None
|
||||
aws_session_token: str | None = None
|
||||
aws_region_name: str | None = None
|
||||
aws_session_name: str | None = None
|
||||
aws_profile_name: str | None = None
|
||||
aws_role_name: str | None = None
|
||||
aws_web_identity_token: str | None = None
|
||||
aws_sts_endpoint: str | None = None
|
||||
s3_region_name: str | None = None
|
||||
s3_endpoint_url: str | None = None
|
||||
|
||||
|
||||
class _TrustedS3ModelCredentials(BaseModel):
|
||||
"""The S3 bucket the server trusts file ids against, from the deployment snapshot."""
|
||||
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
s3_bucket_name: str | None = None
|
||||
|
||||
|
||||
def extract_s3_uri_from_file_id(file_id: str) -> str:
|
||||
"""
|
||||
Resolve a Bedrock file id to its S3 URI.
|
||||
|
||||
Accepts either a base64-encoded LiteLLM unified file id (whose decoded
|
||||
form carries `llm_output_file_id,s3://...`) or a direct `s3://` URI.
|
||||
"""
|
||||
try:
|
||||
padded = file_id + "=" * (-len(file_id) % 4)
|
||||
decoded = base64.urlsafe_b64decode(padded).decode()
|
||||
|
||||
if decoded.startswith(SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value):
|
||||
if "llm_output_file_id," in decoded:
|
||||
return decoded.split("llm_output_file_id,")[1].split(";")[0]
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if file_id.startswith("s3://"):
|
||||
return file_id
|
||||
|
||||
raise ValueError("file_id must be a managed LiteLLM S3 file id")
|
||||
|
||||
|
||||
def get_configured_s3_bucket_name(litellm_params: Mapping[str, object]) -> str:
|
||||
"""
|
||||
Resolve the server-configured S3 bucket for Bedrock file operations.
|
||||
|
||||
Only trusts the immutable server-side credential snapshot or the
|
||||
environment; never a request-supplied param, since the bucket is what
|
||||
`validate_managed_cloud_file_id` checks file ids against.
|
||||
"""
|
||||
trusted_model_credentials = litellm_params.get(
|
||||
"_litellm_internal_model_credentials"
|
||||
)
|
||||
bucket_name: str | None = None
|
||||
if isinstance(trusted_model_credentials, MappingProxyType):
|
||||
snapshot: dict[str, object] = {}
|
||||
snapshot.update(trusted_model_credentials) # any-ok: untyped snapshot
|
||||
bucket_name = _TrustedS3ModelCredentials.model_validate(snapshot).s3_bucket_name
|
||||
bucket_name = bucket_name or os.getenv("AWS_S3_BUCKET_NAME")
|
||||
if not bucket_name:
|
||||
raise ValueError(
|
||||
"S3 bucket_name is required. Set 's3_bucket_name' in proxy config or AWS_S3_BUCKET_NAME for Bedrock file content retrieval."
|
||||
)
|
||||
return bucket_name
|
||||
|
||||
|
||||
class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
||||
"""
|
||||
|
|
@ -63,16 +157,21 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
headers: MutableMapping[str, object],
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
litellm_params: MutableMapping[str, object],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict:
|
||||
# No additional headers needed for S3 uploads - AWS credentials handled by BaseAWSLLM
|
||||
return headers
|
||||
result: dict[str, object] = {}
|
||||
result.update(headers)
|
||||
signed_headers = litellm_params.pop(S3_SIGNED_GET_HEADERS_PARAM, None)
|
||||
if isinstance(signed_headers, Mapping):
|
||||
result.update(signed_headers) # any-ok: untyped handoff headers
|
||||
# otherwise no extra headers - AWS credentials are handled by BaseAWSLLM
|
||||
return result
|
||||
|
||||
def _get_content_from_openai_file(self, openai_file_content: FileTypes) -> str:
|
||||
"""
|
||||
|
|
@ -927,23 +1026,114 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
|
||||
def transform_file_content_request(
|
||||
self,
|
||||
file_content_request,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> tuple[str, dict]:
|
||||
raise NotImplementedError(
|
||||
"BedrockFilesConfig does not support file content retrieval"
|
||||
file_content_request: FileContentRequest,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: MutableMapping[str, object],
|
||||
) -> tuple[str, dict[str, str]]:
|
||||
"""
|
||||
Build a SigV4-signed S3 GetObject request for a Bedrock batch file.
|
||||
|
||||
Bedrock batch file ids are `s3://bucket/key` URIs (or unified ids
|
||||
that decode to one); the bucket and key are validated against the
|
||||
server-configured bucket before any request is signed.
|
||||
"""
|
||||
file_id = file_content_request.get("file_id")
|
||||
if not file_id:
|
||||
raise ValueError("file_id is required for Bedrock file content retrieval")
|
||||
|
||||
s3_uri = extract_s3_uri_from_file_id(file_id)
|
||||
bucket_name, object_key = validate_managed_cloud_file_id(
|
||||
file_id=s3_uri,
|
||||
scheme="s3://",
|
||||
configured_bucket_name=get_configured_s3_bucket_name(litellm_params),
|
||||
allowed_object_prefixes=BEDROCK_MANAGED_S3_PREFIXES,
|
||||
allow_legacy_cloud_file_ids=should_allow_legacy_cloud_file_ids(
|
||||
litellm_params
|
||||
),
|
||||
)
|
||||
|
||||
# The shared file-content handler passes optional_params={}, so AWS
|
||||
# credentials/region arrive via litellm_params here (unlike the upload
|
||||
# path). s3_region_name wins over aws_region_name, same priority as
|
||||
# get_complete_file_url above.
|
||||
merged_params: dict[str, object] = {}
|
||||
merged_params.update(litellm_params)
|
||||
merged_params.update(optional_params)
|
||||
request_params = _BedrockS3RequestParams.model_validate(merged_params)
|
||||
|
||||
region_preference = (
|
||||
request_params.s3_region_name or request_params.aws_region_name
|
||||
)
|
||||
region_params: dict[str, str | None] = {"aws_region_name": region_preference}
|
||||
aws_region_name = self._get_aws_region_name(
|
||||
optional_params=region_params, model=""
|
||||
)
|
||||
|
||||
s3_endpoint_url = (
|
||||
request_params.s3_endpoint_url
|
||||
or f"https://s3.{aws_region_name}.amazonaws.com"
|
||||
).rstrip("/")
|
||||
url = f"{s3_endpoint_url}/{bucket_name}/{encode_s3_object_key_for_url(object_key)}"
|
||||
|
||||
litellm_params[S3_SIGNED_GET_HEADERS_PARAM] = self._sign_s3_get_request(
|
||||
api_base=url,
|
||||
aws_region_name=aws_region_name,
|
||||
request_params=request_params,
|
||||
)
|
||||
return url, {}
|
||||
|
||||
def _sign_s3_get_request(
|
||||
self,
|
||||
api_base: str,
|
||||
aws_region_name: str,
|
||||
request_params: _BedrockS3RequestParams,
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
SigV4-sign an S3 GetObject request, mirroring `_sign_s3_request` (PUT).
|
||||
"""
|
||||
try:
|
||||
import hashlib
|
||||
|
||||
from botocore.auth import SigV4Auth
|
||||
from botocore.awsrequest import AWSRequest
|
||||
except ImportError:
|
||||
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
|
||||
|
||||
credentials = self.get_credentials( # any-ok: boto3 Credentials is untyped
|
||||
aws_access_key_id=request_params.aws_access_key_id,
|
||||
aws_secret_access_key=request_params.aws_secret_access_key,
|
||||
aws_session_token=request_params.aws_session_token,
|
||||
aws_region_name=aws_region_name,
|
||||
aws_session_name=request_params.aws_session_name,
|
||||
aws_profile_name=request_params.aws_profile_name,
|
||||
aws_role_name=request_params.aws_role_name,
|
||||
aws_web_identity_token=request_params.aws_web_identity_token,
|
||||
aws_sts_endpoint=request_params.aws_sts_endpoint,
|
||||
)
|
||||
|
||||
empty_body_hash = hashlib.sha256(b"").hexdigest()
|
||||
aws_request = AWSRequest( # any-ok: botocore AWSRequest is untyped
|
||||
method="GET",
|
||||
url=api_base,
|
||||
headers={"x-amz-content-sha256": empty_body_hash},
|
||||
)
|
||||
auth = SigV4Auth(credentials, "s3", aws_region_name) # any-ok: botocore untyped
|
||||
auth.add_auth(aws_request) # any-ok: botocore request mutation is untyped
|
||||
return dict(aws_request.headers) # any-ok: botocore headers are untyped
|
||||
|
||||
def transform_file_content_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
) -> HttpxBinaryResponseContent:
|
||||
raise NotImplementedError(
|
||||
"BedrockFilesConfig does not support file content retrieval"
|
||||
)
|
||||
if raw_response.status_code >= 400:
|
||||
raise BedrockError(
|
||||
status_code=raw_response.status_code,
|
||||
message=raw_response.text,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
return HttpxBinaryResponseContent(response=raw_response)
|
||||
|
||||
|
||||
class BedrockJsonlFilesTransformation:
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from litellm.secret_managers.main import get_secret_str
|
|||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
from ..common_utils import mantle_base_segment
|
||||
from ...openai_like.chat.transformation import OpenAILikeChatConfig
|
||||
|
||||
|
||||
|
|
@ -48,6 +49,7 @@ class BedrockMantleChatConfig(BedrockMantleAuthMixin, OpenAILikeChatConfig):
|
|||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
litellm_params: Optional[GenericLiteLLMParams] = None,
|
||||
model: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
region = (
|
||||
(litellm_params.aws_region_name if litellm_params else None)
|
||||
|
|
@ -57,10 +59,13 @@ class BedrockMantleChatConfig(BedrockMantleAuthMixin, OpenAILikeChatConfig):
|
|||
or BEDROCK_MANTLE_DEFAULT_REGION
|
||||
)
|
||||
BaseAWSLLM._validate_aws_region_name(region)
|
||||
# The base path segment is data-driven per model (use_openai_responses_path
|
||||
# flag): gemma-4-* and gpt-5.x are served on /openai/v1, everything else on
|
||||
# /v1. An explicit api_base still wins over the derived default.
|
||||
api_base = (
|
||||
api_base
|
||||
or get_secret_str("BEDROCK_MANTLE_API_BASE")
|
||||
or f"https://bedrock-mantle.{region}.api.aws/v1"
|
||||
or f"https://bedrock-mantle.{region}.api.aws/{mantle_base_segment(model, litellm.model_cost)}"
|
||||
)
|
||||
dynamic_api_key = self._resolve_bearer_token(api_key)
|
||||
return api_base, dynamic_api_key
|
||||
|
|
|
|||
|
|
@ -1,5 +1,4 @@
|
|||
"""
|
||||
Shared auth and region resolution for the Amazon Bedrock Mantle backends.
|
||||
"""Shared auth, region resolution, and routing helpers for the Amazon Bedrock Mantle provider.
|
||||
|
||||
Mantle authenticates with a Bearer token when one is available
|
||||
(litellm_params.api_key, BEDROCK_MANTLE_API_KEY, or the standard
|
||||
|
|
@ -7,6 +6,10 @@ AWS_BEARER_TOKEN_BEDROCK); otherwise it falls back to AWS SigV4 (service
|
|||
"bedrock") over the standard credential chain (IAM role / access key / profile /
|
||||
web identity). The Chat Completions and Responses backends share this behaviour
|
||||
through BedrockMantleAuthMixin so the two paths can never drift apart.
|
||||
|
||||
The two routing helpers (mantle_supports_responses, mantle_base_segment) are
|
||||
pure functions of (model, model_cost) so they can be unit-tested without patching
|
||||
global state.
|
||||
"""
|
||||
|
||||
import re
|
||||
|
|
@ -72,13 +75,9 @@ class BedrockMantleAuthMixin:
|
|||
) -> Tuple[dict, bytes | None]:
|
||||
bearer = self._resolve_bearer_token(api_key)
|
||||
if not bearer:
|
||||
# SigV4 path. Pin the credential-scope region to the region of the actual
|
||||
# signing URL so the SigV4 scope and the URL host can never disagree, even
|
||||
# when a stale api_base and aws_region_name point at different regions.
|
||||
# Fall back to _resolve_region only for custom proxy hosts that do not
|
||||
# match the standard Mantle URL pattern. Also drop any caller Authorization
|
||||
# so _sign_request's restore-original-Authorization step cannot override
|
||||
# the SigV4 header.
|
||||
# Pin the credential-scope region to the region of the actual signing URL
|
||||
# so the SigV4 scope and URL host can never disagree, even when a stale
|
||||
# api_base and aws_region_name point at different regions.
|
||||
host_match = MANTLE_HOST_RE.match(api_base.rstrip("/"))
|
||||
optional_params = {
|
||||
**optional_params,
|
||||
|
|
@ -113,3 +112,36 @@ class BedrockMantleAuthMixin:
|
|||
"or pass api_key for Bearer auth, or provide AWS credentials "
|
||||
"(IAM role / access key / profile / web identity) for SigV4."
|
||||
) from e
|
||||
|
||||
|
||||
def mantle_supports_responses(model: str | None, model_cost: dict) -> bool:
|
||||
"""Whether a Bedrock Mantle model can serve the native Responses API.
|
||||
|
||||
Purely data-driven from the model's price-map capability signal -- either
|
||||
/v1/responses in supported_endpoints, or mode=responses -- both overridable
|
||||
via register_model and proxy model_info, so onboarding a model is a JSON
|
||||
change, never a code change. There is deliberately NO model-name match here:
|
||||
capability is per-model, not per-family (openai.gpt-oss-120b supports
|
||||
Responses while openai.gpt-oss-safeguard-120b does not, despite sharing the
|
||||
gpt-oss substring), so a substring gate would be wrong. A model absent from
|
||||
model_cost simply has no signal and returns False (chat-completions emulation).
|
||||
"""
|
||||
entry = model_cost.get(f"bedrock_mantle/{model}", {})
|
||||
if "/v1/responses" in (entry.get("supported_endpoints") or []):
|
||||
return True
|
||||
return entry.get("mode") == "responses"
|
||||
|
||||
|
||||
def mantle_base_segment(model: str | None, model_cost: dict) -> str:
|
||||
"""Return the base path segment for a Bedrock Mantle model's OpenAI surface.
|
||||
|
||||
Data-driven from the model's price-map use_openai_responses_path flag
|
||||
(overridable via register_model / proxy model_info). Per the AWS model cards,
|
||||
gpt-5.x and the google gemma-4-* family carry that flag and are served on the
|
||||
/openai/v1 base (.../openai/v1/responses and .../openai/v1/chat/completions);
|
||||
every other model including gpt-oss uses the standard /v1 base. The segment is
|
||||
the base for the model's whole OpenAI-compatible surface, so both the chat and
|
||||
responses configs derive from it -- there is no separate model-name rule.
|
||||
"""
|
||||
entry = model_cost.get(f"bedrock_mantle/{model}", {})
|
||||
return "openai/v1" if entry.get("use_openai_responses_path") is True else "v1"
|
||||
|
|
|
|||
|
|
@ -159,5 +159,14 @@
|
|||
"max_completion_tokens": "max_tokens"
|
||||
},
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"]
|
||||
},
|
||||
"pinstripes": {
|
||||
"base_url": "https://pinstripes.io/v1",
|
||||
"api_key_env": "PINSTRIPES_API_KEY",
|
||||
"api_base_env": "PINSTRIPES_API_BASE",
|
||||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
},
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses", "/v1/embeddings"]
|
||||
}
|
||||
}
|
||||
|
|
|
|||
3
litellm/llms/tinyfish/search/__init__.py
Normal file
3
litellm/llms/tinyfish/search/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig
|
||||
|
||||
__all__ = ["TinyfishSearchConfig"]
|
||||
164
litellm/llms/tinyfish/search/transformation.py
Normal file
164
litellm/llms/tinyfish/search/transformation.py
Normal file
|
|
@ -0,0 +1,164 @@
|
|||
"""
|
||||
TinyFish Search API.
|
||||
Endpoint: GET https://api.search.tinyfish.ai
|
||||
Docs: https://docs.tinyfish.ai/search-api
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Literal, TypedDict
|
||||
from urllib.parse import urlencode
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.search.transformation import (
|
||||
BaseSearchConfig,
|
||||
SearchResponse,
|
||||
SearchResult,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
class _TinyfishSearchRequestRequired(TypedDict):
|
||||
query: str
|
||||
|
||||
|
||||
class TinyfishSearchRequest(_TinyfishSearchRequestRequired, total=False):
|
||||
location: str
|
||||
language: str
|
||||
page: int
|
||||
include_thumbnail: bool
|
||||
max_results: int
|
||||
|
||||
|
||||
class _TinyfishResultItem(BaseModel, frozen=True):
|
||||
title: str = ""
|
||||
url: str = ""
|
||||
snippet: str = ""
|
||||
|
||||
|
||||
class _TinyfishApiResponse(BaseModel, frozen=True):
|
||||
results: tuple[_TinyfishResultItem, ...] = ()
|
||||
|
||||
|
||||
_UrlEncodableParams = TypeAdapter(dict[str, str | int | bool])
|
||||
_StrList = TypeAdapter(list[str])
|
||||
_StrFrozenSet = TypeAdapter(frozenset[str])
|
||||
|
||||
_TINYFISH_PARAMS_KEY = "_tinyfish_params"
|
||||
|
||||
|
||||
class TinyfishSearchConfig(BaseSearchConfig):
|
||||
TINYFISH_API_BASE = "https://api.search.tinyfish.ai"
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "TinyFish"
|
||||
|
||||
def get_http_method(self) -> Literal["GET", "POST"]:
|
||||
return "GET"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict[str, str],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
**kwargs: object,
|
||||
) -> dict[str, str]:
|
||||
resolved_key = api_key or get_secret_str("TINYFISH_API_KEY")
|
||||
if not resolved_key:
|
||||
raise ValueError(
|
||||
"TINYFISH_API_KEY is not set. Set `TINYFISH_API_KEY` environment variable."
|
||||
)
|
||||
return {**headers, "X-API-Key": resolved_key, "Accept": "application/json"}
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
optional_params: dict[str, object],
|
||||
data: dict[str, object] | list[dict[str, object]] | None = None,
|
||||
**kwargs: object,
|
||||
) -> str:
|
||||
resolved_base = (
|
||||
api_base or get_secret_str("TINYFISH_API_BASE") or self.TINYFISH_API_BASE
|
||||
)
|
||||
if isinstance(data, dict) and _TINYFISH_PARAMS_KEY in data:
|
||||
validated_params = _UrlEncodableParams.validate_python(
|
||||
data[_TINYFISH_PARAMS_KEY]
|
||||
)
|
||||
return f"{resolved_base}?{urlencode(validated_params, doseq=True)}"
|
||||
return resolved_base
|
||||
|
||||
def transform_search_request(
|
||||
self,
|
||||
query: str | list[str],
|
||||
optional_params: dict[str, object],
|
||||
**kwargs: object,
|
||||
) -> dict[str, object]:
|
||||
resolved_query = " ".join(query) if isinstance(query, list) else query
|
||||
|
||||
request_data: TinyfishSearchRequest = {"query": resolved_query}
|
||||
|
||||
country = optional_params.get("country")
|
||||
if isinstance(country, str):
|
||||
request_data["location"] = country
|
||||
|
||||
raw_max = optional_params.get("max_results")
|
||||
if isinstance(raw_max, (int, float, str)):
|
||||
request_data["max_results"] = max(1, min(int(raw_max), 20))
|
||||
|
||||
try:
|
||||
domains = _StrList.validate_python(
|
||||
optional_params.get("search_domain_filter")
|
||||
)
|
||||
except (ValidationError, TypeError):
|
||||
domains = []
|
||||
if domains:
|
||||
request_data["query"] = _append_domain_filters(
|
||||
request_data["query"], domains
|
||||
)
|
||||
|
||||
result_data: dict[str, object] = dict(request_data)
|
||||
|
||||
raw_supported: object = (
|
||||
self.get_supported_perplexity_optional_params() # any-ok: base class returns bare set
|
||||
)
|
||||
supported_perplexity = _StrFrozenSet.validate_python(raw_supported)
|
||||
for param, value in optional_params.items():
|
||||
if param not in supported_perplexity and param not in result_data:
|
||||
result_data[param] = value
|
||||
|
||||
return {_TINYFISH_PARAMS_KEY: result_data}
|
||||
|
||||
def transform_search_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj | None,
|
||||
**kwargs: object,
|
||||
) -> SearchResponse:
|
||||
raw_json: object = raw_response.json() # any-ok: httpx Response.json() -> Any
|
||||
parsed = _TinyfishApiResponse.model_validate(raw_json)
|
||||
|
||||
max_results_str: str = "20"
|
||||
if raw_response.request:
|
||||
raw_param: object = (
|
||||
raw_response.request.url.params.get( # any-ok: httpx QueryParams.get() -> Any
|
||||
"max_results", "20"
|
||||
)
|
||||
)
|
||||
max_results_str = str(raw_param)
|
||||
max_results: int = min(int(max_results_str), 20)
|
||||
|
||||
results = [
|
||||
SearchResult(title=item.title, url=item.url, snippet=item.snippet)
|
||||
for item in parsed.results[:max_results]
|
||||
]
|
||||
|
||||
return SearchResponse(results=results, object="search")
|
||||
|
||||
|
||||
def _append_domain_filters(query: str, domains: list[str]) -> str:
|
||||
domain_clauses = " OR ".join(f"site:{d}" for d in domains)
|
||||
return f"({query}) ({domain_clauses})"
|
||||
|
|
@ -271,7 +271,7 @@ def supports_response_json_schema(model: str) -> bool:
|
|||
|
||||
# Gemini 2.0+ and 2.5+ models support responseJsonSchema
|
||||
# Pattern matches: gemini-2.0-*, gemini-2.5-*, gemini-3-*, etc.
|
||||
gemini_2_plus_pattern = re.compile(r"gemini-([2-9]|[1-9]\d+)\.")
|
||||
gemini_2_plus_pattern = re.compile(r"gemini-(?:[2-9]|[1-9]\d+)(?:\.|\-)")
|
||||
|
||||
return bool(gemini_2_plus_pattern.search(model_lower))
|
||||
|
||||
|
|
|
|||
|
|
@ -42383,6 +42383,7 @@
|
|||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42397,6 +42398,7 @@
|
|||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42411,6 +42413,7 @@
|
|||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions"],
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
|
|
@ -42424,6 +42427,7 @@
|
|||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions"],
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
|
|
@ -42477,6 +42481,8 @@
|
|||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"use_openai_responses_path": true,
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": false,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42491,6 +42497,8 @@
|
|||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"use_openai_responses_path": true,
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": false,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42505,6 +42513,8 @@
|
|||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"use_openai_responses_path": true,
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": false,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42803,6 +42813,19 @@
|
|||
],
|
||||
"supports_audio_input": true
|
||||
},
|
||||
"soniox/stt-async-v5": {
|
||||
"litellm_provider": "soniox",
|
||||
"max_output_tokens": 8000,
|
||||
"max_tokens": 8000,
|
||||
"input_cost_per_second": 0.0,
|
||||
"output_cost_per_second": 0.0000277778,
|
||||
"mode": "audio_transcription",
|
||||
"source": "https://soniox.com/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/audio/transcriptions"
|
||||
],
|
||||
"supports_audio_input": true
|
||||
},
|
||||
"tensormesh/Qwen/Qwen3.5-397B-A17B-FP8": {
|
||||
"litellm_provider": "tensormesh",
|
||||
"mode": "chat",
|
||||
|
|
|
|||
|
|
@ -61,6 +61,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
|
||||
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_safe_get_request_headers,
|
||||
|
|
@ -725,6 +726,7 @@ async def common_checks(
|
|||
user_spend = await get_current_spend(
|
||||
counter_key=f"spend:user:{user_object.user_id}",
|
||||
fallback_spend=user_object.spend or 0.0,
|
||||
max_budget=user_budget,
|
||||
)
|
||||
if math.isfinite(user_budget) and user_spend >= user_budget:
|
||||
raise litellm.BudgetExceededError(
|
||||
|
|
@ -1127,6 +1129,8 @@ async def _check_end_user_budget(
|
|||
end_user_spend = await get_current_spend(
|
||||
counter_key=f"spend:end_user:{end_user_obj.user_id}",
|
||||
fallback_spend=end_user_obj.spend or 0.0,
|
||||
max_budget=end_user_budget,
|
||||
fallback_authoritative=True,
|
||||
)
|
||||
if end_user_spend > end_user_budget:
|
||||
raise litellm.BudgetExceededError(
|
||||
|
|
@ -3615,6 +3619,7 @@ async def _virtual_key_max_budget_check(
|
|||
spend = await get_current_spend(
|
||||
counter_key=counter_key,
|
||||
fallback_spend=fallback_spend,
|
||||
max_budget=valid_token.max_budget,
|
||||
)
|
||||
|
||||
####################################
|
||||
|
|
@ -3684,6 +3689,10 @@ async def _virtual_key_multi_budget_check(
|
|||
window_spend = await get_current_spend(
|
||||
counter_key=counter_key,
|
||||
fallback_spend=0.0,
|
||||
max_budget=w["max_budget"],
|
||||
window_entity_type="Key",
|
||||
window_entity_id=valid_token.token,
|
||||
window_start=get_budget_window_start(w),
|
||||
)
|
||||
if math.isfinite(w["max_budget"]) and window_spend >= w["max_budget"]:
|
||||
raise litellm.BudgetExceededError(
|
||||
|
|
@ -3938,6 +3947,7 @@ async def _check_team_member_budget(
|
|||
team_member_spend = await get_current_spend(
|
||||
counter_key=f"spend:team_member:{valid_token.user_id}:{team_object.team_id}",
|
||||
fallback_spend=team_member_spend,
|
||||
max_budget=team_member_budget,
|
||||
)
|
||||
|
||||
if (
|
||||
|
|
@ -4023,6 +4033,7 @@ async def _team_max_budget_check(
|
|||
spend = await get_current_spend(
|
||||
counter_key=f"spend:team:{team_object.team_id}",
|
||||
fallback_spend=team_object.spend or 0.0,
|
||||
max_budget=team_object.max_budget,
|
||||
)
|
||||
|
||||
if math.isfinite(team_object.max_budget) and spend > team_object.max_budget:
|
||||
|
|
@ -4072,6 +4083,10 @@ async def _team_multi_budget_check(
|
|||
window_spend = await get_current_spend(
|
||||
counter_key=counter_key,
|
||||
fallback_spend=0.0,
|
||||
max_budget=w["max_budget"],
|
||||
window_entity_type="Team",
|
||||
window_entity_id=team_object.team_id,
|
||||
window_start=get_budget_window_start(w),
|
||||
)
|
||||
if math.isfinite(w["max_budget"]) and window_spend >= w["max_budget"]:
|
||||
raise litellm.BudgetExceededError(
|
||||
|
|
@ -4377,6 +4392,7 @@ async def _organization_max_budget_check(
|
|||
org_spend = await get_current_spend(
|
||||
counter_key=f"spend:org:{org_id}",
|
||||
fallback_spend=org_table.spend or 0.0,
|
||||
max_budget=org_max_budget,
|
||||
)
|
||||
|
||||
# Check if organization spend exceeds max budget
|
||||
|
|
@ -4454,6 +4470,8 @@ async def _tag_max_budget_check(
|
|||
tag_spend = await get_current_spend(
|
||||
counter_key=f"spend:tag:{tag_name}",
|
||||
fallback_spend=tag_object.spend or 0.0,
|
||||
max_budget=tag_object.litellm_budget_table.max_budget,
|
||||
fallback_authoritative=True,
|
||||
)
|
||||
if tag_spend <= tag_object.litellm_budget_table.max_budget:
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -1837,6 +1837,7 @@ async def _user_api_key_auth_builder(
|
|||
team_member_spend = await get_current_spend(
|
||||
counter_key=f"spend:team_member:{valid_token.user_id}:{valid_token.team_id}",
|
||||
fallback_spend=team_member_spend,
|
||||
max_budget=team_member_budget,
|
||||
)
|
||||
if team_member_spend > team_member_budget:
|
||||
raise litellm.BudgetExceededError(
|
||||
|
|
|
|||
|
|
@ -2542,6 +2542,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
debug_enabled = verbose_proxy_logger.isEnabledFor(logging.DEBUG)
|
||||
stream_completed = False
|
||||
client_disconnected = False
|
||||
delivered_chunk = False
|
||||
try:
|
||||
str_so_far = ""
|
||||
async for (
|
||||
|
|
@ -2558,36 +2559,38 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"async_data_generator: received streaming chunk - %s", chunk
|
||||
)
|
||||
|
||||
if fast_path:
|
||||
yield serialize_chunk(chunk)
|
||||
continue
|
||||
if not fast_path:
|
||||
chunk = await proxy_logging_obj.async_post_call_streaming_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=chunk,
|
||||
data=request_data,
|
||||
str_so_far=str_so_far,
|
||||
)
|
||||
|
||||
chunk = await proxy_logging_obj.async_post_call_streaming_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=chunk,
|
||||
data=request_data,
|
||||
str_so_far=str_so_far,
|
||||
)
|
||||
if isinstance(chunk, (ModelResponse, ModelResponseStream)):
|
||||
response_str = litellm.get_response_string(response_obj=chunk)
|
||||
str_so_far += response_str
|
||||
elif hasattr(chunk, "model_dump"):
|
||||
try:
|
||||
d = chunk.model_dump(mode="json", exclude_none=True)
|
||||
if isinstance(d, dict):
|
||||
str_so_far += str(d.get("content", ""))
|
||||
except Exception:
|
||||
pass
|
||||
elif isinstance(chunk, dict):
|
||||
str_so_far += str(chunk.get("content", ""))
|
||||
|
||||
if isinstance(chunk, (ModelResponse, ModelResponseStream)):
|
||||
response_str = litellm.get_response_string(response_obj=chunk)
|
||||
str_so_far += response_str
|
||||
elif hasattr(chunk, "model_dump"):
|
||||
try:
|
||||
d = chunk.model_dump(mode="json", exclude_none=True)
|
||||
if isinstance(d, dict):
|
||||
str_so_far += str(d.get("content", ""))
|
||||
except Exception:
|
||||
pass
|
||||
elif isinstance(chunk, dict):
|
||||
str_so_far += str(chunk.get("content", ""))
|
||||
|
||||
model_name = request_data.get("model", "")
|
||||
chunk = (
|
||||
ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
|
||||
model_name = request_data.get("model", "")
|
||||
chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
|
||||
chunk, model_name
|
||||
)
|
||||
)
|
||||
|
||||
# Set before the yield: an async generator suspends at the yield,
|
||||
# so a GeneratorExit on client disconnect is raised there and any
|
||||
# statement after the yield never runs. The slow-path hook is
|
||||
# awaited above, so a cancellation during it still leaves this
|
||||
# False and refunds.
|
||||
delivered_chunk = True
|
||||
yield serialize_chunk(chunk)
|
||||
stream_completed = True
|
||||
except (asyncio.CancelledError, GeneratorExit):
|
||||
|
|
@ -2602,6 +2605,14 @@ class ProxyBaseLLMRequestProcessing:
|
|||
user_api_key_dict
|
||||
)
|
||||
client_disconnected = True
|
||||
if not delivered_chunk:
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
release_budget_reservation_on_cancel,
|
||||
)
|
||||
|
||||
await release_budget_reservation_on_cancel(
|
||||
getattr(user_api_key_dict, "budget_reservation", None)
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
|
|
|
|||
|
|
@ -10,7 +10,10 @@ url_to_redirect_to += "/login"
|
|||
new_ui_login_url = get_custom_url("", "ui/login")
|
||||
|
||||
|
||||
def build_ui_login_form(show_deprecation_banner: bool = False) -> str:
|
||||
def build_ui_login_form(
|
||||
show_deprecation_banner: bool = False,
|
||||
hide_default_credentials_hint: bool = False,
|
||||
) -> str:
|
||||
banner_html = (
|
||||
f"""
|
||||
<div class="deprecation-banner">
|
||||
|
|
@ -23,6 +26,25 @@ def build_ui_login_form(show_deprecation_banner: bool = False) -> str:
|
|||
else ""
|
||||
)
|
||||
|
||||
info_box_html = (
|
||||
""
|
||||
if hide_default_credentials_hint
|
||||
else """
|
||||
<div class="info-box">
|
||||
<div class="info-header">
|
||||
<svg xmlns="http://www.w3.org/2000/svg" width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
|
||||
<circle cx="12" cy="12" r="10"></circle>
|
||||
<line x1="12" y1="16" x2="12" y2="12"></line>
|
||||
<line x1="12" y1="8" x2="12.01" y2="8"></line>
|
||||
</svg>
|
||||
Default Credentials
|
||||
</div>
|
||||
<p>By default, Username is <code>admin</code> and Password is your set LiteLLM Proxy <code>MASTER_KEY</code>.</p>
|
||||
<p>Need to set UI credentials or SSO? <a href="https://docs.litellm.ai/docs/proxy/ui" target="_blank">Check the documentation</a>.</p>
|
||||
</div>
|
||||
"""
|
||||
)
|
||||
|
||||
return f"""
|
||||
<!DOCTYPE html>
|
||||
<html lang="en">
|
||||
|
|
@ -232,18 +254,7 @@ def build_ui_login_form(show_deprecation_banner: bool = False) -> str:
|
|||
</div>
|
||||
<h2>Login</h2>
|
||||
<p class="subtitle">Access your LiteLLM Admin UI.</p>
|
||||
<div class="info-box">
|
||||
<div class="info-header">
|
||||
<svg xmlns="http://www.w3.org/2000/svg" width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
|
||||
<circle cx="12" cy="12" r="10"></circle>
|
||||
<line x1="12" y1="16" x2="12" y2="12"></line>
|
||||
<line x1="12" y1="8" x2="12.01" y2="8"></line>
|
||||
</svg>
|
||||
Default Credentials
|
||||
</div>
|
||||
<p>By default, Username is <code>admin</code> and Password is your set LiteLLM Proxy <code>MASTER_KEY</code>.</p>
|
||||
<p>Need to set UI credentials or SSO? <a href="https://docs.litellm.ai/docs/proxy/ui" target="_blank">Check the documentation</a>.</p>
|
||||
</div>
|
||||
{info_box_html}
|
||||
<label for="username">Username<span class="required">*</span></label>
|
||||
<input type="text" id="username" name="username" required placeholder="Enter your username" autocomplete="username">
|
||||
|
||||
|
|
@ -264,6 +275,3 @@ def build_ui_login_form(show_deprecation_banner: bool = False) -> str:
|
|||
</body>
|
||||
</html>
|
||||
"""
|
||||
|
||||
|
||||
html_form = build_ui_login_form(show_deprecation_banner=True)
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import urllib
|
|||
import urllib.parse
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Dict, Optional, Union
|
||||
from typing import Any, Callable, Union
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
|
@ -31,7 +31,7 @@ class IAMEndpoint:
|
|||
port: str
|
||||
user: str
|
||||
name: str
|
||||
schema: Optional[str] = None
|
||||
schema: str | None = None
|
||||
|
||||
def build_url(self, token: str) -> str:
|
||||
url = f"postgresql://{self.user}:{token}@{self.host}:{self.port}/{self.name}"
|
||||
|
|
@ -53,7 +53,7 @@ def parse_iam_endpoint_from_url(url: str) -> IAMEndpoint:
|
|||
if not name:
|
||||
raise ValueError("Cannot parse IAM endpoint from URL: missing database name")
|
||||
port = str(parsed.port) if parsed.port else "5432"
|
||||
schema: Optional[str] = None
|
||||
schema: str | None = None
|
||||
if parsed.query:
|
||||
qs = urllib.parse.parse_qs(parsed.query)
|
||||
schema_vals = qs.get("schema")
|
||||
|
|
@ -94,7 +94,7 @@ class PrismaWrapper:
|
|||
iam_token_db_auth: bool,
|
||||
*,
|
||||
db_url_env_var: str = "DATABASE_URL",
|
||||
iam_endpoint: Optional[IAMEndpoint] = None,
|
||||
iam_endpoint: IAMEndpoint | None = None,
|
||||
recreate_uses_datasource: bool = False,
|
||||
log_prefix: str = "",
|
||||
):
|
||||
|
|
@ -116,9 +116,25 @@ class PrismaWrapper:
|
|||
self._log_prefix = f"{log_prefix} " if log_prefix else ""
|
||||
|
||||
# Background token refresh task management
|
||||
self._token_refresh_task: Optional[asyncio.Task] = None
|
||||
self._token_refresh_task: asyncio.Task | None = None
|
||||
self._reconnection_lock = asyncio.Lock()
|
||||
self._last_refresh_time: Optional[datetime] = None
|
||||
self._last_refresh_time: datetime | None = None
|
||||
|
||||
# Coordination for planned engine restarts (issue #29176). Every
|
||||
# `recreate_prisma_client` SIGTERMs the running query-engine on
|
||||
# purpose. The engine-death watcher (in `PrismaClient`) must be able
|
||||
# to tell that planned kill apart from a real crash, otherwise it
|
||||
# triggers its own reconnect and kills the freshly-spawned engine.
|
||||
# - `_expected_engine_deaths`: PIDs we intentionally killed; the
|
||||
# watcher consumes these instead of reconnecting.
|
||||
# - `_engine_generation`: monotonic counter bumped on every
|
||||
# successful recreate, used by callers as an optimistic-lock token
|
||||
# so racing/cascading recreates collapse into a single restart.
|
||||
# - `on_engine_replaced`: optional callback fired after a recreate so
|
||||
# the owner (PrismaClient) can re-arm its watcher on the new PID.
|
||||
self._expected_engine_deaths: set[int] = set()
|
||||
self._engine_generation: int = 0
|
||||
self.on_engine_replaced: Callable[[], None] | None = None
|
||||
|
||||
def _get_engine_pid(self) -> int:
|
||||
"""Get the PID of the current Prisma engine subprocess, or 0 if unavailable."""
|
||||
|
|
@ -167,7 +183,7 @@ class PrismaWrapper:
|
|||
except (ProcessLookupError, PermissionError, OSError):
|
||||
pass # Exited after SIGTERM — expected
|
||||
|
||||
def _extract_token_from_db_url(self, db_url: Optional[str]) -> Optional[str]:
|
||||
def _extract_token_from_db_url(self, db_url: str | None) -> str | None:
|
||||
"""
|
||||
Extract the token (password) from the DATABASE_URL.
|
||||
|
||||
|
|
@ -188,7 +204,7 @@ class PrismaWrapper:
|
|||
except Exception:
|
||||
return None
|
||||
|
||||
def _parse_token_expiration(self, token: Optional[str]) -> Optional[datetime]:
|
||||
def _parse_token_expiration(self, token: str | None) -> datetime | None:
|
||||
"""
|
||||
Parse the token to extract its expiration time.
|
||||
|
||||
|
|
@ -255,7 +271,7 @@ class PrismaWrapper:
|
|||
# If already past refresh time, return 0 (refresh immediately)
|
||||
return max(0, seconds_until_refresh)
|
||||
|
||||
def is_token_expired(self, token_url: Optional[str]) -> bool:
|
||||
def is_token_expired(self, token_url: str | None) -> bool:
|
||||
"""Check if the token in the given URL is expired."""
|
||||
if token_url is None:
|
||||
return True
|
||||
|
|
@ -272,7 +288,7 @@ class PrismaWrapper:
|
|||
|
||||
return datetime.utcnow() > expiration_time
|
||||
|
||||
def get_rds_iam_token(self) -> Optional[str]:
|
||||
def get_rds_iam_token(self) -> str | None:
|
||||
"""Generate a new RDS IAM token and update the configured DB URL env var.
|
||||
|
||||
When the wrapper was constructed with an explicit `iam_endpoint`
|
||||
|
|
@ -313,8 +329,12 @@ class PrismaWrapper:
|
|||
return _db_url
|
||||
|
||||
async def recreate_prisma_client(
|
||||
self, new_db_url: str, http_client: Optional[Any] = None
|
||||
):
|
||||
self,
|
||||
new_db_url: str,
|
||||
http_client: Any | None = None,
|
||||
*,
|
||||
expected_generation: int | None = None,
|
||||
) -> bool:
|
||||
"""Disconnect and reconnect the Prisma client with a new database URL.
|
||||
|
||||
Kills the old engine subprocess directly (SIGTERM → SIGKILL) rather than
|
||||
|
|
@ -327,14 +347,70 @@ class PrismaWrapper:
|
|||
the reader wrapper opts into `recreate_uses_datasource=True` so the
|
||||
new URL is passed explicitly via `datasource={"url": ...}` (Prisma
|
||||
does not auto-read alternate env vars like DATABASE_URL_READ_REPLICA).
|
||||
|
||||
Serializes all recreations through `self._reconnection_lock` so the
|
||||
IAM-refresh path and the engine-death/transport-error reconnect paths
|
||||
cannot recreate concurrently (issue #29176). `expected_generation`, if
|
||||
given, is an optimistic-lock token: when it no longer matches
|
||||
`self._engine_generation` once the lock is held, another path already
|
||||
replaced the engine, so this call is a no-op and returns ``False``.
|
||||
|
||||
Returns:
|
||||
bool: ``True`` if the client was actually recreated, ``False`` if
|
||||
the recreate was skipped because the engine generation moved on.
|
||||
"""
|
||||
async with self._reconnection_lock:
|
||||
return await self._recreate_prisma_client_locked(
|
||||
new_db_url,
|
||||
http_client=http_client,
|
||||
expected_generation=expected_generation,
|
||||
)
|
||||
|
||||
async def _recreate_prisma_client_locked(
|
||||
self,
|
||||
new_db_url: str,
|
||||
http_client: Any | None = None,
|
||||
*,
|
||||
expected_generation: int | None = None,
|
||||
) -> bool:
|
||||
"""Core recreate logic. Caller MUST hold `self._reconnection_lock`.
|
||||
|
||||
Split out so callers that already hold the lock (e.g.
|
||||
`_safe_refresh_token`, which double-checks token freshness under the
|
||||
lock) don't re-acquire it — `asyncio.Lock` is not reentrant.
|
||||
"""
|
||||
from prisma import Prisma # type: ignore
|
||||
|
||||
if (
|
||||
expected_generation is not None
|
||||
and expected_generation != self._engine_generation
|
||||
):
|
||||
verbose_proxy_logger.info(
|
||||
"%sSkipping Prisma client recreate: engine already replaced "
|
||||
"(generation %s != expected %s).",
|
||||
self._log_prefix,
|
||||
self._engine_generation,
|
||||
expected_generation,
|
||||
)
|
||||
return False
|
||||
|
||||
old_engine_pid = self._get_engine_pid()
|
||||
if old_engine_pid > 0:
|
||||
# Record BEFORE the kill so the engine-death watcher, which may
|
||||
# fire the instant the process dies, recognizes this as a planned
|
||||
# restart and does not launch its own reconnect.
|
||||
#
|
||||
# A stale entry can linger when the watcher re-arms on the new PID
|
||||
# before the old PID's death callback runs (the callback then
|
||||
# early-returns on PID mismatch without consuming it). Such entries
|
||||
# are harmless but would accumulate on a long-running proxy (~one
|
||||
# per IAM refresh), so cap the set — those old PIDs are long dead.
|
||||
if len(self._expected_engine_deaths) >= 64:
|
||||
self._expected_engine_deaths.clear()
|
||||
self._expected_engine_deaths.add(old_engine_pid)
|
||||
await self._kill_engine_process(old_engine_pid)
|
||||
|
||||
kwargs: Dict[str, Any] = {}
|
||||
kwargs: dict[str, Any] = {}
|
||||
if http_client is not None:
|
||||
kwargs["http"] = http_client
|
||||
if self._recreate_uses_datasource:
|
||||
|
|
@ -342,6 +418,15 @@ class PrismaWrapper:
|
|||
self._original_prisma = Prisma(**kwargs)
|
||||
|
||||
await self._original_prisma.connect()
|
||||
self._engine_generation += 1
|
||||
|
||||
# Let the owner (PrismaClient) re-arm its engine-death watcher on the
|
||||
# newly-spawned engine PID. Scheduled, never awaited, so a slow watcher
|
||||
# can't stall the refresh while we hold the reconnection lock.
|
||||
if self.on_engine_replaced is not None:
|
||||
self.on_engine_replaced()
|
||||
|
||||
return True
|
||||
|
||||
async def start_token_refresh_task(self) -> None:
|
||||
"""
|
||||
|
|
@ -441,9 +526,23 @@ class PrismaWrapper:
|
|||
preventing multiple concurrent reconnection attempts.
|
||||
"""
|
||||
async with self._reconnection_lock:
|
||||
# Double-checked under the lock: another trigger (e.g. the
|
||||
# proactive loop racing a __getattr__ fallback) may have already
|
||||
# refreshed while we waited. Recreating again would needlessly kill
|
||||
# the engine that refresh just spawned (issue #29176), so coalesce
|
||||
# by skipping when the current token still has comfortable runway.
|
||||
if self._token_refresh_not_needed(os.getenv(self._db_url_env_var)):
|
||||
verbose_proxy_logger.debug(
|
||||
"%sRDS IAM token still fresh; skipping redundant refresh.",
|
||||
self._log_prefix,
|
||||
)
|
||||
return
|
||||
|
||||
new_db_url = self.get_rds_iam_token()
|
||||
if new_db_url:
|
||||
await self.recreate_prisma_client(new_db_url)
|
||||
# We already hold `_reconnection_lock`; call the locked core
|
||||
# directly (the public method would re-acquire and deadlock).
|
||||
await self._recreate_prisma_client_locked(new_db_url)
|
||||
self._last_refresh_time = datetime.utcnow()
|
||||
verbose_proxy_logger.info(
|
||||
"%sRDS IAM token refreshed successfully. New token valid for ~15 minutes.",
|
||||
|
|
@ -455,6 +554,23 @@ class PrismaWrapper:
|
|||
self._log_prefix,
|
||||
)
|
||||
|
||||
def _token_refresh_not_needed(self, token_url: str | None) -> bool:
|
||||
"""True iff the token in ``token_url`` has more than the refresh buffer
|
||||
of runway left, so a refresh would be redundant.
|
||||
|
||||
Used to coalesce stacked refresh triggers. Deliberately mirrors the
|
||||
proactive loop's schedule (refresh at ``expiration - buffer``): a token
|
||||
with exactly ``buffer`` seconds left is NOT considered fresh, so the
|
||||
legitimate proactive refresh still fires. Unparseable tokens return
|
||||
``False`` (refresh) — skipping them would mean never refreshing.
|
||||
"""
|
||||
token = self._extract_token_from_db_url(token_url)
|
||||
expiration_time = self._parse_token_expiration(token)
|
||||
if expiration_time is None:
|
||||
return False
|
||||
seconds_left = (expiration_time - datetime.utcnow()).total_seconds()
|
||||
return seconds_left > self.TOKEN_REFRESH_BUFFER_SECONDS
|
||||
|
||||
def __getattr__(self, name: str):
|
||||
"""
|
||||
Proxy attribute access to the underlying Prisma client.
|
||||
|
|
@ -598,7 +714,7 @@ class PrismaManager:
|
|||
|
||||
|
||||
def should_update_prisma_schema(
|
||||
disable_updates: Optional[Union[bool, str]] = None,
|
||||
disable_updates: Union[bool, str] | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Determines if Prisma Schema updates should be applied during startup.
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ otherwise PrismaClient uses the writer-only PrismaWrapper directly.
|
|||
"""
|
||||
|
||||
import os
|
||||
from typing import Any, Callable, Optional
|
||||
from typing import Any, Callable
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||||
|
|
@ -117,7 +117,7 @@ class RoutingPrismaWrapper:
|
|||
)
|
||||
|
||||
async def disconnect(self, *args: Any, **kwargs: Any) -> None:
|
||||
first_error: Optional[BaseException] = None
|
||||
first_error: BaseException | None = None
|
||||
for client in (self._writer, self._reader):
|
||||
try:
|
||||
await client.disconnect(*args, **kwargs)
|
||||
|
|
@ -144,8 +144,12 @@ class RoutingPrismaWrapper:
|
|||
await self._reader.stop_token_refresh_task()
|
||||
|
||||
async def recreate_prisma_client(
|
||||
self, new_db_url: str, http_client: Optional[Any] = None
|
||||
) -> None:
|
||||
self,
|
||||
new_db_url: str,
|
||||
http_client: Any | None = None,
|
||||
*,
|
||||
expected_generation: int | None = None,
|
||||
) -> bool:
|
||||
"""Recreate both writer and reader Prisma clients.
|
||||
|
||||
The writer reconnect path in PrismaClient calls
|
||||
|
|
@ -155,8 +159,19 @@ class RoutingPrismaWrapper:
|
|||
the writer first (its URL is the one passed in), then best-effort
|
||||
recreate the reader. A reader failure flips `_reader_unavailable=True`
|
||||
so reads transparently fall through to the writer.
|
||||
|
||||
`expected_generation` is forwarded to the writer's optimistic-lock
|
||||
guard. If the writer recreate is skipped (another path already replaced
|
||||
the engine — issue #29176), we skip the reader too rather than churning
|
||||
it needlessly, and return ``False``.
|
||||
"""
|
||||
await self._writer.recreate_prisma_client(new_db_url, http_client=http_client)
|
||||
writer_recreated = await self._writer.recreate_prisma_client(
|
||||
new_db_url,
|
||||
http_client=http_client,
|
||||
expected_generation=expected_generation,
|
||||
)
|
||||
if not writer_recreated:
|
||||
return False
|
||||
try:
|
||||
await self._recreate_reader(http_client=http_client)
|
||||
self._reader_unavailable = False
|
||||
|
|
@ -167,8 +182,9 @@ class RoutingPrismaWrapper:
|
|||
"Reads will fall back to the writer until the reader recovers.",
|
||||
e,
|
||||
)
|
||||
return True
|
||||
|
||||
async def _recreate_reader(self, http_client: Optional[Any] = None) -> None:
|
||||
async def _recreate_reader(self, http_client: Any | None = None) -> None:
|
||||
"""Resolve the reader URL and recreate its Prisma client.
|
||||
|
||||
IAM-enabled readers regenerate their token (host/port/user came from
|
||||
|
|
|
|||
|
|
@ -0,0 +1,51 @@
|
|||
from typing import TYPE_CHECKING, Union
|
||||
|
||||
from litellm.types.guardrails import (
|
||||
GuardrailEventHooks,
|
||||
Mode,
|
||||
SupportedGuardrailIntegrations,
|
||||
)
|
||||
|
||||
from .repelloai import RepelloAIGuardrail
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
|
||||
def _event_hook_from_mode(
|
||||
mode: str | list[str] | Mode,
|
||||
) -> Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode]:
|
||||
if isinstance(mode, Mode):
|
||||
return mode
|
||||
if isinstance(mode, list):
|
||||
return [GuardrailEventHooks(item) for item in mode]
|
||||
return GuardrailEventHooks(mode)
|
||||
|
||||
|
||||
def initialize_guardrail(
|
||||
litellm_params: "LitellmParams", guardrail: "Guardrail"
|
||||
) -> RepelloAIGuardrail:
|
||||
import litellm
|
||||
|
||||
_repelloai_callback = RepelloAIGuardrail(
|
||||
guardrail_name=guardrail["guardrail_name"],
|
||||
api_key=litellm_params.api_key,
|
||||
api_base=litellm_params.api_base,
|
||||
asset_id=litellm_params.asset_id,
|
||||
unreachable_fallback=litellm_params.unreachable_fallback,
|
||||
event_hook=_event_hook_from_mode(litellm_params.mode),
|
||||
default_on=litellm_params.default_on or False,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_repelloai_callback)
|
||||
|
||||
return _repelloai_callback
|
||||
|
||||
|
||||
guardrail_initializer_registry = {
|
||||
SupportedGuardrailIntegrations.REPELLOAI.value: initialize_guardrail,
|
||||
}
|
||||
|
||||
|
||||
guardrail_class_registry = {
|
||||
SupportedGuardrailIntegrations.REPELLOAI.value: RepelloAIGuardrail,
|
||||
}
|
||||
613
litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py
Normal file
613
litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py
Normal file
|
|
@ -0,0 +1,613 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import AsyncGenerator, Literal
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import TypeGuard
|
||||
|
||||
from fastapi import HTTPException
|
||||
from httpx import HTTPError, Response as HttpxResponse
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType]
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header, # pyright: ignore[reportUnknownVariableType]
|
||||
)
|
||||
from litellm.proxy.guardrails._content_utils import build_inspection_messages
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.guardrails import GuardrailEventHooks, Mode
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.repelloai import (
|
||||
RepelloAIAnalyzeResponse,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
CallTypesLiteral,
|
||||
GuardrailStatus,
|
||||
LLMResponseTypes,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
)
|
||||
|
||||
DEFAULT_REPELLOAI_API_BASE = "https://argusapi.repello.ai/sdk/v1"
|
||||
DEFAULT_REPELLOAI_TIMEOUT = 30.0
|
||||
BLOCKED_VERDICT = "blocked"
|
||||
FLAGGED_VERDICT = "flagged"
|
||||
PASSED_VERDICT = "passed"
|
||||
|
||||
# Argus returns these for a permanently broken guardrail (bad key, unknown
|
||||
# asset_id, malformed payload), not a transient outage. They must always
|
||||
# block, never honour fail_open.
|
||||
CONFIG_ERROR_STATUS_CODES = frozenset({400, 401, 403, 404, 422})
|
||||
_SCHEMA_SCALAR_KEYS = frozenset(("name", "description", "title", "const", "default"))
|
||||
_SCHEMA_LIST_KEYS = frozenset(("enum", "examples"))
|
||||
_SCHEMA_EXTRACTED_KEYS = _SCHEMA_SCALAR_KEYS | _SCHEMA_LIST_KEYS
|
||||
|
||||
|
||||
class RepelloAIGuardrailMissingSecrets(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def _is_object_dict(value: object) -> TypeGuard[dict[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
|
||||
return isinstance(value, dict)
|
||||
|
||||
|
||||
def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
|
||||
return isinstance(value, list)
|
||||
|
||||
|
||||
class RepelloAIGuardrail(CustomGuardrail):
|
||||
@staticmethod
|
||||
def _get_field(obj: object, key: str) -> object:
|
||||
if _is_object_dict(obj):
|
||||
return obj.get(key)
|
||||
return getattr(obj, key, None)
|
||||
|
||||
@classmethod
|
||||
def _extract_tool_call_args_from_message(cls, message: object) -> list[str]:
|
||||
args: list[str] = []
|
||||
|
||||
tool_calls = cls._get_field(message, "tool_calls")
|
||||
if _is_object_list(tool_calls):
|
||||
for tool_call in tool_calls:
|
||||
function = cls._get_field(tool_call, "function")
|
||||
arguments = cls._get_field(function, "arguments")
|
||||
if isinstance(arguments, str) and arguments.strip():
|
||||
args.append(arguments)
|
||||
|
||||
function_call = cls._get_field(message, "function_call")
|
||||
arguments = cls._get_field(function_call, "arguments")
|
||||
if isinstance(arguments, str) and arguments.strip():
|
||||
args.append(arguments)
|
||||
|
||||
return args
|
||||
|
||||
@staticmethod
|
||||
def _iter_schema_text(node: object) -> list[str]:
|
||||
texts: list[str] = []
|
||||
stack: list[object] = [node]
|
||||
|
||||
while stack:
|
||||
current = stack.pop()
|
||||
if _is_object_dict(current):
|
||||
for key in _SCHEMA_SCALAR_KEYS:
|
||||
value = current.get(key)
|
||||
if isinstance(value, str) and value:
|
||||
texts.append(value)
|
||||
for key in _SCHEMA_LIST_KEYS:
|
||||
items = current.get(key)
|
||||
if _is_object_list(items):
|
||||
for item in items:
|
||||
if isinstance(item, str) and item:
|
||||
texts.append(item)
|
||||
remaining: list[object] = [
|
||||
v for k, v in current.items() if k not in _SCHEMA_EXTRACTED_KEYS
|
||||
]
|
||||
stack.extend(reversed(remaining))
|
||||
elif _is_object_list(current):
|
||||
stack.extend(reversed(current))
|
||||
|
||||
return texts
|
||||
|
||||
@classmethod
|
||||
def _extract_tool_definition_text(cls, data: dict[str, object]) -> list[str]:
|
||||
texts: list[str] = []
|
||||
|
||||
tools = data.get("tools")
|
||||
for tool in tools if _is_object_list(tools) else []:
|
||||
if not _is_object_dict(tool):
|
||||
continue
|
||||
function = tool.get("function")
|
||||
if _is_object_dict(function):
|
||||
texts.extend(cls._iter_schema_text(function))
|
||||
|
||||
functions = data.get("functions")
|
||||
for function in functions if _is_object_list(functions) else []:
|
||||
if _is_object_dict(function):
|
||||
texts.extend(cls._iter_schema_text(function))
|
||||
|
||||
return texts
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
asset_id: str | None = None,
|
||||
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
|
||||
guardrail_name: str | None = None,
|
||||
event_hook: (
|
||||
GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None
|
||||
) = None,
|
||||
default_on: bool = False,
|
||||
):
|
||||
self.repelloai_api_key = (
|
||||
api_key
|
||||
or get_secret_str("ARGUS_API_KEY")
|
||||
or get_secret_str("REPELLOAI_API_KEY")
|
||||
or ""
|
||||
)
|
||||
if not self.repelloai_api_key:
|
||||
raise RepelloAIGuardrailMissingSecrets(
|
||||
"Couldn't get Repello API key. Set `ARGUS_API_KEY` in the environment "
|
||||
"or pass `api_key` to the guardrail in the config file."
|
||||
)
|
||||
|
||||
self.asset_id = asset_id
|
||||
if not self.asset_id:
|
||||
raise ValueError(
|
||||
"Repello guardrail requires an `asset_id`. Create an asset in the Repello "
|
||||
"dashboard and set `asset_id` on the guardrail in the config file."
|
||||
)
|
||||
|
||||
self.api_base = (
|
||||
api_base
|
||||
or get_secret_str("REPELLOAI_API_BASE")
|
||||
or DEFAULT_REPELLOAI_API_BASE
|
||||
)
|
||||
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
|
||||
"fail_open" if unreachable_fallback == "fail_open" else "fail_closed"
|
||||
)
|
||||
self.async_handler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.GuardrailCallback,
|
||||
params={"timeout": DEFAULT_REPELLOAI_TIMEOUT},
|
||||
)
|
||||
super().__init__( # pyright: ignore[reportUnknownMemberType]
|
||||
guardrail_name=guardrail_name,
|
||||
event_hook=event_hook,
|
||||
default_on=default_on,
|
||||
)
|
||||
|
||||
async def _call_analyze(
|
||||
self,
|
||||
text: str,
|
||||
stage: Literal["prompt", "response"],
|
||||
request_data: dict[str, object],
|
||||
event_type: GuardrailEventHooks,
|
||||
) -> RepelloAIAnalyzeResponse | None:
|
||||
endpoint = f"{self.api_base}/analyze/{stage}"
|
||||
request: dict[str, object] = {
|
||||
"asset_id": self.asset_id or "",
|
||||
"scan_data": {stage: text},
|
||||
}
|
||||
|
||||
status: GuardrailStatus = "success"
|
||||
guardrail_json_response: str | dict[str, object] | list[dict[str, object]] = ""
|
||||
start_time: datetime = datetime.now()
|
||||
repelloai_response: RepelloAIAnalyzeResponse | None = None
|
||||
try:
|
||||
verbose_proxy_logger.debug("RepelloAI Argus request: %s", request)
|
||||
raw_response: HttpxResponse | None = (
|
||||
await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
|
||||
url=endpoint,
|
||||
headers={"X-API-Key": self.repelloai_api_key},
|
||||
json=request,
|
||||
)
|
||||
)
|
||||
if raw_response is None:
|
||||
raise ValueError("RepelloAI Argus returned no response")
|
||||
response: HttpxResponse = raw_response
|
||||
self._raise_for_config_error(response)
|
||||
response.raise_for_status()
|
||||
try:
|
||||
repelloai_response = TypeAdapter(
|
||||
RepelloAIAnalyzeResponse
|
||||
).validate_json(response.text)
|
||||
except ValidationError as e:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": "RepelloAI Argus guardrail returned invalid JSON",
|
||||
"status_code": response.status_code,
|
||||
},
|
||||
) from e
|
||||
verbose_proxy_logger.debug(
|
||||
"RepelloAI Argus response: %s", repelloai_response
|
||||
)
|
||||
if self._verdict_blocks(repelloai_response):
|
||||
status = "guardrail_intervened"
|
||||
return repelloai_response
|
||||
except HTTPException as e:
|
||||
status = "guardrail_failed_to_respond"
|
||||
guardrail_json_response = str(e.detail) if not isinstance(e.detail, (dict, list)) else e.detail # type: ignore[assignment]
|
||||
raise
|
||||
except HTTPError as e:
|
||||
status = "guardrail_failed_to_respond"
|
||||
guardrail_json_response = str(e)
|
||||
return self._handle_unreachable(e)
|
||||
except Exception as e:
|
||||
status = "guardrail_failed_to_respond"
|
||||
guardrail_json_response = str(e)
|
||||
raise HTTPException(
|
||||
status_code=500, detail={"error": "RepelloAI Argus guardrail failed"}
|
||||
) from e
|
||||
finally:
|
||||
end_time = datetime.now()
|
||||
if repelloai_response is not None:
|
||||
guardrail_json_response = dict(repelloai_response)
|
||||
self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType]
|
||||
guardrail_json_response=guardrail_json_response,
|
||||
guardrail_status=status,
|
||||
request_data=request_data,
|
||||
start_time=start_time.timestamp(),
|
||||
end_time=end_time.timestamp(),
|
||||
duration=(end_time - start_time).total_seconds(),
|
||||
masked_entity_count={},
|
||||
event_type=event_type,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _raise_for_config_error(response: HttpxResponse) -> None:
|
||||
if response.status_code in CONFIG_ERROR_STATUS_CODES:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": "RepelloAI Argus guardrail is misconfigured",
|
||||
"status_code": response.status_code,
|
||||
},
|
||||
)
|
||||
|
||||
def _verdict_blocks(
|
||||
self, repelloai_response: RepelloAIAnalyzeResponse | None
|
||||
) -> bool:
|
||||
if repelloai_response is None:
|
||||
return False
|
||||
verdict = repelloai_response.get("verdict")
|
||||
if verdict == BLOCKED_VERDICT:
|
||||
return True
|
||||
if verdict in (PASSED_VERDICT, FLAGGED_VERDICT):
|
||||
return False
|
||||
verbose_proxy_logger.warning(
|
||||
"RepelloAI Argus returned an unrecognized verdict (%s) - blocking.",
|
||||
verdict,
|
||||
)
|
||||
return True
|
||||
|
||||
def _handle_unreachable(self, error: Exception) -> RepelloAIAnalyzeResponse | None:
|
||||
verbose_proxy_logger.warning("RepelloAI Argus unreachable: %s", str(error))
|
||||
if self.unreachable_fallback == "fail_closed":
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": "RepelloAI Argus guardrail unreachable"},
|
||||
)
|
||||
return None
|
||||
|
||||
def _raise_if_blocked(
|
||||
self, repelloai_response: RepelloAIAnalyzeResponse | None
|
||||
) -> None:
|
||||
if repelloai_response is None:
|
||||
return
|
||||
if self._verdict_blocks(repelloai_response):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=self._format_blocked_detail(repelloai_response),
|
||||
)
|
||||
self._log_flagged_verdict(repelloai_response)
|
||||
|
||||
@classmethod
|
||||
def _format_blocked_detail(
|
||||
cls, repelloai_response: RepelloAIAnalyzeResponse
|
||||
) -> str:
|
||||
policies = repelloai_response.get("policies_violated")
|
||||
if not isinstance(policies, list) or not policies:
|
||||
return "Blocked by RepelloAI Argus guardrail."
|
||||
|
||||
formatted_policies: list[str] = []
|
||||
for policy in policies:
|
||||
policy_name = policy.get("policy_name") or "unknown_policy"
|
||||
details: list[str] = []
|
||||
action_taken = policy.get("action_taken")
|
||||
if action_taken:
|
||||
details.append(f"action: {action_taken}")
|
||||
policy_details = policy.get("details")
|
||||
if isinstance(policy_details, dict):
|
||||
score = policy_details.get("score")
|
||||
if score is not None:
|
||||
details.append(f"score: {score}")
|
||||
suffix = f" ({', '.join(details)})" if details else ""
|
||||
formatted_policies.append(f"{policy_name}{suffix}")
|
||||
|
||||
if not formatted_policies:
|
||||
return "Blocked by RepelloAI Argus guardrail."
|
||||
return f"Blocked by RepelloAI Argus guardrail. Policies violated: {'; '.join(formatted_policies)}."
|
||||
|
||||
@staticmethod
|
||||
def _log_flagged_verdict(repelloai_response: RepelloAIAnalyzeResponse) -> None:
|
||||
if repelloai_response.get("verdict") == FLAGGED_VERDICT:
|
||||
verbose_proxy_logger.warning(
|
||||
"RepelloAI Argus flagged content (allowed): %s",
|
||||
repelloai_response.get("policies_violated"),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _extract_prompt_message_text(data: dict[str, object]) -> list[str]:
|
||||
messages = build_inspection_messages(data)
|
||||
return [
|
||||
content
|
||||
for message in messages
|
||||
if isinstance(content := message.get("content"), str) and content
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _extract_input_text_parts(content: object) -> list[str]:
|
||||
if not _is_object_list(content):
|
||||
return []
|
||||
return [
|
||||
text
|
||||
for part in content
|
||||
if _is_object_dict(part) and part.get("type") == "input_text"
|
||||
if isinstance(text := part.get("text"), str) and text
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _extract_prompt_field_text(data: dict[str, object]) -> list[str]:
|
||||
prompt = data.get("prompt")
|
||||
if isinstance(prompt, str) and prompt:
|
||||
return [prompt]
|
||||
if _is_object_list(prompt):
|
||||
return [item for item in prompt if isinstance(item, str) and item]
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def _extract_prompt_text(cls, data: dict[str, object]) -> str | None:
|
||||
texts = cls._extract_prompt_message_text(data)
|
||||
texts.extend(cls._extract_prompt_field_text(data))
|
||||
|
||||
instructions = data.get("instructions")
|
||||
if isinstance(instructions, str) and instructions:
|
||||
texts.append(instructions)
|
||||
|
||||
raw_messages = data.get("messages")
|
||||
if _is_object_list(raw_messages):
|
||||
for message in raw_messages:
|
||||
texts.extend(cls._extract_tool_call_args_from_message(message))
|
||||
|
||||
raw_input = data.get("input")
|
||||
if _is_object_list(raw_input):
|
||||
for item in raw_input:
|
||||
if _is_object_dict(item):
|
||||
if "role" not in item:
|
||||
continue
|
||||
texts.extend(cls._extract_tool_call_args_from_message(item))
|
||||
texts.extend(cls._extract_input_text_parts(item.get("content")))
|
||||
|
||||
texts.extend(cls._extract_tool_definition_text(data))
|
||||
return "\n".join(text for text in texts if text) if texts else None
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: litellm.DualCache,
|
||||
data: dict[str, object],
|
||||
call_type: CallTypesLiteral,
|
||||
) -> Exception | str | dict[str, object] | None:
|
||||
verbose_proxy_logger.debug("RepelloAI Argus: pre_call_hook")
|
||||
|
||||
event_type = GuardrailEventHooks.pre_call
|
||||
if (
|
||||
self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType]
|
||||
data=data, event_type=event_type
|
||||
)
|
||||
is not True
|
||||
):
|
||||
return data
|
||||
|
||||
text = self._extract_prompt_text(data)
|
||||
if not text:
|
||||
verbose_proxy_logger.warning(
|
||||
"RepelloAI Argus: no inspectable prompt text in data - skipping."
|
||||
)
|
||||
return data
|
||||
|
||||
repelloai_response = await self._call_analyze(
|
||||
text=text,
|
||||
stage="prompt",
|
||||
request_data=data,
|
||||
event_type=event_type,
|
||||
)
|
||||
self._raise_if_blocked(repelloai_response)
|
||||
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
return data
|
||||
|
||||
async def async_post_call_success_hook(
|
||||
self,
|
||||
data: dict[str, object],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: LLMResponseTypes,
|
||||
):
|
||||
verbose_proxy_logger.debug("RepelloAI Argus: post_call_success_hook")
|
||||
|
||||
event_type = GuardrailEventHooks.post_call
|
||||
if (
|
||||
self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType]
|
||||
data=data, event_type=event_type
|
||||
)
|
||||
is not True
|
||||
):
|
||||
return response
|
||||
|
||||
text = self._extract_response_text(response)
|
||||
if not text:
|
||||
verbose_proxy_logger.warning(
|
||||
"RepelloAI Argus: no inspectable response text - skipping."
|
||||
)
|
||||
return response
|
||||
|
||||
repelloai_response = await self._call_analyze(
|
||||
text=text,
|
||||
stage="response",
|
||||
request_data=data,
|
||||
event_type=event_type,
|
||||
)
|
||||
self._raise_if_blocked(repelloai_response)
|
||||
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
return response
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: AsyncGenerator[ModelResponseStream, None],
|
||||
request_data: dict[str, object],
|
||||
) -> AsyncGenerator[ModelResponseStream, None]:
|
||||
from litellm import main as litellm_main
|
||||
|
||||
event_type = GuardrailEventHooks.post_call
|
||||
if (
|
||||
self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType]
|
||||
data=request_data, event_type=event_type
|
||||
)
|
||||
is not True
|
||||
):
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
chunks: list[ModelResponseStream] = []
|
||||
async for chunk in response:
|
||||
chunks.append(chunk)
|
||||
|
||||
assembled = litellm_main.stream_chunk_builder( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType]
|
||||
chunks=chunks
|
||||
)
|
||||
text = (
|
||||
self._extract_response_text(assembled)
|
||||
if isinstance(assembled, ModelResponse)
|
||||
else None
|
||||
)
|
||||
if text:
|
||||
repelloai_response = await self._call_analyze(
|
||||
text=text,
|
||||
stage="response",
|
||||
request_data=request_data,
|
||||
event_type=event_type,
|
||||
)
|
||||
if repelloai_response is not None:
|
||||
self._log_flagged_verdict(repelloai_response)
|
||||
if self._verdict_blocks(repelloai_response):
|
||||
from litellm.proxy.proxy_server import StreamingCallbackError
|
||||
|
||||
raise StreamingCallbackError("Blocked by RepelloAI Argus guardrail")
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=request_data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.warning(
|
||||
"RepelloAI Argus: no inspectable text in streamed response; skipping scan. "
|
||||
"guardrail=%s assembled_type=%s",
|
||||
self.guardrail_name,
|
||||
type(assembled).__name__,
|
||||
)
|
||||
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
@staticmethod
|
||||
def _extract_response_text(response: object) -> str | None:
|
||||
if _is_object_dict(response):
|
||||
response_dict = response
|
||||
elif isinstance(response, ModelResponse):
|
||||
response_dict = (
|
||||
response.model_dump() # pyright: ignore[reportUnknownMemberType]
|
||||
)
|
||||
else:
|
||||
output_text = getattr(response, "output_text", None)
|
||||
if isinstance(output_text, str) and output_text:
|
||||
return output_text
|
||||
response_dict = {}
|
||||
|
||||
text = RepelloAIGuardrail._extract_chat_completion_text(response_dict)
|
||||
if text:
|
||||
return text
|
||||
return RepelloAIGuardrail._extract_responses_api_text(response_dict)
|
||||
|
||||
@classmethod
|
||||
def _extract_chat_completion_text(
|
||||
cls, response_dict: dict[str, object]
|
||||
) -> str | None:
|
||||
choices = response_dict.get("choices")
|
||||
if not _is_object_list(choices):
|
||||
return None
|
||||
parts: list[str] = []
|
||||
for choice in choices:
|
||||
if not _is_object_dict(choice):
|
||||
continue
|
||||
message = choice.get("message")
|
||||
if _is_object_dict(message):
|
||||
content = message.get("content")
|
||||
if isinstance(content, str) and content:
|
||||
parts.append(content)
|
||||
parts.extend(cls._extract_tool_call_args_from_message(message))
|
||||
text = choice.get("text")
|
||||
if isinstance(text, str) and text:
|
||||
parts.append(text)
|
||||
return "\n".join(parts) if parts else None
|
||||
|
||||
@staticmethod
|
||||
def _extract_responses_api_text(response_dict: dict[str, object]) -> str | None:
|
||||
output = response_dict.get("output")
|
||||
if not _is_object_list(output):
|
||||
return None
|
||||
texts: list[str] = []
|
||||
for output_item in output:
|
||||
if not _is_object_dict(output_item):
|
||||
continue
|
||||
item_type = output_item.get("type")
|
||||
if item_type == "function_call":
|
||||
arguments = output_item.get("arguments")
|
||||
if isinstance(arguments, str) and arguments:
|
||||
texts.append(arguments)
|
||||
continue
|
||||
if item_type != "message":
|
||||
continue
|
||||
content = output_item.get("content")
|
||||
if not _is_object_list(content):
|
||||
continue
|
||||
for content_item in content:
|
||||
if not _is_object_dict(content_item):
|
||||
continue
|
||||
if content_item.get("type") not in ("output_text", "text"):
|
||||
continue
|
||||
text = content_item.get("text")
|
||||
if isinstance(text, str) and text:
|
||||
texts.append(text)
|
||||
return "".join(texts) if texts else None
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> type[GuardrailConfigModel[BaseModel]] | None:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.repelloai import (
|
||||
RepelloAIGuardrailConfigModel,
|
||||
)
|
||||
|
||||
return RepelloAIGuardrailConfigModel
|
||||
|
|
@ -401,7 +401,7 @@ def _resolve_health_check_max_tokens(
|
|||
3. For non-wildcard reasoning routes: BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING
|
||||
from env (if set)
|
||||
4. BACKGROUND_HEALTH_CHECK_MAX_TOKENS (global, any route including wildcards)
|
||||
5. Non-wildcard default: 5
|
||||
5. Non-wildcard default: 16
|
||||
6. Wildcard and nothing from (1)(4): leave unset (caller omits max_tokens)
|
||||
"""
|
||||
explicit = model_info.get("health_check_max_tokens", None)
|
||||
|
|
@ -432,7 +432,7 @@ def _resolve_health_check_max_tokens(
|
|||
return int(BACKGROUND_HEALTH_CHECK_MAX_TOKENS)
|
||||
|
||||
if not is_wildcard:
|
||||
return 5
|
||||
return 16
|
||||
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -162,9 +162,20 @@ class _ProxyDBLogger(CustomLogger):
|
|||
if obj_start is not None:
|
||||
actual_start_time = obj_start
|
||||
|
||||
# A stream that broke mid-flight still billed the provider for the
|
||||
# chunks already delivered. ``post_call_failure_hook`` lifts that
|
||||
# recovered cost onto request_data (the usage rides along in
|
||||
# ``combined_usage_object`` for the token columns), so attribute the
|
||||
# real partial spend to this failure row instead of zero.
|
||||
recovered_response_cost = 0.0
|
||||
if isinstance(request_data.get("combined_usage_object"), litellm.Usage):
|
||||
recovered_response_cost = max(
|
||||
float(request_data.get("response_cost") or 0.0), 0.0
|
||||
)
|
||||
|
||||
await proxy_logging_obj.db_spend_update_writer.update_database(
|
||||
token=user_api_key_dict.api_key,
|
||||
response_cost=0.0,
|
||||
response_cost=recovered_response_cost,
|
||||
user_id=user_api_key_dict.user_id,
|
||||
end_user_id=user_api_key_dict.end_user_id,
|
||||
team_id=user_api_key_dict.team_id,
|
||||
|
|
|
|||
|
|
@ -5122,7 +5122,7 @@ async def list_keys(
|
|||
size: int = Query(10, description="Page size", ge=1, le=100),
|
||||
user_id: Optional[str] = Query(
|
||||
None,
|
||||
description="Filter keys by user ID. Supports partial matching (substring, case-insensitive).",
|
||||
description="Filter keys by user ID. Exact match by default; set substring_matching=true (admin only) for case-insensitive substring matching.",
|
||||
),
|
||||
team_id: Optional[str] = Query(None, description="Filter keys by team ID"),
|
||||
organization_id: Optional[str] = Query(
|
||||
|
|
@ -5131,7 +5131,7 @@ async def list_keys(
|
|||
key_hash: Optional[str] = Query(None, description="Filter keys by key hash"),
|
||||
key_alias: Optional[str] = Query(
|
||||
None,
|
||||
description="Filter keys by key alias. Supports partial matching (substring, case-insensitive).",
|
||||
description="Filter keys by key alias. Exact match by default; set substring_matching=true (admin only) for case-insensitive substring matching.",
|
||||
),
|
||||
return_full_object: bool = Query(False, description="Return full key object"),
|
||||
include_team_keys: bool = Query(
|
||||
|
|
@ -5155,6 +5155,10 @@ async def list_keys(
|
|||
access_group_id: Optional[str] = Query(
|
||||
None, description="Filter keys by access group ID"
|
||||
),
|
||||
substring_matching: bool = Query(
|
||||
False,
|
||||
description="If true (proxy admins only), match user_id/key_alias as case-insensitive substrings instead of exact values. Defaults to false: /key/list matched these exactly before substring search was added, and an exact user_id/key_alias filter must never return another user's keys.",
|
||||
),
|
||||
) -> KeyListResponseObject:
|
||||
"""
|
||||
List all keys for a given user / team / organization.
|
||||
|
|
@ -5236,12 +5240,21 @@ async def list_keys(
|
|||
else:
|
||||
admin_team_ids = None
|
||||
|
||||
use_substring_matching = user_api_key_dict.user_role in [
|
||||
is_proxy_admin = user_api_key_dict.user_role in [
|
||||
LitellmUserRoles.PROXY_ADMIN.value,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value,
|
||||
]
|
||||
|
||||
if not user_id and not use_substring_matching:
|
||||
# Substring matching is opt-in (admin-only). /key/list matched user_id and
|
||||
# key_alias exactly before substring search was added; auto-applying a
|
||||
# substring match to every admin call broke that contract and let a caller
|
||||
# passing an exact user_id (e.g. an integration scoping to one user with an
|
||||
# admin key) receive other users' keys (user_id="alice" -> "alice2"). Exact
|
||||
# by default restores the prior behavior; the dashboard opts in explicitly.
|
||||
use_substring_matching = substring_matching and is_proxy_admin
|
||||
|
||||
# Admins may omit user_id to list all keys; non-admins are scoped to self.
|
||||
if not user_id and not is_proxy_admin:
|
||||
user_id = user_api_key_dict.user_id
|
||||
|
||||
response = await _list_key_helper(
|
||||
|
|
|
|||
|
|
@ -90,7 +90,7 @@ from litellm.proxy.common_utils.admin_ui_utils import (
|
|||
from litellm.proxy.common_utils.html_forms.jwt_display_template import (
|
||||
jwt_display_template,
|
||||
)
|
||||
from litellm.proxy.common_utils.html_forms.ui_login import html_form
|
||||
from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
|
||||
from litellm.proxy.management_endpoints.sso import CustomMicrosoftSSO
|
||||
|
|
@ -902,6 +902,7 @@ async def google_login(
|
|||
Example:
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
premium_user,
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
|
|
@ -948,7 +949,6 @@ async def google_login(
|
|||
missing_env_vars = show_missing_vars_in_env()
|
||||
if missing_env_vars is not None:
|
||||
return missing_env_vars
|
||||
ui_username = os.getenv("UI_USERNAME")
|
||||
|
||||
# get url from request - always use regular callback, but set state for CLI
|
||||
redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(
|
||||
|
|
@ -1009,16 +1009,20 @@ async def google_login(
|
|||
samesite="lax",
|
||||
)
|
||||
return sso_redirect
|
||||
elif ui_username is not None:
|
||||
# No Google, Microsoft SSO
|
||||
# Use UI Credentials set in .env
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
return HTMLResponse(content=html_form, status_code=200)
|
||||
else:
|
||||
from fastapi.responses import HTMLResponse
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
return HTMLResponse(content=html_form, status_code=200)
|
||||
hide_default_credentials_hint = (
|
||||
os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true"
|
||||
or general_settings.get("hide_default_credentials_hint", False) is True
|
||||
)
|
||||
return HTMLResponse(
|
||||
content=build_ui_login_form(
|
||||
show_deprecation_banner=True,
|
||||
hide_default_credentials_hint=hide_default_credentials_hint,
|
||||
),
|
||||
status_code=200,
|
||||
)
|
||||
|
||||
|
||||
def generic_response_convertor(
|
||||
|
|
|
|||
53
litellm/proxy/middleware/security_headers_middleware.py
Normal file
53
litellm/proxy/middleware/security_headers_middleware.py
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
"""
|
||||
Adds anti-framing / content-type security headers to every HTTP response.
|
||||
|
||||
X-Frame-Options and Content-Security-Policy: frame-ancestors 'none' stop the
|
||||
admin UI and login pages from being embedded cross-origin (clickjacking).
|
||||
X-Content-Type-Options: nosniff stops MIME sniffing.
|
||||
|
||||
Strict-Transport-Security is opt-in via LITELLM_ENABLE_HSTS because it only
|
||||
makes sense over HTTPS and would lock browsers out of plain-http deployments.
|
||||
|
||||
Headers are set with setdefault so a route that intentionally sets its own
|
||||
value is never overridden.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
from starlette.datastructures import MutableHeaders
|
||||
from starlette.types import ASGIApp, Message, Receive, Scope, Send
|
||||
|
||||
STATIC_SECURITY_HEADERS = (
|
||||
("X-Frame-Options", "DENY"),
|
||||
("Content-Security-Policy", "frame-ancestors 'none'"),
|
||||
("X-Content-Type-Options", "nosniff"),
|
||||
)
|
||||
HSTS_HEADER = ("Strict-Transport-Security", "max-age=31536000; includeSubDomains")
|
||||
|
||||
|
||||
def _hsts_enabled() -> bool:
|
||||
return os.getenv("LITELLM_ENABLE_HSTS", "false").strip().lower() == "true"
|
||||
|
||||
|
||||
class SecurityHeadersMiddleware:
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] != "http":
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
async def send_with_security_headers(message: Message) -> None:
|
||||
if message["type"] == "http.response.start":
|
||||
headers = MutableHeaders(scope=message)
|
||||
applied = (
|
||||
(*STATIC_SECURITY_HEADERS, HSTS_HEADER)
|
||||
if _hsts_enabled()
|
||||
else STATIC_SECURITY_HEADERS
|
||||
)
|
||||
for name, value in applied:
|
||||
headers.setdefault(name, value)
|
||||
await send(message)
|
||||
|
||||
await self.app(scope, receive, send_with_security_headers)
|
||||
|
|
@ -38,6 +38,9 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
|
|||
get_custom_llm_provider_from_request_headers,
|
||||
get_custom_llm_provider_from_request_query,
|
||||
)
|
||||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
is_managed_cloud_storage_uri,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
_is_base64_encoded_unified_file_id,
|
||||
encode_file_id_with_model,
|
||||
|
|
@ -726,6 +729,15 @@ async def get_file_content(
|
|||
}
|
||||
)
|
||||
else:
|
||||
# A raw cloud-storage URI (s3://, gs://) supplied here would skip the
|
||||
# managed-file owner/team check that only runs for unified ids, letting
|
||||
# a caller read another tenant's object by its key. Such objects are only
|
||||
# reachable through their managed unified id.
|
||||
if is_managed_cloud_storage_uri(file_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Raw cloud storage file ids cannot be retrieved directly. Use the LiteLLM managed file id returned when the file was created.",
|
||||
)
|
||||
# Check for model-based credential routing
|
||||
(
|
||||
should_route,
|
||||
|
|
|
|||
|
|
@ -8,6 +8,9 @@ import litellm
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_content_from_model_response,
|
||||
)
|
||||
from litellm.llms.anthropic import get_anthropic_config
|
||||
from litellm.llms.anthropic.chat.handler import (
|
||||
ModelResponseIterator as AnthropicModelResponseIterator,
|
||||
|
|
@ -136,6 +139,84 @@ class AnthropicPassthroughLoggingHandler:
|
|||
return model
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _stream_was_interrupted(
|
||||
all_chunks: Sequence[Union[str, bytes]],
|
||||
) -> bool:
|
||||
"""
|
||||
Anthropic ends a stream with ``content_block_stop`` -> ``message_delta``
|
||||
-> ``message_stop``; a client disconnect leaves the last event mid
|
||||
``content_block_delta``. Scan from the tail and decide on the first
|
||||
terminal-region event, so the common completed case is O(1) rather than
|
||||
re-deserializing every line of the stream.
|
||||
"""
|
||||
for raw in reversed(all_chunks):
|
||||
text = raw.decode("utf-8") if isinstance(raw, bytes) else raw
|
||||
for line in reversed(text.splitlines()):
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
try:
|
||||
data = json.loads(line[len("data:") :].strip())
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
continue
|
||||
if not isinstance(data, dict):
|
||||
continue
|
||||
etype = data.get("type")
|
||||
if etype == "message_delta":
|
||||
return False
|
||||
if etype in (
|
||||
"content_block_delta",
|
||||
"content_block_stop",
|
||||
"message_start",
|
||||
):
|
||||
return True
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _recover_interrupted_stream_output_tokens(
|
||||
response: Union[ModelResponse, TextCompletionResponse],
|
||||
all_chunks: Sequence[Union[str, bytes]],
|
||||
model: str,
|
||||
) -> None:
|
||||
"""
|
||||
An Anthropic stream interrupted before its terminal ``message_delta``
|
||||
(client disconnect) carries only the ``message_start`` ``output_tokens``
|
||||
placeholder (typically 1-3), so completion tokens and spend are
|
||||
undercounted ~20x. Re-tokenize the buffered output text to recover a
|
||||
realistic ``output_tokens`` for usage/cost. Completed streams are
|
||||
untouched because their terminal ``message_delta`` short-circuits here.
|
||||
"""
|
||||
if not isinstance(response, ModelResponse):
|
||||
return
|
||||
if not AnthropicPassthroughLoggingHandler._stream_was_interrupted(all_chunks):
|
||||
return
|
||||
usage = getattr(response, "usage", None)
|
||||
if usage is None:
|
||||
return
|
||||
output_text = get_content_from_model_response(response)
|
||||
if not output_text:
|
||||
return
|
||||
try:
|
||||
recovered_output_tokens = litellm.token_counter(
|
||||
model=model, text=output_text, count_response_tokens=True
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.warning(
|
||||
"Could not re-tokenize interrupted stream output; "
|
||||
"keeping placeholder completion token count."
|
||||
)
|
||||
return
|
||||
if recovered_output_tokens <= (usage.completion_tokens or 0):
|
||||
return
|
||||
usage.completion_tokens = recovered_output_tokens
|
||||
usage.total_tokens = (usage.prompt_tokens or 0) + recovered_output_tokens
|
||||
# Anthropic costing reads completion_tokens_details.text_tokens, so the
|
||||
# stale message_start placeholder there must be corrected too or spend
|
||||
# stays undercounted even after completion_tokens is fixed.
|
||||
details = getattr(usage, "completion_tokens_details", None)
|
||||
if details is not None and getattr(details, "text_tokens", None) is not None:
|
||||
details.text_tokens = recovered_output_tokens
|
||||
|
||||
@staticmethod
|
||||
def _create_anthropic_response_logging_payload(
|
||||
litellm_model_response: Union[ModelResponse, TextCompletionResponse],
|
||||
|
|
@ -277,6 +358,11 @@ class AnthropicPassthroughLoggingHandler:
|
|||
"result": None,
|
||||
"kwargs": {},
|
||||
}
|
||||
AnthropicPassthroughLoggingHandler._recover_interrupted_stream_output_tokens(
|
||||
response=complete_streaming_response,
|
||||
all_chunks=all_chunks,
|
||||
model=model,
|
||||
)
|
||||
kwargs = AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload(
|
||||
litellm_model_response=complete_streaming_response,
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -426,6 +426,9 @@ from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMi
|
|||
from litellm.proxy.middleware.request_size_limit_middleware import (
|
||||
RequestSizeLimitMiddleware,
|
||||
)
|
||||
from litellm.proxy.middleware.security_headers_middleware import (
|
||||
SecurityHeadersMiddleware,
|
||||
)
|
||||
from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
router as openai_files_router,
|
||||
|
|
@ -840,37 +843,6 @@ async def proxy_startup_event(app: FastAPI):
|
|||
if isinstance(worker_config, dict):
|
||||
await initialize(**worker_config)
|
||||
|
||||
## V2 OTEL: now that config (and therefore the callbacks) is loaded, publish
|
||||
## the chosen V2 logger's TracerProvider as the OTel global. The FastAPI
|
||||
## instrumentation mounted at app-creation binds to the global provider, so
|
||||
## this is what makes server spans and gen-ai spans share one provider and
|
||||
## land in the same trace. Prefer an already-registered preset logger
|
||||
## (arize, langfuse, …) so server spans export to that backend too; otherwise
|
||||
## build a generic one from OTEL_* envs. ``set_tracer_provider`` only takes
|
||||
## effect once, so the first configured logger wins.
|
||||
try:
|
||||
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
||||
|
||||
if is_otel_v2_enabled():
|
||||
from opentelemetry import trace as _otel_trace
|
||||
|
||||
from litellm.integrations.otel.logger import OpenTelemetryV2
|
||||
|
||||
_otel_v2_logger = (
|
||||
next(
|
||||
(
|
||||
cb
|
||||
for cb in litellm.service_callback
|
||||
if isinstance(cb, OpenTelemetryV2)
|
||||
),
|
||||
None,
|
||||
)
|
||||
or OpenTelemetryV2()
|
||||
)
|
||||
_otel_trace.set_tracer_provider(_otel_v2_logger._tracer_provider)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Skipping OTel V2 provider setup: %s", e)
|
||||
|
||||
# check if DATABASE_URL in environment - load from there
|
||||
if prisma_client is None:
|
||||
_db_url: Optional[str] = get_secret("DATABASE_URL", None) # type: ignore
|
||||
|
|
@ -907,6 +879,42 @@ async def proxy_startup_event(app: FastAPI):
|
|||
redis_usage_cache=transaction_buffer_redis_cache,
|
||||
)
|
||||
|
||||
## V2 OTEL: publish the chosen V2 logger's TracerProvider as the OTel global.
|
||||
## This MUST run after callback initialization above: a preset (arize, langfuse,
|
||||
## …) builds its logger there, folding the OTEL_* base exporter and its own
|
||||
## exporter into one logger. The FastAPI instrumentation mounted at app-creation
|
||||
## binds to the global provider, so reusing that one logger is what makes the
|
||||
## server span and the gen-ai spans share one provider and land in the same
|
||||
## trace, exporting to every configured backend. Running before callback init
|
||||
## (when no logger exists yet) would build a second, generic logger whose
|
||||
## provider became the global, orphaning the gen-ai spans onto a different
|
||||
## backend than the server span. A generic logger is built only when none was
|
||||
## configured.
|
||||
try:
|
||||
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
||||
|
||||
if is_otel_v2_enabled():
|
||||
from opentelemetry import trace as _otel_trace
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import _in_memory_loggers
|
||||
from litellm.integrations.otel.logger import (
|
||||
OpenTelemetryV2,
|
||||
publish_global_otel_v2_provider,
|
||||
)
|
||||
|
||||
registered = (
|
||||
open_telemetry_logger
|
||||
if isinstance(open_telemetry_logger, OpenTelemetryV2)
|
||||
else None
|
||||
)
|
||||
publish_global_otel_v2_provider(
|
||||
_in_memory_loggers, # any-ok: pre-existing untyped List[Any] global
|
||||
_otel_trace.set_tracer_provider,
|
||||
registered=registered,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Skipping OTel V2 provider setup: %s", e)
|
||||
|
||||
## Validate use_redis_transaction_buffer requires Redis cache ##
|
||||
ProxyStartupEvent._validate_redis_transaction_buffer_config(
|
||||
general_settings=general_settings,
|
||||
|
|
@ -1757,6 +1765,7 @@ app.add_middleware(
|
|||
|
||||
app.add_middleware(PrometheusAuthMiddleware)
|
||||
app.add_middleware(InFlightRequestsMiddleware)
|
||||
app.add_middleware(SecurityHeadersMiddleware)
|
||||
|
||||
|
||||
def mount_swagger_ui():
|
||||
|
|
@ -2026,7 +2035,43 @@ def cost_tracking():
|
|||
)
|
||||
|
||||
|
||||
async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
|
||||
# Bounds authoritative DB re-reads when enforcing a budget against a
|
||||
# stale-low spend counter: at most one DB read per counter per window.
|
||||
SPEND_DB_FLOOR_CACHE_TTL_SECONDS = 5
|
||||
|
||||
|
||||
def _fail_closed_budget_enforcement() -> bool:
|
||||
return general_settings.get("fail_closed_budget_enforcement") is True
|
||||
|
||||
|
||||
def _raise_budget_unverifiable(counter_key: str) -> None:
|
||||
verbose_proxy_logger.warning(
|
||||
"fail_closed_budget_enforcement: rejecting request — spend for %s could "
|
||||
"not be verified against Redis or the database",
|
||||
counter_key,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail={
|
||||
"error": (
|
||||
"Budget enforcement unavailable: current spend could not be "
|
||||
"verified against Redis or the database, and "
|
||||
"fail_closed_budget_enforcement is enabled, so the request was "
|
||||
"rejected to avoid exceeding the configured budget. Retry shortly."
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def get_current_spend(
|
||||
counter_key: str,
|
||||
fallback_spend: float,
|
||||
max_budget: float | None = None,
|
||||
window_entity_type: str | None = None,
|
||||
window_entity_id: str | None = None,
|
||||
window_start: datetime | None = None,
|
||||
fallback_authoritative: bool = False,
|
||||
) -> float:
|
||||
"""
|
||||
Read current spend from the cross-pod spend counter.
|
||||
|
||||
|
|
@ -2040,7 +2085,168 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
|
|||
2. In-memory counter (single-instance or Redis failure)
|
||||
3. Reseed from authoritative DB spend (counter expired, cross-pod stale)
|
||||
4. Caller-supplied fallback (DB unavailable, cold start)
|
||||
|
||||
When ``max_budget`` is supplied, the counter is re-checked against the
|
||||
authoritative recorded spend before a request is admitted. A Redis counter
|
||||
that survived a Redis restart can return a stale-low value loaded from an
|
||||
older RDB snapshot; that read is a hit (not a clean miss), so step 3 never
|
||||
runs and a key can leak spend past ``max_budget`` indefinitely. The
|
||||
authoritative source depends on the counter: primary key/team/user/org
|
||||
counters read the DB row; per-window counters (``window_start`` supplied)
|
||||
aggregate spend logs; end-user/tag counters have no DB row, so the caller's
|
||||
``fallback_spend`` (loaded fresh in auth) is authoritative. The DB read is
|
||||
skipped for healthy primary counters (counter at or above recorded spend)
|
||||
and cached in-process for a few seconds, so a persistently stale counter
|
||||
drives at most one read per counter per window rather than one per request.
|
||||
"""
|
||||
current, verified = await _read_spend_counter_estimate(
|
||||
counter_key=counter_key, fallback_spend=fallback_spend
|
||||
)
|
||||
if fallback_authoritative:
|
||||
verified = True
|
||||
|
||||
if max_budget is None or current >= max_budget:
|
||||
return current
|
||||
|
||||
# Cheap staleness signal for primary counters: the counter reads below the
|
||||
# spend this caller already knows about. Window counters have no such signal
|
||||
# (fallback is 0), so they always re-check, bounded by the cache. Strict mode
|
||||
# (fail_closed_budget_enforcement) always re-checks against the authoritative
|
||||
# source too, so a counter that is stale-low at the same time as the caller's
|
||||
# cached spend cannot slip through; the 5s cache keeps that bounded.
|
||||
is_window = window_start is not None
|
||||
if fallback_spend > current or is_window or _fail_closed_budget_enforcement():
|
||||
authoritative = await _authoritative_floor_spend(
|
||||
counter_key=counter_key,
|
||||
window_entity_type=window_entity_type,
|
||||
window_entity_id=window_entity_id,
|
||||
window_start=window_start,
|
||||
)
|
||||
if authoritative is not None:
|
||||
verified = True
|
||||
if authoritative > current:
|
||||
await _repair_stale_spend_counter(
|
||||
counter_key=counter_key, db_spend=authoritative
|
||||
)
|
||||
return authoritative
|
||||
elif fallback_spend > current:
|
||||
# end-user / tag counters have no DB row; fallback_spend is the
|
||||
# authoritative recorded value loaded in auth.
|
||||
return fallback_spend
|
||||
|
||||
# Opt-in hard guarantee: when the spend backing this admit decision came
|
||||
# only from a per-pod cache (Redis and DB both unreadable), reject rather
|
||||
# than admit on an unverifiable budget. No-op unless the flag is set, so
|
||||
# default behavior is unchanged.
|
||||
if not verified and _fail_closed_budget_enforcement():
|
||||
_raise_budget_unverifiable(counter_key)
|
||||
|
||||
return current
|
||||
|
||||
|
||||
async def _repair_stale_spend_counter(counter_key: str, db_spend: float) -> None:
|
||||
"""Raise a counter that has fallen below the authoritative DB spend (e.g.
|
||||
Redis restarted and reloaded an older snapshot) so every worker reads the
|
||||
corrected value directly instead of re-deriving it per request, and so a
|
||||
worker whose own cached spend is also stale still sees the true total.
|
||||
|
||||
The write is monotonic: it only ever raises the counter, so a repair that
|
||||
carries a slightly-stale DB total cannot clobber a concurrent increment that
|
||||
already pushed the counter higher (which would let racing requests
|
||||
under-count). Redis enforces this atomically via async_set_max; the
|
||||
in-memory copy is guarded by a read-compare-write with no await in between,
|
||||
so it is atomic within the worker.
|
||||
"""
|
||||
cached = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
|
||||
needs_update = True
|
||||
if cached is not None:
|
||||
try:
|
||||
needs_update = float(cached) < db_spend
|
||||
except (TypeError, ValueError):
|
||||
needs_update = True
|
||||
if needs_update:
|
||||
spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=db_spend)
|
||||
if spend_counter_cache.redis_cache is not None:
|
||||
try:
|
||||
await spend_counter_cache.redis_cache.async_set_max(
|
||||
key=counter_key, value=db_spend
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.debug(
|
||||
"Unable to repair stale spend counter %s in Redis",
|
||||
counter_key,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
async def reseed_spend_counter_from_db(counter_key: str) -> None:
|
||||
"""Recover a counter that the reservation reconcile found in an inconsistent
|
||||
state (missing, or where applying the reconcile delta would drive it
|
||||
negative) by reseeding it from the DB instead of deleting it.
|
||||
|
||||
The DB row is a LAGGING authoritative floor, not post-request truth: the
|
||||
entity .spend column is flushed in batches (every PROXY_BATCH_WRITE_AT), so
|
||||
it can exclude this request's just-recorded cost and other buffered spend.
|
||||
That is fine here: the monotonic set-max can only RAISE a stale-low counter
|
||||
toward that floor (never lowers it or clobbers a concurrent increment), and
|
||||
the read-time floor (_authoritative_floor_spend) converges to the true total
|
||||
as the buffer flushes. The point is to restore enforcement to a real floor
|
||||
rather than leave the counter deleted and unenforced (the prior fail-open).
|
||||
Counters with no DB row (window/end-user/tag) are left untouched rather than
|
||||
deleted, so enforcement keeps reading whatever value they hold.
|
||||
"""
|
||||
db_spend = await SpendCounterReseed.from_db(
|
||||
prisma_client=prisma_client, counter_key=counter_key
|
||||
)
|
||||
if db_spend is None:
|
||||
return
|
||||
await _repair_stale_spend_counter(counter_key=counter_key, db_spend=db_spend)
|
||||
|
||||
|
||||
async def _authoritative_floor_spend(
|
||||
counter_key: str,
|
||||
window_entity_type: str | None = None,
|
||||
window_entity_id: str | None = None,
|
||||
window_start: datetime | None = None,
|
||||
) -> float | None:
|
||||
marker_key = f"spend_db_floor:{counter_key}"
|
||||
cached = spend_counter_cache.in_memory_cache.get_cache(key=marker_key)
|
||||
if cached is not None:
|
||||
return float(cached)
|
||||
|
||||
db_spend = await SpendCounterReseed.from_db(
|
||||
prisma_client=prisma_client, counter_key=counter_key
|
||||
)
|
||||
if (
|
||||
db_spend is None
|
||||
and window_entity_type is not None
|
||||
and window_entity_id is not None
|
||||
and window_start is not None
|
||||
):
|
||||
db_spend = await SpendCounterReseed.window_from_spend_logs(
|
||||
prisma_client=prisma_client,
|
||||
entity_type=window_entity_type,
|
||||
entity_id=window_entity_id,
|
||||
window_start=window_start,
|
||||
)
|
||||
if db_spend is None:
|
||||
return None
|
||||
|
||||
spend_counter_cache.in_memory_cache.set_cache(
|
||||
key=marker_key,
|
||||
value=db_spend,
|
||||
ttl=SPEND_DB_FLOOR_CACHE_TTL_SECONDS,
|
||||
)
|
||||
return db_spend
|
||||
|
||||
|
||||
async def _read_spend_counter_estimate(
|
||||
counter_key: str, fallback_spend: float
|
||||
) -> tuple[float, bool]:
|
||||
"""Return (spend, authoritative). ``authoritative`` is True when the value
|
||||
came from Redis or a fresh DB read (cross-pod truth), False when it came
|
||||
from the per-pod in-memory copy or the caller's fallback. Only the
|
||||
fail-closed path reads the flag; normal callers ignore it."""
|
||||
# 1. Redis first (cross-pod authoritative). On clean miss, skip
|
||||
# in-memory: per-pod in-memory only has this pod's writes, so it
|
||||
# would mask cross-pod increments.
|
||||
|
|
@ -2049,7 +2255,7 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
|
|||
try:
|
||||
val = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key)
|
||||
if val is not None:
|
||||
return float(val)
|
||||
return float(val), True
|
||||
redis_clean_miss = True
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -2062,7 +2268,7 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
|
|||
if not redis_clean_miss:
|
||||
val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
|
||||
if val is not None:
|
||||
return float(val)
|
||||
return float(val), False
|
||||
|
||||
# 3. Reseed from DB - fallback_spend lags cross-pod, would allow bypass.
|
||||
db_spend = await SpendCounterReseed.coalesced(
|
||||
|
|
@ -2071,10 +2277,10 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
|
|||
counter_key=counter_key,
|
||||
)
|
||||
if db_spend is not None:
|
||||
return db_spend
|
||||
return db_spend, True
|
||||
|
||||
# 4. Caller-supplied fallback (DB unavailable).
|
||||
return fallback_spend
|
||||
return fallback_spend, False
|
||||
|
||||
|
||||
async def increment_spend_counters(
|
||||
|
|
@ -8758,7 +8964,7 @@ async def chat_completion(
|
|||
completion_stream=_iterator,
|
||||
model=e.model,
|
||||
custom_llm_provider="cached_response",
|
||||
logging_obj=data.get("litellm_logging_obj", None),
|
||||
logging_obj=_data.get("litellm_logging_obj", None),
|
||||
)
|
||||
selected_data_generator = select_data_generator(
|
||||
response=_streaming_response,
|
||||
|
|
@ -8793,7 +8999,7 @@ async def chat_completion(
|
|||
completion_stream=_iterator,
|
||||
model=data.get("model", ""),
|
||||
custom_llm_provider="cached_response",
|
||||
logging_obj=data.get("litellm_logging_obj", None),
|
||||
logging_obj=_data.get("litellm_logging_obj", None),
|
||||
)
|
||||
selected_data_generator = select_data_generator(
|
||||
response=_streaming_response,
|
||||
|
|
@ -12940,6 +13146,9 @@ async def model_info_v1(
|
|||
# use internal routing keys (model_name_{team_id}_{uuid}) and were omitted
|
||||
# when v1 resolved models only via public model_name strings.
|
||||
all_models: List[dict] = copy.deepcopy(llm_router.model_list)
|
||||
alias_models = copy.deepcopy(llm_router.get_model_list_from_model_alias())
|
||||
all_models.extend(alias_models)
|
||||
|
||||
allowed_model_names = _get_v1_model_info_allowed_model_names(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
llm_router=llm_router,
|
||||
|
|
@ -13507,26 +13716,24 @@ async def fallback_login(request: Request):
|
|||
|
||||
# get url from request
|
||||
redirect_url = get_custom_url(str(request.base_url))
|
||||
ui_username = os.getenv("UI_USERNAME")
|
||||
if redirect_url.endswith("/"):
|
||||
redirect_url += "sso/callback"
|
||||
else:
|
||||
redirect_url += "/sso/callback"
|
||||
|
||||
if ui_username is not None:
|
||||
# No Google, Microsoft SSO
|
||||
# Use UI Credentials set in .env
|
||||
from fastapi.responses import HTMLResponse
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
return HTMLResponse(
|
||||
content=build_ui_login_form(show_deprecation_banner=False), status_code=200
|
||||
)
|
||||
else:
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
return HTMLResponse(
|
||||
content=build_ui_login_form(show_deprecation_banner=False), status_code=200
|
||||
)
|
||||
hide_default_credentials_hint = (
|
||||
os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true"
|
||||
or general_settings.get("hide_default_credentials_hint", False) is True
|
||||
)
|
||||
return HTMLResponse(
|
||||
content=build_ui_login_form(
|
||||
show_deprecation_banner=False,
|
||||
hide_default_credentials_hint=hide_default_credentials_hint,
|
||||
),
|
||||
status_code=200,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
|
|
|
|||
|
|
@ -2658,7 +2658,7 @@
|
|||
"default_value": null
|
||||
}
|
||||
],
|
||||
"default_model_placeholder": "soniox/stt-async-v4"
|
||||
"default_model_placeholder": "soniox/stt-async-v5"
|
||||
},
|
||||
{
|
||||
"provider": "TEXT_COMPLETION_CODESTRAL",
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
|
@ -162,10 +163,14 @@ async def reserve_budget_for_request(
|
|||
if not applied_entries:
|
||||
return None
|
||||
|
||||
input_cost = estimate_request_input_cost(
|
||||
request_body=request_body, route=route, llm_router=llm_router
|
||||
)
|
||||
return {
|
||||
"reserved_cost": reservation_cost,
|
||||
"entries": applied_entries,
|
||||
"finalized": False,
|
||||
"input_cost": min(float(input_cost or 0.0), reservation_cost),
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -195,6 +200,41 @@ async def release_budget_reservation(budget_reservation: Optional[dict]) -> None
|
|||
)
|
||||
|
||||
|
||||
async def release_budget_reservation_on_cancel(
|
||||
budget_reservation: dict | None,
|
||||
) -> None:
|
||||
"""Reconcile a still-open reservation when the request is cancelled mid-flight.
|
||||
|
||||
A client disconnect or timeout cancels the request task, which surfaces as
|
||||
CancelledError / GeneratorExit rather than a normal exception, so neither the
|
||||
success cost callback nor the failure hook runs and the pre-call reservation
|
||||
is never reconciled. Left alone it pins the spend counter above real spend
|
||||
and 429s subsequent requests until the counter's TTL expires.
|
||||
|
||||
Reconcile to the request's input-token cost rather than refunding to zero:
|
||||
by the time a request is cancelled in-flight the provider call was already
|
||||
dispatched, so the input tokens were billed even if no chunk reached the
|
||||
client. Refunding to zero would let a caller abort pre-token to dodge that
|
||||
charge; the worst-case output portion of the reservation is still released.
|
||||
|
||||
asyncio.shield keeps the reconcile running to completion even though the
|
||||
surrounding task is being cancelled. The `finalized` guard makes this a no-op
|
||||
when success/failure handling already reconciled, so calling it on every
|
||||
cancellation path is safe.
|
||||
"""
|
||||
if not budget_reservation or budget_reservation.get("finalized") is True:
|
||||
return
|
||||
incurred_cost = float(budget_reservation.get("input_cost") or 0.0)
|
||||
try:
|
||||
await asyncio.shield(
|
||||
reconcile_budget_reservation(
|
||||
budget_reservation=budget_reservation, actual_cost=incurred_cost
|
||||
)
|
||||
)
|
||||
except (asyncio.CancelledError, Exception):
|
||||
pass
|
||||
|
||||
|
||||
async def invalidate_budget_reservation_counters(
|
||||
budget_reservation: Optional[dict],
|
||||
) -> None:
|
||||
|
|
@ -628,12 +668,14 @@ async def _set_reserved_entries_actual_cost(
|
|||
entries: List[dict],
|
||||
actual_cost: float,
|
||||
default_reserved_cost: float,
|
||||
reseed_on_inconsistent: bool = True,
|
||||
) -> None:
|
||||
for entry in entries:
|
||||
await _set_reserved_entry_actual_cost(
|
||||
entry=entry,
|
||||
actual_cost=actual_cost,
|
||||
default_reserved_cost=default_reserved_cost,
|
||||
reseed_on_inconsistent=reseed_on_inconsistent,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -641,8 +683,12 @@ async def _set_reserved_entry_actual_cost(
|
|||
entry: dict,
|
||||
actual_cost: float,
|
||||
default_reserved_cost: float,
|
||||
reseed_on_inconsistent: bool = True,
|
||||
) -> None:
|
||||
from litellm.proxy.proxy_server import _increment_spend_counter_cache
|
||||
from litellm.proxy.proxy_server import (
|
||||
_increment_spend_counter_cache,
|
||||
reseed_spend_counter_from_db,
|
||||
)
|
||||
|
||||
counter_key = entry.get("counter_key")
|
||||
if counter_key is None:
|
||||
|
|
@ -656,46 +702,49 @@ async def _set_reserved_entry_actual_cost(
|
|||
adjustment = target_adjustment - applied_adjustment
|
||||
if adjustment == 0:
|
||||
return
|
||||
await _ensure_counter_can_apply_adjustment(
|
||||
if await _counter_can_apply_adjustment(
|
||||
counter_key=counter_key,
|
||||
adjustment=adjustment,
|
||||
)
|
||||
await _increment_spend_counter_cache(
|
||||
counter_key=counter_key,
|
||||
increment=adjustment,
|
||||
)
|
||||
):
|
||||
await _increment_spend_counter_cache(
|
||||
counter_key=counter_key,
|
||||
increment=adjustment,
|
||||
)
|
||||
elif reseed_on_inconsistent:
|
||||
# Post-call reconcile / release: the counter was flushed or reseeded
|
||||
# between reservation and reconcile (Redis restart / cross-pod reset),
|
||||
# so the optimistic delta no longer applies. Recover by reseeding from
|
||||
# the DB's lagging authoritative floor rather than deleting the counter
|
||||
# and failing open — deleting it is what left budgets unenforced after a
|
||||
# Redis reload.
|
||||
await reseed_spend_counter_from_db(counter_key=counter_key)
|
||||
else:
|
||||
# Pre-call admission resize: the in-flight reservation cost is not yet
|
||||
# persisted, so the DB floor would discard it. Keep the original
|
||||
# fail-closed behavior (raise -> reserve_budget_for_request releases and
|
||||
# denies) rather than admitting against an inconsistent counter.
|
||||
raise RuntimeError(
|
||||
f"Cannot resize budget reservation against inconsistent counter {counter_key}"
|
||||
)
|
||||
entry["applied_adjustment"] = target_adjustment
|
||||
|
||||
|
||||
async def _ensure_counter_can_apply_adjustment(
|
||||
async def _counter_can_apply_adjustment(
|
||||
counter_key: str,
|
||||
adjustment: float,
|
||||
) -> None:
|
||||
from litellm.proxy.proxy_server import (
|
||||
_invalidate_spend_counter,
|
||||
spend_counter_cache,
|
||||
)
|
||||
) -> bool:
|
||||
from litellm.proxy.proxy_server import spend_counter_cache
|
||||
|
||||
current_value = await spend_counter_cache.async_get_cache(key=counter_key)
|
||||
if current_value is None:
|
||||
await _invalidate_spend_counter(counter_key=counter_key)
|
||||
raise RuntimeError(
|
||||
f"Cannot apply budget reservation adjustment to missing counter {counter_key}"
|
||||
)
|
||||
return False
|
||||
|
||||
try:
|
||||
current_float = float(current_value)
|
||||
except (TypeError, ValueError):
|
||||
await _invalidate_spend_counter(counter_key=counter_key)
|
||||
raise RuntimeError(
|
||||
f"Cannot apply budget reservation adjustment to non-numeric counter {counter_key}"
|
||||
)
|
||||
return False
|
||||
|
||||
if adjustment < 0 and current_float + adjustment < -1e-12:
|
||||
await _invalidate_spend_counter(counter_key=counter_key)
|
||||
raise RuntimeError(
|
||||
f"Budget reservation adjustment would make counter negative {counter_key}"
|
||||
)
|
||||
return not (adjustment < 0 and current_float + adjustment < -1e-12)
|
||||
|
||||
|
||||
async def _release_applied_entries_best_effort(
|
||||
|
|
@ -735,6 +784,7 @@ async def _resize_applied_reservation(
|
|||
entries=entries,
|
||||
actual_cost=new_reserved_cost,
|
||||
default_reserved_cost=current_reserved_cost,
|
||||
reseed_on_inconsistent=False,
|
||||
)
|
||||
for entry in entries:
|
||||
entry["reserved_cost"] = new_reserved_cost
|
||||
|
|
@ -817,6 +867,61 @@ def estimate_request_max_cost(
|
|||
return max(cast(List[float], estimates))
|
||||
|
||||
|
||||
def estimate_request_input_cost(
|
||||
request_body: dict,
|
||||
route: str,
|
||||
llm_router: Router | None,
|
||||
) -> float | None:
|
||||
"""Cost of the request's input tokens alone.
|
||||
|
||||
Once the provider request is dispatched the input tokens are billed even if
|
||||
the client disconnects before the first chunk, so this is the cost floor a
|
||||
cancelled in-flight request has already incurred. A cancelled reservation is
|
||||
reconciled to this instead of being refunded to zero.
|
||||
"""
|
||||
model = get_model_from_request(request_body, route, llm_router=llm_router)
|
||||
if model is None:
|
||||
return None
|
||||
|
||||
models = [model] if isinstance(model, str) else model
|
||||
estimates = [
|
||||
_estimate_request_input_cost_for_model(
|
||||
request_body=request_body,
|
||||
route=route,
|
||||
model=model_name,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
for model_name in models
|
||||
]
|
||||
estimates = [estimate for estimate in estimates if estimate is not None]
|
||||
if not estimates:
|
||||
return None
|
||||
return max(cast("list[float]", estimates))
|
||||
|
||||
|
||||
def _estimate_request_input_cost_for_model(
|
||||
request_body: dict,
|
||||
route: str,
|
||||
model: str,
|
||||
llm_router: Router | None,
|
||||
) -> float | None:
|
||||
model_info = _get_model_cost_info(model=model, llm_router=llm_router)
|
||||
if model_info is None:
|
||||
return None
|
||||
input_cost_per_token = _to_float(model_info.get("input_cost_per_token"))
|
||||
if input_cost_per_token is None:
|
||||
return None
|
||||
input_tokens = _estimate_input_tokens(
|
||||
request_body=request_body,
|
||||
route=route,
|
||||
model=model,
|
||||
model_info=model_info,
|
||||
)
|
||||
if input_tokens is None:
|
||||
return None
|
||||
return input_tokens * input_cost_per_token
|
||||
|
||||
|
||||
def _estimate_request_max_cost_for_model(
|
||||
request_body: dict,
|
||||
route: str,
|
||||
|
|
|
|||
|
|
@ -263,6 +263,13 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
|
|||
elif isinstance(_usage, dict):
|
||||
usage = _usage
|
||||
|
||||
# A request that failed mid-stream has no usable response_obj usage, but the
|
||||
# streaming handler may have recovered the usage from the chunks already
|
||||
# delivered. Honor that override so the partial usage lands in spend tracking.
|
||||
_combined_usage = kwargs.get("combined_usage_object")
|
||||
if not usage and isinstance(_combined_usage, litellm.Usage):
|
||||
usage = _combined_usage.model_dump()
|
||||
|
||||
id = get_spend_logs_id(call_type or "acompletion", response_obj_dict, kwargs)
|
||||
standard_logging_payload = cast(
|
||||
Optional[StandardLoggingPayload], kwargs.get("standard_logging_object", None)
|
||||
|
|
|
|||
|
|
@ -14,13 +14,11 @@ from litellm.proxy._types import *
|
|||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.repositories.config_repository import ConfigRepository
|
||||
from litellm.repositories.table_repositories import (
|
||||
DailyTagSpendRepository,
|
||||
SSOConfigRepository,
|
||||
UISettingsRepository,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.ui_sso import (
|
||||
DefaultTeamSSOParams,
|
||||
InProductNudgeResponse,
|
||||
SSOConfig,
|
||||
)
|
||||
|
||||
|
|
@ -178,11 +176,6 @@ class UISettings(BaseModel):
|
|||
description="If true, org admins cannot generate API keys via /key/generate.",
|
||||
)
|
||||
|
||||
disable_ui_nudges: bool = Field(
|
||||
default=False,
|
||||
description="If true, suppresses in-product UI nudges (survey and Claude Code feedback popups) for all users.",
|
||||
)
|
||||
|
||||
|
||||
class UISettingsResponse(SettingsResponse):
|
||||
"""Response model for UI settings"""
|
||||
|
|
@ -206,7 +199,6 @@ ALLOWED_UI_SETTINGS_FIELDS = {
|
|||
"scope_user_search_to_org",
|
||||
"disable_custom_api_keys",
|
||||
"disable_key_generate_for_org_admin",
|
||||
"disable_ui_nudges",
|
||||
}
|
||||
|
||||
# Flags that must be synced from the persisted UISettings into
|
||||
|
|
@ -1117,34 +1109,6 @@ async def update_mcp_semantic_filter_settings(
|
|||
return result
|
||||
|
||||
|
||||
@router.get(
|
||||
"/in_product_nudges",
|
||||
tags=["UI Settings"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=InProductNudgeResponse,
|
||||
)
|
||||
async def get_in_product_nudges():
|
||||
"""
|
||||
Get in-product nudges configuration.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": "Database not connected. Please connect a database."},
|
||||
)
|
||||
|
||||
db_record = await DailyTagSpendRepository(prisma_client).table.find_first(
|
||||
where={"tag": "User-Agent: claude-cli"}
|
||||
)
|
||||
|
||||
if db_record:
|
||||
return InProductNudgeResponse(is_claude_code_enabled=True)
|
||||
|
||||
return InProductNudgeResponse(is_claude_code_enabled=False)
|
||||
|
||||
|
||||
UI_SETTINGS_CACHE_KEY = "ui_settings:settings_dict"
|
||||
UI_SETTINGS_CACHE_TTL = 600 # 10 minutes
|
||||
|
||||
|
|
|
|||
|
|
@ -2128,12 +2128,21 @@ class ProxyLogging:
|
|||
# compute preprocessing latency after the logging object is popped.
|
||||
_logging_obj = request_data.get("litellm_logging_obj")
|
||||
if _logging_obj is not None:
|
||||
_first_handoff = getattr(_logging_obj, "model_call_details", {}).get(
|
||||
"first_api_call_start_time"
|
||||
)
|
||||
_model_call_details = getattr(_logging_obj, "model_call_details", {})
|
||||
_first_handoff = _model_call_details.get("first_api_call_start_time")
|
||||
if _first_handoff is not None:
|
||||
request_data["first_api_call_start_time"] = _first_handoff
|
||||
|
||||
# A stream that broke mid-flight still billed the provider for the
|
||||
# chunks already delivered; the streaming handler stashes that
|
||||
# recovered usage and cost here. Lift them onto request_data so the
|
||||
# failure-path spend callbacks (which run after the logging object
|
||||
# is popped) record the real partial spend instead of zero.
|
||||
_recovered_usage = _model_call_details.get("combined_usage_object")
|
||||
if _recovered_usage is not None:
|
||||
request_data["combined_usage_object"] = _recovered_usage
|
||||
request_data["response_cost"] = _model_call_details.get("response_cost")
|
||||
|
||||
# Remove before callbacks iterate — not serialisable
|
||||
request_data.pop("litellm_logging_obj", None)
|
||||
|
||||
|
|
@ -4453,6 +4462,14 @@ class PrismaClient:
|
|||
"prisma-query-engine PID %s already dead at watch start.",
|
||||
pid,
|
||||
)
|
||||
if self._consume_expected_death(pid):
|
||||
verbose_proxy_logger.info(
|
||||
"PID %s death was planned (engine already replaced); "
|
||||
"not reconnecting.",
|
||||
pid,
|
||||
)
|
||||
self._cleanup_engine_watcher()
|
||||
return True
|
||||
self._engine_confirmed_dead = True
|
||||
self._reap_all_zombies()
|
||||
self._cleanup_engine_watcher()
|
||||
|
|
@ -4497,12 +4514,39 @@ class PrismaClient:
|
|||
except RuntimeError:
|
||||
pass
|
||||
|
||||
def _consume_expected_death(self, pid: int) -> bool:
|
||||
"""True iff ``pid`` was killed on purpose by a planned recreate.
|
||||
|
||||
`PrismaWrapper.recreate_prisma_client` records the old engine PID in
|
||||
`_expected_engine_deaths` before SIGTERM-ing it (IAM token refresh,
|
||||
guarded reconnect). When the watcher then sees that PID die, this lets
|
||||
it recognize the death as planned and skip its own reconnect, which
|
||||
would otherwise kill the engine the recreate just spawned (#29176).
|
||||
|
||||
Consumes (removes) the PID so a later real crash of a reused PID is
|
||||
still handled. Tolerant of `self.db` stand-ins (tests / older clients)
|
||||
that don't expose a real set.
|
||||
"""
|
||||
expected = getattr(self.db, "_expected_engine_deaths", None)
|
||||
if isinstance(expected, set) and pid in expected:
|
||||
expected.discard(pid)
|
||||
return True
|
||||
return False
|
||||
|
||||
def _on_engine_death_from_thread(self, dead_pid: int) -> None:
|
||||
"""Called on the event loop thread when the waitpid thread detects engine death."""
|
||||
if self._engine_confirmed_dead:
|
||||
return
|
||||
if dead_pid != self._engine_pid:
|
||||
return
|
||||
if self._consume_expected_death(dead_pid):
|
||||
verbose_proxy_logger.info(
|
||||
"prisma-query-engine PID %s exited as part of a planned restart; "
|
||||
"not reconnecting (engine already replaced).",
|
||||
dead_pid,
|
||||
)
|
||||
self._cleanup_engine_watcher()
|
||||
return
|
||||
verbose_proxy_logger.error(
|
||||
"prisma-query-engine PID %s exited (waitpid thread); triggering reconnect.",
|
||||
dead_pid,
|
||||
|
|
@ -4557,6 +4601,14 @@ class PrismaClient:
|
|||
self._engine_pidfd = -1
|
||||
return
|
||||
dead_pid = self._engine_pid
|
||||
if self._consume_expected_death(dead_pid):
|
||||
verbose_proxy_logger.info(
|
||||
"prisma-query-engine PID %s exited (pidfd event) as part of a "
|
||||
"planned restart; not reconnecting (engine already replaced).",
|
||||
dead_pid,
|
||||
)
|
||||
self._cleanup_engine_watcher()
|
||||
return
|
||||
verbose_proxy_logger.error(
|
||||
"prisma-query-engine PID %s exited (pidfd event); triggering reconnect.",
|
||||
dead_pid,
|
||||
|
|
@ -4580,9 +4632,18 @@ class PrismaClient:
|
|||
try:
|
||||
os.kill(self._engine_pid, 0)
|
||||
except ProcessLookupError:
|
||||
dead_pid = self._engine_pid
|
||||
if self._consume_expected_death(dead_pid):
|
||||
verbose_proxy_logger.info(
|
||||
"prisma-query-engine PID %s gone as part of a planned "
|
||||
"restart; not reconnecting (engine already replaced).",
|
||||
dead_pid,
|
||||
)
|
||||
self._cleanup_engine_watcher()
|
||||
return
|
||||
verbose_proxy_logger.error(
|
||||
"prisma-query-engine PID %s gone; triggering reconnect.",
|
||||
self._engine_pid,
|
||||
dead_pid,
|
||||
)
|
||||
self._engine_confirmed_dead = True
|
||||
self._reap_all_zombies()
|
||||
|
|
@ -4669,6 +4730,22 @@ class PrismaClient:
|
|||
self._engine_confirmed_dead = False
|
||||
verbose_proxy_logger.debug("Stopped engine process watcher.")
|
||||
|
||||
def _handle_writer_engine_replaced(self) -> None:
|
||||
"""Re-arm the engine watcher after a planned writer-engine restart.
|
||||
|
||||
Wired as `PrismaWrapper.on_engine_replaced` and invoked from inside
|
||||
`recreate_prisma_client` once the new engine is connected (IAM token
|
||||
refresh, guarded reconnect). The old watcher was tracking the engine
|
||||
we just intentionally killed, so we tear it down and re-arm on the new
|
||||
PID. Scheduling `_start_engine_watcher` as a task (rather than awaiting)
|
||||
keeps us from blocking the recreate while it still holds the wrapper's
|
||||
reconnection lock. Without this re-arm, a planned restart would leave
|
||||
the proxy with no engine-death detection until the next reconnect.
|
||||
"""
|
||||
self._engine_confirmed_dead = False
|
||||
self._cleanup_engine_watcher()
|
||||
asyncio.create_task(self._start_engine_watcher())
|
||||
|
||||
async def _run_reconnect_cycle(
|
||||
self, timeout_seconds: Optional[float] = None
|
||||
) -> None:
|
||||
|
|
@ -4689,6 +4766,17 @@ class PrismaClient:
|
|||
else self._db_watchdog_reconnect_timeout_seconds
|
||||
)
|
||||
|
||||
# Snapshot the writer's engine generation BEFORE any await. Both
|
||||
# reconnect branches forward it to recreate_prisma_client as an
|
||||
# optimistic-lock token: if a concurrent IAM token refresh replaces the
|
||||
# engine after this point, the generation moves and the recreate becomes
|
||||
# a no-op instead of killing the engine the refresh just spawned
|
||||
# (#29176). Captured here — atomically with the dead-engine decision
|
||||
# below — rather than inside the reconnect closures, because those run
|
||||
# after an `asyncio.wait_for(...)` yield during which a refresh could
|
||||
# otherwise slip in and bump the very generation the closure then reads.
|
||||
expected_generation = getattr(self.writer_db, "_engine_generation", None)
|
||||
|
||||
engine_is_dead = self._engine_confirmed_dead or (
|
||||
self._engine_pid > 0 and not self._is_engine_alive()
|
||||
)
|
||||
|
|
@ -4709,7 +4797,16 @@ class PrismaClient:
|
|||
"DATABASE_URL not set; cannot recreate Prisma client."
|
||||
)
|
||||
raise RuntimeError("DATABASE_URL not set")
|
||||
await self.db.recreate_prisma_client(db_url)
|
||||
# Forward the entry-snapshot generation. The engine was
|
||||
# confirmed dead, but a concurrent IAM refresh may have already
|
||||
# respawned it; the guard makes this recreate a no-op in that
|
||||
# case rather than killing the fresh engine (#29176). Unlike the
|
||||
# direct path there is no SELECT 1 probe here, so the generation
|
||||
# guard is the only thing standing between a crash-reconnect and
|
||||
# a refresh that raced it.
|
||||
await self.db.recreate_prisma_client(
|
||||
db_url, expected_generation=expected_generation
|
||||
)
|
||||
await self._start_engine_watcher()
|
||||
|
||||
await asyncio.wait_for(_do_heavy_reconnect(), timeout=effective_timeout)
|
||||
|
|
@ -4731,13 +4828,36 @@ class PrismaClient:
|
|||
"DATABASE_URL not set; cannot reconnect Prisma client."
|
||||
)
|
||||
raise RuntimeError("DATABASE_URL not set")
|
||||
# Probe the writer BEFORE recreating. A concurrent IAM token
|
||||
# refresh may have just replaced the engine (issue #29176); if
|
||||
# the writer answers SELECT 1 the connection is already healthy
|
||||
# and recreating would needlessly kill that fresh engine. If we
|
||||
# do recreate, the entry-snapshot generation lets the wrapper
|
||||
# detect a refresh that landed since cycle entry and skip the
|
||||
# redundant restart.
|
||||
writer = self.writer_db
|
||||
try:
|
||||
await writer.query_raw("SELECT 1")
|
||||
verbose_proxy_logger.info(
|
||||
"Writer healthy on probe; skipping recreate (engine "
|
||||
"likely already replaced by a token refresh)."
|
||||
)
|
||||
await self._start_engine_watcher()
|
||||
return
|
||||
except Exception as probe_err:
|
||||
verbose_proxy_logger.warning(
|
||||
"Writer probe failed (%s); recreating Prisma client.",
|
||||
probe_err,
|
||||
)
|
||||
# Fresh Prisma client + new engine subprocess. The previous
|
||||
# "lightweight" path called `disconnect()` which blocks the
|
||||
# event loop on `subprocess.Popen.wait()`; since that call
|
||||
# ends up killing the engine anyway, we do it non-blockingly
|
||||
# via `_kill_engine_process` inside `recreate_prisma_client`.
|
||||
self._cleanup_engine_watcher()
|
||||
await self.db.recreate_prisma_client(db_url)
|
||||
await self.db.recreate_prisma_client(
|
||||
db_url, expected_generation=expected_generation
|
||||
)
|
||||
await self._start_engine_watcher()
|
||||
# Smoke-test the writer specifically; query_raw on the routing
|
||||
# wrapper sends to the reader, which would not validate the
|
||||
|
|
@ -4898,6 +5018,11 @@ class PrismaClient:
|
|||
return
|
||||
if self._db_health_watchdog_task is not None:
|
||||
return
|
||||
# Let planned writer-engine restarts (IAM token refresh, guarded
|
||||
# reconnect) re-arm the watcher on the new PID instead of being
|
||||
# mistaken for a crash (issue #29176). Set on the writer wrapper since
|
||||
# the watcher tracks the writer engine.
|
||||
self.writer_db.on_engine_replaced = self._handle_writer_engine_replaced
|
||||
self._db_health_watchdog_task = asyncio.create_task(
|
||||
self._db_health_watchdog_loop()
|
||||
)
|
||||
|
|
|
|||
|
|
@ -42,6 +42,8 @@ class BaseRAGIngestion(ABC):
|
|||
vector stores, so it overrides the embedding step to be a no-op.
|
||||
"""
|
||||
|
||||
supports_existing_file_id: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ingest_options: RAGIngestOptions,
|
||||
|
|
@ -280,6 +282,7 @@ class BaseRAGIngestion(ABC):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in vector store.
|
||||
|
|
@ -292,6 +295,7 @@ class BaseRAGIngestion(ABC):
|
|||
content_type: MIME type
|
||||
chunks: Text chunks (if chunking was done locally)
|
||||
embeddings: Embeddings (if embedding was done locally)
|
||||
existing_file_id: Provider file ID supplied by the caller, if any
|
||||
|
||||
Returns:
|
||||
Tuple of (vector_store_id, file_id)
|
||||
|
|
@ -326,6 +330,12 @@ class BaseRAGIngestion(ABC):
|
|||
)
|
||||
|
||||
try:
|
||||
if existing_file_id and not self.supports_existing_file_id:
|
||||
raise ValueError(
|
||||
f"{self.__class__.__name__} does not support ingesting an existing file_id. "
|
||||
"Upload file data or provide file_url instead."
|
||||
)
|
||||
|
||||
# Step 2: OCR (optional)
|
||||
extracted_text = await self.ocr(
|
||||
file_content=file_content,
|
||||
|
|
@ -349,6 +359,7 @@ class BaseRAGIngestion(ABC):
|
|||
content_type=content_type,
|
||||
chunks=chunks,
|
||||
embeddings=embeddings,
|
||||
existing_file_id=existing_file_id,
|
||||
)
|
||||
|
||||
return RAGIngestResponse(
|
||||
|
|
|
|||
|
|
@ -685,6 +685,7 @@ class BedrockRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in Bedrock Knowledge Base.
|
||||
|
|
@ -701,6 +702,7 @@ class BedrockRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
content_type: MIME type
|
||||
chunks: Ignored - Bedrock handles chunking
|
||||
embeddings: Ignored - Bedrock handles embedding
|
||||
existing_file_id: Existing provider file ID, unsupported for Bedrock
|
||||
|
||||
Returns:
|
||||
Tuple of (knowledge_base_id, file_key)
|
||||
|
|
|
|||
|
|
@ -61,6 +61,7 @@ class GeminiRAGIngestion(BaseRAGIngestion):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in Gemini File Search store.
|
||||
|
|
@ -75,6 +76,7 @@ class GeminiRAGIngestion(BaseRAGIngestion):
|
|||
content_type: MIME type
|
||||
chunks: Ignored - Gemini handles chunking
|
||||
embeddings: Ignored - Gemini handles embedding
|
||||
existing_file_id: Existing provider file ID, unsupported for Gemini
|
||||
|
||||
Returns:
|
||||
Tuple of (vector_store_id, file_id)
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ so this implementation skips the embedding step and directly uploads files.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast
|
||||
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, cast
|
||||
|
||||
import litellm
|
||||
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
|
||||
|
|
@ -29,6 +29,8 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
- Chunking is done by OpenAI's vector store (uses 'auto' strategy)
|
||||
"""
|
||||
|
||||
supports_existing_file_id = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ingest_options: "RAGIngestOptions",
|
||||
|
|
@ -56,6 +58,7 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in OpenAI vector store.
|
||||
|
|
@ -71,6 +74,7 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
content_type: MIME type
|
||||
chunks: Ignored - OpenAI handles chunking
|
||||
embeddings: Ignored - OpenAI handles embedding
|
||||
existing_file_id: Existing OpenAI file ID to attach
|
||||
|
||||
Returns:
|
||||
Tuple of (vector_store_id, file_id)
|
||||
|
|
@ -82,6 +86,11 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
api_key = self.vector_store_config.get("api_key")
|
||||
api_base = self.vector_store_config.get("api_base")
|
||||
|
||||
if existing_file_id and not vector_store_id:
|
||||
raise ValueError(
|
||||
"vector_store_id is required when ingesting an existing file_id"
|
||||
)
|
||||
|
||||
# Create vector store if not provided
|
||||
if not vector_store_id:
|
||||
expires_after = (
|
||||
|
|
@ -96,9 +105,20 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
)
|
||||
vector_store_id = create_response.get("id")
|
||||
|
||||
if existing_file_id and vector_store_id:
|
||||
await vector_store_file_acreate(
|
||||
vector_store_id=vector_store_id,
|
||||
file_id=existing_file_id,
|
||||
custom_llm_provider="openai",
|
||||
chunking_strategy=cast(dict[str, Any] | None, self.chunking_strategy),
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
return vector_store_id, existing_file_id
|
||||
|
||||
# Upload file and attach to vector store
|
||||
result_file_id = None
|
||||
if file_content and filename and vector_store_id:
|
||||
if file_content is not None and filename and vector_store_id:
|
||||
# Upload file to OpenAI
|
||||
file_response = await litellm.acreate_file(
|
||||
file=(
|
||||
|
|
@ -118,9 +138,7 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
vector_store_id=vector_store_id,
|
||||
file_id=result_file_id,
|
||||
custom_llm_provider="openai",
|
||||
chunking_strategy=cast(
|
||||
Optional[Dict[str, Any]], self.chunking_strategy
|
||||
),
|
||||
chunking_strategy=cast(dict[str, Any] | None, self.chunking_strategy),
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -464,6 +464,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store vectors in S3 Vectors using PutVectors API.
|
||||
|
|
@ -480,6 +481,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
content_type: MIME type (not used for S3 Vectors)
|
||||
chunks: Text chunks
|
||||
embeddings: Vector embeddings
|
||||
existing_file_id: Existing provider file ID, unsupported for S3 Vectors
|
||||
|
||||
Returns:
|
||||
Tuple of (index_name, filename)
|
||||
|
|
|
|||
|
|
@ -74,6 +74,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in Vertex AI RAG corpus.
|
||||
|
|
@ -88,6 +89,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase):
|
|||
content_type: MIME type
|
||||
chunks: Ignored - Vertex AI handles chunking
|
||||
embeddings: Ignored - Vertex AI handles embedding
|
||||
existing_file_id: Existing provider file ID, unsupported for Vertex AI
|
||||
|
||||
Returns:
|
||||
Tuple of (rag_corpus_id, file_id)
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ class LiteLLMCacheType(str, Enum):
|
|||
LOCAL = "local"
|
||||
REDIS = "redis"
|
||||
REDIS_SEMANTIC = "redis-semantic"
|
||||
VALKEY_SEMANTIC = "valkey-semantic"
|
||||
S3 = "s3"
|
||||
DISK = "disk"
|
||||
QDRANT_SEMANTIC = "qdrant-semantic"
|
||||
|
|
|
|||
|
|
@ -44,6 +44,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import (
|
|||
from litellm.types.proxy.guardrails.guardrail_hooks.qohash import (
|
||||
QostodianNexusConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.repelloai import (
|
||||
RepelloAIGuardrailConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import (
|
||||
VigilGuardGuardrailConfigModel,
|
||||
)
|
||||
|
|
@ -115,6 +118,7 @@ class SupportedGuardrailIntegrations(Enum):
|
|||
QOSTODIAN_NEXUS = "qostodian_nexus"
|
||||
RUBRIK = "rubrik"
|
||||
VIGIL_GUARD = "vigil_guard"
|
||||
REPELLOAI = "repelloai"
|
||||
|
||||
|
||||
class Role(Enum):
|
||||
|
|
@ -758,7 +762,7 @@ class BaseLitellmParams(
|
|||
default="fail_closed",
|
||||
description=(
|
||||
"Behavior when a guardrail endpoint is unreachable due to network errors. "
|
||||
"NOTE: This is currently only implemented by guardrail='generic_guardrail_api'. "
|
||||
"Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', and 'repelloai'. "
|
||||
"'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed."
|
||||
),
|
||||
)
|
||||
|
|
@ -856,6 +860,7 @@ class LitellmParams(
|
|||
PresidioConfigModel,
|
||||
BedrockGuardrailConfigModel,
|
||||
LakeraV2GuardrailConfigModel,
|
||||
RepelloAIGuardrailConfigModel,
|
||||
LassoGuardrailConfigModel,
|
||||
PillarGuardrailConfigModel,
|
||||
GraySwanGuardrailConfigModel,
|
||||
|
|
|
|||
65
litellm/types/proxy/guardrails/guardrail_hooks/repelloai.py
Normal file
65
litellm/types/proxy/guardrails/guardrail_hooks/repelloai.py
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
from typing import List, Literal, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from .base import GuardrailConfigModel
|
||||
|
||||
|
||||
class RepelloAIGuardrailConfigModel(GuardrailConfigModel[BaseModel]):
|
||||
"""Config model for the RepelloAI Argus guardrail."""
|
||||
|
||||
api_key: Optional[str] = Field(
|
||||
default=None,
|
||||
description="API key for the RepelloAI Argus service. Falls back to ARGUS_API_KEY or REPELLOAI_API_KEY.",
|
||||
)
|
||||
api_base: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Base URL for the RepelloAI Argus API. Defaults to https://argusapi.repello.ai/sdk/v1",
|
||||
)
|
||||
asset_id: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Repello asset ID whose dashboard policies are enforced. Required; the guardrail raises at init if it is missing.",
|
||||
)
|
||||
unreachable_fallback: Literal["fail_closed", "fail_open"] = Field(
|
||||
default="fail_closed",
|
||||
description="What to do when the RepelloAI Argus API is unreachable. 'fail_closed' = block (default), 'fail_open' = allow.",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "RepelloAI Argus"
|
||||
|
||||
|
||||
class RepelloAIScanData(TypedDict, total=False):
|
||||
"""The text payload sent to the RepelloAI Argus analyze endpoints.
|
||||
Only one of 'prompt' or 'response' is set per request.
|
||||
"""
|
||||
|
||||
prompt: Optional[str]
|
||||
response: Optional[str]
|
||||
|
||||
|
||||
class RepelloAIAnalyzeRequest(TypedDict, total=False):
|
||||
"""Request body for POST {api_base}/analyze/{prompt|response}."""
|
||||
|
||||
asset_id: str
|
||||
scan_data: RepelloAIScanData
|
||||
|
||||
|
||||
class RepelloAIViolatedPolicy(TypedDict, total=False):
|
||||
policy_name: Optional[str]
|
||||
policy_id: Optional[str]
|
||||
action_taken: Optional[str]
|
||||
scope: Optional[str]
|
||||
details: Optional[dict[str, object]]
|
||||
masked_result: Optional[str]
|
||||
|
||||
|
||||
class RepelloAIAnalyzeResponse(TypedDict, total=False):
|
||||
"""Response body returned by the RepelloAI Argus analyze endpoints."""
|
||||
|
||||
verdict: Optional[str] # "blocked" | "flagged" | "passed"
|
||||
request_id: Optional[str]
|
||||
policies_violated: Optional[List[RepelloAIViolatedPolicy]]
|
||||
policies_applied: Optional[List[dict[str, object]]]
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
from typing import Dict, List, Literal, Optional, Union
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import Field
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from litellm.proxy._types import KeyManagementRoutes, LitellmUserRoles
|
||||
|
|
@ -209,10 +209,3 @@ class DefaultTeamSSOParams(LiteLLMPydanticObjectBase):
|
|||
default=None,
|
||||
description="Default permissions granted to members of newly created teams (e.g. /key/generate, /key/update, /key/delete). /key/info and /key/health are always included.",
|
||||
)
|
||||
|
||||
|
||||
class InProductNudgeResponse(BaseModel):
|
||||
is_claude_code_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Whether the Claude Code nudge should be shown.",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -186,6 +186,7 @@ class CredentialLiteLLMParams(BaseModel):
|
|||
aws_region_name: Optional[str] = None
|
||||
aws_bedrock_runtime_endpoint: Optional[str] = None
|
||||
aws_bedrock_project_id: Optional[str] = None
|
||||
s3_bucket_name: Optional[str] = None
|
||||
## IBM WATSONX ##
|
||||
watsonx_region_name: Optional[str] = None
|
||||
|
||||
|
|
|
|||
|
|
@ -3439,6 +3439,7 @@ class LlmProviders(str, Enum):
|
|||
XIAOMI_MIMO = "xiaomi_mimo"
|
||||
TENSORMESH = "tensormesh"
|
||||
LIBERTAI = "libertai"
|
||||
PINSTRIPES = "pinstripes"
|
||||
LITELLM_AGENT = "litellm_agent"
|
||||
CURSOR = "cursor"
|
||||
BEDROCK_MANTLE = "bedrock_mantle"
|
||||
|
|
@ -3482,6 +3483,7 @@ class SearchProviders(str, Enum):
|
|||
SERPER = "serper"
|
||||
YOU_COM = "you_com"
|
||||
APISERPENT = "apiserpent"
|
||||
TINYFISH = "tinyfish"
|
||||
|
||||
|
||||
# Create a set of all search provider values for quick lookup
|
||||
|
|
|
|||
|
|
@ -9070,34 +9070,26 @@ class ProviderConfigManager:
|
|||
elif litellm.LlmProviders.HOSTED_VLLM == provider:
|
||||
return litellm.HostedVLLMResponsesAPIConfig()
|
||||
elif litellm.LlmProviders.BEDROCK_MANTLE == provider:
|
||||
# Mantle serves Responses on two upstream paths. A model takes the
|
||||
# /openai/v1/responses path when its price-map entry declares
|
||||
# use_openai_responses_path (data-driven, so a non-gpt-named frontier
|
||||
# model can be onboarded by JSON alone), or, as a fallback needing no
|
||||
# price-map entry, when its name matches the openai.gpt- frontier
|
||||
# convention (minus gpt-oss) -- this keeps a future gpt-6 routing
|
||||
# correctly before its entry loads. Any other model declared
|
||||
# mode=responses takes the standard /v1/responses path. Everything
|
||||
# else returns None and keeps the chat-completions emulation (see
|
||||
# responses/main.py "config is None").
|
||||
if not model:
|
||||
return None
|
||||
model_lower = model.lower()
|
||||
entry = litellm.model_cost.get(f"bedrock_mantle/{model}", {})
|
||||
on_openai_path = entry.get("use_openai_responses_path") is True
|
||||
name_is_frontier = (
|
||||
"openai.gpt-" in model_lower and "gpt-oss" not in model_lower
|
||||
# Both decisions are data-driven from the model's price-map entry, with
|
||||
# no model-name logic. Capability (can it serve Responses?) comes from
|
||||
# mantle_supports_responses (supported_endpoints / mode);
|
||||
# chat-only models (gpt-oss safeguard, nvidia, ...) return None and keep
|
||||
# the chat-completions emulation (responses/main.py "config is None").
|
||||
# The wire path comes from mantle_base_segment, which reads the
|
||||
# use_openai_responses_path flag: gpt-5.x and gemma-4-* on
|
||||
# /openai/v1/responses, everything else (incl. gpt-oss) on
|
||||
# /v1/responses.
|
||||
from litellm.llms.bedrock_mantle.common_utils import (
|
||||
mantle_base_segment,
|
||||
mantle_supports_responses,
|
||||
)
|
||||
|
||||
if not model or not mantle_supports_responses(model, litellm.model_cost):
|
||||
return None
|
||||
return litellm.BedrockMantleResponsesAPIConfig(
|
||||
use_openai_path=mantle_base_segment(model, litellm.model_cost)
|
||||
== "openai/v1"
|
||||
)
|
||||
if on_openai_path or name_is_frontier:
|
||||
return litellm.BedrockMantleResponsesAPIConfig(use_openai_path=True)
|
||||
try:
|
||||
if get_model_info(model, "bedrock_mantle").get("mode") == "responses":
|
||||
return litellm.BedrockMantleResponsesAPIConfig(
|
||||
use_openai_path=False
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -9710,6 +9702,7 @@ class ProviderConfigManager:
|
|||
from litellm.llms.searxng.search.transformation import SearXNGSearchConfig
|
||||
from litellm.llms.serper.search.transformation import SerperSearchConfig
|
||||
from litellm.llms.tavily.search.transformation import TavilySearchConfig
|
||||
from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig
|
||||
from litellm.llms.you_com.search.transformation import YouComSearchConfig
|
||||
|
||||
PROVIDER_TO_CONFIG_MAP = {
|
||||
|
|
@ -9729,6 +9722,7 @@ class ProviderConfigManager:
|
|||
SearchProviders.SERPER: SerperSearchConfig,
|
||||
SearchProviders.YOU_COM: YouComSearchConfig,
|
||||
SearchProviders.APISERPENT: APISerpentSearchConfig,
|
||||
SearchProviders.TINYFISH: TinyfishSearchConfig,
|
||||
}
|
||||
config_class = PROVIDER_TO_CONFIG_MAP.get(provider, None)
|
||||
if config_class is None:
|
||||
|
|
|
|||
|
|
@ -13876,6 +13876,14 @@
|
|||
"notes": "APISerpent deep search (/api/search), multi-engine (Google, Bing, Yahoo, DuckDuckGo). Pricing: $0.60/1k searches."
|
||||
}
|
||||
},
|
||||
"tinyfish/search": {
|
||||
"input_cost_per_query": 0.0,
|
||||
"litellm_provider": "tinyfish",
|
||||
"mode": "search",
|
||||
"metadata": {
|
||||
"notes": "TinyFish Search API"
|
||||
}
|
||||
},
|
||||
"elevenlabs/scribe_v1": {
|
||||
"input_cost_per_second": 6.11e-05,
|
||||
"litellm_provider": "elevenlabs",
|
||||
|
|
@ -42594,6 +42602,7 @@
|
|||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42608,6 +42617,7 @@
|
|||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42622,6 +42632,7 @@
|
|||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions"],
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
|
|
@ -42635,6 +42646,7 @@
|
|||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions"],
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
|
|
@ -42688,6 +42700,8 @@
|
|||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"use_openai_responses_path": true,
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": false,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42702,6 +42716,8 @@
|
|||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"use_openai_responses_path": true,
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": false,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42716,6 +42732,8 @@
|
|||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"use_openai_responses_path": true,
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": false,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -43186,6 +43204,17 @@
|
|||
"supported_endpoints": ["/v1/audio/transcriptions"],
|
||||
"supports_audio_input": true
|
||||
},
|
||||
"soniox/stt-async-v5": {
|
||||
"litellm_provider": "soniox",
|
||||
"max_output_tokens": 8000,
|
||||
"max_tokens": 8000,
|
||||
"input_cost_per_second": 0.0,
|
||||
"output_cost_per_second": 0.0000277778,
|
||||
"mode": "audio_transcription",
|
||||
"source": "https://soniox.com/pricing",
|
||||
"supported_endpoints": ["/v1/audio/transcriptions"],
|
||||
"supports_audio_input": true
|
||||
},
|
||||
"tensormesh/Qwen/Qwen3.5-397B-A17B-FP8": {
|
||||
"litellm_provider": "tensormesh",
|
||||
"mode": "chat",
|
||||
|
|
@ -43445,5 +43474,83 @@
|
|||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"pinstripes/ps/glm-4.5-air": {
|
||||
"max_tokens": 128000,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 128000,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.00000045,
|
||||
"litellm_provider": "pinstripes",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_reasoning": true,
|
||||
"source": "https://pinstripes.io/pricing"
|
||||
},
|
||||
"pinstripes/ps/qwen3.6-35b-a3b": {
|
||||
"max_tokens": 131072,
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 131072,
|
||||
"input_cost_per_token": 0.00000014,
|
||||
"output_cost_per_token": 0.00000045,
|
||||
"litellm_provider": "pinstripes",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_reasoning": true,
|
||||
"source": "https://pinstripes.io/pricing"
|
||||
},
|
||||
"pinstripes/ps/qwen3-30b-a3b": {
|
||||
"max_tokens": 131072,
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 131072,
|
||||
"input_cost_per_token": 0.00000009,
|
||||
"output_cost_per_token": 0.0000002,
|
||||
"litellm_provider": "pinstripes",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_reasoning": true,
|
||||
"source": "https://pinstripes.io/pricing"
|
||||
},
|
||||
"pinstripes/ps/qwen3-coder-30b-a3b": {
|
||||
"max_tokens": 131072,
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 131072,
|
||||
"input_cost_per_token": 0.0000003,
|
||||
"output_cost_per_token": 0.0000006,
|
||||
"litellm_provider": "pinstripes",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_reasoning": false,
|
||||
"source": "https://pinstripes.io/pricing"
|
||||
},
|
||||
"pinstripes/ps/deepseek-v4-flash": {
|
||||
"max_tokens": 163840,
|
||||
"max_input_tokens": 163840,
|
||||
"max_output_tokens": 163840,
|
||||
"input_cost_per_token": 0.0000001,
|
||||
"output_cost_per_token": 0.0000002,
|
||||
"litellm_provider": "pinstripes",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_reasoning": true,
|
||||
"source": "https://pinstripes.io/pricing"
|
||||
},
|
||||
"pinstripes/ps/minimax-m2.7": {
|
||||
"max_tokens": 1000192,
|
||||
"max_input_tokens": 1000192,
|
||||
"max_output_tokens": 1000192,
|
||||
"input_cost_per_token": 0.000000255,
|
||||
"output_cost_per_token": 0.00000055,
|
||||
"litellm_provider": "pinstripes",
|
||||
"mode": "chat",
|
||||
"supports_function_calling": true,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_reasoning": false,
|
||||
"source": "https://pinstripes.io/pricing"
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1940,6 +1940,23 @@
|
|||
"interactions": true
|
||||
}
|
||||
},
|
||||
"pinstripes": {
|
||||
"display_name": "Pinstripes (`pinstripes`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/pinstripes",
|
||||
"endpoints": {
|
||||
"chat_completions": true,
|
||||
"messages": true,
|
||||
"responses": true,
|
||||
"embeddings": true,
|
||||
"image_generations": false,
|
||||
"audio_transcriptions": false,
|
||||
"audio_speech": false,
|
||||
"moderations": false,
|
||||
"batches": false,
|
||||
"rerank": false,
|
||||
"a2a": false
|
||||
}
|
||||
},
|
||||
"poe": {
|
||||
"display_name": "Poe (`poe`)",
|
||||
"endpoints": {
|
||||
|
|
@ -2303,6 +2320,13 @@
|
|||
"search": true
|
||||
}
|
||||
},
|
||||
"tinyfish": {
|
||||
"display_name": "TinyFish (`tinyfish`)",
|
||||
"url": "https://docs.tinyfish.ai/search-api",
|
||||
"endpoints": {
|
||||
"search": true
|
||||
}
|
||||
},
|
||||
"triton": {
|
||||
"display_name": "Triton (`triton`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/triton-inference-server",
|
||||
|
|
|
|||
|
|
@ -167,6 +167,8 @@ dev = [
|
|||
"types-setuptools==75.8.0.20250225",
|
||||
"types-redis==4.6.0.20241004",
|
||||
"types-PyYAML==6.0.12.20250915",
|
||||
"botocore-stubs==1.43.14",
|
||||
"types-boto3[bedrock,bedrock-agent,bedrock-runtime,kms,s3,sagemaker-runtime,sts]==1.43.30",
|
||||
"opentelemetry-api==1.28.0",
|
||||
"opentelemetry-sdk==1.28.0",
|
||||
"opentelemetry-exporter-otlp==1.28.0",
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ SEARCH_PROVIDERS = [
|
|||
"searchapi",
|
||||
"serper",
|
||||
"apiserpent",
|
||||
"tinyfish",
|
||||
]
|
||||
|
||||
ALLOWED_FILES_IN_LLMS_FOLDER = [
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ import asyncio
|
|||
import os
|
||||
import threading
|
||||
import time
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -219,7 +219,7 @@ async def test_run_reconnect_cycle_uses_heavy_path_when_engine_dead(
|
|||
await engine_client._run_reconnect_cycle(timeout_seconds=5.0)
|
||||
|
||||
engine_client.db.recreate_prisma_client.assert_awaited_once_with(
|
||||
"postgresql://test"
|
||||
"postgresql://test", expected_generation=ANY
|
||||
)
|
||||
engine_client._start_engine_watcher.assert_awaited_once()
|
||||
engine_client.db.connect.assert_not_awaited()
|
||||
|
|
@ -246,7 +246,7 @@ async def test_run_reconnect_cycle_uses_heavy_path_when_confirmed_dead(
|
|||
await engine_client._run_reconnect_cycle(timeout_seconds=5.0)
|
||||
|
||||
engine_client.db.recreate_prisma_client.assert_awaited_once_with(
|
||||
"postgresql://test"
|
||||
"postgresql://test", expected_generation=ANY
|
||||
)
|
||||
engine_client._start_engine_watcher.assert_awaited_once()
|
||||
engine_client.db.connect.assert_not_awaited()
|
||||
|
|
@ -257,12 +257,16 @@ async def test_run_reconnect_cycle_uses_heavy_path_when_confirmed_dead(
|
|||
async def test_run_reconnect_cycle_uses_direct_path_when_engine_alive(
|
||||
engine_client,
|
||||
) -> None:
|
||||
"""Direct reconnect (engine alive) calls recreate_prisma_client + SELECT 1.
|
||||
"""Direct reconnect (engine alive) probes the writer first and skips the
|
||||
recreate when the probe is healthy.
|
||||
|
||||
The old "lightweight" path called `disconnect()` + `connect()`, which
|
||||
blocks the event loop on the sync `process.wait()` inside aclose().
|
||||
The fix routes both engine-alive and engine-dead paths through
|
||||
`recreate_prisma_client`, which non-blockingly kills the old engine.
|
||||
The engine-alive path now runs a SELECT 1 probe before recreating. A
|
||||
healthy probe means the connection is fine — e.g. an IAM token refresh
|
||||
already replaced the engine (issue #29176) — so recreating would kill a
|
||||
working engine. Recreate happens only when the probe fails (covered in
|
||||
test_prisma_client_reconnect.py::
|
||||
test_run_reconnect_cycle_direct_path_recreates_when_probe_fails). Either
|
||||
way the blocking `disconnect()` is never called.
|
||||
"""
|
||||
engine_client._engine_pid = 1234
|
||||
engine_client._start_engine_watcher = AsyncMock()
|
||||
|
|
@ -273,29 +277,28 @@ async def test_run_reconnect_cycle_uses_direct_path_when_engine_alive(
|
|||
):
|
||||
await engine_client._run_reconnect_cycle(timeout_seconds=5.0)
|
||||
|
||||
engine_client.db.recreate_prisma_client.assert_awaited_once_with(
|
||||
"postgresql://test"
|
||||
)
|
||||
engine_client.db.recreate_prisma_client.assert_not_awaited()
|
||||
engine_client.db.query_raw.assert_awaited_once_with("SELECT 1")
|
||||
engine_client.db.disconnect.assert_not_awaited()
|
||||
engine_client._start_engine_watcher.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_reconnect_cycle_uses_direct_path_when_pid_unknown(
|
||||
engine_client,
|
||||
) -> None:
|
||||
"""When the engine PID is not tracked, direct reconnect still runs."""
|
||||
"""When the engine PID is not tracked, direct reconnect still runs and a
|
||||
healthy probe likewise skips the recreate."""
|
||||
engine_client._engine_pid = 0
|
||||
engine_client._start_engine_watcher = AsyncMock()
|
||||
|
||||
with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}):
|
||||
await engine_client._run_reconnect_cycle(timeout_seconds=5.0)
|
||||
|
||||
engine_client.db.recreate_prisma_client.assert_awaited_once_with(
|
||||
"postgresql://test"
|
||||
)
|
||||
engine_client.db.recreate_prisma_client.assert_not_awaited()
|
||||
engine_client.db.query_raw.assert_awaited_once_with("SELECT 1")
|
||||
engine_client.db.disconnect.assert_not_awaited()
|
||||
engine_client._start_engine_watcher.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -497,7 +500,10 @@ async def test_escalation_after_consecutive_direct_reconnect_failures(engine_cli
|
|||
engine_client._db_reconnect_cooldown_seconds = 0 # disable cooldown for test
|
||||
engine_client._start_engine_watcher = AsyncMock(return_value=None)
|
||||
|
||||
# Make direct reconnect fail every time
|
||||
# Make the direct path's writer probe fail so it proceeds to recreate
|
||||
# (a healthy probe would correctly skip recreate), then make recreate
|
||||
# fail every time.
|
||||
engine_client.db.query_raw = AsyncMock(side_effect=Exception("probe failed"))
|
||||
engine_client.db.recreate_prisma_client = AsyncMock(
|
||||
side_effect=Exception("recreate failed")
|
||||
)
|
||||
|
|
|
|||
|
|
@ -81,8 +81,11 @@ async def _list_hashes(proxy_client, caller_cleartext: str, query: str) -> set:
|
|||
|
||||
|
||||
async def test_key_list_admin_key_alias_substring_match(proxy_client, scratch, world):
|
||||
"""A PROXY_ADMIN's key_alias filter is a case-insensitive substring match;
|
||||
a narrower fragment selects the subset whose alias contains it."""
|
||||
"""A PROXY_ADMIN's key_alias filter is a case-insensitive substring match
|
||||
when substring_matching=true is requested (the dashboard search box); a
|
||||
narrower fragment selects the subset whose alias contains it. Substring
|
||||
matching is opt-in: without the flag the filter is exact (see
|
||||
test_key_list_admin_key_alias_exact_without_substring_flag)."""
|
||||
admin = world.keys[Actor.PROXY_ADMIN]
|
||||
a = await create_scratch_key(
|
||||
proxy_client,
|
||||
|
|
@ -101,16 +104,46 @@ async def test_key_list_admin_key_alias_substring_match(proxy_client, scratch, w
|
|||
seeded = {hash_token(a), hash_token(b)}
|
||||
|
||||
broad = await _list_hashes(
|
||||
proxy_client, admin.cleartext, f"key_alias={scratch.prefix}-sub"
|
||||
proxy_client,
|
||||
admin.cleartext,
|
||||
f"key_alias={scratch.prefix}-sub&substring_matching=true",
|
||||
)
|
||||
assert broad & seeded == seeded
|
||||
|
||||
narrow = await _list_hashes(
|
||||
proxy_client, admin.cleartext, f"key_alias={scratch.prefix}-sub-a"
|
||||
proxy_client,
|
||||
admin.cleartext,
|
||||
f"key_alias={scratch.prefix}-sub-a&substring_matching=true",
|
||||
)
|
||||
assert narrow & seeded == {hash_token(a)}
|
||||
|
||||
|
||||
async def test_key_list_admin_key_alias_exact_without_substring_flag(
|
||||
proxy_client, scratch, world
|
||||
):
|
||||
"""Regression guard for the prior exact-match contract: without
|
||||
substring_matching, even a PROXY_ADMIN's key_alias filter is exact, so a
|
||||
fragment of a seeded alias does not select it."""
|
||||
admin = world.keys[Actor.PROXY_ADMIN]
|
||||
full_alias = f"{scratch.prefix}-exactflag"
|
||||
key = await create_scratch_key(
|
||||
proxy_client,
|
||||
admin.cleartext,
|
||||
scratch.prefix,
|
||||
user_id=admin.user_id,
|
||||
key_alias=full_alias,
|
||||
)
|
||||
key_hash = hash_token(key)
|
||||
|
||||
exact = await _list_hashes(proxy_client, admin.cleartext, f"key_alias={full_alias}")
|
||||
assert key_hash in exact
|
||||
|
||||
fragment = await _list_hashes(
|
||||
proxy_client, admin.cleartext, f"key_alias={scratch.prefix}-exactfla"
|
||||
)
|
||||
assert key_hash not in fragment
|
||||
|
||||
|
||||
async def test_key_list_non_admin_key_alias_is_exact_match(
|
||||
proxy_client, scratch, world
|
||||
):
|
||||
|
|
|
|||
|
|
@ -2085,6 +2085,48 @@ async def test_gemini_pass_through_endpoint():
|
|||
print(resp.body)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("hidden", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_info_alias_without_prisma(hidden):
|
||||
from litellm.proxy.proxy_server import model_info_v1
|
||||
|
||||
_model_list = [
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {"model": "gpt-3.5-turbo"},
|
||||
}
|
||||
]
|
||||
|
||||
model_alias = "gpt-4"
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=_model_list,
|
||||
model_group_alias={
|
||||
model_alias: {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"hidden": hidden,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
setattr(litellm.proxy.proxy_server, "llm_router", router)
|
||||
setattr(litellm.proxy.proxy_server, "llm_model_list", _model_list)
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", None)
|
||||
|
||||
resp = await model_info_v1(
|
||||
user_api_key_dict=UserAPIKeyAuth(models=[]),
|
||||
)
|
||||
|
||||
models = resp["data"]
|
||||
|
||||
alias_found = any(
|
||||
m["model_name"] == model_alias
|
||||
for m in models
|
||||
)
|
||||
|
||||
assert alias_found is (not hidden)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("hidden", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skip(reason="Requires reliable external DB connection (prisma).")
|
||||
|
|
|
|||
224
tests/search_tests/test_tinyfish_search.py
Normal file
224
tests/search_tests/test_tinyfish_search.py
Normal file
|
|
@ -0,0 +1,224 @@
|
|||
"""
|
||||
Tests for TinyFish Search API integration.
|
||||
"""
|
||||
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
MOCK_TINYFISH_RESPONSE = {
|
||||
"query": "web automation tools",
|
||||
"results": [
|
||||
{
|
||||
"position": 1,
|
||||
"site_name": "tinyfish.ai",
|
||||
"title": "TinyFish - AI Web Automation",
|
||||
"snippet": "Automate any website with natural language.",
|
||||
"url": "https://tinyfish.ai",
|
||||
},
|
||||
{
|
||||
"position": 2,
|
||||
"site_name": "github.com",
|
||||
"title": "Top Web Automation Tools",
|
||||
"snippet": "A curated list of browser automation frameworks.",
|
||||
"url": "https://github.com/example/web-automation",
|
||||
},
|
||||
],
|
||||
"total_results": 2,
|
||||
"page": 0,
|
||||
}
|
||||
|
||||
|
||||
def _make_mock_response(
|
||||
json_data: dict, status_code: int = 200, request_url: str | None = None
|
||||
) -> MagicMock:
|
||||
mock = MagicMock()
|
||||
mock.status_code = status_code
|
||||
mock.json.return_value = json_data
|
||||
if request_url:
|
||||
mock.request = MagicMock()
|
||||
mock.request.url = httpx.URL(request_url)
|
||||
else:
|
||||
mock.request = None
|
||||
return mock
|
||||
|
||||
|
||||
class TestTinyfishSearch:
|
||||
@pytest.mark.asyncio
|
||||
async def test_basic_search(self):
|
||||
os.environ["TINYFISH_API_KEY"] = "sk-tinyfish-test"
|
||||
|
||||
mock_response = _make_mock_response(MOCK_TINYFISH_RESPONSE)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get:
|
||||
mock_get.return_value = mock_response
|
||||
|
||||
response = await litellm.asearch(
|
||||
query="web automation tools",
|
||||
search_provider="tinyfish",
|
||||
)
|
||||
|
||||
assert mock_get.call_count == 1
|
||||
|
||||
call_args = mock_get.call_args
|
||||
parsed_url = urlparse(call_args.kwargs["url"])
|
||||
assert parsed_url.scheme == "https"
|
||||
assert parsed_url.netloc == "api.search.tinyfish.ai"
|
||||
assert parsed_url.path == ""
|
||||
|
||||
query_params = parse_qs(parsed_url.query)
|
||||
assert query_params["query"] == ["web automation tools"]
|
||||
|
||||
headers = call_args.kwargs.get("headers", {})
|
||||
assert headers["X-API-Key"] == "sk-tinyfish-test"
|
||||
|
||||
assert hasattr(response, "results")
|
||||
assert response.object == "search"
|
||||
assert len(response.results) == 2
|
||||
|
||||
first = response.results[0]
|
||||
assert first.title == "TinyFish - AI Web Automation"
|
||||
assert first.url == "https://tinyfish.ai"
|
||||
assert first.snippet == "Automate any website with natural language."
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_country_maps_to_location(self):
|
||||
os.environ["TINYFISH_API_KEY"] = "sk-tinyfish-test"
|
||||
|
||||
mock_response = _make_mock_response(MOCK_TINYFISH_RESPONSE)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get:
|
||||
mock_get.return_value = mock_response
|
||||
|
||||
await litellm.asearch(
|
||||
query="test",
|
||||
search_provider="tinyfish",
|
||||
country="US",
|
||||
)
|
||||
|
||||
call_args = mock_get.call_args
|
||||
parsed_url = urlparse(call_args.kwargs["url"])
|
||||
query_params = parse_qs(parsed_url.query)
|
||||
assert query_params["location"] == ["US"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_domain_filter_injection(self):
|
||||
os.environ["TINYFISH_API_KEY"] = "sk-tinyfish-test"
|
||||
|
||||
mock_response = _make_mock_response(MOCK_TINYFISH_RESPONSE)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get:
|
||||
mock_get.return_value = mock_response
|
||||
|
||||
await litellm.asearch(
|
||||
query="python tutorials",
|
||||
search_provider="tinyfish",
|
||||
search_domain_filter=["arxiv.org", "github.com"],
|
||||
)
|
||||
|
||||
call_args = mock_get.call_args
|
||||
parsed_url = urlparse(call_args.kwargs["url"])
|
||||
query_params = parse_qs(parsed_url.query)
|
||||
query_value = query_params["query"][0]
|
||||
assert "site:arxiv.org" in query_value
|
||||
assert "site:github.com" in query_value
|
||||
assert "python tutorials" in query_value
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_language_passthrough(self):
|
||||
os.environ["TINYFISH_API_KEY"] = "sk-tinyfish-test"
|
||||
|
||||
mock_response = _make_mock_response(MOCK_TINYFISH_RESPONSE)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get:
|
||||
mock_get.return_value = mock_response
|
||||
|
||||
await litellm.asearch(
|
||||
query="test",
|
||||
search_provider="tinyfish",
|
||||
language="en",
|
||||
)
|
||||
|
||||
call_args = mock_get.call_args
|
||||
parsed_url = urlparse(call_args.kwargs["url"])
|
||||
query_params = parse_qs(parsed_url.query)
|
||||
assert query_params["language"] == ["en"]
|
||||
|
||||
def test_max_results_truncates_response(self):
|
||||
from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig
|
||||
|
||||
config = TinyfishSearchConfig()
|
||||
many_results = {
|
||||
"results": [
|
||||
{
|
||||
"title": f"Result {i}",
|
||||
"url": f"https://example.com/{i}",
|
||||
"snippet": f"Snippet {i}",
|
||||
}
|
||||
for i in range(10)
|
||||
]
|
||||
}
|
||||
mock_response = _make_mock_response(
|
||||
many_results,
|
||||
request_url="https://api.search.tinyfish.ai?query=test&max_results=3",
|
||||
)
|
||||
|
||||
result = config.transform_search_response(
|
||||
raw_response=mock_response,
|
||||
logging_obj=None,
|
||||
)
|
||||
assert len(result.results) == 3
|
||||
assert result.results[0].title == "Result 0"
|
||||
assert result.results[2].title == "Result 2"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_results(self):
|
||||
os.environ["TINYFISH_API_KEY"] = "sk-tinyfish-test"
|
||||
|
||||
empty_response = {
|
||||
"query": "xyznonexistent",
|
||||
"results": [],
|
||||
"total_results": 0,
|
||||
"page": 0,
|
||||
}
|
||||
mock_response = _make_mock_response(empty_response)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get:
|
||||
mock_get.return_value = mock_response
|
||||
|
||||
response = await litellm.asearch(
|
||||
query="xyznonexistent",
|
||||
search_provider="tinyfish",
|
||||
)
|
||||
|
||||
assert response.object == "search"
|
||||
assert len(response.results) == 0
|
||||
|
||||
def test_missing_api_key(self):
|
||||
os.environ.pop("TINYFISH_API_KEY", None)
|
||||
|
||||
from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig
|
||||
|
||||
config = TinyfishSearchConfig()
|
||||
with pytest.raises(ValueError, match="TINYFISH_API_KEY"):
|
||||
config.validate_environment(headers={})
|
||||
|
|
@ -76,3 +76,73 @@ def test_get_per_item_prompt_tokens_distributes_with_remainder():
|
|||
per_item = [cache._get_per_item_prompt_tokens(result, i) for i in range(3)]
|
||||
assert sum(per_item) == 10 # 4 + 3 + 3
|
||||
assert per_item == [4, 3, 3]
|
||||
|
||||
|
||||
def _semantic_cache():
|
||||
return Cache(
|
||||
type=LiteLLMCacheType.VALKEY_SEMANTIC,
|
||||
host="localhost",
|
||||
port="6379",
|
||||
similarity_threshold=0.8,
|
||||
)
|
||||
|
||||
|
||||
def test_semantic_cache_key_excludes_prompt_so_paraphrases_share_a_bucket():
|
||||
cache = _semantic_cache()
|
||||
tenant = {"user_api_key": "hash-abc"}
|
||||
key_a = cache.get_cache_key(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "What color is the sky?"}],
|
||||
metadata=dict(tenant),
|
||||
)
|
||||
key_b = cache.get_cache_key(
|
||||
model="gpt-4o-mini",
|
||||
messages=[
|
||||
{"role": "user", "content": "Tell me the colour of the daytime sky."}
|
||||
],
|
||||
metadata=dict(tenant),
|
||||
)
|
||||
assert key_a == key_b
|
||||
|
||||
|
||||
def test_semantic_cache_key_isolates_tenants():
|
||||
messages = [{"role": "user", "content": "What color is the sky?"}]
|
||||
cache = _semantic_cache()
|
||||
key_a = cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=messages, metadata={"user_api_key": "hash-A"}
|
||||
)
|
||||
key_b = cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=messages, metadata={"user_api_key": "hash-B"}
|
||||
)
|
||||
key_team = cache.get_cache_key(
|
||||
model="gpt-4o-mini",
|
||||
messages=messages,
|
||||
metadata={"user_api_key": "hash-A", "user_api_key_team_id": "team-1"},
|
||||
)
|
||||
assert key_a != key_b
|
||||
assert key_a != key_team
|
||||
|
||||
|
||||
def test_semantic_cache_key_still_separates_models_and_params():
|
||||
cache = _semantic_cache()
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
tenant = {"user_api_key": "hash-A"}
|
||||
assert cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=messages, metadata=dict(tenant)
|
||||
) != cache.get_cache_key(model="gpt-4o", messages=messages, metadata=dict(tenant))
|
||||
assert cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=messages, temperature=0, metadata=dict(tenant)
|
||||
) != cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=messages, temperature=1, metadata=dict(tenant)
|
||||
)
|
||||
|
||||
|
||||
def test_exact_cache_key_still_includes_prompt():
|
||||
cache = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
key_a = cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=[{"role": "user", "content": "a"}]
|
||||
)
|
||||
key_b = cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=[{"role": "user", "content": "b"}]
|
||||
)
|
||||
assert key_a != key_b
|
||||
|
|
|
|||
|
|
@ -44,3 +44,64 @@ async def test_gcs_cache_async_set_and_get(mock_gcs_dependencies):
|
|||
mock_gcs_dependencies["async_client"].get.return_value.text = '{"foo": "bar"}'
|
||||
result = await cache.async_get_cache("key")
|
||||
assert result == {"foo": "bar"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gcs_cache_async_get_encodes_object_name_in_path(mock_gcs_dependencies):
|
||||
"""
|
||||
Regression test for https://github.com/BerriAI/litellm/issues/30377
|
||||
|
||||
When gcs_path is set, the object name contains a '/' (e.g. "my_cache/<hash>").
|
||||
The GCS JSON API requires the object name in the GET path to be URL-encoded,
|
||||
so the '/' must be sent as '%2F'. Otherwise GCS returns 404 and every read
|
||||
silently misses.
|
||||
"""
|
||||
cache = GCSCache(bucket_name="test-bucket", gcs_path="my_cache/")
|
||||
|
||||
mock_gcs_dependencies["async_client"].get.return_value.status_code = 200
|
||||
mock_gcs_dependencies["async_client"].get.return_value.text = '{"foo": "bar"}'
|
||||
|
||||
result = await cache.async_get_cache("abc123")
|
||||
assert result == {"foo": "bar"}
|
||||
|
||||
called_url = mock_gcs_dependencies["async_client"].get.call_args.kwargs["url"]
|
||||
# The slash from gcs_path must be percent-encoded in the path segment.
|
||||
assert "/o/my_cache%2Fabc123?alt=media" in called_url
|
||||
assert "/o/my_cache/abc123" not in called_url
|
||||
|
||||
|
||||
def test_gcs_cache_get_encodes_object_name_in_path(mock_gcs_dependencies):
|
||||
"""Sync counterpart of the regression test for issue #30377."""
|
||||
cache = GCSCache(bucket_name="test-bucket", gcs_path="my_cache/")
|
||||
|
||||
mock_gcs_dependencies["sync_client"].get.return_value.status_code = 200
|
||||
mock_gcs_dependencies["sync_client"].get.return_value.text = '{"foo": "bar"}'
|
||||
|
||||
result = cache.get_cache("abc123")
|
||||
assert result == {"foo": "bar"}
|
||||
|
||||
called_url = mock_gcs_dependencies["sync_client"].get.call_args.kwargs["url"]
|
||||
assert "/o/my_cache%2Fabc123?alt=media" in called_url
|
||||
assert "/o/my_cache/abc123" not in called_url
|
||||
|
||||
|
||||
def test_gcs_cache_set_encodes_object_name_in_query(mock_gcs_dependencies):
|
||||
"""
|
||||
The set path uses the object name as a query parameter. Encoding it keeps
|
||||
both sides symmetric so the key written matches the key read back.
|
||||
"""
|
||||
cache = GCSCache(bucket_name="test-bucket", gcs_path="my_cache/")
|
||||
cache.set_cache("abc123", {"foo": "bar"})
|
||||
|
||||
called_url = mock_gcs_dependencies["sync_client"].post.call_args.kwargs["url"]
|
||||
assert "name=my_cache%2Fabc123" in called_url
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gcs_cache_async_set_encodes_object_name_in_query(mock_gcs_dependencies):
|
||||
"""Async counterpart of test_gcs_cache_set_encodes_object_name_in_query."""
|
||||
cache = GCSCache(bucket_name="test-bucket", gcs_path="my_cache/")
|
||||
await cache.async_set_cache("abc123", {"foo": "bar"})
|
||||
|
||||
called_url = mock_gcs_dependencies["async_client"].post.call_args.kwargs["url"]
|
||||
assert "name=my_cache%2Fabc123" in called_url
|
||||
|
|
|
|||
473
tests/test_litellm/caching/test_valkey_semantic_cache.py
Normal file
473
tests/test_litellm/caching/test_valkey_semantic_cache.py
Normal file
|
|
@ -0,0 +1,473 @@
|
|||
import hashlib
|
||||
import os
|
||||
import struct
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
from litellm.caching.valkey_semantic_cache import ValkeySemanticCache
|
||||
|
||||
_REPO_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
|
||||
|
||||
|
||||
def _make_cache(sync_client=None, async_client=None, similarity_threshold=0.8):
|
||||
return ValkeySemanticCache(
|
||||
similarity_threshold=similarity_threshold,
|
||||
index_name="test_index",
|
||||
sync_client=sync_client or MagicMock(),
|
||||
async_client=async_client or AsyncMock(),
|
||||
)
|
||||
|
||||
|
||||
def _search_result(distance, response='{"content": "Paris"}'):
|
||||
return SimpleNamespace(
|
||||
docs=[SimpleNamespace(response=response, vector_distance=str(distance))]
|
||||
)
|
||||
|
||||
|
||||
def test_build_valkey_url_prefers_valkey_env(monkeypatch):
|
||||
monkeypatch.setenv("REDIS_HOST", "redis-host")
|
||||
monkeypatch.setenv("REDIS_PORT", "6379")
|
||||
monkeypatch.setenv("REDIS_PASSWORD", "rpass")
|
||||
monkeypatch.setenv("VALKEY_HOST", "valkey-host")
|
||||
monkeypatch.setenv("VALKEY_PORT", "6380")
|
||||
monkeypatch.setenv("VALKEY_PASSWORD", "vpass")
|
||||
|
||||
assert (
|
||||
ValkeySemanticCache._build_valkey_url(None, None, None)
|
||||
== "redis://:vpass@valkey-host:6380"
|
||||
)
|
||||
|
||||
|
||||
def test_build_valkey_url_supports_passwordless(monkeypatch):
|
||||
monkeypatch.delenv("REDIS_PASSWORD", raising=False)
|
||||
monkeypatch.delenv("VALKEY_PASSWORD", raising=False)
|
||||
monkeypatch.setenv("VALKEY_HOST", "valkey-host")
|
||||
monkeypatch.setenv("VALKEY_PORT", "6380")
|
||||
|
||||
assert (
|
||||
ValkeySemanticCache._build_valkey_url(None, None, None)
|
||||
== "redis://valkey-host:6380"
|
||||
)
|
||||
|
||||
|
||||
def test_build_valkey_url_falls_back_to_redis_env(monkeypatch):
|
||||
monkeypatch.delenv("VALKEY_HOST", raising=False)
|
||||
monkeypatch.delenv("VALKEY_PORT", raising=False)
|
||||
monkeypatch.delenv("VALKEY_PASSWORD", raising=False)
|
||||
monkeypatch.setenv("REDIS_HOST", "redis-host")
|
||||
monkeypatch.setenv("REDIS_PORT", "6379")
|
||||
monkeypatch.setenv("REDIS_PASSWORD", "rpass")
|
||||
|
||||
assert (
|
||||
ValkeySemanticCache._build_valkey_url(None, None, None)
|
||||
== "redis://:rpass@redis-host:6379"
|
||||
)
|
||||
|
||||
|
||||
def test_build_valkey_url_requires_host_and_port(monkeypatch):
|
||||
for var in (
|
||||
"VALKEY_HOST",
|
||||
"VALKEY_PORT",
|
||||
"VALKEY_PASSWORD",
|
||||
"REDIS_HOST",
|
||||
"REDIS_PORT",
|
||||
"REDIS_PASSWORD",
|
||||
):
|
||||
monkeypatch.delenv(var, raising=False)
|
||||
|
||||
with pytest.raises(ValueError, match="Missing required Valkey configuration"):
|
||||
ValkeySemanticCache._build_valkey_url(None, None, None)
|
||||
|
||||
|
||||
def test_build_valkey_url_uses_rediss_scheme_when_ssl(monkeypatch):
|
||||
monkeypatch.setenv("VALKEY_HOST", "valkey-host")
|
||||
monkeypatch.setenv("VALKEY_PORT", "6379")
|
||||
monkeypatch.setenv("VALKEY_PASSWORD", "vpass")
|
||||
|
||||
assert (
|
||||
ValkeySemanticCache._build_valkey_url(None, None, None, ssl=True)
|
||||
== "rediss://:vpass@valkey-host:6379"
|
||||
)
|
||||
assert ValkeySemanticCache._build_valkey_url(
|
||||
"h", "6379", None, ssl=False
|
||||
).startswith("redis://")
|
||||
|
||||
|
||||
def test_init_requires_similarity_threshold():
|
||||
with pytest.raises(ValueError, match="similarity_threshold must be provided"):
|
||||
ValkeySemanticCache(sync_client=MagicMock(), async_client=AsyncMock())
|
||||
|
||||
|
||||
def test_init_rejects_cluster_startup_nodes():
|
||||
with pytest.raises(ValueError, match="cluster-mode-enabled"):
|
||||
ValkeySemanticCache(
|
||||
similarity_threshold=0.8,
|
||||
startup_nodes=[{"host": "shard1", "port": 6379}],
|
||||
)
|
||||
|
||||
|
||||
def test_cache_dispatch_rejects_cluster_for_valkey_semantic():
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.caching import LiteLLMCacheType
|
||||
|
||||
with pytest.raises(ValueError, match="cluster-mode-enabled"):
|
||||
Cache(
|
||||
type=LiteLLMCacheType.VALKEY_SEMANTIC,
|
||||
host="valkey-host",
|
||||
port="6379",
|
||||
similarity_threshold=0.8,
|
||||
redis_startup_nodes=[{"host": "shard1", "port": 6379}],
|
||||
)
|
||||
|
||||
|
||||
def test_scope_tag_is_deterministic_hex():
|
||||
tag = ValkeySemanticCache._scope_tag("model:gpt-4o::abc-123")
|
||||
assert tag == hashlib.sha256(b"model:gpt-4o::abc-123").hexdigest()
|
||||
assert len(tag) == 64
|
||||
assert ValkeySemanticCache._scope_tag("a") != ValkeySemanticCache._scope_tag("b")
|
||||
|
||||
|
||||
def test_embedding_to_bytes_is_little_endian_float32():
|
||||
assert ValkeySemanticCache._embedding_to_bytes([1.0, 0.0]) == struct.pack(
|
||||
"<2f", 1.0, 0.0
|
||||
)
|
||||
|
||||
|
||||
def test_set_cache_stores_scoped_doc_with_embedding(monkeypatch):
|
||||
sync_client = MagicMock()
|
||||
cache = _make_cache(sync_client=sync_client)
|
||||
cache._get_embedding = MagicMock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
cache.set_cache(
|
||||
key="cache-key",
|
||||
value={"content": "Paris"},
|
||||
messages=[{"role": "user", "content": "What is the capital of France?"}],
|
||||
)
|
||||
|
||||
sync_client.ft.return_value.create_index.assert_called_once()
|
||||
assert sync_client.hset.call_count == 1
|
||||
doc_key, kwargs = (
|
||||
sync_client.hset.call_args.args[0],
|
||||
sync_client.hset.call_args.kwargs,
|
||||
)
|
||||
mapping = kwargs["mapping"]
|
||||
scope = ValkeySemanticCache._scope_tag("cache-key")
|
||||
assert mapping[ValkeySemanticCache.CACHE_KEY_FIELD_NAME] == scope
|
||||
assert mapping["prompt"] == "What is the capital of France?"
|
||||
assert mapping["response"] == "{'content': 'Paris'}"
|
||||
assert mapping["embedding"] == struct.pack("<3f", 0.1, 0.2, 0.3)
|
||||
assert doc_key.startswith(f"test_index:{scope}:")
|
||||
|
||||
|
||||
def test_set_cache_applies_ttl():
|
||||
sync_client = MagicMock()
|
||||
cache = _make_cache(sync_client=sync_client)
|
||||
cache._get_embedding = MagicMock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
cache.set_cache(
|
||||
key="cache-key",
|
||||
value={"content": "Paris"},
|
||||
messages=[{"role": "user", "content": "What is the capital of France?"}],
|
||||
ttl=60,
|
||||
)
|
||||
|
||||
sync_client.expire.assert_called_once()
|
||||
assert sync_client.expire.call_args.args[1] == 60
|
||||
|
||||
|
||||
def test_set_cache_skips_ttl_when_absent():
|
||||
sync_client = MagicMock()
|
||||
cache = _make_cache(sync_client=sync_client)
|
||||
cache._get_embedding = MagicMock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
cache.set_cache(
|
||||
key="cache-key",
|
||||
value={"content": "Paris"},
|
||||
messages=[{"role": "user", "content": "What is the capital of France?"}],
|
||||
)
|
||||
|
||||
sync_client.expire.assert_not_called()
|
||||
|
||||
|
||||
def test_get_cache_returns_hit_above_threshold():
|
||||
sync_client = MagicMock()
|
||||
sync_client.ft.return_value.search.return_value = _search_result(0.1)
|
||||
cache = _make_cache(sync_client=sync_client, similarity_threshold=0.8)
|
||||
cache._get_embedding = MagicMock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
metadata = {}
|
||||
result = cache.get_cache(
|
||||
key="cache-key",
|
||||
messages=[{"role": "user", "content": "capital of France?"}],
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
assert result == {"content": "Paris"}
|
||||
assert metadata["semantic-similarity"] == pytest.approx(0.9)
|
||||
|
||||
|
||||
def test_get_cache_misses_below_threshold():
|
||||
sync_client = MagicMock()
|
||||
sync_client.ft.return_value.search.return_value = _search_result(0.5)
|
||||
cache = _make_cache(sync_client=sync_client, similarity_threshold=0.8)
|
||||
cache._get_embedding = MagicMock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
metadata = {}
|
||||
result = cache.get_cache(
|
||||
key="cache-key",
|
||||
messages=[{"role": "user", "content": "capital of Germany?"}],
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
assert result is None
|
||||
assert metadata["semantic-similarity"] == pytest.approx(0.5)
|
||||
|
||||
|
||||
def test_get_cache_misses_when_no_docs():
|
||||
sync_client = MagicMock()
|
||||
sync_client.ft.return_value.search.return_value = SimpleNamespace(docs=[])
|
||||
cache = _make_cache(sync_client=sync_client)
|
||||
cache._get_embedding = MagicMock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
metadata = {}
|
||||
result = cache.get_cache(
|
||||
key="cache-key",
|
||||
messages=[{"role": "user", "content": "capital of France?"}],
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
assert result is None
|
||||
assert metadata["semantic-similarity"] == 0.0
|
||||
|
||||
|
||||
def test_get_cache_query_filters_by_scope_tag():
|
||||
sync_client = MagicMock()
|
||||
sync_client.ft.return_value.search.return_value = _search_result(0.1)
|
||||
cache = _make_cache(sync_client=sync_client)
|
||||
cache._get_embedding = MagicMock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
cache.get_cache(
|
||||
key="cache-key",
|
||||
messages=[{"role": "user", "content": "capital of France?"}],
|
||||
metadata={},
|
||||
)
|
||||
|
||||
query = sync_client.ft.return_value.search.call_args.args[0]
|
||||
scope = ValkeySemanticCache._scope_tag("cache-key")
|
||||
assert scope in query.query_string()
|
||||
assert "KNN 1 @embedding" in query.query_string()
|
||||
|
||||
|
||||
def _async_ft(search_distance):
|
||||
search_obj = SimpleNamespace(
|
||||
search=AsyncMock(return_value=_search_result(search_distance)),
|
||||
create_index=AsyncMock(),
|
||||
)
|
||||
return MagicMock(return_value=search_obj)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_set_and_get_roundtrip():
|
||||
async_client = AsyncMock()
|
||||
async_client.ft = _async_ft(0.05)
|
||||
cache = _make_cache(async_client=async_client, similarity_threshold=0.8)
|
||||
cache._get_async_embedding = AsyncMock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
await cache.async_set_cache(
|
||||
key="cache-key",
|
||||
value={"content": "Paris"},
|
||||
messages=[{"role": "user", "content": "What is the capital of France?"}],
|
||||
ttl=30,
|
||||
)
|
||||
async_client.hset.assert_awaited_once()
|
||||
async_client.expire.assert_awaited_once()
|
||||
assert async_client.expire.call_args.args[1] == 30
|
||||
|
||||
metadata = {}
|
||||
result = await cache.async_get_cache(
|
||||
key="cache-key",
|
||||
messages=[{"role": "user", "content": "capital city of France"}],
|
||||
metadata=metadata,
|
||||
)
|
||||
assert result == {"content": "Paris"}
|
||||
assert metadata["semantic-similarity"] == pytest.approx(0.95)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_get_cache_misses_below_threshold():
|
||||
async_client = AsyncMock()
|
||||
async_client.ft = _async_ft(0.4)
|
||||
cache = _make_cache(async_client=async_client, similarity_threshold=0.8)
|
||||
cache._get_async_embedding = AsyncMock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
metadata = {}
|
||||
result = await cache.async_get_cache(
|
||||
key="cache-key",
|
||||
messages=[{"role": "user", "content": "capital of Germany?"}],
|
||||
metadata=metadata,
|
||||
)
|
||||
assert result is None
|
||||
assert metadata["semantic-similarity"] == pytest.approx(0.6)
|
||||
|
||||
|
||||
def test_ensure_index_swallows_already_exists():
|
||||
sync_client = MagicMock()
|
||||
sync_client.ft.return_value.create_index.side_effect = Exception(
|
||||
"Index test_index already exists."
|
||||
)
|
||||
cache = _make_cache(sync_client=sync_client)
|
||||
|
||||
cache._ensure_index_sync(3)
|
||||
assert cache._index_dim == 3
|
||||
|
||||
|
||||
def test_ensure_index_reraises_unexpected_error():
|
||||
sync_client = MagicMock()
|
||||
sync_client.ft.return_value.create_index.side_effect = Exception(
|
||||
"connection refused"
|
||||
)
|
||||
cache = _make_cache(sync_client=sync_client)
|
||||
|
||||
with pytest.raises(Exception, match="connection refused"):
|
||||
cache._ensure_index_sync(3)
|
||||
|
||||
|
||||
_FT_INFO_ATTRS_DIM_1536 = [
|
||||
[b"identifier", b"litellm_cache_key", b"type", b"TAG"],
|
||||
[
|
||||
b"identifier",
|
||||
b"embedding",
|
||||
b"type",
|
||||
b"VECTOR",
|
||||
b"index",
|
||||
[b"capacity", 10240, b"dimensions", 1536, b"distance_metric", b"COSINE"],
|
||||
],
|
||||
]
|
||||
|
||||
|
||||
def test_extract_index_dim_parses_nested_ft_info():
|
||||
info = {"attributes": _FT_INFO_ATTRS_DIM_1536}
|
||||
assert ValkeySemanticCache._extract_index_dim(info) == 1536
|
||||
assert ValkeySemanticCache._extract_index_dim({"attributes": []}) is None
|
||||
|
||||
|
||||
def test_ensure_index_raises_on_dimension_mismatch():
|
||||
sync_client = MagicMock()
|
||||
sync_client.ft.return_value.create_index.side_effect = Exception(
|
||||
"Index test_index already exists."
|
||||
)
|
||||
sync_client.ft.return_value.info.return_value = {
|
||||
"attributes": _FT_INFO_ATTRS_DIM_1536
|
||||
}
|
||||
cache = _make_cache(sync_client=sync_client)
|
||||
|
||||
with pytest.raises(
|
||||
ValueError, match="already exists with embedding dimension 1536"
|
||||
):
|
||||
cache._ensure_index_sync(768)
|
||||
assert cache._index_dim is None
|
||||
|
||||
|
||||
def test_ensure_index_accepts_matching_existing_dimension():
|
||||
sync_client = MagicMock()
|
||||
sync_client.ft.return_value.create_index.side_effect = Exception(
|
||||
"Index test_index already exists."
|
||||
)
|
||||
sync_client.ft.return_value.info.return_value = {
|
||||
"attributes": _FT_INFO_ATTRS_DIM_1536
|
||||
}
|
||||
cache = _make_cache(sync_client=sync_client)
|
||||
|
||||
cache._ensure_index_sync(1536)
|
||||
assert cache._index_dim == 1536
|
||||
|
||||
|
||||
def test_init_builds_only_missing_client_from_url():
|
||||
sync_client = MagicMock()
|
||||
cache = ValkeySemanticCache(
|
||||
similarity_threshold=0.8,
|
||||
redis_url="redis://valkey-host:6380",
|
||||
sync_client=sync_client,
|
||||
)
|
||||
assert cache.sync_client is sync_client
|
||||
assert cache.async_client is not None and cache.async_client is not sync_client
|
||||
|
||||
|
||||
def test_init_uses_both_injected_clients_without_connection_info(monkeypatch):
|
||||
for var in ("VALKEY_HOST", "VALKEY_PORT", "REDIS_HOST", "REDIS_PORT"):
|
||||
monkeypatch.delenv(var, raising=False)
|
||||
sync_client = MagicMock()
|
||||
async_client = AsyncMock()
|
||||
|
||||
cache = ValkeySemanticCache(
|
||||
similarity_threshold=0.8,
|
||||
sync_client=sync_client,
|
||||
async_client=async_client,
|
||||
)
|
||||
|
||||
assert cache.sync_client is sync_client
|
||||
assert cache.async_client is async_client
|
||||
|
||||
|
||||
def test_cache_dispatches_valkey_semantic_type():
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.caching import LiteLLMCacheType
|
||||
|
||||
cache = Cache(
|
||||
type=LiteLLMCacheType.VALKEY_SEMANTIC,
|
||||
host="valkey-host",
|
||||
port="6380",
|
||||
similarity_threshold=0.8,
|
||||
)
|
||||
|
||||
assert isinstance(cache.cache, ValkeySemanticCache)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_index_info_uses_valkey_ft_info():
|
||||
# The /health/readiness endpoint calls _index_info() on any
|
||||
# RedisSemanticCache instance; since ValkeySemanticCache subclasses it,
|
||||
# the inherited RedisVL implementation (which reads self.llmcache) would
|
||||
# break. This override must query valkey-search FT.INFO instead.
|
||||
async_client = AsyncMock()
|
||||
info_namespace = SimpleNamespace(info=AsyncMock(return_value={"num_docs": 3}))
|
||||
async_client.ft = MagicMock(return_value=info_namespace)
|
||||
cache = _make_cache(async_client=async_client)
|
||||
|
||||
result = await cache._index_info()
|
||||
|
||||
assert result == {"num_docs": 3}
|
||||
async_client.ft.assert_called_once_with("test_index")
|
||||
|
||||
|
||||
def test_importing_caching_does_not_require_redis():
|
||||
# redis is an optional dependency (extra_proxy), so the base SDK can be
|
||||
# installed without it. Selecting valkey-semantic needs redis, but merely
|
||||
# importing litellm.caching.caching must not, or `import litellm` breaks for
|
||||
# every base-SDK user. This runs in a subprocess with redis blocked so the
|
||||
# check is not polluted by redis already being imported in this session.
|
||||
code = textwrap.dedent("""
|
||||
import sys
|
||||
for name in ("redis", "redis.asyncio", "redis.commands",
|
||||
"redis.commands.search"):
|
||||
sys.modules[name] = None
|
||||
import litellm.caching.caching # must not import redis at module top
|
||||
from litellm.types.caching import LiteLLMCacheType
|
||||
assert LiteLLMCacheType.VALKEY_SEMANTIC == "valkey-semantic"
|
||||
print("ok")
|
||||
""")
|
||||
result = subprocess.run(
|
||||
[sys.executable, "-c", code],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env={**os.environ, "PYTHONPATH": _REPO_ROOT},
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert "ok" in result.stdout
|
||||
|
|
@ -255,3 +255,133 @@ async def test_should_skip_non_file_unified_id_on_output_file_id():
|
|||
assert batch_response.output_file_id == batch_unified
|
||||
mock_afile_retrieve.assert_not_called()
|
||||
managed_files.store_unified_file_id.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_passes_trusted_model_credentials_to_router():
|
||||
"""
|
||||
afile_content must hand the deployment's credential snapshot to the router
|
||||
call as an immutable server-side mapping. Cloud-storage providers (Bedrock
|
||||
S3) validate file ids against the bucket in that snapshot, so without it
|
||||
unified-id content retrieval only works when AWS_S3_BUCKET_NAME is set.
|
||||
"""
|
||||
from types import MappingProxyType
|
||||
|
||||
managed_files = _make_managed_files_instance()
|
||||
unified_file_id = "unified-file-id"
|
||||
s3_uri = "s3://my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
|
||||
managed_files.get_model_file_id_mapping = AsyncMock(
|
||||
return_value={unified_file_id: {"model-123": s3_uri}}
|
||||
)
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_deployment_credentials_with_provider = MagicMock(
|
||||
return_value={
|
||||
"custom_llm_provider": "bedrock",
|
||||
"s3_bucket_name": "my-bucket",
|
||||
"aws_region_name": "us-west-2",
|
||||
}
|
||||
)
|
||||
mock_router.afile_content = AsyncMock(return_value=MagicMock())
|
||||
|
||||
await managed_files.afile_content(
|
||||
file_id=unified_file_id,
|
||||
litellm_parent_otel_span=None,
|
||||
llm_router=mock_router,
|
||||
)
|
||||
|
||||
call_kwargs = mock_router.afile_content.call_args.kwargs
|
||||
assert call_kwargs["model"] == "model-123"
|
||||
assert call_kwargs["file_id"] == s3_uri
|
||||
trusted_credentials = call_kwargs["_litellm_internal_model_credentials"]
|
||||
assert isinstance(trusted_credentials, MappingProxyType)
|
||||
assert trusted_credentials["s3_bucket_name"] == "my-bucket"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_bedrock_unified_id_end_to_end(monkeypatch):
|
||||
"""
|
||||
Proxy repro for Bedrock batch output retrieval: a unified file id that
|
||||
resolves to an s3:// output object must be fetched via a SigV4-signed S3
|
||||
GET using the deployment's s3_bucket_name (no AWS_S3_BUCKET_NAME env).
|
||||
|
||||
Regression test for "BedrockFilesConfig does not support file content
|
||||
retrieval" raised on this path.
|
||||
"""
|
||||
import httpx
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
|
||||
monkeypatch.delenv("AWS_S3_BUCKET_NAME", raising=False)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "bedrock-claude",
|
||||
"litellm_params": {
|
||||
"model": "bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
"aws_access_key_id": "AKIAEXAMPLE",
|
||||
"aws_secret_access_key": "secret",
|
||||
"aws_region_name": "us-west-2",
|
||||
"s3_bucket_name": "my-bucket",
|
||||
},
|
||||
"model_info": {"id": "model-123"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
managed_files = _make_managed_files_instance()
|
||||
unified_file_id = "unified-file-id"
|
||||
s3_uri = "s3://my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
|
||||
managed_files.get_model_file_id_mapping = AsyncMock(
|
||||
return_value={unified_file_id: {"model-123": s3_uri}}
|
||||
)
|
||||
|
||||
expected_url = "https://s3.us-west-2.amazonaws.com/my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
|
||||
with respx.mock:
|
||||
route = respx.get(expected_url).mock(
|
||||
return_value=httpx.Response(200, content=b'{"recordId": "x"}')
|
||||
)
|
||||
|
||||
response = await managed_files.afile_content(
|
||||
file_id=unified_file_id,
|
||||
litellm_parent_otel_span=None,
|
||||
llm_router=router,
|
||||
)
|
||||
|
||||
assert route.called
|
||||
assert (
|
||||
route.calls[0].request.headers["Authorization"].startswith("AWS4-HMAC-SHA256")
|
||||
)
|
||||
assert response.content == b'{"recordId": "x"}'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_error_reports_unified_id_not_provider_uri():
|
||||
"""When every model attempt fails, the error must name the caller's unified
|
||||
file id, never the resolved internal s3:// URI (no internal-path leak)."""
|
||||
managed_files = _make_managed_files_instance()
|
||||
unified_file_id = "litellm_proxy_unified_id_abc"
|
||||
s3_uri = "s3://my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
|
||||
managed_files.get_model_file_id_mapping = AsyncMock(
|
||||
return_value={unified_file_id: {"model-123": s3_uri}}
|
||||
)
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_deployment_credentials_with_provider = MagicMock(return_value=None)
|
||||
mock_router.afile_content = AsyncMock(side_effect=Exception("deployment failed"))
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
await managed_files.afile_content(
|
||||
file_id=unified_file_id,
|
||||
litellm_parent_otel_span=None,
|
||||
llm_router=mock_router,
|
||||
)
|
||||
|
||||
message = str(exc_info.value)
|
||||
assert unified_file_id in message
|
||||
assert s3_uri not in message
|
||||
|
|
|
|||
|
|
@ -123,9 +123,53 @@ def test_non_participating_callback_uses_default_tracer():
|
|||
def test_dynamic_headers_applied_to_otlp_exporter_only():
|
||||
cache = _cache(
|
||||
"arize",
|
||||
exporters=[ExporterSpec(kind="otlp_http"), ExporterSpec(kind="in_memory")],
|
||||
exporters=[
|
||||
ExporterSpec(kind="otlp_http", owner="arize"),
|
||||
ExporterSpec(kind="in_memory", owner="arize"),
|
||||
],
|
||||
)
|
||||
new_cfg = cache._config_with_headers({"arize-space-id": "S", "api_key": "K"})
|
||||
otlp, in_mem = new_cfg.exporters
|
||||
assert otlp.headers == "arize-space-id=S,api_key=K"
|
||||
assert in_mem.headers is None # console/in_memory left untouched
|
||||
|
||||
|
||||
def test_dynamic_headers_do_not_leak_to_other_owners_exporter():
|
||||
"""A tenant's Arize credentials must never be stamped onto a co-configured
|
||||
exporter owned by a different backend (a self-hosted collector, Langfuse).
|
||||
|
||||
Regression for the cross-backend credential leak: ``_config_with_headers``
|
||||
used to rewrite the headers of every OTLP exporter, so one request carrying
|
||||
a team's Arize key clobbered the base collector's and Langfuse's headers
|
||||
with that key.
|
||||
"""
|
||||
cache = _cache(
|
||||
"arize",
|
||||
exporters=[
|
||||
ExporterSpec(
|
||||
kind="otlp_http",
|
||||
endpoint="http://self-hosted-collector:4318",
|
||||
headers="x=base-collector",
|
||||
owner=None,
|
||||
),
|
||||
ExporterSpec(
|
||||
kind="otlp_http",
|
||||
endpoint="https://cloud.langfuse.com/api/public/otel",
|
||||
headers="Authorization=Basic base-langfuse",
|
||||
owner="langfuse_otel",
|
||||
),
|
||||
ExporterSpec(
|
||||
kind="otlp_grpc",
|
||||
endpoint="https://otlp.arize.com/v1",
|
||||
headers="space_id=base,api_key=base",
|
||||
owner="arize",
|
||||
),
|
||||
],
|
||||
)
|
||||
new_cfg = cache._config_with_headers(
|
||||
{"arize-space-id": "TEAMX", "api_key": "TEAMX_KEY"}
|
||||
)
|
||||
by_owner = {e.owner: e.headers for e in new_cfg.exporters}
|
||||
assert by_owner["arize"] == "arize-space-id=TEAMX,api_key=TEAMX_KEY"
|
||||
assert by_owner[None] == "x=base-collector"
|
||||
assert by_owner["langfuse_otel"] == "Authorization=Basic base-langfuse"
|
||||
|
|
|
|||
|
|
@ -1058,6 +1058,97 @@ def test_proxy_global_first_registered_wins(monkeypatch):
|
|||
assert second is not first
|
||||
|
||||
|
||||
def test_select_global_otel_v2_logger_reuses_existing_preset_logger():
|
||||
"""The global-provider selection must reuse the logger the callback factory
|
||||
already built (e.g. an arize preset logger that folds the OTEL_* base exporter
|
||||
and its own exporter into one logger), not mint a second generic one.
|
||||
|
||||
Regression for the orphan span: the startup publish used to search
|
||||
``service_callback`` (which a preset logger does not always reach), miss the
|
||||
existing logger, and build a second generic ``OpenTelemetryV2`` whose provider
|
||||
became the OTel global. The server span then exported through that generic
|
||||
provider while the preset logger's gen-ai spans exported to the preset backend,
|
||||
so on that backend the LLM span had no parent. Selecting from the loggers the
|
||||
factory registered keeps one logger, one provider, one connected trace.
|
||||
"""
|
||||
from litellm.integrations.otel.logger import select_global_otel_v2_logger
|
||||
|
||||
cfg = OpenTelemetryV2Config(exporter="in_memory")
|
||||
tp = providers.build_tracer_provider(cfg)
|
||||
preset_logger = OpenTelemetryV2(
|
||||
config=cfg, callback_name="arize", tracer_provider=tp
|
||||
)
|
||||
|
||||
chosen = select_global_otel_v2_logger([object(), preset_logger, object()])
|
||||
assert chosen is preset_logger
|
||||
|
||||
|
||||
def test_select_global_otel_v2_logger_prefers_registered_owner_over_list_scan():
|
||||
"""Selection reuses the canonical owner the factory registered, not whatever
|
||||
the ``in_memory_loggers`` scan happens to reach first.
|
||||
|
||||
The factory designates one logger as ``proxy_server.open_telemetry_logger`` the
|
||||
moment it builds the first one, and every other v2 path (guardrail, seed,
|
||||
phase spans) routes through that owner. With two presets configured, the list
|
||||
scan's "first ``OpenTelemetryV2``" is order-dependent and could disagree with
|
||||
that owner, publishing one backend's provider as the global while the rest of
|
||||
the v2 code emits through another. Passing the registered owner pins the global
|
||||
provider to the same logger the rest of the code already uses.
|
||||
"""
|
||||
from litellm.integrations.otel.logger import select_global_otel_v2_logger
|
||||
|
||||
cfg = OpenTelemetryV2Config(exporter="in_memory")
|
||||
owner = OpenTelemetryV2(
|
||||
config=cfg,
|
||||
callback_name="arize",
|
||||
tracer_provider=providers.build_tracer_provider(cfg),
|
||||
)
|
||||
other = OpenTelemetryV2(
|
||||
config=cfg,
|
||||
callback_name="langfuse_otel",
|
||||
tracer_provider=providers.build_tracer_provider(cfg),
|
||||
)
|
||||
|
||||
chosen = select_global_otel_v2_logger([other, owner], registered=owner)
|
||||
assert chosen is owner
|
||||
|
||||
|
||||
def test_select_global_otel_v2_logger_builds_one_when_none_registered():
|
||||
"""With no logger registered, selection builds exactly one generic logger so
|
||||
the proxy still publishes a provider; it must not return ``None``."""
|
||||
from litellm.integrations.otel.logger import select_global_otel_v2_logger
|
||||
|
||||
chosen = select_global_otel_v2_logger([])
|
||||
assert isinstance(chosen, OpenTelemetryV2)
|
||||
|
||||
|
||||
def test_publish_global_otel_v2_provider_sets_selected_logger_provider():
|
||||
"""The startup publish must set the OTel global provider to the *selected*
|
||||
logger's provider (the preset logger that owns every exporter), so the FastAPI
|
||||
server span and the gen-ai spans share one provider and one trace.
|
||||
|
||||
Drives the publish step the proxy runs at startup, with the global-setter
|
||||
injected so no real global OTel state is mutated. Guards the wiring that a unit
|
||||
test would otherwise miss: that the published provider is the selected logger's,
|
||||
not some other.
|
||||
"""
|
||||
from litellm.integrations.otel.logger import publish_global_otel_v2_provider
|
||||
|
||||
cfg = OpenTelemetryV2Config(exporter="in_memory")
|
||||
tp = providers.build_tracer_provider(cfg)
|
||||
preset_logger = OpenTelemetryV2(
|
||||
config=cfg, callback_name="arize", tracer_provider=tp
|
||||
)
|
||||
|
||||
published = []
|
||||
chosen = publish_global_otel_v2_provider(
|
||||
[object(), preset_logger], published.append
|
||||
)
|
||||
|
||||
assert chosen is preset_logger
|
||||
assert published == [preset_logger._tracer_provider]
|
||||
|
||||
|
||||
def test_registers_into_litellm_service_callback(monkeypatch):
|
||||
"""The logger must mutate ``litellm.service_callback`` in place. An empty
|
||||
list is falsy, so a ``getattr(..) or []`` would append to a throwaway local
|
||||
|
|
|
|||
|
|
@ -44,6 +44,37 @@ def test_agentops_exporter_factory_is_registered():
|
|||
assert _AGENTOPS_EXPORTER_KIND in providers._EXPORTER_FACTORIES
|
||||
|
||||
|
||||
def test_dynamic_cred_presets_tag_exporter_with_matching_owner(monkeypatch):
|
||||
"""Each dynamic-credential preset must tag the exporter it contributes with
|
||||
its own callback name, so per-request tenant routing
|
||||
(``TenantTracerCache``) applies that integration's credentials only to its
|
||||
own exporter and never bleeds them onto a co-configured backend.
|
||||
"""
|
||||
from litellm.integrations.otel.presets import (
|
||||
DYNAMIC_HEADERS_BY_CALLBACK,
|
||||
PRESET_BY_CALLBACK,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("ARIZE_SPACE_ID", "S")
|
||||
monkeypatch.setenv("ARIZE_API_KEY", "K")
|
||||
monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk")
|
||||
monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk")
|
||||
monkeypatch.setenv("LANGFUSE_HOST", "https://cloud.langfuse.com")
|
||||
monkeypatch.setenv("WANDB_API_KEY", "w")
|
||||
monkeypatch.setenv("WANDB_PROJECT_ID", "entity/project")
|
||||
|
||||
from litellm.integrations.otel.model.config import ExporterOwner
|
||||
|
||||
for callback_name in DYNAMIC_HEADERS_BY_CALLBACK:
|
||||
cfg = PRESET_BY_CALLBACK[callback_name]()
|
||||
owners = {e.owner for e in cfg.exporters}
|
||||
assert ExporterOwner(callback_name) in owners, (
|
||||
f"{callback_name} preset did not tag its exporter with "
|
||||
f"owner={callback_name!r}; tenant credentials would leak across "
|
||||
f"exporters. owners present: {owners}"
|
||||
)
|
||||
|
||||
|
||||
def test_agentops_exporter_mints_jwt_lazily(monkeypatch):
|
||||
pytest.importorskip("opentelemetry.exporter.otlp.proto.http.trace_exporter")
|
||||
monkeypatch.setattr(
|
||||
|
|
|
|||
|
|
@ -0,0 +1,15 @@
|
|||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
is_managed_cloud_storage_uri,
|
||||
)
|
||||
|
||||
|
||||
def test_is_managed_cloud_storage_uri_detects_raw_object_uris():
|
||||
assert is_managed_cloud_storage_uri("s3://bucket/litellm-batch-outputs/x.jsonl.out")
|
||||
assert is_managed_cloud_storage_uri("gs://bucket/litellm-vertex-files/x")
|
||||
|
||||
|
||||
def test_is_managed_cloud_storage_uri_ignores_provider_and_unified_ids():
|
||||
# Plain provider ids and base64 unified ids carry no storage scheme.
|
||||
assert not is_managed_cloud_storage_uri("file-abc123")
|
||||
assert not is_managed_cloud_storage_uri("bGl0ZWxsbV9wcm94eQ==")
|
||||
assert not is_managed_cloud_storage_uri("")
|
||||
|
|
@ -3406,3 +3406,46 @@ def test_handle_anthropic_messages_response_logging_degrades_on_unparseable_resp
|
|||
assert isinstance(result, ModelResponse)
|
||||
assert result.model == "openai/my-local"
|
||||
assert result.usage.prompt_tokens == 4 # type: ignore[attr-defined]
|
||||
|
||||
|
||||
def test_failure_handler_records_recovered_partial_spend(logging_obj):
|
||||
"""A stream interrupted mid-flight still billed the provider for the chunks
|
||||
already delivered. When the router stashes that recovered usage as
|
||||
``combined_usage_object`` and pre-computes ``response_cost``, the failure
|
||||
handler must preserve them so the failure row carries the real partial
|
||||
spend instead of zero.
|
||||
"""
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
logging_obj.model_call_details["combined_usage_object"] = Usage(
|
||||
prompt_tokens=17, completion_tokens=9, total_tokens=26
|
||||
)
|
||||
logging_obj.model_call_details["response_cost"] = 0.00012
|
||||
|
||||
logging_obj._failure_handler_helper_fn(
|
||||
exception=Exception("Connection lost"),
|
||||
traceback_exception="Traceback ...",
|
||||
)
|
||||
|
||||
payload = logging_obj.model_call_details["standard_logging_object"]
|
||||
assert payload["status"] == "failure"
|
||||
assert payload["response_cost"] == 0.00012
|
||||
assert payload["prompt_tokens"] == 17
|
||||
assert payload["completion_tokens"] == 9
|
||||
assert payload["total_tokens"] == 26
|
||||
|
||||
|
||||
def test_failure_handler_zeroes_spend_without_recovered_usage(logging_obj):
|
||||
"""A failure with no recovered partial usage keeps the existing behavior of
|
||||
recording zero spend, so the partial-spend preservation does not leak into
|
||||
ordinary failures.
|
||||
"""
|
||||
logging_obj._failure_handler_helper_fn(
|
||||
exception=Exception("boom"),
|
||||
traceback_exception="Traceback ...",
|
||||
)
|
||||
|
||||
payload = logging_obj.model_call_details["standard_logging_object"]
|
||||
assert payload["status"] == "failure"
|
||||
assert payload["response_cost"] == 0
|
||||
assert payload["total_tokens"] == 0
|
||||
|
|
|
|||
|
|
@ -2325,3 +2325,79 @@ def test_chunk_creator_tool_calls_not_dropped_on_finish(
|
|||
assert result.choices[0].delta.tool_calls is not None
|
||||
assert result.choices[0].finish_reason is None
|
||||
assert initialized_custom_stream_wrapper.received_finish_reason == "tool_calls"
|
||||
|
||||
|
||||
def test_record_partial_usage_for_failure_stashes_usage_and_cost():
|
||||
"""A stream that breaks mid-flight must surface the usage assembled from the
|
||||
chunks already delivered, plus its cost, on the logging object so the
|
||||
failure handler records the real partial spend instead of zero.
|
||||
"""
|
||||
logging_obj = Logging(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "Hey"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=time.time(),
|
||||
litellm_call_id="partial-usage-1",
|
||||
function_id="1245",
|
||||
)
|
||||
logging_obj.model_call_details["custom_llm_provider"] = "openai"
|
||||
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=None,
|
||||
model="gpt-4o-mini",
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
wrapper.chunks = [
|
||||
ModelResponseStream(
|
||||
id="chatcmpl-partial-1",
|
||||
created=1742056047,
|
||||
model="gpt-4o-mini",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(
|
||||
content="The Roman Empire began when", role="assistant"
|
||||
),
|
||||
)
|
||||
],
|
||||
usage=Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31),
|
||||
)
|
||||
]
|
||||
|
||||
wrapper._record_partial_usage_for_failure()
|
||||
|
||||
stashed = logging_obj.model_call_details["combined_usage_object"]
|
||||
assert stashed.prompt_tokens == 30
|
||||
assert stashed.completion_tokens == 1
|
||||
assert stashed.total_tokens == 31
|
||||
assert isinstance(logging_obj.model_call_details["response_cost"], float)
|
||||
|
||||
|
||||
def test_record_partial_usage_for_failure_noop_without_chunks():
|
||||
"""With no chunks delivered there is nothing billed to recover, so the
|
||||
failure stash must stay absent and not force a zero-usage row.
|
||||
"""
|
||||
logging_obj = Logging(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "Hey"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=time.time(),
|
||||
litellm_call_id="partial-usage-2",
|
||||
function_id="1245",
|
||||
)
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=None,
|
||||
model="gpt-4o-mini",
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
wrapper.chunks = []
|
||||
|
||||
wrapper._record_partial_usage_for_failure()
|
||||
|
||||
assert "combined_usage_object" not in logging_obj.model_call_details
|
||||
|
|
|
|||
|
|
@ -2747,3 +2747,23 @@ def test_translate_openai_response_to_anthropic_with_polyfill_both_compaction_an
|
|||
cm = result.get("context_management")
|
||||
assert cm is not None
|
||||
assert cm["applied_edits"][0]["type"] == "compact_20260112"
|
||||
|
||||
|
||||
def test_translate_anthropic_tools_to_openai_preserves_parameters_type():
|
||||
"""Regression for #30557: the Anthropic tool `type` ("custom") must not be
|
||||
merged into the OpenAI function `parameters`, overwriting parameters.type."""
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
tools = [
|
||||
{
|
||||
"type": "custom",
|
||||
"name": "get_weather",
|
||||
"description": "Get weather",
|
||||
"input_schema": {"type": "object", "properties": {}},
|
||||
}
|
||||
]
|
||||
|
||||
new_tools, _ = adapter.translate_anthropic_tools_to_openai(tools=tools)
|
||||
|
||||
params = new_tools[0]["function"]["parameters"]
|
||||
assert params["type"] == "object"
|
||||
assert new_tools[0]["type"] == "function"
|
||||
|
|
|
|||
|
|
@ -4,8 +4,11 @@ Test bedrock files transformation functionality
|
|||
|
||||
import json
|
||||
import os
|
||||
from unittest.mock import MagicMock
|
||||
from urllib.parse import unquote, urlparse
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.bedrock.files.transformation import BedrockJsonlFilesTransformation
|
||||
|
||||
|
||||
|
|
@ -1173,3 +1176,314 @@ class TestBedrockFilesEmbeddingTransformation:
|
|||
assert not BedrockFilesConfig._is_embedding_record(
|
||||
{"url": "/v1/responses", "body": {"input": "x"}}
|
||||
)
|
||||
|
||||
|
||||
class TestBedrockFileContentTransformation:
|
||||
"""SigV4-signed S3 GetObject retrieval of Bedrock batch output files."""
|
||||
|
||||
S3_URI = "s3://my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
|
||||
EXPECTED_URL = "https://s3.us-west-2.amazonaws.com/my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
|
||||
|
||||
def _litellm_params(self) -> dict:
|
||||
return {
|
||||
"aws_access_key_id": "AKIAEXAMPLE",
|
||||
"aws_secret_access_key": "secret",
|
||||
"aws_region_name": "us-west-2",
|
||||
}
|
||||
|
||||
def test_transform_file_content_request_signs_s3_get(self, monkeypatch):
|
||||
"""The request transform must produce the S3 object URL plus SigV4 GET headers."""
|
||||
import hashlib
|
||||
|
||||
from litellm.llms.bedrock.files.transformation import (
|
||||
S3_SIGNED_GET_HEADERS_PARAM,
|
||||
BedrockFilesConfig,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
litellm_params = self._litellm_params()
|
||||
|
||||
url, params = BedrockFilesConfig().transform_file_content_request(
|
||||
file_content_request={"file_id": self.S3_URI},
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert url == self.EXPECTED_URL
|
||||
assert params == {}
|
||||
|
||||
signed_headers = litellm_params[S3_SIGNED_GET_HEADERS_PARAM]
|
||||
assert (
|
||||
signed_headers["x-amz-content-sha256"] == hashlib.sha256(b"").hexdigest()
|
||||
), "GET has no payload, so the content hash must be the empty-body hash"
|
||||
authorization = signed_headers["Authorization"]
|
||||
assert authorization.startswith("AWS4-HMAC-SHA256 Credential=AKIAEXAMPLE/")
|
||||
assert "/us-west-2/s3/aws4_request" in authorization
|
||||
assert "x-amz-content-sha256" in authorization
|
||||
assert "X-Amz-Date" in signed_headers
|
||||
|
||||
def test_transform_file_content_request_decodes_unified_file_id(self, monkeypatch):
|
||||
"""Base64 unified ids carrying llm_output_file_id must resolve to their S3 object."""
|
||||
import base64
|
||||
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
from litellm.types.utils import SpecialEnums
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
unified_file_id = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
|
||||
"application/json", "unified-id", "", self.S3_URI, "model-id"
|
||||
)
|
||||
encoded_file_id = (
|
||||
base64.urlsafe_b64encode(unified_file_id.encode()).decode().rstrip("=")
|
||||
)
|
||||
|
||||
url, _ = BedrockFilesConfig().transform_file_content_request(
|
||||
file_content_request={"file_id": encoded_file_id},
|
||||
optional_params={},
|
||||
litellm_params=self._litellm_params(),
|
||||
)
|
||||
|
||||
assert url == self.EXPECTED_URL
|
||||
|
||||
def test_transform_file_content_request_rejects_foreign_bucket(self, monkeypatch):
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
|
||||
with pytest.raises(ValueError, match="configured storage bucket"):
|
||||
BedrockFilesConfig().transform_file_content_request(
|
||||
file_content_request={
|
||||
"file_id": "s3://other-bucket/litellm-batch-outputs/job/x.jsonl.out"
|
||||
},
|
||||
optional_params={},
|
||||
litellm_params=self._litellm_params(),
|
||||
)
|
||||
|
||||
def test_transform_file_content_request_rejects_unmanaged_key(self, monkeypatch):
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
|
||||
with pytest.raises(ValueError, match="LiteLLM-managed"):
|
||||
BedrockFilesConfig().transform_file_content_request(
|
||||
file_content_request={"file_id": "s3://my-bucket/private/x.jsonl"},
|
||||
optional_params={},
|
||||
litellm_params=self._litellm_params(),
|
||||
)
|
||||
|
||||
def test_extract_s3_uri_rejects_non_managed_file_id(self):
|
||||
"""A file id that is neither an s3:// URI nor a unified id must be rejected."""
|
||||
from litellm.llms.bedrock.files.transformation import (
|
||||
extract_s3_uri_from_file_id,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="managed LiteLLM S3 file id"):
|
||||
extract_s3_uri_from_file_id("file-1234567890")
|
||||
|
||||
def test_transform_file_content_request_requires_configured_bucket(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""Without a server-configured bucket (env or snapshot), the request must fail
|
||||
before any S3 call rather than guessing a bucket from the file id."""
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.delenv("AWS_S3_BUCKET_NAME", raising=False)
|
||||
|
||||
with pytest.raises(ValueError, match="S3 bucket_name is required"):
|
||||
BedrockFilesConfig().transform_file_content_request(
|
||||
file_content_request={"file_id": self.S3_URI},
|
||||
optional_params={},
|
||||
litellm_params=self._litellm_params(),
|
||||
)
|
||||
|
||||
def test_transform_file_content_request_requires_file_id(self, monkeypatch):
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
|
||||
with pytest.raises(ValueError, match="file_id is required"):
|
||||
BedrockFilesConfig().transform_file_content_request(
|
||||
file_content_request={},
|
||||
optional_params={},
|
||||
litellm_params=self._litellm_params(),
|
||||
)
|
||||
|
||||
def test_sign_request_without_botocore_raises_helpful_error(self, monkeypatch):
|
||||
"""A missing botocore must surface an actionable 'install boto3' error
|
||||
rather than a raw import failure."""
|
||||
import sys
|
||||
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
monkeypatch.setitem(sys.modules, "botocore.auth", None)
|
||||
|
||||
with pytest.raises(ImportError, match="boto3"):
|
||||
BedrockFilesConfig().transform_file_content_request(
|
||||
file_content_request={"file_id": self.S3_URI},
|
||||
optional_params={},
|
||||
litellm_params=self._litellm_params(),
|
||||
)
|
||||
|
||||
def test_bucket_resolved_from_trusted_model_credentials(self, monkeypatch):
|
||||
"""Per-model s3_bucket_name must be honored via the server-side credential snapshot."""
|
||||
from types import MappingProxyType
|
||||
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.delenv("AWS_S3_BUCKET_NAME", raising=False)
|
||||
litellm_params = self._litellm_params()
|
||||
litellm_params["_litellm_internal_model_credentials"] = MappingProxyType(
|
||||
{"s3_bucket_name": "my-bucket"}
|
||||
)
|
||||
|
||||
url, _ = BedrockFilesConfig().transform_file_content_request(
|
||||
file_content_request={"file_id": self.S3_URI},
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert url == self.EXPECTED_URL
|
||||
|
||||
def test_s3_region_name_wins_for_content_signing(self, monkeypatch):
|
||||
"""s3_region_name must override aws_region_name for both the URL and the signature."""
|
||||
from litellm.llms.bedrock.files.transformation import (
|
||||
S3_SIGNED_GET_HEADERS_PARAM,
|
||||
BedrockFilesConfig,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
litellm_params = self._litellm_params()
|
||||
litellm_params["s3_region_name"] = "eu-west-1"
|
||||
|
||||
url, _ = BedrockFilesConfig().transform_file_content_request(
|
||||
file_content_request={"file_id": self.S3_URI},
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert url.startswith("https://s3.eu-west-1.amazonaws.com/")
|
||||
authorization = litellm_params[S3_SIGNED_GET_HEADERS_PARAM]["Authorization"]
|
||||
assert "/eu-west-1/s3/aws4_request" in authorization
|
||||
|
||||
def test_validate_environment_merges_and_pops_signed_get_headers(self):
|
||||
from litellm.llms.bedrock.files.transformation import (
|
||||
S3_SIGNED_GET_HEADERS_PARAM,
|
||||
BedrockFilesConfig,
|
||||
)
|
||||
|
||||
litellm_params = {
|
||||
S3_SIGNED_GET_HEADERS_PARAM: {"Authorization": "AWS4-HMAC-SHA256 test"}
|
||||
}
|
||||
|
||||
headers = BedrockFilesConfig().validate_environment(
|
||||
headers={"x-custom": "kept"},
|
||||
model="",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert headers == {
|
||||
"x-custom": "kept",
|
||||
"Authorization": "AWS4-HMAC-SHA256 test",
|
||||
}
|
||||
assert S3_SIGNED_GET_HEADERS_PARAM not in litellm_params
|
||||
|
||||
def test_transform_file_content_response_wraps_binary_content(self):
|
||||
import httpx
|
||||
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
|
||||
raw_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=b'{"recordId": "CALL0000001"}',
|
||||
request=httpx.Request("GET", self.EXPECTED_URL),
|
||||
)
|
||||
|
||||
result = BedrockFilesConfig().transform_file_content_response(
|
||||
raw_response=raw_response,
|
||||
logging_obj=MagicMock(),
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
assert result.response.content == b'{"recordId": "CALL0000001"}'
|
||||
|
||||
def test_transform_file_content_response_raises_on_s3_error(self):
|
||||
import httpx
|
||||
|
||||
from litellm.llms.bedrock.common_utils import BedrockError
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
raw_response = httpx.Response(
|
||||
status_code=403,
|
||||
content=b"<Error><Code>AccessDenied</Code></Error>",
|
||||
request=httpx.Request("GET", self.EXPECTED_URL),
|
||||
)
|
||||
|
||||
with pytest.raises(BedrockError, match="AccessDenied"):
|
||||
BedrockFilesConfig().transform_file_content_response(
|
||||
raw_response=raw_response,
|
||||
logging_obj=MagicMock(),
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
def test_file_content_end_to_end_sends_signed_get(self, monkeypatch):
|
||||
"""litellm.file_content must issue a SigV4-signed GET and return the S3 object bytes."""
|
||||
import httpx
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
|
||||
with respx.mock:
|
||||
route = respx.get(self.EXPECTED_URL).mock(
|
||||
return_value=httpx.Response(200, content=b'{"recordId": "x"}')
|
||||
)
|
||||
|
||||
response = litellm.file_content(
|
||||
file_id=self.S3_URI,
|
||||
custom_llm_provider="bedrock",
|
||||
**self._litellm_params(),
|
||||
)
|
||||
|
||||
assert route.called
|
||||
request = route.calls[0].request
|
||||
assert request.headers["Authorization"].startswith("AWS4-HMAC-SHA256")
|
||||
assert "x-amz-content-sha256" in request.headers
|
||||
assert response.content == b'{"recordId": "x"}'
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_end_to_end_sends_signed_get(self, monkeypatch):
|
||||
"""Async variant: litellm.afile_content over the same signed GET path."""
|
||||
import httpx
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
# respx can only intercept httpx transports
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
with respx.mock:
|
||||
route = respx.get(self.EXPECTED_URL).mock(
|
||||
return_value=httpx.Response(200, content=b'{"recordId": "x"}')
|
||||
)
|
||||
|
||||
response = await litellm.afile_content(
|
||||
file_id=self.S3_URI,
|
||||
custom_llm_provider="bedrock",
|
||||
**self._litellm_params(),
|
||||
)
|
||||
|
||||
assert route.called
|
||||
assert (
|
||||
route.calls[0]
|
||||
.request.headers["Authorization"]
|
||||
.startswith("AWS4-HMAC-SHA256")
|
||||
)
|
||||
assert response.content == b'{"recordId": "x"}'
|
||||
|
|
|
|||
|
|
@ -174,7 +174,9 @@ class TestBedrockMantleResponsesURL:
|
|||
|
||||
|
||||
class TestBedrockMantleGetLlmProviderRegion:
|
||||
def test_get_llm_provider_uses_supplemental_litellm_params(self, monkeypatch):
|
||||
def test_get_llm_provider_uses_supplemental_litellm_params(
|
||||
self, monkeypatch, local_cost_map
|
||||
):
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_REGION", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
|
||||
monkeypatch.delenv("AWS_REGION", raising=False)
|
||||
|
|
@ -187,9 +189,13 @@ class TestBedrockMantleGetLlmProviderRegion:
|
|||
litellm_params=GenericLiteLLMParams(aws_region_name="us-east-2"),
|
||||
)
|
||||
assert provider == "bedrock_mantle"
|
||||
assert api_base == "https://bedrock-mantle.us-east-2.api.aws/v1"
|
||||
# gpt-5.x carries use_openai_responses_path, so its whole surface (incl.
|
||||
# the resolved chat base) is on the /openai/v1 base per the AWS card.
|
||||
assert api_base == "https://bedrock-mantle.us-east-2.api.aws/openai/v1"
|
||||
|
||||
def test_get_llm_provider_uses_aws_region_from_litellm_params(self, monkeypatch):
|
||||
def test_get_llm_provider_uses_aws_region_from_litellm_params(
|
||||
self, monkeypatch, local_cost_map
|
||||
):
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_REGION", raising=False)
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
|
||||
monkeypatch.delenv("AWS_REGION", raising=False)
|
||||
|
|
@ -205,7 +211,7 @@ class TestBedrockMantleGetLlmProviderRegion:
|
|||
litellm_params=params,
|
||||
)
|
||||
assert provider == "bedrock_mantle"
|
||||
assert api_base == "https://bedrock-mantle.us-east-2.api.aws/v1"
|
||||
assert api_base == "https://bedrock-mantle.us-east-2.api.aws/openai/v1"
|
||||
|
||||
|
||||
class TestBedrockMantleResponsesAuth:
|
||||
|
|
@ -368,7 +374,10 @@ class TestBedrockMantleResponsesTools:
|
|||
|
||||
|
||||
class TestBedrockMantleResponsesRegistry:
|
||||
def test_registry_returns_config_for_gpt_5_5(self):
|
||||
def test_registry_returns_config_for_gpt_5_5(self, local_cost_map):
|
||||
# gpt-5.x advertises /v1/responses in supported_endpoints (capability)
|
||||
# and use_openai_responses_path (wire path), so it gets the native config
|
||||
# on the /openai/v1/responses path. local_cost_map loads the entry.
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
cfg = ProviderConfigManager.get_provider_responses_api_config(
|
||||
|
|
@ -378,7 +387,7 @@ class TestBedrockMantleResponsesRegistry:
|
|||
assert isinstance(cfg, BedrockMantleResponsesAPIConfig)
|
||||
assert cfg.use_openai_path is True
|
||||
|
||||
def test_registry_returns_config_for_gpt_5_4_enum(self):
|
||||
def test_registry_returns_config_for_gpt_5_4_enum(self, local_cost_map):
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
cfg = ProviderConfigManager.get_provider_responses_api_config(
|
||||
|
|
@ -388,39 +397,76 @@ class TestBedrockMantleResponsesRegistry:
|
|||
assert isinstance(cfg, BedrockMantleResponsesAPIConfig)
|
||||
assert cfg.use_openai_path is True
|
||||
|
||||
def test_registry_returns_none_for_gpt_oss(self):
|
||||
# Regression guard: gpt-oss must NOT get the native Responses config; it
|
||||
# keeps the chat-completions emulation path (responses/main.py ~line 1109).
|
||||
def test_registry_returns_native_config_for_gpt_oss(self, local_cost_map):
|
||||
# Core regression: gpt-oss-120b supports the native Responses API (AWS
|
||||
# model card), so it must get a BedrockMantleResponsesAPIConfig on the
|
||||
# STANDARD /v1/responses path -- NOT fall through to None / chat-completions
|
||||
# emulation. Driven by /v1/responses in its price-map supported_endpoints.
|
||||
# Fails on the old gate, which had no responses entry for gpt-oss.
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
cfg = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="bedrock_mantle",
|
||||
model="openai.gpt-oss-120b",
|
||||
)
|
||||
assert cfg is None
|
||||
assert isinstance(cfg, BedrockMantleResponsesAPIConfig)
|
||||
assert cfg.use_openai_path is False
|
||||
|
||||
def test_registry_returns_none_for_gpt_oss_safeguard(self):
|
||||
def test_registry_returns_native_config_for_gpt_oss_20b(self, local_cost_map):
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
cfg = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="bedrock_mantle",
|
||||
model="openai.gpt-oss-safeguard-20b",
|
||||
model="openai.gpt-oss-20b",
|
||||
)
|
||||
assert cfg is None
|
||||
assert isinstance(cfg, BedrockMantleResponsesAPIConfig)
|
||||
assert cfg.use_openai_path is False
|
||||
|
||||
def test_registry_returns_config_for_future_frontier_model(self):
|
||||
# Forward-compatibility: an unseen OpenAI gpt frontier model (e.g. gpt-6),
|
||||
# not yet in the price map, must get the openai-path Responses config with
|
||||
# no code or JSON change. The name-convention fallback (openai.gpt- minus
|
||||
# gpt-oss) catches it before any price-map entry exists.
|
||||
def test_registry_returns_none_for_gpt_oss_safeguard(self, local_cost_map):
|
||||
# Key discriminator: gpt-oss-safeguard shares the "gpt-oss" substring with
|
||||
# gpt-oss-120b but does NOT support Responses (AWS card), so it must return
|
||||
# None. Proves the gate is per-model (supported_endpoints) and not a naive
|
||||
# gpt-oss substring match. local_cost_map loads the chat-only entry.
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
for model in ("openai.gpt-oss-safeguard-120b", "openai.gpt-oss-safeguard-20b"):
|
||||
cfg = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="bedrock_mantle",
|
||||
model=model,
|
||||
)
|
||||
assert cfg is None, model
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
["google.gemma-4-31b", "google.gemma-4-26b-a4b", "google.gemma-4-e2b"],
|
||||
)
|
||||
def test_registry_returns_native_config_for_gemma_4(self, local_cost_map, model):
|
||||
# All three gemma-4 models support Responses (AWS cards) on the /openai/v1
|
||||
# base, so each must get the native config with the openai path.
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
cfg = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="bedrock_mantle",
|
||||
model=model,
|
||||
)
|
||||
assert isinstance(cfg, BedrockMantleResponsesAPIConfig)
|
||||
assert cfg.use_openai_path is True
|
||||
|
||||
def test_unmapped_frontier_model_falls_through_to_none(self, restore_model_cost):
|
||||
# The gate is data-driven, not name-based: an unseen model not yet in the
|
||||
# price map (e.g. a future gpt-6) has no capability signal, so it falls
|
||||
# through to None (chat-completions emulation) rather than being routed
|
||||
# natively by a model-name guess. Onboarding it is a JSON / register_model
|
||||
# change, never a code change (see the register_model tests below).
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
litellm.model_cost.pop("bedrock_mantle/openai.gpt-6", None)
|
||||
litellm.get_model_info.cache_clear()
|
||||
cfg = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="bedrock_mantle",
|
||||
model="openai.gpt-6",
|
||||
)
|
||||
assert isinstance(cfg, BedrockMantleResponsesAPIConfig)
|
||||
assert cfg.use_openai_path is True
|
||||
assert cfg is None
|
||||
|
||||
def test_price_map_flag_routes_non_gpt_name_to_openai_path(
|
||||
self, restore_model_cost
|
||||
|
|
@ -542,8 +588,8 @@ class TestBedrockMantleResponsesRegistry:
|
|||
assert cfg.use_openai_path is False
|
||||
|
||||
def test_unmapped_model_degrades_to_none_without_crashing(self, restore_model_cost):
|
||||
# A non-frontier model that is not in model_cost makes get_model_info
|
||||
# raise; the gate must swallow it and return None rather than crash.
|
||||
# A model absent from model_cost has no capability signal, so the gate
|
||||
# returns None (chat-completions emulation) rather than crashing.
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
litellm.model_cost.pop("bedrock_mantle/somelab.unmapped-model", None)
|
||||
|
|
@ -560,6 +606,9 @@ class TestBedrockMantleResponsesRegistry:
|
|||
# place, so the snapshot must be a deepcopy: a shallow dict() copy would
|
||||
# share that nested dict and leave mode=responses after restore, making
|
||||
# the final assertion fail. The in-place clear+update mirrors the fixture.
|
||||
# gpt-oss-safeguard is the right vehicle here: it is chat-only, so without
|
||||
# the registered mode=responses it resolves to None, isolating the effect
|
||||
# of the register/restore from the model's own (lack of) capability.
|
||||
from litellm.utils import ProviderConfigManager, register_model
|
||||
|
||||
snapshot = copy.deepcopy(litellm.model_cost)
|
||||
|
|
@ -567,14 +616,14 @@ class TestBedrockMantleResponsesRegistry:
|
|||
try:
|
||||
register_model(
|
||||
{
|
||||
"bedrock_mantle/openai.gpt-oss-120b": {
|
||||
"bedrock_mantle/openai.gpt-oss-safeguard-120b": {
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"mode": "responses",
|
||||
}
|
||||
}
|
||||
)
|
||||
during = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="bedrock_mantle", model="openai.gpt-oss-120b"
|
||||
provider="bedrock_mantle", model="openai.gpt-oss-safeguard-120b"
|
||||
)
|
||||
assert isinstance(during, BedrockMantleResponsesAPIConfig)
|
||||
finally:
|
||||
|
|
@ -582,11 +631,151 @@ class TestBedrockMantleResponsesRegistry:
|
|||
litellm.model_cost.update(snapshot)
|
||||
litellm.get_model_info.cache_clear()
|
||||
after = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="bedrock_mantle", model="openai.gpt-oss-120b"
|
||||
provider="bedrock_mantle", model="openai.gpt-oss-safeguard-120b"
|
||||
)
|
||||
assert after is None
|
||||
|
||||
|
||||
class TestMantleBaseSegment:
|
||||
"""The wire-path helper is data-driven from the price-map
|
||||
use_openai_responses_path flag (NOT a model-name match): flagged models are on
|
||||
the /openai/v1 base, everything else on /v1. An unmapped model defaults to /v1.
|
||||
"""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,model_cost,expected",
|
||||
[
|
||||
(
|
||||
"openai.gpt-5.5",
|
||||
{"bedrock_mantle/openai.gpt-5.5": {"use_openai_responses_path": True}},
|
||||
"openai/v1",
|
||||
),
|
||||
(
|
||||
"google.gemma-4-31b",
|
||||
{
|
||||
"bedrock_mantle/google.gemma-4-31b": {
|
||||
"use_openai_responses_path": True
|
||||
}
|
||||
},
|
||||
"openai/v1",
|
||||
),
|
||||
(
|
||||
"openai.gpt-oss-120b",
|
||||
{"bedrock_mantle/openai.gpt-oss-120b": {}},
|
||||
"v1",
|
||||
),
|
||||
("openai.gpt-oss-120b", {}, "v1"),
|
||||
(None, {}, "v1"),
|
||||
],
|
||||
)
|
||||
def test_base_segment(self, model, model_cost, expected):
|
||||
from litellm.llms.bedrock_mantle.common_utils import mantle_base_segment
|
||||
|
||||
assert mantle_base_segment(model, model_cost) == expected
|
||||
|
||||
|
||||
class TestMantleSupportsResponses:
|
||||
"""The capability helper is data-driven (supported_endpoints / mode), with no
|
||||
model-name match: per-model, so gpt-oss-120b is supported but the safeguard
|
||||
variant is not despite the shared substring."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,model_cost,expected",
|
||||
[
|
||||
# supported_endpoints lists responses -> supported
|
||||
(
|
||||
"openai.gpt-oss-120b",
|
||||
{
|
||||
"bedrock_mantle/openai.gpt-oss-120b": {
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"]
|
||||
}
|
||||
},
|
||||
True,
|
||||
),
|
||||
# chat-only supported_endpoints -> not supported (the discriminator)
|
||||
(
|
||||
"openai.gpt-oss-safeguard-120b",
|
||||
{
|
||||
"bedrock_mantle/openai.gpt-oss-safeguard-120b": {
|
||||
"supported_endpoints": ["/v1/chat/completions"]
|
||||
}
|
||||
},
|
||||
False,
|
||||
),
|
||||
# mode=responses (no supported_endpoints) -> supported
|
||||
(
|
||||
"somelab.future-model",
|
||||
{"bedrock_mantle/somelab.future-model": {"mode": "responses"}},
|
||||
True,
|
||||
),
|
||||
# mode=chat, no responses endpoint -> not supported
|
||||
(
|
||||
"google.gemma-3-27b-it",
|
||||
{"bedrock_mantle/google.gemma-3-27b-it": {"mode": "chat"}},
|
||||
False,
|
||||
),
|
||||
# absent from model_cost -> no signal -> not supported
|
||||
("somelab.unmapped", {}, False),
|
||||
(None, {}, False),
|
||||
],
|
||||
)
|
||||
def test_supports_responses(self, model, model_cost, expected):
|
||||
from litellm.llms.bedrock_mantle.common_utils import mantle_supports_responses
|
||||
|
||||
assert mantle_supports_responses(model, model_cost) is expected
|
||||
|
||||
|
||||
class TestBedrockMantlePerModelResponsesURL:
|
||||
"""End-to-end: the registry-selected config must build the correct wire URL
|
||||
per model. gpt-oss on /v1/responses, gpt-5.x and gemma-4 on
|
||||
/openai/v1/responses."""
|
||||
|
||||
def _url_for(self, model, region="us-east-2"):
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
cfg = ProviderConfigManager.get_provider_responses_api_config(
|
||||
provider="bedrock_mantle",
|
||||
model=model,
|
||||
)
|
||||
assert isinstance(cfg, BedrockMantleResponsesAPIConfig)
|
||||
return cfg.get_complete_url(
|
||||
api_base=None, litellm_params={"aws_region_name": region}
|
||||
)
|
||||
|
||||
def test_gpt_oss_uses_standard_responses_path(self, local_cost_map):
|
||||
url = self._url_for("openai.gpt-oss-120b")
|
||||
assert url == "https://bedrock-mantle.us-east-2.api.aws/v1/responses"
|
||||
assert "/openai/v1/responses" not in url
|
||||
|
||||
def test_gpt_5_5_uses_openai_responses_path(self, local_cost_map):
|
||||
url = self._url_for("openai.gpt-5.5")
|
||||
assert url == "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
["google.gemma-4-31b", "google.gemma-4-26b-a4b", "google.gemma-4-e2b"],
|
||||
)
|
||||
def test_gemma_4_uses_openai_responses_path(self, local_cost_map, model):
|
||||
url = self._url_for(model)
|
||||
assert url == "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses"
|
||||
|
||||
|
||||
class TestBedrockMantleEndpointHonoring:
|
||||
def test_plain_chat_call_to_gpt_oss_is_not_bridged(self, local_cost_map):
|
||||
# Adding native Responses support to gpt-oss must NOT reroute its plain
|
||||
# chat-completions traffic. responses_api_bridge_check keys off mode, and
|
||||
# gpt-oss stays mode=chat, so a completion() call is not flipped to the
|
||||
# Responses API. Guards the dual-capability contract.
|
||||
from litellm.main import responses_api_bridge_check
|
||||
|
||||
model_info, resolved_model = responses_api_bridge_check(
|
||||
model="openai.gpt-oss-120b",
|
||||
custom_llm_provider="bedrock_mantle",
|
||||
)
|
||||
assert model_info.get("mode") != "responses"
|
||||
assert resolved_model == "openai.gpt-oss-120b"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def restore_model_cost():
|
||||
"""Snapshot litellm.model_cost so register_model edits don't leak across tests.
|
||||
|
|
|
|||
|
|
@ -131,7 +131,9 @@ class TestBedrockMantleConfig:
|
|||
),
|
||||
)
|
||||
|
||||
def test_get_llm_provider_uses_aws_region_name_for_responses(self, monkeypatch):
|
||||
def test_get_llm_provider_uses_aws_region_name_for_responses(
|
||||
self, monkeypatch, local_cost_map
|
||||
):
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_REGION", raising=False)
|
||||
|
|
@ -143,7 +145,9 @@ class TestBedrockMantleConfig:
|
|||
litellm_params=GenericLiteLLMParams(aws_region_name="us-east-2"),
|
||||
)
|
||||
assert provider == "bedrock_mantle"
|
||||
assert api_base == "https://bedrock-mantle.us-east-2.api.aws/v1"
|
||||
# gpt-5.x carries use_openai_responses_path, so it is served on the
|
||||
# /openai/v1 base per the AWS model card.
|
||||
assert api_base == "https://bedrock-mantle.us-east-2.api.aws/openai/v1"
|
||||
|
||||
def test_default_api_base_fallback_to_us_east_1(self, monkeypatch):
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_REGION", raising=False)
|
||||
|
|
@ -159,6 +163,50 @@ class TestBedrockMantleConfig:
|
|||
api_base, _ = cfg._get_openai_compatible_provider_info(custom_base, None)
|
||||
assert api_base == custom_base
|
||||
|
||||
def test_chat_base_for_gpt_oss_uses_v1(self, monkeypatch):
|
||||
# gpt-oss carries no use_openai_responses_path flag, so it stays on the
|
||||
# standard /v1 base; no regression for existing chat usage now that the
|
||||
# segment is data-driven.
|
||||
monkeypatch.setenv("BEDROCK_MANTLE_REGION", "us-east-2")
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
|
||||
cfg = BedrockMantleChatConfig()
|
||||
api_base, _ = cfg._get_openai_compatible_provider_info(
|
||||
None, None, model="openai.gpt-oss-120b"
|
||||
)
|
||||
assert api_base == "https://bedrock-mantle.us-east-2.api.aws/v1"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_id",
|
||||
["google.gemma-4-31b", "google.gemma-4-26b-a4b", "google.gemma-4-e2b"],
|
||||
)
|
||||
def test_chat_base_for_gemma_4_uses_openai_v1(
|
||||
self, monkeypatch, local_cost_map, model_id
|
||||
):
|
||||
# The chat-config bug the Gemma 4 cards exposed: gemma-4-* is served on the
|
||||
# /openai/v1 base, not the hardcoded /v1. Driven by the price-map
|
||||
# use_openai_responses_path flag (loaded by local_cost_map). Fails before
|
||||
# the data-driven segment lands.
|
||||
monkeypatch.setenv("BEDROCK_MANTLE_REGION", "us-east-2")
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
|
||||
cfg = BedrockMantleChatConfig()
|
||||
api_base, _ = cfg._get_openai_compatible_provider_info(
|
||||
None, None, model=model_id
|
||||
)
|
||||
assert api_base == "https://bedrock-mantle.us-east-2.api.aws/openai/v1"
|
||||
|
||||
def test_chat_base_explicit_api_base_wins_over_derived(
|
||||
self, monkeypatch, local_cost_map
|
||||
):
|
||||
# An explicit api_base must not be overridden by the data-driven default,
|
||||
# even for a model whose default differs (gemma-4 -> openai/v1).
|
||||
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
|
||||
custom_base = "https://bedrock-mantle.us-west-2.api.aws/v1"
|
||||
cfg = BedrockMantleChatConfig()
|
||||
api_base, _ = cfg._get_openai_compatible_provider_info(
|
||||
custom_base, None, model="google.gemma-4-31b"
|
||||
)
|
||||
assert api_base == custom_base
|
||||
|
||||
def test_api_key_from_env(self, monkeypatch):
|
||||
monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "test-key-123")
|
||||
cfg = BedrockMantleChatConfig()
|
||||
|
|
|
|||
|
|
@ -175,6 +175,75 @@ class TestJSONProviderLoader:
|
|||
assert config.custom_llm_provider == "publicai"
|
||||
|
||||
|
||||
class TestPinstripes:
|
||||
"""Tests for Pinstripes JSON-configured provider"""
|
||||
|
||||
def test_pinstripes_json_config_exists(self):
|
||||
"""Test that pinstripes is configured in providers.json"""
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
assert JSONProviderRegistry.exists("pinstripes")
|
||||
|
||||
pinstripes = JSONProviderRegistry.get("pinstripes")
|
||||
assert pinstripes is not None
|
||||
assert pinstripes.base_url == "https://pinstripes.io/v1"
|
||||
assert pinstripes.api_key_env == "PINSTRIPES_API_KEY"
|
||||
assert pinstripes.param_mappings.get("max_completion_tokens") == "max_tokens"
|
||||
|
||||
def test_pinstripes_provider_resolution(self):
|
||||
"""Test that provider resolution finds pinstripes and returns the default base URL"""
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="pinstripes/ps/glm-4.5-air",
|
||||
custom_llm_provider=None,
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert model == "ps/glm-4.5-air"
|
||||
assert provider == "pinstripes"
|
||||
assert api_base == "https://pinstripes.io/v1"
|
||||
|
||||
def test_pinstripes_dynamic_config(self):
|
||||
"""Test dynamic config class creation for pinstripes"""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("pinstripes")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
|
||||
assert api_base == "https://pinstripes.io/v1"
|
||||
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(
|
||||
"https://custom.pinstripes.io/v1", "test-key"
|
||||
)
|
||||
assert api_base == "https://custom.pinstripes.io/v1"
|
||||
assert api_key == "test-key"
|
||||
|
||||
def test_pinstripes_parameter_mapping(self):
|
||||
"""Test that max_completion_tokens is mapped to max_tokens for pinstripes"""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("pinstripes")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
optional_params = {}
|
||||
non_default_params = {"max_completion_tokens": 100, "temperature": 0.7}
|
||||
result = config.map_openai_params(
|
||||
non_default_params, optional_params, "ps/glm-4.5-air", False
|
||||
)
|
||||
|
||||
assert "max_tokens" in result
|
||||
assert result["max_tokens"] == 100
|
||||
assert "max_completion_tokens" not in result
|
||||
assert result["temperature"] == 0.7
|
||||
|
||||
|
||||
class TestPublicAIIntegration:
|
||||
"""Integration tests for PublicAI provider"""
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,97 @@
|
|||
"""
|
||||
Tests for Pinstripes provider configuration and integration.
|
||||
"""
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
class TestPinstripeProviderConfig:
|
||||
"""Test Pinstripes provider configuration"""
|
||||
|
||||
def test_pinstripes_in_provider_list(self):
|
||||
"""Test that pinstripes is in the provider list"""
|
||||
from litellm import LlmProviders
|
||||
|
||||
assert hasattr(LlmProviders, "PINSTRIPES")
|
||||
assert LlmProviders.PINSTRIPES.value == "pinstripes"
|
||||
assert "pinstripes" in litellm.provider_list
|
||||
|
||||
def test_pinstripes_json_config_exists(self):
|
||||
"""Test that pinstripes is configured in providers.json"""
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
assert JSONProviderRegistry.exists("pinstripes")
|
||||
|
||||
pinstripes = JSONProviderRegistry.get("pinstripes")
|
||||
assert pinstripes is not None
|
||||
assert pinstripes.base_url == "https://pinstripes.io/v1"
|
||||
assert pinstripes.api_key_env == "PINSTRIPES_API_KEY"
|
||||
assert pinstripes.param_mappings.get("max_completion_tokens") == "max_tokens"
|
||||
|
||||
def test_pinstripes_in_openai_compatible_providers(self):
|
||||
"""Test that pinstripes is in the openai_compatible_providers list"""
|
||||
from litellm.constants import openai_compatible_providers
|
||||
|
||||
assert "pinstripes" in openai_compatible_providers
|
||||
|
||||
def test_pinstripes_provider_resolution(self):
|
||||
"""Test that provider resolution finds pinstripes and returns the default base URL"""
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="pinstripes/ps/glm-4.5-air",
|
||||
custom_llm_provider=None,
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert model == "ps/glm-4.5-air"
|
||||
assert provider == "pinstripes"
|
||||
assert api_base == "https://pinstripes.io/v1"
|
||||
|
||||
def test_pinstripes_api_base_override(self):
|
||||
"""Test that an explicit api_base / api_key overrides the default"""
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="pinstripes/ps/glm-4.5-air",
|
||||
custom_llm_provider=None,
|
||||
api_base="https://custom.pinstripes.io/v1",
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
assert provider == "pinstripes"
|
||||
assert api_base == "https://custom.pinstripes.io/v1"
|
||||
assert api_key == "sk-test"
|
||||
|
||||
def test_pinstripes_url_autodetection(self):
|
||||
"""Test that api_base=pinstripes.io/v1 auto-sets custom_llm_provider=pinstripes"""
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="ps/glm-4.5-air",
|
||||
custom_llm_provider=None,
|
||||
api_base="https://pinstripes.io/v1",
|
||||
api_key=None,
|
||||
)
|
||||
assert provider == "pinstripes"
|
||||
assert api_base == "https://pinstripes.io/v1"
|
||||
|
||||
def test_pinstripes_router_config(self):
|
||||
"""Test that pinstripes can be used in Router configuration"""
|
||||
from litellm import Router
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "pinstripes-chat",
|
||||
"litellm_params": {
|
||||
"model": "pinstripes/ps/glm-4.5-air",
|
||||
"api_key": "test-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
assert len(router.model_list) == 1
|
||||
assert router.model_list[0]["model_name"] == "pinstripes-chat"
|
||||
|
|
@ -31,6 +31,16 @@ class TestProviderRegistration:
|
|||
assert api_key == "test-key"
|
||||
assert api_base == "https://api.soniox.com"
|
||||
|
||||
def test_should_resolve_soniox_v5_via_get_llm_provider(self, monkeypatch):
|
||||
monkeypatch.setenv("SONIOX_API_KEY", "test-key")
|
||||
model, provider, api_key, api_base = litellm.get_llm_provider(
|
||||
model="soniox/stt-async-v5"
|
||||
)
|
||||
assert provider == "soniox"
|
||||
assert model == "stt-async-v5"
|
||||
assert api_key == "test-key"
|
||||
assert api_base == "https://api.soniox.com"
|
||||
|
||||
def test_should_return_soniox_config_from_provider_config_manager(self):
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue