Merge branch 'litellm_internal_staging' into fix-thinking-block-duplication-stream-chunk-builder

This commit is contained in:
Vineeth Sai Varikuntla 2026-08-25 12:47:41 -07:00 • committed by GitHub
commit 9704d37311
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
164 changed files with 10753 additions and 667 deletions

3
.github/CODEOWNERS vendored
View file

@ -1,6 +1,9 @@
/ui/ @yuneng-berri @ryan-crabbe-berri
/litellm/proxy/_experimental/out/ @yuneng-berri @ryan-crabbe-berri
/ui/Dockerfile
/ui/nginx.conf
/ui/litellm-dashboard/src/lib/http/schema.d.ts
/ui/litellm-dashboard/tsconfig.tsbuildinfo
/model_prices_and_context_window.json @mateo-berri
/litellm/model_prices_and_context_window_backup.json @mateo-berri
/litellm-proxy-extras/litellm_proxy_extras/migrations/ @yuneng-berri @ryan-crabbe-berri

198
.github/scripts/e2e_egress_sentinel.py vendored Executable file
View file

@ -0,0 +1,198 @@
"""Prove an e2e replay run makes zero outbound provider calls, by counting them.
`serve` pins each provider host (`--host`) to a local sink address in the hosts
file and binds a counting listener on that address, so any connection the proxy
or the record/replay edge opens to a real provider is redirected to the sink,
recorded as one line in `--hits-file`, and never leaves the box. The record and
replay edge only ever dials `127.0.0.1:<edge-port>` (a different host than the
pinned provider names), so in a clean replay the sink sees nothing; a single hit
means a provider call escaped the bundle. `assert-empty` turns that hit file into
the pass/fail check.
Stdlib only, so CI runs it under the system interpreter as root (binding :443 and
editing the hosts file both need root); `--sink-address`, `--port`, and
`--hosts-file` are injectable so it runs unprivileged against a temp hosts file on
a high port under test.
"""
# ruff: noqa: T201 # CLI script: its stdout/stderr progress and results are the interface
from __future__ import annotations
import argparse
import json
import os
import signal
import socket
import sys
import threading
import time
from dataclasses import dataclass
from pathlib import Path
from types import FrameType
from typing import Final
_BLOCK_BEGIN: Final = "# BEGIN e2e-egress-sentinel"
_BLOCK_END: Final = "# END e2e-egress-sentinel"
@dataclass(frozen=True, slots=True)
class ServeConfig:
hosts: tuple[str, ...]
sink_address: str
ports: tuple[int, ...]
hits_file: Path
hosts_file: Path
ready_file: Path | None
pid_file: Path | None
def _pin_block(sink_address: str, hosts: tuple[str, ...]) -> str:
lines = "\n".join(f"{sink_address}\t{host}" for host in hosts)
return f"\n{_BLOCK_BEGIN}\n{lines}\n{_BLOCK_END}\n"
def _install_pins(hosts_file: Path, sink_address: str, hosts: tuple[str, ...]) -> bytes:
original = hosts_file.read_bytes() if hosts_file.exists() else b""
hosts_file.write_bytes(original + _pin_block(sink_address, hosts).encode())
return original
def _restore_pins(hosts_file: Path, original: bytes) -> None:
hosts_file.write_bytes(original)
def _bind(sink_address: str, port: int) -> socket.socket:
listener = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
listener.bind((sink_address, port))
listener.listen(128)
return listener
@dataclass(frozen=True, slots=True)
class _HitLog:
path: Path
_lock: threading.Lock
def record(self, *, port: int, peer: tuple[str, int]) -> None:
entry = json.dumps({"ts": time.time(), "port": port, "peer": list(peer)})
with self._lock:
with self.path.open("a", encoding="utf-8") as handle:
handle.write(entry + "\n")
def _serve_socket(listener: socket.socket, port: int, hits: _HitLog, stop: threading.Event) -> None:
while not stop.is_set():
try:
conn, peer = listener.accept()
except OSError:
return
hits.record(port=port, peer=(peer[0], peer[1]))
try:
conn.close()
except OSError:
pass
def serve(config: ServeConfig) -> int:
config.hits_file.write_text("", encoding="utf-8")
original_hosts = _install_pins(config.hosts_file, config.sink_address, config.hosts)
try:
listeners = tuple(_bind(config.sink_address, port) for port in config.ports)
except OSError as exc:
_restore_pins(config.hosts_file, original_hosts)
print(f"egress sentinel could not bind a sink: {exc}", file=sys.stderr)
return 1
stop = threading.Event()
hits = _HitLog(path=config.hits_file, _lock=threading.Lock())
threads = tuple(
threading.Thread(target=_serve_socket, args=(listener, port, hits, stop), daemon=True)
for listener, port in zip(listeners, config.ports)
)
for thread in threads:
thread.start()
def _handle(_signum: int, _frame: FrameType | None) -> None:
stop.set()
for listener in listeners:
try:
listener.close()
except OSError:
pass
signal.signal(signal.SIGTERM, _handle)
signal.signal(signal.SIGINT, _handle)
if config.pid_file is not None:
config.pid_file.write_text(str(os.getpid()), encoding="utf-8")
if config.ready_file is not None:
config.ready_file.write_text("ready", encoding="utf-8")
print(
f"egress sentinel up: pinned {', '.join(config.hosts)} to {config.sink_address} "
f"on port(s) {', '.join(str(p) for p in config.ports)}",
flush=True,
)
stop.wait()
_restore_pins(config.hosts_file, original_hosts)
if config.ready_file is not None and config.ready_file.exists():
config.ready_file.unlink()
if config.pid_file is not None and config.pid_file.exists():
config.pid_file.unlink()
return 0
def assert_empty(hits_file: Path) -> int:
if not hits_file.exists():
print(f"egress sentinel recorded no provider calls ({hits_file} absent): zero egress")
return 0
hits = [line for line in hits_file.read_text(encoding="utf-8").splitlines() if line.strip()]
if not hits:
print("egress sentinel recorded no provider calls: zero egress")
return 0
print(f"egress sentinel recorded {len(hits)} provider call(s); replay was not hermetic:", file=sys.stderr)
for line in hits:
print(f" {line}", file=sys.stderr)
return 1
def _serve_from_args(args: argparse.Namespace) -> int:
config = ServeConfig(
hosts=tuple(args.host),
sink_address=args.sink_address,
ports=tuple(args.port),
hits_file=Path(args.hits_file),
hosts_file=Path(args.hosts_file),
ready_file=Path(args.ready_file) if args.ready_file else None,
pid_file=Path(args.pid_file) if args.pid_file else None,
)
return serve(config)
def main(argv: tuple[str, ...]) -> int:
parser = argparse.ArgumentParser(description="count outbound provider calls during an e2e replay")
sub = parser.add_subparsers(dest="command", required=True)
serve_parser = sub.add_parser("serve", help="pin provider hosts and count connection attempts")
serve_parser.add_argument("--host", action="append", required=True, help="provider host to pin and watch")
serve_parser.add_argument("--sink-address", default="127.0.0.1")
serve_parser.add_argument("--port", action="append", type=int, default=None)
serve_parser.add_argument("--hits-file", required=True)
serve_parser.add_argument("--hosts-file", default="/etc/hosts")
serve_parser.add_argument("--ready-file", default=None)
serve_parser.add_argument("--pid-file", default=None)
assert_parser = sub.add_parser("assert-empty", help="exit non-zero if any provider call was recorded")
assert_parser.add_argument("--hits-file", required=True)
args = parser.parse_args(argv)
if args.command == "serve":
if args.port is None:
args.port = [443]
return _serve_from_args(args)
return assert_empty(Path(args.hits_file))
if __name__ == "__main__":
raise SystemExit(main(tuple(sys.argv[1:])))

55
.github/scripts/e2e_fetch_fixture_bundle.sh vendored Executable file
View file

@ -0,0 +1,55 @@
#!/usr/bin/env bash
set -euo pipefail
REPO="${1:-${GITHUB_REPOSITORY:?REPO required}}"
ARTIFACT_NAME="${2:-e2e-fixtures-bundle}"
BASE_BRANCH="${3:?base branch required}"
DEST_DIR="${4:?destination bundle dir required}"
: "${GH_TOKEN:?GH_TOKEN required to query and download artifacts}"
WORKDIR="$(mktemp -d)"
trap 'rm -rf "${WORKDIR}"' EXIT
echo "resolving newest non-expired '${ARTIFACT_NAME}' artifact on ${REPO}@${BASE_BRANCH}"
SELECTED="$(
gh api "repos/${REPO}/actions/artifacts" -X GET -f per_page=100 --paginate \
--jq ".artifacts[] | select(.name == \"${ARTIFACT_NAME}\" and .expired == false and .workflow_run.head_branch == \"${BASE_BRANCH}\") | {id, digest, created_at, run_id: .workflow_run.id, run_number: .workflow_run.run_number}" \
| jq -s 'sort_by(.created_at) | reverse | .[0] // empty'
)"
if [[ -z "${SELECTED}" ]]; then
echo "no usable '${ARTIFACT_NAME}' artifact on ${BASE_BRANCH}: the last record run produced none (a red Saturday), so there is nothing fresh to replay; failing loudly instead of replaying a stale bundle" >&2
exit 1
fi
RUN_ID="$(echo "${SELECTED}" | jq -r '.run_id')"
RUN_NUMBER="$(echo "${SELECTED}" | jq -r '.run_number')"
ARTIFACT_ID="$(echo "${SELECTED}" | jq -r '.id')"
GH_DIGEST="$(echo "${SELECTED}" | jq -r '.digest // "unknown"')"
CREATED_AT="$(echo "${SELECTED}" | jq -r '.created_at')"
echo "pinned bundle: run #${RUN_NUMBER} (run_id=${RUN_ID}, artifact_id=${ARTIFACT_ID}), recorded ${CREATED_AT}, github digest ${GH_DIGEST}"
gh run download "${RUN_ID}" --repo "${REPO}" -n "${ARTIFACT_NAME}" -D "${WORKDIR}"
TARBALL="$(find "${WORKDIR}" -name '*.tar.gz' -type f | head -n 1)"
if [[ -z "${TARBALL}" ]]; then
echo "downloaded artifact contained no tarball" >&2
exit 1
fi
SIDECAR="${TARBALL}.sha256"
if [[ ! -f "${SIDECAR}" ]]; then
echo "downloaded artifact has no ${SIDECAR}: cannot verify the bundle digest" >&2
exit 1
fi
echo "verifying bundle against its recorded sha256 digest"
( cd "$(dirname "${TARBALL}")" && sha256sum -c "$(basename "${SIDECAR}")" )
mkdir -p "${DEST_DIR}"
tar xzf "${TARBALL}" -C "${DEST_DIR}"
echo "extracted bundle into ${DEST_DIR}"
python3 -c "import json,sys; m=json.load(open(sys.argv[1])); print(' recorded_at', m['recorded_at'], 'harness', m['harness_version'], 'format_version', m['format_version'])" "${DEST_DIR}/manifest.json"

36
.github/scripts/e2e_pack_fixture_bundle.sh vendored Executable file
View file

@ -0,0 +1,36 @@
#!/usr/bin/env bash
set -euo pipefail
if [[ $# -ne 2 ]]; then
echo "usage: $0 <bundle-dir> <out-tarball>" >&2
exit 2
fi
BUNDLE_DIR="$1"
OUT_TARBALL="$2"
MANIFEST="${BUNDLE_DIR}/manifest.json"
if [[ ! -f "${MANIFEST}" ]]; then
echo "no ${MANIFEST}: refusing to publish a bundle with no manifest (record produced nothing)" >&2
exit 1
fi
echo "packing fixture bundle from ${BUNDLE_DIR}"
python3 -c "import json,sys; m=json.load(open(sys.argv[1])); print(' format_version', m['format_version'], 'recorded_at', m['recorded_at'], 'harness', m['harness_version'])" "${MANIFEST}"
TEST_DIRS=$(find "${BUNDLE_DIR}" -mindepth 1 -maxdepth 1 -type d | wc -l | tr -d ' ')
if [[ "${TEST_DIRS}" -eq 0 ]]; then
echo "bundle at ${BUNDLE_DIR} has a manifest but no recorded interactions; refusing to publish an empty bundle" >&2
exit 1
fi
echo " ${TEST_DIRS} recorded test director(ies)"
mkdir -p "$(dirname "${OUT_TARBALL}")"
tar czf "${OUT_TARBALL}" -C "${BUNDLE_DIR}" .
OUT_DIR="$(cd "$(dirname "${OUT_TARBALL}")" && pwd)"
OUT_BASE="$(basename "${OUT_TARBALL}")"
( cd "${OUT_DIR}" && sha256sum "${OUT_BASE}" > "${OUT_BASE}.sha256" )
echo "wrote ${OUT_TARBALL} ($(du -h "${OUT_TARBALL}" | cut -f1)) and ${OUT_BASE}.sha256"
cat "${OUT_DIR}/${OUT_BASE}.sha256"

237
.github/workflows/e2e_record_replay.yml vendored Normal file
View file

@ -0,0 +1,237 @@
name: "E2E Record and Replay"
on:
schedule:
- cron: "0 8 * * 6"
- cron: "0 8 * * 1-5"
workflow_dispatch:
inputs:
mode:
description: "record (hits real providers and publishes a fresh bundle) or replay (bundle only, zero provider egress)"
type: choice
options:
- record
- replay
default: record
permissions:
contents: read
jobs:
record:
name: "Record the e2e suite against real providers"
if: >-
(github.event_name != 'schedule' || github.repository == 'BerriAI/litellm') &&
(github.event.schedule == '0 8 * * 6' ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.mode == 'record'))
runs-on: ubuntu-latest
timeout-minutes: 45
services:
postgres:
image: postgres:16.6
env:
POSTGRES_USER: llmproxy
POSTGRES_PASSWORD: dbpassword9090
POSTGRES_DB: litellm
ports:
- 5432:5432
options: >-
--health-cmd "pg_isready -U llmproxy"
--health-interval 5s
--health-timeout 5s
--health-retries 10
env:
DATABASE_URL: postgresql://llmproxy:dbpassword9090@localhost:5432/litellm
LITELLM_MASTER_KEY: sk-e2e-record-replay
LITELLM_LOCAL_MODEL_COST_MAP: "True"
OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }}
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache the Rust build
uses: ./.github/actions/cache-cargo-build
- name: Install dependencies
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra proxy
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Start the proxy
run: |
nohup uv run --no-sync litellm --config tests/e2e/gateway/record_replay_ci_config.yml --port 4000 > proxy.log 2>&1 &
for _ in $(seq 1 90); do
if curl -fs http://localhost:4000/health/liveliness > /dev/null; then
exit 0
fi
sleep 2
done
echo "proxy never became live"
tail -n 100 proxy.log
exit 1
- name: Record the replayable e2e lane
env:
E2E_FIXTURE_MODE: record
run: |
uv run --no-sync pytest tests/e2e -m replayable --reruns 0 -v --tb=short -rA
- name: Pack the fixture bundle
run: |
.github/scripts/e2e_pack_fixture_bundle.sh tests/e2e/.fixtures "${RUNNER_TEMP}/bundle/e2e-fixtures.tar.gz"
- name: Publish the fixture bundle
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
with:
name: e2e-fixtures-bundle
path: |
${{ runner.temp }}/bundle/e2e-fixtures.tar.gz
${{ runner.temp }}/bundle/e2e-fixtures.tar.gz.sha256
if-no-files-found: error
retention-days: 30
- name: Show proxy log on failure
if: failure()
run: tail -n 300 proxy.log
replay:
name: "Replay the e2e suite from the pinned bundle with zero egress"
if: >-
(github.event_name != 'schedule' || github.repository == 'BerriAI/litellm') &&
(github.event.schedule == '0 8 * * 1-5' ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.mode == 'replay'))
runs-on: ubuntu-latest
timeout-minutes: 45
permissions:
contents: read
actions: read
services:
postgres:
image: postgres:16.6
env:
POSTGRES_USER: llmproxy
POSTGRES_PASSWORD: dbpassword9090
POSTGRES_DB: litellm
ports:
- 5432:5432
options: >-
--health-cmd "pg_isready -U llmproxy"
--health-interval 5s
--health-timeout 5s
--health-retries 10
env:
DATABASE_URL: postgresql://llmproxy:dbpassword9090@localhost:5432/litellm
LITELLM_MASTER_KEY: sk-e2e-record-replay
LITELLM_LOCAL_MODEL_COST_MAP: "True"
GH_TOKEN: ${{ github.token }}
OPENAI_API_KEY: sk-replay-must-never-reach-a-provider
ANTHROPIC_API_KEY: sk-ant-replay-must-never-reach-a-provider
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache the Rust build
uses: ./.github/actions/cache-cargo-build
- name: Install dependencies
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra proxy
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Fetch the pinned fixture bundle by digest
env:
BASE_BRANCH: ${{ github.ref_name }}
run: |
.github/scripts/e2e_fetch_fixture_bundle.sh \
"${GITHUB_REPOSITORY}" \
e2e-fixtures-bundle \
"${BASE_BRANCH}" \
tests/e2e/.fixtures
- name: Start the proxy
run: |
nohup uv run --no-sync litellm --config tests/e2e/gateway/record_replay_ci_config.yml --port 4000 > proxy.log 2>&1 &
for _ in $(seq 1 90); do
if curl -fs http://localhost:4000/health/liveliness > /dev/null; then
exit 0
fi
sleep 2
done
echo "proxy never became live"
tail -n 100 proxy.log
exit 1
- name: Start the egress sentinel
run: |
# shellcheck disable=SC2024 # the log redirect is deliberately the runner user's, so a later non-sudo cat can read it
sudo python3 .github/scripts/e2e_egress_sentinel.py serve \
--host api.openai.com \
--host api.anthropic.com \
--hits-file "${RUNNER_TEMP}/egress-hits.jsonl" \
--ready-file "${RUNNER_TEMP}/egress-ready" \
--pid-file "${RUNNER_TEMP}/egress.pid" \
> "${RUNNER_TEMP}/egress-sentinel.log" 2>&1 &
for _ in $(seq 1 30); do
if [[ -f "${RUNNER_TEMP}/egress-ready" ]]; then
cat "${RUNNER_TEMP}/egress-sentinel.log"
exit 0
fi
sleep 1
done
echo "egress sentinel never became ready"
cat "${RUNNER_TEMP}/egress-sentinel.log"
exit 1
- name: Replay the replayable e2e lane
env:
E2E_FIXTURE_MODE: replay
run: |
uv run --no-sync pytest tests/e2e -m replayable --reruns 0 -v --tb=short -rA
- name: Stop the egress sentinel and assert zero provider egress
if: always()
run: |
if [[ -f "${RUNNER_TEMP}/egress.pid" ]]; then
sudo kill -TERM "$(cat "${RUNNER_TEMP}/egress.pid")" 2>/dev/null || true
sleep 2
fi
python3 .github/scripts/e2e_egress_sentinel.py assert-empty --hits-file "${RUNNER_TEMP}/egress-hits.jsonl"
- name: Show proxy log on failure
if: failure()
run: tail -n 300 proxy.log

View file

@ -131,6 +131,9 @@ jobs:
- name: check_e2e_no_raw_requests
run: uv run --no-sync python ./tests/code_coverage_tests/check_e2e_no_raw_requests.py
- name: check_migrations_no_data_rewrites
run: uv run --no-sync python ./tests/code_coverage_tests/check_migrations_no_data_rewrites.py
- name: memory_test
run: uv run --no-sync python ./tests/code_coverage_tests/memory_test.py

View file

@ -79,6 +79,8 @@ Do not put names of customers or customer company names in code, PR descriptions
CI supply-chain safety: Never pipe a remote script into a shell (`curl ... | bash`, `wget ... | sh`); download the artifact to a file, verify its SHA-256 checksum, then install. Pin every external tool to a specific version with a full URL (not `latest` or `stable`). Verify checksums for all downloaded binaries, using the provider's official `.sha256` / `.sha256sum` sidecar when available. These rules apply to every download in CI
Prisma migrations apply synchronously at proxy boot, before it serves traffic, so a migration must only change schema, never rewrite rows. No `UPDATE`, `DELETE` or `MERGE`, and no `INSERT ... SELECT`: on a spend-log-sized table any of those is minutes of downtime plus a doubled heap that plain autovacuum won't give back. `tests/code_coverage_tests/check_migrations_no_data_rewrites.py` enforces this. When a rewrite is genuinely bounded and has to ship inside the migration, mark the statement `-- data-migration-ok: <what bounds it>`
Follow these coding conventions for new/updated code (a three-line fix in a legacy file shouldn't trigger huge drive-by refactors):
- Composition over inheritance

View file

@ -1,6 +1,6 @@
{
"reportAny": {
"limit": 19955
"limit": 19949
},
"reportArgumentType": {
"limit": 2566
@ -54,7 +54,7 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5663
"limit": 5661
},
"reportMissingTypeArgument": {
"limit": 15555
@ -105,10 +105,10 @@
"limit": 109
},
"reportUnknownMemberType": {
"limit": 39011
"limit": 39009
},
"reportUnknownParameterType": {
"limit": 19885
"limit": 19883
},
"reportUnknownVariableType": {
"limit": 30569

View file

@ -199,6 +199,7 @@ standard_logging_payload_excluded_fields: Optional[List[str]] = (
None # Fields to exclude from StandardLoggingPayload before callbacks receive it
)
log_raw_request_response: bool = False
log_client_error_tracebacks: bool = False
request_correlation_in_logs: bool = False
redact_messages_in_exceptions: Optional[bool] = False
redact_user_api_key_info: Optional[bool] = False
@ -1801,6 +1802,9 @@ if TYPE_CHECKING:
from .llms.gemini.interactions.transformation import (
GoogleAIStudioInteractionsConfig as GoogleAIStudioInteractionsConfig,
)
from .llms.vertex_ai.interactions.transformation import (
VertexAIInteractionsConfig as VertexAIInteractionsConfig,
)
from .llms.openai.chat.o_series_transformation import (
OpenAIOSeriesConfig as OpenAIOSeriesConfig,
OpenAIOSeriesConfig as OpenAIO1Config,

View file

@ -242,6 +242,7 @@ LLM_CONFIG_NAMES: Final = (
"OpenRouterResponsesAPIConfig",
"BedrockMantleResponsesAPIConfig",
"GoogleAIStudioInteractionsConfig",
"VertexAIInteractionsConfig",
"OpenAIOSeriesConfig",
"AnthropicSkillsConfig",
"BaseSkillsAPIConfig",
@ -977,6 +978,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
".llms.gemini.interactions.transformation",
"GoogleAIStudioInteractionsConfig",
),
"VertexAIInteractionsConfig": (
".llms.vertex_ai.interactions.transformation",
"VertexAIInteractionsConfig",
),
"OpenAIOSeriesConfig": (
".llms.openai.chat.o_series_transformation",
"OpenAIOSeriesConfig",

View file

@ -18,6 +18,14 @@ already does when one of its pooled connections errors), leaving every other nod
connections untouched. Every other branch (MOVED, ASK, CLUSTERDOWN, slot-not-covered,
retry-exhaustion) is unchanged from upstream, since those already carry real evidence the
topology changed.
redis-py 8.x fixed this upstream with gentler machinery than this override's
``node.disconnect()`` (which also kills connections other coroutines are mid-operation
on, so one timeout cascades into a reconnect storm and, with TLS, a fresh handshake per
killed connection): it marks in-use connections for reconnect only after their current
operation completes, disconnects only the idle pooled ones, and defers reinitialization
to the outer retry loop. When the installed ``ClusterNode`` has that per-connection
recovery API, the factory returns the base ``RedisCluster`` unmodified.
"""
import asyncio
@ -72,8 +80,16 @@ class _ClusterAttrs(Protocol):
_VERIFIED_REDIS_VERSIONS: Final = frozenset({"5.3.1"})
def get_litellm_async_redis_cluster_class() -> type["_AsyncRedisClusterType"]:
"""Builds the ``RedisCluster`` subclass with the per-node isolation fix.
def get_litellm_async_redis_cluster_class(
cluster_node_class: type | None = None,
) -> type["_AsyncRedisClusterType"]:
"""Returns the base ``RedisCluster`` when the installed redis-py already recovers a
node-level connection error per-connection (8.x+), else builds the ``RedisCluster``
subclass with the per-node isolation fix for older versions whose upstream branch
tears down the whole cluster client.
``cluster_node_class`` exists for dependency injection in tests; production callers
leave it unset and the installed ``ClusterNode`` is used.
Imported lazily because this module is reachable from a base ``import litellm`` while
redis is not a base dependency. Cheap to call repeatedly: the underlying redis
@ -81,7 +97,10 @@ def get_litellm_async_redis_cluster_class() -> type["_AsyncRedisClusterType"]:
"""
import redis
from redis.asyncio.cluster import (
RedisCluster as _BaseAsyncRedisCluster, # pyright: ignore[reportUnknownVariableType] # redis-py ships no resolvable stub for this class under the repo's current (stale) types-redis pin
ClusterNode as _AsyncClusterNode, # pyright: ignore[reportUnknownVariableType] # redis-py ships no resolvable stub for this class under the repo's current (stale) types-redis pin
)
from redis.asyncio.cluster import (
RedisCluster as _BaseAsyncRedisCluster, # pyright: ignore[reportUnknownVariableType] # same stale-stub gap as the import above
)
from redis.cluster import get_node_name
from redis.commands import READ_COMMANDS
@ -98,6 +117,15 @@ def get_litellm_async_redis_cluster_class() -> type["_AsyncRedisClusterType"]:
from redis.exceptions import ConnectionError as _RedisConnectionError
from redis.exceptions import TimeoutError as _RedisTimeoutError
node_class: Final = cluster_node_class if cluster_node_class is not None else _AsyncClusterNode
if hasattr(node_class, "update_active_connections_for_reconnect"):
verbose_logger.debug(
"redis-py %s recovers a node-level connection error per-connection upstream; "
"using the base RedisCluster without litellm's node-isolation override.",
redis.__version__,
)
return _BaseAsyncRedisCluster
if redis.__version__ not in _VERIFIED_REDIS_VERSIONS:
verbose_logger.warning(
"redis-py %s is not in the set this cluster-teardown-storm fix was verified "

View file

@ -5,7 +5,7 @@ Handler for transforming /chat/completions api requests to litellm.responses req
import json
import os
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, Union, cast
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, Union, cast, get_args
from openai.types.responses.custom_tool_param import CustomToolParam
from openai.types.responses.response_input_param import (
@ -35,6 +35,7 @@ from litellm.responses.sse_output_recovery import (
)
from litellm.responses.utils import normalize_responses_api_stream_options
from litellm.types.llms.openai import (
REASONING_EFFORT,
ChatCompletionAnnotation,
ChatCompletionReasoningItem,
ChatCompletionToolCallChunk,
@ -1113,22 +1114,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
litellm.reasoning_auto_summary or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true"
)
# If string is passed, map with optional summary based on flag/env var
if reasoning_effort == "none":
return Reasoning(effort="none", summary="detailed") if auto_summary_enabled else Reasoning(effort="none")
elif reasoning_effort == "high":
return Reasoning(effort="high", summary="detailed") if auto_summary_enabled else Reasoning(effort="high")
elif reasoning_effort == "xhigh":
return Reasoning(effort="xhigh", summary="detailed") if auto_summary_enabled else Reasoning(effort="xhigh")
elif reasoning_effort == "medium":
if reasoning_effort in get_args(REASONING_EFFORT):
return (
Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium")
)
elif reasoning_effort == "low":
return Reasoning(effort="low", summary="detailed") if auto_summary_enabled else Reasoning(effort="low")
elif reasoning_effort == "minimal":
return (
Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal")
Reasoning(effort=reasoning_effort, summary="detailed")
if auto_summary_enabled
else Reasoning(effort=reasoning_effort)
)
return None

View file

@ -48,6 +48,7 @@ LITELLM_MAX_STREAMING_DURATION_SECONDS: Final = (
# Data URIs exceeding this are replaced with a size placeholder.
# Set to 0 to disable truncation.
MAX_BASE64_LENGTH_FOR_LOGGING: Final = int(os.getenv("MAX_BASE64_LENGTH_FOR_LOGGING", 64))
REDACTED_BY_LITELLM: Final = "redacted-by-litellm"
MAX_STRING_LENGTH_STDOUT_LOG: Final = get_env_int("MAX_STRING_LENGTH_STDOUT_LOG", 4096)
@ -749,6 +750,7 @@ openai_compatible_endpoints: Final[list] = [
"api.groq.com/openai/v1",
"https://integrate.api.nvidia.com/v1",
"api.deepseek.com/v1",
"api.together.ai/v1",
"api.together.xyz/v1",
"app.empower.dev/api/v1",
"https://api.friendli.ai/serverless/v1",

View file

@ -168,17 +168,20 @@ class LangsmithLogger(CustomBatchLogger):
return outputs
def _ensure_required_ids(self, data: dict, run_id: str | None):
resolved_id: Final = run_id or str(uuid.uuid4())
if "id" not in data or data["id"] is None:
run_id = str(uuid.uuid4())
data["id"] = run_id
data["id"] = resolved_id
if "trace_id" not in data or data["trace_id"] is None:
if run_id is not None and isinstance(run_id, str):
data["trace_id"] = run_id
# LangSmith rejects the whole ingest batch unless a root run's trace_id
# equals the run id embedded in the first segment of dotted_order
posts_as_root: Final = ("parent_run_id" not in data or data["parent_run_id"] is None) and (
"dotted_order" not in data or data["dotted_order"] is None
)
if posts_as_root or "trace_id" not in data or data["trace_id"] is None:
data["trace_id"] = resolved_id
if "dotted_order" not in data or data["dotted_order"] is None:
if run_id is not None and isinstance(run_id, str):
data["dotted_order"] = self.make_dot_order(run_id=run_id)
data["dotted_order"] = self.make_dot_order(run_id=resolved_id)
def _prepare_log_data(
self,
@ -193,6 +196,11 @@ class LangsmithLogger(CustomBatchLogger):
metadata = _litellm_params.get("metadata", {}) or {}
fields: Final = self._extract_metadata_fields(metadata, credentials)
# the proxy header fan-out mirrors one value into both keys, and LangSmith
# rejects the whole ingest batch when run-body session_id is not an
# existing tracer-session uuid
if fields["session_id"] == fields["trace_id"]:
fields["session_id"] = None
verbose_logger.debug(
"Langsmith Logging - project_name: %s, run_name %s", fields["project_name"], fields["run_name"]
)

View file

@ -11,6 +11,7 @@ import time
from collections.abc import Mapping
from datetime import datetime
from typing import Final, cast
from urllib.parse import quote
import litellm
from litellm._logging import print_verbose, verbose_logger
@ -206,6 +207,23 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
params.get("s3_sse_kms_key_id") or s3_sse_kms_key_id,
)
def _build_object_url(self, s3_object_key: str) -> str:
"""
Build the exact URL that is both signed and sent, with the key percent-encoded once.
S3SigV4Auth signs the path verbatim while S3 canonicalizes the received path with reserved
characters encoded, so an unencoded `=`, `+`, `&`, `#`, `?`, `%` or space in the key makes
the two signatures disagree (403 SignatureDoesNotMatch).
"""
encoded_key: Final = quote(s3_object_key, safe="/")
if self.s3_endpoint_url and self.s3_bucket_name:
if self.s3_use_virtual_hosted_style:
endpoint_host: Final = self.s3_endpoint_url.replace("https://", "").replace("http://", "")
protocol: Final = "https://" if self.s3_endpoint_url.startswith("https://") else "http://"
return f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{encoded_key}"
return f"{self.s3_endpoint_url}/{self.s3_bucket_name}/{encoded_key}"
return f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{encoded_key}"
def _sse_headers(self) -> Mapping[str, str]:
candidates: Final = {
"x-amz-server-side-encryption": self.s3_server_side_encryption,
@ -292,7 +310,6 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
import base64
import hashlib
import requests
from botocore.auth import S3SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
@ -316,18 +333,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
verbose_logger.debug("s3_v2 logger - uploading data to s3 - %s", batch_logging_element.s3_object_key)
verbose_logger.debug("s3_v2 logger - s3_verify setting: %s", self.s3_verify)
# Prepare the URL
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{batch_logging_element.s3_object_key}"
if self.s3_endpoint_url and self.s3_bucket_name:
if self.s3_use_virtual_hosted_style:
# Virtual-hosted-style: bucket.endpoint/key
endpoint_host: Final = self.s3_endpoint_url.replace("https://", "").replace("http://", "")
protocol: Final = "https://" if self.s3_endpoint_url.startswith("https://") else "http://"
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{batch_logging_element.s3_object_key}"
else:
# Path-style: endpoint/bucket/key
url = self.s3_endpoint_url + "/" + self.s3_bucket_name + "/" + batch_logging_element.s3_object_key
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
# Convert JSON to string
json_string: Final = safe_dumps(batch_logging_element.payload)
@ -348,29 +354,19 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
**self._sse_headers(),
}
req: Final = requests.Request("PUT", url, data=json_string, headers=headers)
prepped: Final = req.prepare()
# Sign the request
aws_request: Final = AWSRequest(
method=prepped.method,
url=prepped.url,
data=prepped.body,
headers=prepped.headers,
)
aws_request: Final = AWSRequest(method="PUT", url=url, data=json_string, headers=headers)
aws_region_name: Final = self.get_aws_region_name_for_non_llm_api_calls(aws_region_name=self.s3_region_name)
S3SigV4Auth(credentials, "s3", aws_region_name).add_auth(aws_request)
# Prepare the signed headers
signed_headers: Final = dict(aws_request.headers.items())
# Use prepared URL so path segments match SigV4 canonical request (e.g. %20 for spaces).
request_url: Final = prepped.url or url
# Make the request with retry for transient S3 errors (500/503)
max_retries: Final = 3
for attempt in range(max_retries):
response = await self.async_httpx_client.put(request_url, data=json_string, headers=signed_headers)
response = await self.async_httpx_client.put(url, data=json_string, headers=signed_headers)
if response.status_code in (500, 503) and attempt < max_retries - 1:
wait_time = 2**attempt # 1s, 2s
verbose_logger.warning(
@ -478,7 +474,6 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
import base64
import hashlib
import requests
from botocore.auth import S3SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.credentials import Credentials
@ -493,18 +488,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
aws_region_name=self.s3_region_name,
)
# Prepare the URL
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{batch_logging_element.s3_object_key}"
if self.s3_endpoint_url and self.s3_bucket_name:
if self.s3_use_virtual_hosted_style:
# Virtual-hosted-style: bucket.endpoint/key
endpoint_host: Final = self.s3_endpoint_url.replace("https://", "").replace("http://", "")
protocol: Final = "https://" if self.s3_endpoint_url.startswith("https://") else "http://"
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{batch_logging_element.s3_object_key}"
else:
# Path-style: endpoint/bucket/key
url = self.s3_endpoint_url + "/" + self.s3_bucket_name + "/" + batch_logging_element.s3_object_key
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
# Convert JSON to string
json_string: Final = safe_dumps(batch_logging_element.payload)
@ -525,32 +509,22 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
**self._sse_headers(),
}
req: Final = requests.Request("PUT", url, data=json_string, headers=headers)
prepped: Final = req.prepare()
# Sign the request
aws_request: Final = AWSRequest(
method=prepped.method,
url=prepped.url,
data=prepped.body,
headers=prepped.headers,
)
aws_request: Final = AWSRequest(method="PUT", url=url, data=json_string, headers=headers)
aws_region_name: Final = self.get_aws_region_name_for_non_llm_api_calls(aws_region_name=self.s3_region_name)
S3SigV4Auth(credentials, "s3", aws_region_name).add_auth(aws_request)
# Prepare the signed headers
signed_headers: Final = dict(aws_request.headers.items())
# Use prepared URL so path segments match SigV4 canonical request (e.g. %20 for spaces).
request_url: Final = prepped.url or url
httpx_client: Final = _get_httpx_client(
params=({"ssl_verify": self.s3_verify} if self.s3_verify is not None else None)
)
# Make the request with retry for transient S3 errors (500/503)
max_retries: Final = 3
for attempt in range(max_retries):
response = httpx_client.put(request_url, data=json_string, headers=signed_headers)
response = httpx_client.put(url, data=json_string, headers=signed_headers)
if response.status_code in (500, 503) and attempt < max_retries - 1:
wait_time = 2**attempt # 1s, 2s
verbose_logger.warning(
@ -582,7 +556,6 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
try:
import hashlib
import requests
from botocore.auth import S3SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
@ -607,18 +580,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
verbose_logger.debug("s3_v2 logger - downloading data from s3 - %s", s3_object_key)
# Prepare the URL
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{s3_object_key}"
if self.s3_endpoint_url and self.s3_bucket_name:
if self.s3_use_virtual_hosted_style:
# Virtual-hosted-style: bucket.endpoint/key
endpoint_host: Final = self.s3_endpoint_url.replace("https://", "").replace("http://", "")
protocol: Final = "https://" if self.s3_endpoint_url.startswith("https://") else "http://"
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{s3_object_key}"
else:
# Path-style: endpoint/bucket/key
url = self.s3_endpoint_url + "/" + self.s3_bucket_name + "/" + s3_object_key
url: Final = self._build_object_url(s3_object_key)
# Prepare the request for GET operation
# For GET requests, we need x-amz-content-sha256 with hash of empty string
@ -626,22 +588,15 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
headers: Final = {
"x-amz-content-sha256": empty_string_hash,
}
req: Final = requests.Request("GET", url, headers=headers)
prepped: Final = req.prepare()
# Sign the request
aws_request: Final = AWSRequest(
method=prepped.method,
url=prepped.url,
headers=prepped.headers,
)
aws_request: Final = AWSRequest(method="GET", url=url, headers=headers)
S3SigV4Auth(credentials, "s3", self.s3_region_name).add_auth(aws_request)
# Prepare the signed headers
signed_headers: Final = dict(aws_request.headers.items())
request_url: Final = prepped.url or url
response: Final = await self.async_httpx_client.get(request_url, headers=signed_headers)
response: Final = await self.async_httpx_client.get(url, headers=signed_headers)
if response.status_code != 200:
verbose_logger.exception("S3 object not found, saw response=", response.text)

View file

@ -47,6 +47,13 @@ def get_provider_interactions_api_config(
return GoogleAIStudioInteractionsConfig()
if provider in (LlmProviders.VERTEX_AI.value, LlmProviders.VERTEX_AI_BETA.value):
from litellm.llms.vertex_ai.interactions.transformation import (
VertexAIInteractionsConfig,
)
return VertexAIInteractionsConfig()
return None

View file

@ -58,6 +58,26 @@ def safe_divide(
return numerator / denominator
def is_expected_client_error(exception: BaseException | None) -> bool:
"""
True when the exception maps to an HTTP 4xx status.
ProxyException stores the status on .code (as a str), HTTPException and
litellm exceptions on .status_code.
"""
if exception is None:
return False
code: Final[object] = getattr(exception, "code", None)
status_code: Final[object] = code if code is not None else getattr(exception, "status_code", None)
if status_code is None or isinstance(status_code, bool):
return False
try:
status: Final = int(str(status_code))
except ValueError:
return False
return 400 <= status < 500
def coerce_token_limit(value: object) -> int | None:
"""
Coerce a max_input_tokens / max_output_tokens value to an int, treating a

View file

@ -272,6 +272,14 @@ def get_llm_provider(
elif endpoint == "api.deepseek.com/v1":
custom_llm_provider = "deepseek"
dynamic_api_key = get_secret_str("DEEPSEEK_API_KEY")
elif endpoint == "api.together.ai/v1" or endpoint == "api.together.xyz/v1":
custom_llm_provider = "together_ai"
dynamic_api_key = api_key or (
get_secret_str("TOGETHER_API_KEY")
or get_secret_str("TOGETHER_AI_API_KEY")
or get_secret_str("TOGETHERAI_API_KEY")
or get_secret_str("TOGETHER_AI_TOKEN")
)
elif endpoint == "ollama.com":
custom_llm_provider = "ollama"
dynamic_api_key = get_secret_str("OLLAMA_API_KEY")
@ -707,7 +715,7 @@ def _get_openai_compatible_provider_info(
dynamic_api_key,
) = litellm.ZAIChatConfig()._get_openai_compatible_provider_info(api_base, api_key)
elif custom_llm_provider == "together_ai":
api_base = api_base or get_secret_str("TOGETHER_AI_API_BASE") or "https://api.together.xyz/v1"
api_base = api_base or get_secret_str("TOGETHER_AI_API_BASE") or "https://api.together.ai/v1"
dynamic_api_key = api_key or (
get_secret_str("TOGETHER_API_KEY")
or get_secret_str("TOGETHER_AI_API_KEY")

View file

@ -62,7 +62,7 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.deepeval.deepeval import DeepEvalLogger
from litellm.integrations.mlflow import MlflowLogger
from litellm.integrations.sqs import SQSLogger
from litellm.litellm_core_utils.core_helpers import reconstruct_model_name
from litellm.litellm_core_utils.core_helpers import is_expected_client_error, reconstruct_model_name
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import (
cost_breakdown_with_guardrail,
@ -3124,6 +3124,13 @@ class Logging(LiteLLMLoggingBaseClass):
if not hasattr(self, "model_call_details"):
self.model_call_details = {}
if (
self.model_call_details.get("log_event_type") == "failed_api_call"
and self.model_call_details.get("exception") is exception
and self.model_call_details.get("standard_logging_object") is not None
):
return start_time, self.model_call_details["end_time"]
self.model_call_details["log_event_type"] = "failed_api_call"
self.model_call_details["exception"] = exception
self.model_call_details["traceback_exception"] = (
@ -5455,9 +5462,10 @@ class StandardLoggingPayloadSetup:
error_class: Final[str] = str(original_exception.__class__.__name__) if original_exception else ""
_llm_provider_in_exception: Final = getattr(original_exception, "llm_provider", "")
# Get traceback information (first 100 lines)
traceback_info = traceback_str or ""
if original_exception:
if original_exception and (
litellm.log_client_error_tracebacks or not is_expected_client_error(original_exception)
):
tb: Final[TracebackType | None] = getattr(original_exception, "__traceback__", None)
if tb:
tb_lines: Final = traceback.format_tb(tb)
@ -5930,11 +5938,15 @@ def get_standard_logging_object_payload(
response_model_name = final_response_obj.get("model")
# For Azure Model Router, preserve the actual model in the top-level standard
# logging payload only when the user has opted in.
# logging payload.
from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo
requested_model: Final = kwargs.get("model")
if (
isinstance(requested_model, str)
and ("model_router" in requested_model.lower() or "model-router" in requested_model.lower())
stamped_selected_model: Final = AzureFoundryModelInfo.get_model_router_selected_model(hidden_params)
if stamped_selected_model is not None:
model_name = stamped_selected_model
elif (
AzureFoundryModelInfo.is_model_router_call(model=requested_model, hidden_params=hidden_params)
and isinstance(response_model_name, str)
and response_model_name
):

View file

@ -470,6 +470,8 @@ def update_messages_with_model_file_ids(
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
convert_b64_uid_to_unified_uid,
get_original_file_id,
is_model_embedded_id,
)
for message in messages:
@ -508,6 +510,11 @@ def update_messages_with_model_file_ids(
unified_file_id = convert_b64_uid_to_unified_uid(file_id)
if "llm_output_file_id," in unified_file_id:
provider_file_id = unified_file_id.split("llm_output_file_id,")[1].split(";")[0]
if not provider_file_id and is_model_embedded_id(file_id):
# `litellm:<raw_id>;model,<m>` encoding from the
# x-litellm-model upload path. Strip the wrapper
# so the provider sees its own ID.
provider_file_id = get_original_file_id(file_id)
file_object_file_field["file_id"] = provider_file_id or file_id
if format:
file_object_file_field["format"] = format
@ -535,6 +542,8 @@ def update_responses_input_with_model_file_ids(
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
convert_b64_uid_to_unified_uid,
get_original_file_id,
is_model_embedded_id,
)
if isinstance(input, str):
@ -578,6 +587,13 @@ def update_responses_input_with_model_file_ids(
updated_content_item = content_item.copy()
updated_content_item["file_id"] = provider_file_id
updated_content.append(updated_content_item)
elif is_model_embedded_id(file_id):
# `litellm:<raw_id>;model,<m>` encoding from the
# x-litellm-model upload path. Strip the wrapper
# so the provider sees its own ID.
updated_content_item = content_item.copy()
updated_content_item["file_id"] = get_original_file_id(file_id)
updated_content.append(updated_content_item)
else:
# Not a managed file, keep as-is
updated_content.append(content_item)

View file

@ -16,6 +16,7 @@ import litellm.types
import litellm.types.llms
from litellm import verbose_logger
from litellm._uuid import uuid
from litellm.constants import REDACTED_BY_LITELLM
from litellm.litellm_core_utils.url_utils import async_safe_get, safe_get
from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_async_httpx_client
from litellm.types.files import get_file_extension_from_mime_type
@ -642,49 +643,6 @@ def claude_2_1_pt(
return prompt
### TOGETHER AI
def get_model_info(token, model):
try:
headers: Final = {"Authorization": f"Bearer {token}"}
client: Final = HTTPHandler(concurrent_limit=1)
response: Final = client.get("https://api.together.xyz/models/info", headers=headers)
if response.status_code == 200:
model_info: Final = response.json()
for m in model_info:
if m["name"].lower().strip() == model.strip():
return m["config"].get("prompt_format", None), m["config"].get("chat_template", None)
return None, None
else:
return None, None
except Exception: # safely fail a prompt template request
return None, None
## OLD TOGETHER AI FLOW
# def format_prompt_togetherai(messages, prompt_format, chat_template):
# if prompt_format is None:
# return default_pt(messages)
# human_prompt, assistant_prompt = prompt_format.split("{prompt}")
# if chat_template is not None:
# prompt = hf_chat_template(
# model=None, messages=messages, chat_template=chat_template
# )
# elif prompt_format is not None:
# prompt = custom_prompt(
# role_dict={},
# messages=messages,
# initial_prompt_value=human_prompt,
# final_prompt_value=assistant_prompt,
# )
# else:
# prompt = default_pt(messages)
# return prompt
### IBM Granite
@ -5383,12 +5341,13 @@ def _parse_tool_call_arguments(raw: Any, tool_name: str | None, context: str) ->
return raw
if not isinstance(raw, str):
return {}
normalized_raw: Final = "{}" if raw == REDACTED_BY_LITELLM else raw
from litellm.litellm_core_utils.prompt_templates.common_utils import (
parse_tool_call_arguments,
)
try:
parsed: Final = parse_tool_call_arguments(raw, tool_name=tool_name, context=context)
parsed: Final = parse_tool_call_arguments(normalized_raw, tool_name=tool_name, context=context)
except ValueError as e:
verbose_logger.warning("Failed to parse tool call arguments: %s", e)
return {}

View file

@ -13,6 +13,7 @@ import inspect
from typing import TYPE_CHECKING, Any, Final
import litellm
from litellm.constants import REDACTED_BY_LITELLM
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
@ -84,29 +85,29 @@ def _redact_tool_calls(tool_calls) -> None:
for tool_call in tool_calls:
function = getattr(tool_call, "function", None)
if function is not None and hasattr(function, "arguments"):
function.arguments = "redacted-by-litellm"
function.arguments = REDACTED_BY_LITELLM
def _redact_function_call(function_call) -> None:
"""Redact legacy assistant function_call arguments."""
if function_call is not None and hasattr(function_call, "arguments"):
function_call.arguments = "redacted-by-litellm"
function_call.arguments = REDACTED_BY_LITELLM
def _redact_choice_content(choice):
"""Helper to redact content in a choice (message or delta)."""
if isinstance(choice, litellm.Choices):
choice.message.content = "redacted-by-litellm"
choice.message.content = REDACTED_BY_LITELLM
if hasattr(choice.message, "reasoning_content"):
choice.message.reasoning_content = "redacted-by-litellm"
choice.message.reasoning_content = REDACTED_BY_LITELLM
if hasattr(choice.message, "thinking_blocks"):
choice.message.thinking_blocks = None
_redact_tool_calls(getattr(choice.message, "tool_calls", None))
_redact_function_call(getattr(choice.message, "function_call", None))
elif isinstance(choice, litellm.utils.StreamingChoices):
choice.delta.content = "redacted-by-litellm"
choice.delta.content = REDACTED_BY_LITELLM
if hasattr(choice.delta, "reasoning_content"):
choice.delta.reasoning_content = "redacted-by-litellm"
choice.delta.reasoning_content = REDACTED_BY_LITELLM
if hasattr(choice.delta, "thinking_blocks"):
choice.delta.thinking_blocks = None
_redact_tool_calls(getattr(choice.delta, "tool_calls", None))
@ -117,22 +118,22 @@ def _redact_responses_api_output(output_items):
"""Helper to redact ResponsesAPIResponse output items."""
for output_item in output_items:
if hasattr(output_item, "text"):
output_item.text = "redacted-by-litellm"
output_item.text = REDACTED_BY_LITELLM
if hasattr(output_item, "content") and isinstance(output_item.content, list):
for content_part in output_item.content:
if hasattr(content_part, "text"):
content_part.text = "redacted-by-litellm"
content_part.text = REDACTED_BY_LITELLM
# Redact reasoning items in output array
if hasattr(output_item, "type") and output_item.type == "reasoning":
if hasattr(output_item, "summary") and isinstance(output_item.summary, list):
for summary_item in output_item.summary:
if hasattr(summary_item, "text"):
summary_item.text = "redacted-by-litellm"
summary_item.text = REDACTED_BY_LITELLM
if hasattr(output_item, "type") and output_item.type == "function_call" and hasattr(output_item, "arguments"):
output_item.arguments = "redacted-by-litellm"
output_item.arguments = REDACTED_BY_LITELLM
def _redact_responses_api_output_dict(output_items, redacted_str: str):
@ -164,7 +165,7 @@ def _redact_standard_logging_object(model_call_details: dict):
if standard_logging_object is None:
return
redacted_str: Final = "redacted-by-litellm"
redacted_str: Final = REDACTED_BY_LITELLM
if standard_logging_object.get("messages") is not None:
standard_logging_object["messages"] = [{"role": "user", "content": redacted_str}]
@ -235,7 +236,7 @@ def perform_redaction(model_call_details: dict, result, redact_streaming_respons
copy via redact_streaming_responses_for_custom_logger instead.
"""
# Redact model_call_details
model_call_details["messages"] = [{"role": "user", "content": "redacted-by-litellm"}]
model_call_details["messages"] = [{"role": "user", "content": REDACTED_BY_LITELLM}]
model_call_details["prompt"] = ""
model_call_details["input"] = ""
_redact_standard_logging_object(model_call_details)
@ -256,7 +257,7 @@ def perform_redaction(model_call_details: dict, result, redact_streaming_respons
or hasattr(result, "__anext__") # async generator
): # async iterator
# For async objects, return a simple redacted response without deepcopy
return {"text": "redacted-by-litellm"}
return {"text": REDACTED_BY_LITELLM}
if not (
isinstance(result, (litellm.ModelResponse, litellm.ResponsesAPIResponse, litellm.EmbeddingResponse))
@ -273,11 +274,11 @@ def perform_redaction(model_call_details: dict, result, redact_streaming_respons
elif isinstance(_result, dict) and "choices" in _result:
# Handle dict representation of ModelResponse (e.g., from model_dump())
if _result.get("choices") is not None:
_redact_model_response_dict_choices(_result["choices"], "redacted-by-litellm")
_redact_model_response_dict_choices(_result["choices"], REDACTED_BY_LITELLM)
redact_vertex_ai_metadata_from_logged_object(_result)
elif isinstance(_result, dict) and "output" in _result:
if isinstance(_result.get("output"), list):
_redact_responses_api_output_dict(_result["output"], "redacted-by-litellm")
_redact_responses_api_output_dict(_result["output"], REDACTED_BY_LITELLM)
elif isinstance(_result, litellm.ResponsesAPIResponse):
if hasattr(_result, "output"):
_redact_responses_api_output(_result.output)
@ -288,7 +289,7 @@ def perform_redaction(model_call_details: dict, result, redact_streaming_respons
if hasattr(_result, "data") and _result.data is not None:
_result.data = []
else:
return {"text": "redacted-by-litellm"}
return {"text": REDACTED_BY_LITELLM}
return _result

View file

@ -1215,8 +1215,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if reasoning_effort is None or reasoning_effort == "none":
return None
if AnthropicConfig._is_adaptive_thinking_model(model, custom_llm_provider):
# without display, Anthropic defaults adaptive thinking to
# display="omitted" and returns a blank thinking block
return AnthropicThinkingParam(
type="adaptive",
display="summarized",
)
elif reasoning_effort == "low":
return AnthropicThinkingParam(
@ -2144,7 +2147,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
@staticmethod
def _thinking_tokens_from_usage(usage_object: Mapping[str, object]) -> int | None:
def thinking_tokens_from_usage(usage_object: Mapping[str, object]) -> int | None:
details: Final = usage_object.get("output_tokens_details")
if not isinstance(details, Mapping):
return None
@ -2176,7 +2179,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
reported_thinking_tokens: Final = (
iteration_thinking_tokens
if iteration_thinking_tokens is not None
else self._thinking_tokens_from_usage(usage_object)
else self.thinking_tokens_from_usage(usage_object)
)
if reported_thinking_tokens is not None:
capped_reported: Final = min(max(0, reported_thinking_tokens), completion_tokens)
@ -2199,7 +2202,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _sum_iteration_thinking_tokens(self, iterations: Sequence[object]) -> int | None:
per_iteration: Final = tuple(
self._thinking_tokens_from_usage(iteration) if isinstance(iteration, Mapping) else None
self.thinking_tokens_from_usage(iteration) if isinstance(iteration, Mapping) else None
for iteration in iterations
)
reported: Final = tuple(tokens for tokens in per_iteration if tokens is not None)

View file

@ -0,0 +1,3 @@
from litellm.llms.azure.search.transformation import BingGroundingSearchConfig
__all__ = ("BingGroundingSearchConfig",)

View file

@ -0,0 +1,442 @@
"""
Calls the Microsoft Foundry Responses API with the `bing_grounding` or `web_search`
tool to search the web (Grounding with Bing Search).
Microsoft docs: https://learn.microsoft.com/en-us/azure/ai-foundry/agents/how-to/tools/bing-grounding
Setup:
1. Set BING_GROUNDING_PROJECT_ENDPOINT to the Foundry project endpoint, e.g.
https://<account>.services.ai.azure.com/api/projects/<project>
2. Set BING_GROUNDING_MODEL to a model deployment in that project (e.g. gpt-4.1);
it runs the grounded search and its tokens are billed on that deployment
3. Optional: set BING_GROUNDING_CONNECTION_ID to a Grounding with Bing Search
project connection id to use the `bing_grounding` tool; without it the
project's built-in `web_search` tool is used
4. Auth: pass api_key (an Azure API key, sent in the api-key header), or set
BING_GROUNDING_TOKEN to an Entra bearer token for scope
https://ai.azure.com/.default, or configure azure-identity (AZURE_CLIENT_ID /
AZURE_CLIENT_SECRET / AZURE_TENANT_ID, managed identity, or any
DefaultAzureCredential source) and the token is minted automatically
Usage:
response = litellm.search(
query="latest AI developments",
search_provider="bing_grounding",
max_results=5,
)
"""
from __future__ import annotations
from collections.abc import Callable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal
import httpx
from pydantic import BaseModel, ConfigDict, ValidationError
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.search.transformation import (
BaseSearchConfig,
SearchResponse,
SearchResult,
)
from litellm.secret_managers.main import get_secret_str
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
_DOCS_URL: Final = "https://learn.microsoft.com/en-us/azure/ai-foundry/agents/how-to/tools/bing-grounding"
PROJECT_ENDPOINT_ENV: Final = "BING_GROUNDING_PROJECT_ENDPOINT"
MODEL_ENV: Final = "BING_GROUNDING_MODEL"
CONNECTION_ID_ENV: Final = "BING_GROUNDING_CONNECTION_ID"
TOKEN_ENV: Final = "BING_GROUNDING_TOKEN"
ENTRA_SCOPE: Final = "https://ai.azure.com/.default"
_RESPONSES_PATH: Final = "/openai/v1/responses"
_SNIPPET_FALLBACK_LENGTH: Final = 300
_UPSTREAM_ERROR_STATUS: Final = 502
_RESPONSE_COST_HEADER: Final = "llm_provider-x-litellm-response-cost"
class _Annotation(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
type: str = ""
url: str | None = None
title: str | None = None
start_index: int | None = None
end_index: int | None = None
class _ContentPart(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
type: str = ""
text: str = ""
annotations: tuple[_Annotation, ...] = ()
class _OutputItem(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
type: str = ""
content: tuple[_ContentPart, ...] = ()
class _ErrorBody(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
message: str | None = None
class _IncompleteDetails(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
reason: str | None = None
class _ResponsesEnvelope(BaseModel):
"""A Foundry Responses API body. `output` is required: a body without it is not a
Responses API response and must not be reported as a successful empty search.
A 200 body can still carry `status` `failed` or `incomplete`; those are surfaced as
errors rather than reported as a successful empty search."""
model_config = ConfigDict(extra="ignore", frozen=True)
output: tuple[_OutputItem, ...]
status: str | None = None
error: _ErrorBody | None = None
incomplete_details: _IncompleteDetails | None = None
class _ErrorEnvelope(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
error: _ErrorBody | None = None
def _unwrap_error_detail(error_message: str) -> str:
"""
Surface the human-readable message inside Foundry's error envelope.
Tool failures nest a second JSON document as a string inside `error.message`
(observed live for `bing_grounding` connection errors), so the unwrap runs twice.
Falls back to the raw body for anything else.
"""
try:
envelope: Final = _ErrorEnvelope.model_validate_json(error_message)
except ValidationError:
return error_message
message: Final = envelope.error.message if envelope.error else None
if message is None:
return error_message
try:
nested: Final = _ErrorBody.model_validate_json(message)
except ValidationError:
return message
return nested.message or message
def _snippet(text: str, annotation: _Annotation) -> str:
"""
The text a citation supports, not the citation marker itself.
A url_citation's start/end indices span the inline marker ("([host](url))"),
which follows the claim it backs, so the snippet is the marker's own line up
to where the marker starts.
"""
start: Final = annotation.start_index
marker_start: Final = start if start is not None and 0 <= start <= len(text) else len(text)
claim: Final = text[:marker_start].rsplit("\n", 1)[-1].strip()
if claim:
return claim[-_SNIPPET_FALLBACK_LENGTH:]
return text[:_SNIPPET_FALLBACK_LENGTH]
def _citation_results(envelope: _ResponsesEnvelope) -> tuple[SearchResult, ...]:
"""One result per cited URL: first occurrence wins, order preserved as answered."""
cited: Final = tuple(
SearchResult(
title=annotation.title or "",
url=annotation.url or "",
snippet=_snippet(part.text, annotation),
date=None,
last_updated=None,
)
for item in envelope.output
if item.type == "message"
for part in item.content
if part.type == "output_text"
for annotation in part.annotations
if annotation.type == "url_citation" and annotation.url
)
first_by_url: Final = MappingProxyType({result.url: result for result in reversed(cited)})
return tuple(first_by_url[url] for url in dict.fromkeys(result.url for result in cited))
def _valid_max_results(max_results: object) -> int | None:
"""A positive-int `max_results`, else None. Rejects bools, an `int` subclass, and
non-positive values so neither the request-side `count` nor the response-side cap
forwards a value the other would silently ignore.
"""
if isinstance(max_results, bool) or not isinstance(max_results, int):
return None
return max_results if max_results > 0 else None
def _requested_max_results(response_kwargs: Mapping[str, object]) -> int | None:
"""The unified `max_results` cap the caller asked for, if any.
The built-in web_search tool has no server-side result-count knob, so the cap is
enforced here after the fact; connection mode also honors it as a hard ceiling on
top of the tool's `count` hint.
"""
optional_params: Final = response_kwargs.get("optional_params")
if not isinstance(optional_params, Mapping):
return None
return _valid_max_results(optional_params.get("max_results"))
def _capped(results: tuple[SearchResult, ...], max_results: int | None) -> tuple[SearchResult, ...]:
return results[:max_results] if max_results is not None else results
class _SearchConfiguration(BaseModel):
model_config = ConfigDict(frozen=True)
project_connection_id: str
count: int | None = None
class _BingGroundingParams(BaseModel):
model_config = ConfigDict(frozen=True)
search_configurations: tuple[_SearchConfiguration, ...]
class _BingGroundingTool(BaseModel):
model_config = ConfigDict(frozen=True)
type: Literal["bing_grounding"] = "bing_grounding"
bing_grounding: _BingGroundingParams
class _UserLocation(BaseModel):
model_config = ConfigDict(frozen=True)
type: Literal["approximate"] = "approximate"
country: str
class _WebSearchTool(BaseModel):
model_config = ConfigDict(frozen=True)
type: Literal["web_search"] = "web_search"
user_location: _UserLocation | None = None
class _ResponsesRequest(BaseModel):
model_config = ConfigDict(frozen=True)
model: str
input: str
tools: tuple[_BingGroundingTool | _WebSearchTool, ...]
def _search_tool(optional_params: Mapping[str, object]) -> _BingGroundingTool | _WebSearchTool:
connection_id: Final = get_secret_str(CONNECTION_ID_ENV)
max_results: Final = optional_params.get("max_results")
country: Final = optional_params.get("country")
if connection_id:
configuration: Final = _SearchConfiguration(
project_connection_id=connection_id,
count=_valid_max_results(max_results),
)
return _BingGroundingTool(bing_grounding=_BingGroundingParams(search_configurations=(configuration,)))
location: Final = _UserLocation(country=country.upper()) if isinstance(country, str) else None
return _WebSearchTool(user_location=location)
def _default_entra_token_minter() -> str:
from litellm.secret_managers.get_azure_ad_token_provider import get_azure_ad_token_provider
return get_azure_ad_token_provider(azure_scope=ENTRA_SCOPE)()
class BingGroundingSearchConfig(BaseSearchConfig):
def __init__(self, entra_token_minter: Callable[[], str] | None = None) -> None:
super().__init__()
self._entra_token_minter = entra_token_minter
@staticmethod
def ui_friendly_name() -> str:
return "Grounding with Bing Search"
def validate_environment(
self,
headers: dict[str, str], # mutable-ok: BaseSearchConfig.validate_environment signature
api_key: str | None = None,
api_base: str | None = None,
**kwargs: object, # kwargs-ok: BaseSearchConfig.validate_environment signature
) -> dict[str, str]: # mutable-ok: the http handler passes this straight to httpx as headers
"""
Validate environment and return headers.
Returns a new dict rather than mutating ``headers``: the http handler calls this
a second time after ``litellm/search/main.py`` already did, so it has to be idempotent.
"""
return { # mutable-ok: httpx requires a plain dict of headers
**headers,
**self._auth_header(api_key, api_base),
"Content-Type": "application/json",
}
def _auth_header(self, api_key: str | None, api_base: str | None) -> Mapping[str, str]:
"""
A caller-supplied ``api_key`` is an Azure API key and rides the ``api-key`` header;
an Entra bearer token (``BING_GROUNDING_TOKEN`` or one minted via azure-identity)
rides ``Authorization: Bearer``. Foundry rejects the wrong scheme for each.
"""
if api_key:
return MappingProxyType({"api-key": api_key})
token: Final = self.resolve_server_api_key(
caller_api_key=None,
caller_api_base=api_base,
key_env_vars=(TOKEN_ENV,),
base_env_var=PROJECT_ENDPOINT_ENV,
default_api_base=None,
) or self._mint_entra_token(api_base)
return MappingProxyType({"Authorization": f"Bearer {token}"})
def _mint_entra_token(self, caller_api_base: str | None) -> str:
self._assert_trusted_api_base_for_server_credential(
caller_api_base, None, PROJECT_ENDPOINT_ENV, "Azure AD token"
)
minter: Final = self._entra_token_minter or _default_entra_token_minter
try:
return minter()
except Exception as e:
raise ValueError(
f"Grounding with Bing Search: no credential available. Pass api_key, set {TOKEN_ENV} "
f"to an Entra bearer token, or configure azure-identity (AZURE_CLIENT_ID / "
f"AZURE_CLIENT_SECRET / AZURE_TENANT_ID or any DefaultAzureCredential source) "
f"for scope {ENTRA_SCOPE}. Underlying error: {e}"
) from e
def get_complete_url(
self,
api_base: str | None,
optional_params: dict[str, object], # mutable-ok: BaseSearchConfig.get_complete_url signature
data: dict[str, object] | list[dict[str, object]] | None = None, # mutable-ok: base signature
**kwargs: object, # kwargs-ok: BaseSearchConfig.get_complete_url signature
) -> str:
resolved_base: Final = api_base or get_secret_str(PROJECT_ENDPOINT_ENV)
if not resolved_base:
raise ValueError(
f"{PROJECT_ENDPOINT_ENV} is not set. Set it to your Microsoft Foundry project "
f"endpoint, e.g. https://<account>.services.ai.azure.com/api/projects/<project>."
)
trimmed: Final = resolved_base.rstrip("/")
if trimmed.endswith(_RESPONSES_PATH):
return trimmed
return f"{trimmed}{_RESPONSES_PATH}"
def transform_search_request(
self,
query: str | list[str], # mutable-ok: BaseSearchConfig.transform_search_request signature
optional_params: dict[str, object], # mutable-ok: base signature
**kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_request signature
) -> dict[str, object]: # mutable-ok: the http handler passes this straight to httpx as the JSON body
"""
Transform Search request to the Foundry Responses API format.
The unified params map as far as the API allows:
- max_results -> the bing_grounding search configuration's `count`; the built-in
web_search tool has no result-count knob, so that mode instead caps the returned
results after the fact (see transform_search_response)
- country -> web_search's approximate `user_location` (bing_grounding's `market`
wants a full locale like en-US, which a bare country code cannot fill)
- search_domain_filter, max_tokens_per_page -> no API equivalent, dropped
"""
model: Final = get_secret_str(MODEL_ENV)
if not model:
raise ValueError(
f"{MODEL_ENV} is not set. Set it to a model deployment in the Foundry project "
f"that runs the grounded search, e.g. gpt-4.1."
)
request: Final = _ResponsesRequest(
model=model,
input=" ".join(query) if isinstance(query, list) else query,
tools=(_search_tool(optional_params),),
)
return request.model_dump(mode="json", exclude_none=True)
def transform_search_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_response signature
) -> SearchResponse:
try:
parsed: Final = _ResponsesEnvelope.model_validate_json(raw_response.content)
except ValidationError as e:
raise self.get_error_class(
error_message=f"response does not match the Foundry Responses API schema: {e}",
status_code=raw_response.status_code,
headers=dict(raw_response.headers), # mutable-ok: BaseSearchConfig.get_error_class signature
)
if parsed.status == "failed":
detail: Final = (
parsed.error.message if parsed.error and parsed.error.message else "the grounded search failed"
)
raise self._upstream_error(detail, raw_response)
results: Final = _capped(_citation_results(parsed), _requested_max_results(kwargs))
if not results and parsed.status == "incomplete":
reason: Final = (
parsed.incomplete_details.reason
if parsed.incomplete_details and parsed.incomplete_details.reason
else "unknown reason"
)
raise self._upstream_error(f"the grounded search was incomplete: {reason}", raw_response)
return self._priced(results)
def _upstream_error(self, detail: str, raw_response: httpx.Response) -> Exception:
return self.get_error_class(
error_message=detail,
status_code=_UPSTREAM_ERROR_STATUS,
headers=dict(raw_response.headers), # mutable-ok: BaseSearchConfig.get_error_class signature
)
def _priced(self, results: tuple[SearchResult, ...]) -> SearchResponse:
"""web_search mode runs no paid Grounding with Bing transaction, so it must not
inherit the connection-mode ``bing_grounding/search`` price; zero its per-query
cost while leaving connection mode to the cost map."""
response: Final = SearchResponse(
results=list(results), # mutable-ok: SearchResponse.results is list[SearchResult]
object="search",
)
if get_secret_str(CONNECTION_ID_ENV):
return response
response._hidden_params[
"additional_headers"
] = { # mutable-ok: response_cost_calculator writes into _hidden_params
_RESPONSE_COST_HEADER: 0.0
}
return response
def get_error_class(
self,
error_message: str,
status_code: int,
headers: dict[str, str], # mutable-ok: BaseSearchConfig.get_error_class signature
) -> Exception:
detail: Final = _unwrap_error_detail(error_message).rstrip(". ")
return BaseLLMException(
status_code=status_code,
message=f"Grounding with Bing Search: {detail}. See {_DOCS_URL} for details.",
headers=headers,
)

View file

@ -65,15 +65,24 @@ class AzureModelRouterConfig(AzureAIStudioConfig):
Extracts the actual model used from the Azure response (e.g., gpt-5-nano-2025-08-07)
and returns it with the azure_ai/ prefix for proper display and cost tracking.
Also stamps that model onto ``_hidden_params`` so downstream consumers (spend logs,
response restamping) can read it instead of guessing the route from the model string.
"""
from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo
from litellm.llms.azure_ai.common_utils import (
AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY,
AzureFoundryModelInfo,
)
from litellm.router_utils.add_retry_fallback_headers import (
get_hidden_params_dict,
)
# Get base model for the parent call (strips routing prefixes for API compatibility)
base_model: Final[str] = AzureFoundryModelInfo.get_base_model(model)
# Call parent transform_response first - this will extract the actual model
# from the raw response (e.g., "gpt-5-nano-2025-08-07")
model_response = super().transform_response(
transformed_response: Final = super().transform_response(
model=base_model,
raw_response=raw_response,
model_response=model_response,
@ -86,7 +95,15 @@ class AzureModelRouterConfig(AzureAIStudioConfig):
api_key=api_key,
json_mode=json_mode,
)
return model_response
selected_model: Final = transformed_response.model
if selected_model:
# Rebuilt rather than mutated in place: ModelResponseBase declares _hidden_params as a
# class-level dict, so an in-place write can bleed into unrelated responses.
transformed_response._hidden_params = { # pyright: ignore[reportPrivateUsage] # ModelResponse exposes no public hidden-params setter # mutable-ok: ModelResponse requires _hidden_params to be a plain dict
**get_hidden_params_dict(transformed_response),
AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY: selected_model,
}
return transformed_response
def calculate_additional_costs(self, model: str, prompt_tokens: int, completion_tokens: int) -> dict | None:
"""

View file

@ -51,6 +51,9 @@ def get_azure_ai_auth_headers(
)
AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY: Final = "azure_model_router_selected_model"
class AzureFoundryModelInfo(BaseLLMModelInfo):
"""Model info for Azure AI / Azure Foundry models."""
@ -82,6 +85,41 @@ class AzureFoundryModelInfo(BaseLLMModelInfo):
return "model_router"
return "default"
@staticmethod
def get_model_router_selected_model(hidden_params: Mapping[str, object] | None) -> str | None:
"""The model Azure Model Router actually served, stamped by ``AzureModelRouterConfig``.
Reading this beats re-deriving the route from a model string: the stamp is set on the
code path that was actually taken, so it holds no matter what the caller named the model.
"""
if not hidden_params:
return None
selected: Final = hidden_params.get(AZURE_MODEL_ROUTER_SELECTED_MODEL_KEY)
if isinstance(selected, str) and selected:
return selected
return None
@staticmethod
def is_model_router_call(
model: str | None = None,
hidden_params: Mapping[str, object] | None = None,
) -> bool:
"""Whether a request went down the Azure Model Router route.
Prefers the response stamp, then the deployment's litellm model path, and only then the
caller-supplied name. The last two go through ``get_azure_ai_route`` so the model-router
name heuristic lives in exactly one place.
"""
if AzureFoundryModelInfo.get_model_router_selected_model(hidden_params) is not None:
return True
deployment_model: Final = (
hidden_params.get("litellm_model_name") or hidden_params.get("model") if hidden_params is not None else None
)
return any(
isinstance(candidate, str) and AzureFoundryModelInfo.get_azure_ai_route(candidate) == "model_router"
for candidate in (deployment_model, model)
)
@staticmethod
def get_api_base(api_base: str | None = None) -> str | None:
return api_base or litellm.api_base or get_secret_str("AZURE_AI_API_BASE")

View file

@ -1,9 +1,10 @@
import types
from abc import ABC, abstractmethod
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any
import httpx
from httpx._types import RequestFiles
from httpx._types import FileContent, RequestFiles
from litellm.types.responses.main import *
from litellm.types.router import GenericLiteLLMParams
@ -340,14 +341,18 @@ class BaseVideoConfig(ABC):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
video_file: FileContent | None = None,
extra_body: dict[str, Any] | None = None,
prefetched_source_data: dict[str, Any] | None = None,
) -> tuple[str, dict]:
) -> tuple[str, Mapping[str, object], RequestFiles | None]:
"""
Transform the video edit request into a URL and JSON data.
Transform the video edit request into a URL plus either JSON data or
multipart form fields and files.
Returns:
Tuple[str, Dict]: (url, data) for the POST request
tuple[str, Mapping[str, object], RequestFiles | None]: (url, data,
files). When files is None the handler sends data as JSON; otherwise
data holds the form fields and files holds the uploaded source video.
"""
raise NotImplementedError("video edit is not supported for this provider")

View file

@ -1617,6 +1617,8 @@ class AmazonConverseConfig(BaseConfig):
}
if additional_request_params:
data["additionalModelRequestFields"] = additional_request_params
if "thinking" in additional_request_params:
data["additionalModelResponseFieldPaths"] = ("/usage/output_tokens_details",)
if system_content_blocks:
data["system"] = system_content_blocks
@ -1801,6 +1803,17 @@ class AmazonConverseConfig(BaseConfig):
thinking_blocks_list.append(_redacted_block)
return thinking_blocks_list
@staticmethod
def thinking_tokens_from_additional_fields(additional_fields: object) -> int | None:
"""Converse omits thinking tokens from its usage block; they only arrive under
``additionalModelResponseFields`` when ``/usage/output_tokens_details`` is requested."""
if not isinstance(additional_fields, Mapping):
return None
usage: Final = additional_fields.get("usage")
if not isinstance(usage, Mapping):
return None
return AnthropicConfig.thinking_tokens_from_usage(usage)
@staticmethod
def is_converse_usage_shape(usage_object: Mapping[str, object]) -> bool:
"""Converse-family models report camelCase token counts, not Anthropic's snake_case."""
@ -1842,6 +1855,7 @@ class AmazonConverseConfig(BaseConfig):
usage: ConverseTokenUsageBlock,
reasoning_content: str | None = None,
thinking_ran: bool = False,
provider_reasoning_tokens: int | None = None,
) -> Usage:
input_tokens = usage["inputTokens"]
output_tokens: Final = usage["outputTokens"]
@ -1862,9 +1876,14 @@ class AmazonConverseConfig(BaseConfig):
cache_creation_tokens=cache_creation_input_tokens,
text_tokens=raw_input_tokens,
)
reasoning_tokens: Final = (
estimated_reasoning_tokens: Final = (
token_counter(text=reasoning_content, count_response_tokens=True) if reasoning_content else 0
)
reasoning_tokens: Final = (
min(max(0, provider_reasoning_tokens), output_tokens)
if provider_reasoning_tokens is not None
else estimated_reasoning_tokens
)
completion_tokens_details: Final = (
CompletionTokensDetailsWrapper(
reasoning_tokens=reasoning_tokens,
@ -2272,6 +2291,9 @@ class AmazonConverseConfig(BaseConfig):
completion_response["usage"],
reasoning_content=chat_completion_message.get("reasoning_content"),
thinking_ran=reasoningContentBlocks is not None,
provider_reasoning_tokens=self.thinking_tokens_from_additional_fields(
completion_response.get("additionalModelResponseFields")
),
)
## HANDLE TOOL CALLS

View file

@ -331,6 +331,7 @@ class AWSEventStreamDecoder:
self.json_mode = json_mode
self._current_tool_name: str | None = None
self._thinking_ran = False
self._provider_reasoning_tokens: int | None = None
def check_empty_tool_call_args(self) -> bool:
"""
@ -559,10 +560,14 @@ class AWSEventStreamDecoder:
tool_use = self._handle_converse_stop_event(content_block_index)
elif "stopReason" in chunk_data:
finish_reason = map_finish_reason(chunk_data.get("stopReason", "stop"))
self._provider_reasoning_tokens = AmazonConverseConfig.thinking_tokens_from_additional_fields(
chunk_data.get("additionalModelResponseFields")
)
elif "usage" in chunk_data:
usage = converse_config.transform_usage(
chunk_data.get("usage", {}),
thinking_ran=self._thinking_ran,
provider_reasoning_tokens=self._provider_reasoning_tokens,
)
if thinking_blocks:
self._thinking_ran = True

View file

@ -1,4 +1,5 @@
import json
from collections.abc import Mapping
from typing import TYPE_CHECKING, Final, Optional, cast
from httpx import Response
@ -93,6 +94,9 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD
endpoint_url,
)
def get_bedrock_bearer_token(self, litellm_params: Mapping[str, object]) -> str | None:
return None
def sign_request(
self,
headers: dict,
@ -109,6 +113,7 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD
request_data=request_data or {},
api_base=api_base,
model=model,
api_key=self.get_bedrock_bearer_token(optional_params),
)
def logging_non_streaming_response(

View file

@ -13,6 +13,7 @@ global state.
"""
import re
from collections.abc import Mapping
from typing import Final
from botocore.exceptions import (
@ -31,30 +32,39 @@ BEDROCK_MANTLE_DEFAULT_REGION: Final = "us-east-1"
MANTLE_HOST_RE: Final = re.compile(r"^https?://bedrock-mantle\.([^/.]+)\.api\.aws", re.IGNORECASE)
def resolve_mantle_bearer_token(api_key: str | None) -> str | None:
return api_key or get_secret_str("BEDROCK_MANTLE_API_KEY") or get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
def resolve_mantle_region(params: Mapping[str, object]) -> str:
region: Final = params.get("aws_region_name")
if isinstance(region, str) and region:
BaseAWSLLM._validate_aws_region_name(region)
return region
api_base: Final = params.get("api_base")
base: Final = (api_base if isinstance(api_base, str) else None) or get_secret_str("BEDROCK_MANTLE_API_BASE")
if base:
match: Final = MANTLE_HOST_RE.match(base.rstrip("/"))
if match:
return match.group(1)
return (
get_secret_str("BEDROCK_MANTLE_REGION")
or get_secret_str("AWS_REGION_NAME")
or get_secret_str("AWS_REGION")
or BEDROCK_MANTLE_DEFAULT_REGION
)
class BedrockMantleAuthMixin:
_aws_signer: BaseAWSLLM
@staticmethod
def _resolve_bearer_token(api_key: str | None) -> str | None:
return api_key or get_secret_str("BEDROCK_MANTLE_API_KEY") or get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
return resolve_mantle_bearer_token(api_key)
@staticmethod
def _resolve_region(params: dict) -> str:
region: Final = params.get("aws_region_name")
if region:
BaseAWSLLM._validate_aws_region_name(region)
return region
base: Final = params.get("api_base") or get_secret_str("BEDROCK_MANTLE_API_BASE")
if base:
match: Final = MANTLE_HOST_RE.match(base.rstrip("/"))
if match:
return match.group(1)
return (
get_secret_str("BEDROCK_MANTLE_REGION")
or get_secret_str("AWS_REGION_NAME")
or get_secret_str("AWS_REGION")
or BEDROCK_MANTLE_DEFAULT_REGION
)
return resolve_mantle_region(params)
def sign_request(
self,

View file

@ -0,0 +1,71 @@
from collections.abc import Mapping
from typing import TYPE_CHECKING, Final, Literal, Optional
from httpx import Response
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.bedrock.passthrough.transformation import BedrockPassthroughConfig
from litellm.llms.bedrock_mantle.common_utils import (
MANTLE_HOST_RE,
resolve_mantle_bearer_token,
resolve_mantle_region,
)
from litellm.types.utils import LlmProviders
if TYPE_CHECKING:
from litellm.types.utils import CostResponseTypes
class BedrockMantlePassthroughConfig(BedrockPassthroughConfig):
"""Native Bedrock runtime passthrough (InvokeModel, Converse) for deployments declared as bedrock_mantle.
The Mantle host only serves the OpenAI-compatible surface, so a Mantle api_base lends its region and the
request itself goes to bedrock-runtime, signed with the deployment's Bearer token or SigV4 credentials.
"""
def _get_aws_region_name(
self,
optional_params: Mapping[str, object],
model: str | None = None,
model_id: str | None = None,
) -> str:
return resolve_mantle_region(optional_params)
def get_runtime_endpoint(
self,
api_base: str | None,
aws_bedrock_runtime_endpoint: str | None,
aws_region_name: str,
endpoint_type: Literal["runtime", "agent", "agentcore"] | None = "runtime",
) -> tuple[str, str]:
is_mantle_host: Final = api_base is not None and MANTLE_HOST_RE.match(api_base.rstrip("/")) is not None
return super().get_runtime_endpoint(
api_base=None if is_mantle_host else api_base,
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
aws_region_name=aws_region_name,
endpoint_type=endpoint_type,
)
def get_bedrock_bearer_token(self, litellm_params: Mapping[str, object]) -> str | None:
api_key: Final = litellm_params.get("api_key")
return resolve_mantle_bearer_token(api_key if isinstance(api_key, str) else None)
def logging_non_streaming_response(
self,
model: str,
custom_llm_provider: str,
httpx_response: Response,
request_data: dict, # mutable-ok: mirrors the inherited BedrockPassthroughConfig signature
logging_obj: Logging,
endpoint: str,
) -> Optional["CostResponseTypes"]:
is_converse: Final = "invoke" not in endpoint and "converse" in endpoint
shape_provider: Final = LlmProviders.BEDROCK.value if is_converse else custom_llm_provider
return super().logging_non_streaming_response(
model=model,
custom_llm_provider=shape_provider,
httpx_response=httpx_response,
request_data=request_data,
logging_obj=logging_obj,
endpoint=endpoint,
)

View file

@ -15,8 +15,12 @@ role / access key / profile / web identity), signed via the shared
BaseAWSLLM._sign_request after the request body is finalized.
"""
import json
from collections.abc import Mapping
from typing import Any, Final
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_logger
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
@ -50,6 +54,33 @@ _BEDROCK_MANTLE_SUPPORTED_SERVICE_TIERS: Final = frozenset({"auto", "default"})
_CODEX_ADDITIONAL_TOOLS_INPUT_ITEM_TYPE: Final = "additional_tools"
_CODEX_AGENT_MESSAGE_INPUT_ITEM_TYPE: Final = "agent_message"
_CODEX_CONTEXT_COMPACTION_INPUT_ITEM_TYPE: Final = "context_compaction"
_CODEX_LOCAL_SHELL_CALL_INPUT_ITEM_TYPE: Final = "local_shell_call"
class _RewrittenOutputTextBlock(TypedDict):
type: ReadOnly[str]
text: ReadOnly[str]
class _RewrittenAssistantMessageItem(TypedDict):
type: ReadOnly[str]
role: ReadOnly[str]
content: ReadOnly[tuple[_RewrittenOutputTextBlock, ...]]
class _RewrittenCompactionItem(TypedDict):
type: ReadOnly[str]
encrypted_content: ReadOnly[str]
class _RewrittenFunctionCallItem(TypedDict):
type: ReadOnly[str]
call_id: ReadOnly[str]
name: ReadOnly[str]
arguments: ReadOnly[str]
class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPIConfig):
def __init__(
@ -155,6 +186,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI
headers: dict,
) -> dict:
remaining_input, hoisted_tools = self._hoist_codex_additional_tools(input)
normalized_input: Final = self._normalize_codex_input_items(remaining_input)
request_params: Final = (
{
**response_api_optional_request_params,
@ -168,7 +200,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI
)
return super().transform_responses_api_request(
model=model,
input=remaining_input,
input=normalized_input,
response_api_optional_request_params=request_params,
litellm_params=litellm_params,
headers=headers,
@ -210,6 +242,91 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI
)
return remaining_input, cls._filter_unsupported_tools(hoisted_tools)
@staticmethod
def _agent_message_text(item: "Mapping[str, Any]") -> str:
content: Final = item.get("content")
if not isinstance(content, list):
return ""
return "".join(
str(block.get("text") or block.get("encrypted_content") or "")
for block in content
if isinstance(block, dict)
)
@classmethod
def _normalize_agent_message_item(cls, item: "Mapping[str, Any]") -> "_RewrittenAssistantMessageItem | None":
text: Final = cls._agent_message_text(item)
if not text:
return None
rewritten: Final[_RewrittenAssistantMessageItem] = {
"type": "message",
"role": "assistant",
"content": ({"type": "output_text", "text": text},),
}
return rewritten
@staticmethod
def _normalize_context_compaction_item(item: "Mapping[str, Any]") -> "_RewrittenCompactionItem | None":
encrypted_content: Final = item.get("encrypted_content")
if not isinstance(encrypted_content, str) or not encrypted_content:
return None
rewritten: Final[_RewrittenCompactionItem] = {"type": "compaction", "encrypted_content": encrypted_content}
return rewritten
@staticmethod
def _normalize_local_shell_call_item(item: "Mapping[str, Any]") -> "_RewrittenFunctionCallItem | None":
call_id: Final = item.get("call_id")
if not isinstance(call_id, str) or not call_id:
return None
action: Final = item.get("action")
rewritten: Final[_RewrittenFunctionCallItem] = {
"type": "function_call",
"call_id": call_id,
"name": "local_shell",
"arguments": json.dumps(action) if isinstance(action, dict) else "{}",
}
return rewritten
@classmethod
def _normalize_codex_input_item(cls, item: object) -> "tuple[object, str | None]":
"""Returns (normalized item or None to drop it, original type when rewritten)."""
if not isinstance(item, dict):
return item, None
item_type: Final = item.get("type")
if item_type == _CODEX_AGENT_MESSAGE_INPUT_ITEM_TYPE:
return cls._normalize_agent_message_item(item), item_type
if item_type == _CODEX_CONTEXT_COMPACTION_INPUT_ITEM_TYPE:
return cls._normalize_context_compaction_item(item), item_type
if item_type == _CODEX_LOCAL_SHELL_CALL_INPUT_ITEM_TYPE:
return cls._normalize_local_shell_call_item(item), item_type
return item, None
@classmethod
def _normalize_codex_input_items(
cls,
input: "str | ResponseInputParam",
) -> "str | ResponseInputParam":
"""Rewrite Codex history item types Mantle rejects with 400 "Invalid
'input': value did not match any expected variant" into supported
equivalents. `agent_message` (Codex multi-agent traffic; its
encrypted_content slot carries the plaintext payload when the model
never issued encrypted args) becomes an assistant message,
`context_compaction` becomes the `compaction` spelling Mantle accepts,
and `local_shell_call` becomes the function_call its recorded
function_call_output already pairs with.
"""
if not isinstance(input, list):
return input
normalized: Final = tuple(cls._normalize_codex_input_item(item) for item in input)
rewritten_types: Final = sorted(frozenset(item_type for _, item_type in normalized if item_type is not None))
if rewritten_types:
verbose_logger.warning(
"Bedrock Mantle Responses API: rewrote Codex input item type(s) %s that Mantle rejects.",
rewritten_types,
)
kept: Final = [item for item, _ in normalized if item is not None] # mutable-ok: ResponseInputParam is a list
return kept # pyright: ignore[reportReturnType] # Codex passthrough items sit outside the OpenAI input union
def map_openai_params(
self,
response_api_optional_params: ResponsesAPIOptionalRequestParams,

View file

@ -9,7 +9,7 @@ import threading
import time
from collections.abc import AsyncIterable, Callable, Iterable, Mapping
from http.cookiejar import CookieJar, DefaultCookiePolicy
from typing import TYPE_CHECKING, Any, Final, Optional, TypeAlias, TypedDict
from typing import TYPE_CHECKING, Any, ClassVar, Final, Optional, TypeAlias, TypedDict
import certifi
import httpx
@ -933,11 +933,83 @@ class AsyncHTTPHandler:
response.raise_for_status()
return response
# Strong references to finalizer-scheduled client-close tasks. A bare
# create_task() result may be garbage-collected before it runs, leaving
# the underlying aiohttp session unclosed ("Unclosed client session").
# Mirrors LiteLLMAiohttpTransport._background_close_tasks.
_finalizer_close_tasks: ClassVar[set["asyncio.Task[None]"]] = set() # mutable-ok: strong refs for pending closes
@classmethod
def _on_finalizer_close_done(cls, task: "asyncio.Task[None]") -> None:
cls._finalizer_close_tasks.discard(task)
if task.cancelled():
return
exc: Final = task.exception()
if exc is not None:
verbose_logger.debug("Error closing client at finalization: %s", exc)
def _aiohttp_session_bound_elsewhere(self, loop: asyncio.AbstractEventLoop) -> bool:
"""True when the wrapped aiohttp session is bound to a loop other than
``loop`` — awaiting ``aclose()`` here would touch that loop's internals."""
from litellm.llms.custom_httpx.aiohttp_transport import (
LiteLLMAiohttpTransport,
)
transport: Final = getattr(self._client, "_transport", None)
if not isinstance(transport, LiteLLMAiohttpTransport):
return False
session: Final = transport.client
if not isinstance(session, ClientSession) or session.closed:
return False
return getattr(session, "_loop", None) is not loop
def _dispose_wrapped_aiohttp_session(self) -> None:
"""Dispose the wrapped aiohttp session when ``aclose()`` cannot run here.
Finalization either has no running loop, or a loop the session is not
bound to. Delegating to the transport's lifecycle-aware disposal picks
the safe path per session state (async close on its own loop, threadsafe
handoff to a loop running elsewhere, or the synchronous connector
teardown that flips the flags ``ClientSession.__del__`` checks), so no
"Unclosed client session" / "Unclosed connector" warnings fire at
garbage collection.
"""
from litellm.llms.custom_httpx.aiohttp_transport import (
LiteLLMAiohttpTransport,
)
transport: Final = getattr(self._client, "_transport", None)
if not isinstance(transport, LiteLLMAiohttpTransport):
return
# A shared session (e.g. the proxy's) is never this handler's to close.
if not getattr(transport, "_owns_session", False):
return
session: Final = transport.client
if isinstance(session, ClientSession) and not session.closed:
transport._close_recycled_session(session) # pyright: ignore[reportPrivateUsage] # deliberate reuse of the transport's lifecycle-aware disposal; an async close can never run in this context
def __del__(self) -> None:
try:
if not _handler_may_close_client(sys.getrefcount(self._client), self._owns_client):
return
asyncio.get_running_loop().create_task(self._client.aclose())
try:
loop: Final = asyncio.get_running_loop()
except RuntimeError:
# No running loop at finalization time (worker threads after
# their loop closed, interpreter/worker shutdown, GC in a
# sync context). An async close can never run here.
self._dispose_wrapped_aiohttp_session()
return
if self._aiohttp_session_bound_elsewhere(loop):
# GC ran on a live loop (e.g. the app's) but the session
# belongs to another, possibly dead, loop — awaiting aclose()
# here is the cross-loop path the transport refuses.
self._dispose_wrapped_aiohttp_session()
return
task: Final = loop.create_task(self._client.aclose())
cls: Final = type(self)
cls._finalizer_close_tasks.add(task)
task.add_done_callback(cls._on_finalizer_close_done)
except Exception:
pass

View file

@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, Type
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
import httpx
from httpx._types import FileContent
from openai.types.file_deleted import FileDeleted
import litellm
@ -1846,6 +1847,7 @@ class BaseLLMHTTPHandler:
return provider_config.transform_search_response(
raw_response=response,
logging_obj=logging_obj,
optional_params=optional_params,
)
async def async_search(
@ -1944,6 +1946,7 @@ class BaseLLMHTTPHandler:
return provider_config.transform_search_response(
raw_response=response,
logging_obj=logging_obj,
optional_params=optional_params,
)
async def _async_post_anthropic_messages_with_http_error_retry(
@ -7838,6 +7841,7 @@ class BaseLLMHTTPHandler:
custom_llm_provider: str,
litellm_params,
logging_obj,
video_file: FileContent | None = None,
extra_headers: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
timeout: float | None = None,
@ -7849,6 +7853,7 @@ class BaseLLMHTTPHandler:
return self.async_video_edit_handler(
prompt=prompt,
video_id=video_id,
video_file=video_file,
video_provider_config=video_provider_config,
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
@ -7902,9 +7907,10 @@ class BaseLLMHTTPHandler:
prefetched_source_data = prefetch_resp.json()
try:
url, data = video_provider_config.transform_video_edit_request(
url, data, files = video_provider_config.transform_video_edit_request(
prompt=prompt,
video_id=video_id,
video_file=video_file,
api_base=api_base,
litellm_params=litellm_params,
headers=headers,
@ -7923,11 +7929,10 @@ class BaseLLMHTTPHandler:
},
)
response: Final = sync_httpx_client.post(
url=url,
headers=headers,
json=data,
timeout=timeout,
response: Final = (
sync_httpx_client.post(url=url, headers=headers, data=data, files=files, timeout=timeout)
if files
else sync_httpx_client.post(url=url, headers=headers, json=data, timeout=timeout)
)
response.raise_for_status()
return video_provider_config.transform_video_edit_response(
@ -7947,6 +7952,7 @@ class BaseLLMHTTPHandler:
custom_llm_provider: str,
litellm_params,
logging_obj,
video_file: FileContent | None = None,
extra_headers: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
timeout: float | None = None,
@ -7998,9 +8004,10 @@ class BaseLLMHTTPHandler:
prefetched_source_data = prefetch_resp.json()
try:
url, data = video_provider_config.transform_video_edit_request(
url, data, files = video_provider_config.transform_video_edit_request(
prompt=prompt,
video_id=video_id,
video_file=video_file,
api_base=api_base,
litellm_params=litellm_params,
headers=headers,
@ -8019,11 +8026,10 @@ class BaseLLMHTTPHandler:
},
)
response: Final = await async_httpx_client.post(
url=url,
headers=headers,
json=data,
timeout=timeout,
response: Final = await (
async_httpx_client.post(url=url, headers=headers, data=data, files=files, timeout=timeout)
if files
else async_httpx_client.post(url=url, headers=headers, json=data, timeout=timeout)
)
response.raise_for_status()
return video_provider_config.transform_video_edit_response(

View file

@ -566,6 +566,7 @@ class GeminiVideoConfig(BaseVideoConfig):
api_base,
litellm_params,
headers,
video_file=None,
extra_body=None,
prefetched_source_data=None,
):

View file

@ -16,7 +16,7 @@ def _normalize_reasoning_effort_for_chat_completion(
) -> str | None:
"""Convert reasoning_effort to the string format expected by OpenAI chat completion API.
The chat completion API expects a simple string: 'none', 'low', 'medium', 'high', or 'xhigh'.
The chat completion API expects an effort string such as 'low' or 'high'.
Config/deployments may pass the Responses API format: {'effort': 'high', 'summary': 'detailed'}.
"""
if value is None:

View file

@ -1,5 +1,7 @@
import mimetypes
from collections.abc import Mapping
from io import BufferedReader, BytesIO
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, cast
from urllib.parse import quote
@ -502,15 +504,26 @@ class OpenAIVideoConfig(BaseVideoConfig):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
video_file: FileContent | None = None,
extra_body: dict[str, object] | None = None,
prefetched_source_data: dict[str, object] | None = None,
) -> tuple[str, dict]:
original_video_id: Final = extract_original_video_id(video_id)
) -> tuple[str, Mapping[str, object], RequestFiles | None]:
url: Final = f"{api_base.rstrip('/')}/edits"
if video_file is not None:
files: Final[RequestFiles] = (self._video_file_tuple(video_file, "video"),)
form_data: Final = (
MappingProxyType({"prompt": prompt, **extra_body})
if extra_body
else MappingProxyType({"prompt": prompt})
)
return url, form_data, files
original_video_id: Final = extract_original_video_id(video_id)
data: Final[dict[str, object]] = {"prompt": prompt, "video": {"id": original_video_id}}
if extra_body:
data.update(extra_body)
return url, data
return url, data, None
def transform_video_edit_response(
self,
@ -570,21 +583,22 @@ class OpenAIVideoConfig(BaseVideoConfig):
else:
files_list.append((field_name, ("input_reference.png", image, image_content_type)))
def _video_file_tuple(self, video: FileContent, field_name: str) -> tuple[str, FileTypes]:
"""
Build a multipart field tuple for a video upload with proper video MIME
type detection: these paths must send video/mp4, not image/* content types.
"""
filename: Final = getattr(video, "name", None) or "input_video.mp4"
content_type: Final = self._get_video_content_type(video=video, filename=filename)
return (field_name, (filename, video, content_type))
def _add_video_to_files(
self,
files_list: list[tuple[str, FileTypes]],
video: FileContent,
field_name: str,
) -> None:
"""
Add a video to files with proper video MIME type detection.
This path is used by POST /videos/characters and must send video/mp4,
not image/* content types.
"""
filename: Final = getattr(video, "name", None) or "input_video.mp4"
content_type: Final = self._get_video_content_type(video=video, filename=filename)
files_list.append((field_name, (filename, video, content_type)))
files_list.append(self._video_file_tuple(video, field_name))
def _get_video_content_type(self, video: FileContent, filename: str) -> str:
guessed_content_type, _ = mimetypes.guess_type(filename)

View file

@ -672,6 +672,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
api_base,
litellm_params,
headers,
video_file=None,
extra_body=None,
prefetched_source_data=None,
):

View file

@ -16,11 +16,16 @@ from litellm.llms.together_ai.rerank.transformation import TogetherAIRerankConfi
from litellm.types.rerank import RerankRequest, RerankResponse
def _rerank_url(api_base: str) -> str:
return f"{api_base.rstrip('/')}/rerank"
class TogetherAIRerank(BaseLLM):
def rerank(
self,
model: str,
api_key: str,
api_base: str,
query: str,
documents: list[str | dict[str, Any]],
top_n: int | None = None,
@ -46,10 +51,10 @@ class TogetherAIRerank(BaseLLM):
raise ValueError("TogetherAI does not support max_chunks_per_doc")
if _is_async:
return self.async_rerank(request_data_dict, api_key) # Call async method
return self.async_rerank(request_data_dict, api_key, api_base)
response: Final = client.post(
"https://api.together.xyz/v1/rerank",
_rerank_url(api_base),
headers={
"accept": "application/json",
"content-type": "application/json",
@ -69,11 +74,12 @@ class TogetherAIRerank(BaseLLM):
self,
request_data_dict: dict[str, Any],
api_key: str,
api_base: str,
) -> RerankResponse:
client: Final = get_async_httpx_client(llm_provider=litellm.LlmProviders.TOGETHER_AI) # Use async client
response: Final = await client.post(
"https://api.together.xyz/v1/rerank",
_rerank_url(api_base),
headers={
"accept": "application/json",
"content-type": "application/json",

View file

@ -0,0 +1,149 @@
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from typing import Final
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.gemini.interactions.transformation import GoogleAIStudioInteractionsConfig
from litellm.llms.vertex_ai.common_utils import validate_vertex_location
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
VERTEX_INTERACTIONS_API_VERSION: Final = "v1beta1"
VERTEX_INTERACTIONS_DEFAULT_LOCATION: Final = "global"
@dataclass(frozen=True, slots=True)
class VertexInteractionsTarget:
base_url: str
project_id: str
location: str
@property
def collection_url(self) -> str:
return (
f"{self.base_url}/{VERTEX_INTERACTIONS_API_VERSION}"
f"/projects/{self.project_id}/locations/{self.location}/interactions"
)
def interaction_url(self, interaction_id: str) -> str:
encoded_interaction_id: Final = encode_url_path_segment(interaction_id, field_name="interaction_id")
return f"{self.collection_url}/{encoded_interaction_id}"
class VertexAIInteractionsConfig(VertexBase, GoogleAIStudioInteractionsConfig):
def __init__(
self,
mint_access_token: Callable[[VERTEX_CREDENTIALS_TYPES | None, str | None], tuple[str, str]] | None = None,
) -> None:
super().__init__()
self._mint_access_token: Final[Callable[[VERTEX_CREDENTIALS_TYPES | None, str | None], tuple[str, str]]] = (
mint_access_token or self._mint_access_token_with_vertex_base
)
def _mint_access_token_with_vertex_base(
self,
credentials: VERTEX_CREDENTIALS_TYPES | None,
project_id: str | None,
) -> tuple[str, str]:
return self._ensure_access_token(
credentials=credentials, project_id=project_id, custom_llm_provider="vertex_ai"
)
@property
def custom_llm_provider(self) -> LlmProviders:
return LlmProviders.VERTEX_AI
@property
def api_version(self) -> str:
return VERTEX_INTERACTIONS_API_VERSION
def get_default_vertex_location(self) -> str:
return VERTEX_INTERACTIONS_DEFAULT_LOCATION
def _mint(self, litellm_params: GenericLiteLLMParams) -> tuple[str, str]:
raw_params: Final = litellm_params.model_dump()
return self._mint_access_token(
self.safe_get_vertex_ai_credentials(raw_params),
self.safe_get_vertex_ai_project(raw_params),
)
def _target(self, api_base: str | None, litellm_params: GenericLiteLLMParams) -> VertexInteractionsTarget:
_, project_id = self._mint(litellm_params)
if not project_id:
raise ValueError(
"Vertex AI project is required. Set vertex_project, litellm.vertex_project, or VERTEXAI_PROJECT"
)
location: Final = validate_vertex_location(
self.explicit_vertex_ai_location(litellm_params.model_dump()) or VERTEX_INTERACTIONS_DEFAULT_LOCATION
)
return VertexInteractionsTarget(
base_url=self.get_api_base(api_base or None, location),
project_id=project_id,
location=location,
)
def validate_environment(
self,
headers: Mapping[str, str],
model: str,
litellm_params: GenericLiteLLMParams | None,
) -> dict: # mutable-ok: BaseInteractionsAPIConfig declares plain-dict headers
access_token, _ = self._mint(litellm_params or GenericLiteLLMParams())
return { # mutable-ok: BaseInteractionsAPIConfig declares plain-dict headers
"Content-Type": "application/json",
"Authorization": f"Bearer {access_token}",
**headers,
}
def get_complete_url(
self,
api_base: str | None,
model: str | None,
agent: str | None = None,
litellm_params: Mapping[str, object] | None = None,
stream: bool | None = None,
) -> str:
params: Final = (
GenericLiteLLMParams.model_validate(litellm_params) if litellm_params else GenericLiteLLMParams()
)
collection_url: Final = self._target(api_base, params).collection_url
return f"{collection_url}?alt=sse" if stream else collection_url
def _interaction_by_id_request(
self,
interaction_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
url_suffix: str = "",
) -> tuple[str, dict]: # mutable-ok: BaseInteractionsAPIConfig declares a plain-dict request body
target: Final = self._target(api_base or None, litellm_params)
return f"{target.interaction_url(interaction_id)}{url_suffix}", {} # mutable-ok: same base contract
def transform_get_interaction_request(
self,
interaction_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
) -> tuple[str, dict]: # mutable-ok: BaseInteractionsAPIConfig declares a plain-dict request body
return self._interaction_by_id_request(interaction_id, api_base, litellm_params)
def transform_delete_interaction_request(
self,
interaction_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
) -> tuple[str, dict]: # mutable-ok: BaseInteractionsAPIConfig declares a plain-dict request body
return self._interaction_by_id_request(interaction_id, api_base, litellm_params)
def transform_cancel_interaction_request(
self,
interaction_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
) -> tuple[str, dict]: # mutable-ok: BaseInteractionsAPIConfig declares a plain-dict request body
return self._interaction_by_id_request(interaction_id, api_base, litellm_params, url_suffix=":cancel")

View file

@ -7,11 +7,11 @@ Based on: https://docs.cloud.google.com/vertex-ai/generative-ai/docs/model-refer
import base64
import time
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
import httpx
from httpx._types import RequestFiles
from httpx._types import FileContent, RequestFiles
from typing_extensions import ReadOnly
from litellm.constants import DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS
@ -677,9 +677,10 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
video_file: FileContent | None = None,
extra_body: dict[str, object] | None = None,
prefetched_source_data: dict[str, Any] | None = None,
) -> tuple[str, dict]:
) -> tuple[str, Mapping[str, object], RequestFiles | None]:
"""
Build a predictLongRunning edit request from the pre-fetched source video.
@ -727,7 +728,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
request_data["parameters"] = vertex_params
edit_url: Final = f"{api_base.rstrip('/')}/{model}:predictLongRunning"
return edit_url, request_data
return edit_url, request_data, None
def transform_video_edit_response(
self,

View file

@ -416,7 +416,7 @@ async def acompletion(
logprobs: bool | None = None,
top_logprobs: int | None = None,
deployment_id=None,
reasoning_effort: Literal["none", "minimal", "low", "medium", "high", "xhigh", "default"] | None = None,
reasoning_effort: Literal["none", "minimal", "low", "medium", "high", "xhigh", "max", "default"] | None = None,
verbosity: Literal["low", "medium", "high"] | None = None,
safety_identifier: str | None = None,
service_tier: str | None = None,
@ -602,7 +602,7 @@ async def acompletion(
_, custom_llm_provider, _, _ = get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=base_url,
api_base=kwargs.get("api_base") or base_url,
)
fallbacks = fallbacks or litellm.model_fallbacks
@ -4920,7 +4920,7 @@ def completion(
logit_bias: dict | None = None,
user: str | None = None,
# openai v1.0+ new params
reasoning_effort: Literal["none", "minimal", "low", "medium", "high", "xhigh", "default"] | None = None,
reasoning_effort: Literal["none", "minimal", "low", "medium", "high", "xhigh", "max", "default"] | None = None,
verbosity: Literal["low", "medium", "high"] | None = None,
response_format: dict | type[BaseModel] | None = None,
seed: int | None = None,

View file

@ -17155,6 +17155,14 @@
"notes": "Web Search on Amazon Bedrock AgentCore, billed by AWS on the gateway"
}
},
"bing_grounding/search": {
"input_cost_per_query": 0.035,
"litellm_provider": "bing_grounding",
"mode": "search",
"metadata": {
"notes": "Grounding with Bing Search (G1 SKU): $35 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment."
}
},
"tinyfish/search": {
"input_cost_per_query": 0.0,
"litellm_provider": "tinyfish",
@ -49008,12 +49016,13 @@
"output_cost_per_token": 3.3e-05,
"output_cost_per_token_above_272k_tokens": 4.95e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1000000,
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"use_openai_responses_path": true,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
@ -49040,12 +49049,13 @@
"output_cost_per_token": 1.32e-05,
"output_cost_per_token_above_272k_tokens": 1.98e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1000000,
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"use_openai_responses_path": true,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
@ -49072,12 +49082,13 @@
"output_cost_per_token": 1.32e-06,
"output_cost_per_token_above_272k_tokens": 1.98e-06,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1000000,
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"use_openai_responses_path": true,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [

View file

@ -199,7 +199,7 @@ def llm_passthrough_route(
api_key=api_key,
)
litellm_params_dict: Final = get_litellm_params(**kwargs)
litellm_params_dict: Final = get_litellm_params(api_key=api_key, api_base=api_base, **kwargs)
if client is None:
from litellm.llms.custom_httpx.http_handler import (

View file

@ -815,6 +815,7 @@ class LiteLLMRoutes(enum.Enum):
"/team/member_add",
"/team/member_delete",
"/team/member_update",
"/team/{team_id}/member/{user_id}/reset_spend",
"/team/permissions_list",
"/team/permissions_update",
"/team/daily/activity",
@ -1287,6 +1288,16 @@ class RegenerateKeyRequest(GenerateKeyRequest):
class ResetSpendRequest(LiteLLMPydanticObjectBase):
reset_to: float
@field_validator("reset_to", mode="before")
@classmethod
def reject_bool_reset_to(cls, v):
# bool is a subclass of int, so pydantic silently coerces True/False into
# 1.0/0.0 for a `float` field: a caller who accidentally sends a boolean
# would otherwise get an unintended spend reset instead of a 422.
if isinstance(v, bool):
raise ValueError("reset_to must be a number, not a boolean") # noqa: TRY004 # pydantic needs ValueError
return v
class KeyRequest(LiteLLMPydanticObjectBase):
keys: list[str] | None = None

View file

@ -34,10 +34,18 @@ class CacheActivityFilterOptions(BaseModel):
models: list[str]
class CacheActivityErrorBucket(BaseModel):
call_type: str
error_code: str
error_class: str
count: int
class CacheActivityResponse(BaseModel):
groups: list[CacheActivityGroup]
totals: CacheActivityTotals
filter_options: CacheActivityFilterOptions
error_breakdown: tuple[CacheActivityErrorBucket, ...]
GROUPS_SQL: Final = """
@ -65,6 +73,26 @@ GROUPS_SQL: Final = """
ORDER BY (COUNT(*)) DESC
"""
ERROR_BREAKDOWN_SQL: Final = """
SELECT
CASE WHEN sl."call_type" = '' THEN 'Unknown' ELSE sl."call_type" END AS call_type,
COALESCE(NULLIF(sl."metadata"->'error_information'->>'error_code', ''), 'Unknown') AS error_code,
COALESCE(NULLIF(sl."metadata"->'error_information'->>'error_class', ''), 'Unknown') AS error_class,
COUNT(*)::int AS count
FROM "LiteLLM_SpendLogs" sl
LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token"
WHERE
sl."status" = 'failure'
AND sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
AND ($3::jsonb = '[]'::jsonb
OR COALESCE(vt."key_alias", 'Unnamed Key') IN (SELECT jsonb_array_elements_text($3::jsonb)))
AND ($4::jsonb = '[]'::jsonb
OR sl."model" IN (SELECT jsonb_array_elements_text($4::jsonb)))
GROUP BY 1, 2, 3
ORDER BY (COUNT(*)) DESC
"""
KEY_ALIAS_OPTIONS_SQL: Final = """
SELECT DISTINCT COALESCE(vt."key_alias", 'Unnamed Key') AS key_alias
FROM "LiteLLM_SpendLogs" sl
@ -95,6 +123,7 @@ class _ModelRow(BaseModel):
_groups_adapter: Final = TypeAdapter(list[CacheActivityGroup])
_error_buckets_adapter: Final = TypeAdapter(tuple[CacheActivityErrorBucket, ...])
_key_alias_rows_adapter: Final = TypeAdapter(list[_KeyAliasRow])
_model_rows_adapter: Final = TypeAdapter(list[_ModelRow])
@ -120,10 +149,11 @@ async def get_cache_activity(
key_aliases: Sequence[str],
models: Sequence[str],
) -> CacheActivityResponse:
group_rows, key_alias_rows, model_rows = await asyncio.gather(
prisma_client.db.query_raw(
GROUPS_SQL, start_date, end_date, json.dumps(list(key_aliases)), json.dumps(list(models))
),
key_aliases_json: Final = json.dumps(list(key_aliases))
models_json: Final = json.dumps(list(models))
group_rows, error_rows, key_alias_rows, model_rows = await asyncio.gather(
prisma_client.db.query_raw(GROUPS_SQL, start_date, end_date, key_aliases_json, models_json),
prisma_client.db.query_raw(ERROR_BREAKDOWN_SQL, start_date, end_date, key_aliases_json, models_json),
prisma_client.db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date),
prisma_client.db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date),
)
@ -135,4 +165,5 @@ async def get_cache_activity(
key_aliases=[row.key_alias for row in _key_alias_rows_adapter.validate_python(key_alias_rows or [])],
models=[row.model for row in _model_rows_adapter.validate_python(model_rows or [])],
),
error_breakdown=_error_buckets_adapter.validate_python(error_rows or []),
)

View file

@ -71,7 +71,6 @@ from litellm.proxy.auth.budget_throttle import (
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import publish_auth_cache_invalidation
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_headers,
_safe_get_request_query_params,
@ -87,6 +86,8 @@ from litellm.proxy.common_utils.user_api_key_cache import (
object_permission_cache_key,
tag_cache_key,
tag_registry_cache_key,
team_membership_auth_cache_key,
team_membership_reservation_cache_key,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.guardrails.tool_name_extraction import (
@ -1967,7 +1968,7 @@ async def get_team_membership(
if user_id is None or team_id is None:
return None
_key: Final = f"team_membership:{user_id}:{team_id}"
_key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id)
# check if in cache
cached_membership_obj: Final = await user_api_key_cache.async_get_cache(
@ -2402,6 +2403,116 @@ async def _cache_team_object(
)
async def invalidate_team_member_spend_state(
user_id: str,
team_id: str,
user_api_key_cache: UserApiKeyCache,
new_spend: float | None = None,
) -> None:
"""
Clear every cached read path for one team member's budget so a spend
reset or a raised cap takes effect on the next request instead of
waiting on the membership cache's TTL.
Two independently-keyed cache entries hold the same LiteLLM_TeamMembership
row: user_api_key_auth.py's admission check writes ``{team_id}_{user_id}``,
while budget_reservation.py's pre-call reservation and auth_checks.py's own
get_team_membership() (used by _check_team_member_budget) both write
``team_membership:{user_id}:{team_id}``. Both formats must be invalidated
explicitly; writing one does not refresh the other. All keys are also
broadcast (LIT-3803): each worker's own in-memory copy (membership object,
spend counter, or the counter's own short-TTL DB-floor marker) survives
eviction elsewhere until its TTL, so the handling worker alone clearing its
copy leaves every other worker still enforcing the pre-reset budget.
``new_spend`` is only passed by reset_team_member_spend_fn, which knows the
exact post-reset value: it is SET everywhere (matching /key/{key}/reset_spend's
own precedent) rather than deleted, so a worker's next read reflects it
directly instead of re-deriving it through a DB reseed. team_member_update
only changes the budget cap, not the tracked spend, so it passes no
new_spend; the live spend counter is untouched in that case (deleting it
would force a reseed from the DB's own spend column, which lags the live
counter via periodic batch writes, briefly under-enforcing the raised cap
against a spend value lower than what was actually tracked) and only the
membership caches carrying the new cap are invalidated.
The floor marker (``spend_db_floor:``, proxy_server.py's
_authoritative_floor_spend) caches the pre-reset DB spend for
SPEND_DB_FLOOR_CACHE_TTL_SECONDS; left stale after a real reset, a request
landing on the pod that cached it can read that higher floor and raise the
counter right back above the just-reset spend. It is overwritten here with
the post-reset floor (not merely deleted) and _authoritative_floor_spend
re-checks the marker after its DB read, so a floor read already in flight
on this pod when the reset commits cannot clobber it with the pre-reset
value. Both keys are broadcast as SETs carrying new_spend, not deletes:
every subscriber (remote pods AND this pod's own, which receives its own
message) writes the post-reset value, so the self-delivered message cannot
erase the guard just written here.
Raises HTTPException(503) if Redis still holds the stale pre-reset counter
after both the SET and the fallback DELETE fail: budget checks read Redis
first, so returning success would leave the old value authoritative for
every worker despite the DB write having committed.
"""
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
evict_and_broadcast,
publish_auth_cache_invalidation,
)
if new_spend is not None:
from litellm.proxy.proxy_server import SPEND_DB_FLOOR_CACHE_TTL_SECONDS, spend_counter_cache
spend_counter_key: Final = f"spend:team_member:{user_id}:{team_id}"
spend_db_floor_key: Final = f"spend_db_floor:{spend_counter_key}"
spend_counter_cache.in_memory_cache.set_cache(key=spend_counter_key, value=new_spend, ttl=60)
if spend_counter_cache.redis_cache is not None:
try:
await spend_counter_cache.redis_cache.async_set_cache(key=spend_counter_key, value=new_spend, ttl=60)
except Exception as e: # noqa: BLE001 # fall back to deleting the stale entry before giving up
verbose_proxy_logger.warning(
"Failed to set spend counter %s in Redis after reset: %s; deleting it instead so the next "
"read reseeds from the DB rather than keeping the stale pre-reset value authoritative",
spend_counter_key,
e,
)
try:
await spend_counter_cache.redis_cache.async_delete_cache(key=spend_counter_key)
except Exception: # noqa: BLE001 # stale value now authoritative in Redis; surface instead of reporting success
verbose_proxy_logger.warning(
"Failed to delete stale spend counter %s in Redis after a failed reset write",
spend_counter_key,
exc_info=True,
)
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail={ # mutable-ok: HTTPException.detail takes a dict
"error": "Spend was reset in the database, but Redis is unreachable and still "
"holds the pre-reset counter. Retry once Redis is reachable."
},
) from e
spend_counter_cache.in_memory_cache.set_cache(
key=spend_db_floor_key,
value=new_spend,
ttl=SPEND_DB_FLOOR_CACHE_TTL_SECONDS,
)
await publish_auth_cache_invalidation(cache_key=spend_counter_key, new_value=new_spend, ttl=60)
await publish_auth_cache_invalidation(
cache_key=spend_db_floor_key,
new_value=new_spend,
ttl=SPEND_DB_FLOOR_CACHE_TTL_SECONDS,
)
await evict_and_broadcast(
cache_keys=(
team_membership_auth_cache_key(team_id=team_id, user_id=user_id),
team_membership_reservation_cache_key(user_id=user_id, team_id=team_id),
),
user_api_key_cache=user_api_key_cache,
)
async def delete_cache_team_object(
team_id: str,
team_alias: str | None,
@ -2629,20 +2740,9 @@ async def _get_team_object_from_user_api_key_cache(
async def _get_team_object_from_cache(
key: str,
proxy_logging_obj: ProxyLogging | None,
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None,
) -> LiteLLM_TeamTableCachedObj | None:
## INTERNAL USAGE CACHE (plain DualCache) — checked before UserApiKeyCache stores ##
if proxy_logging_obj is not None and proxy_logging_obj.internal_usage_cache.dual_cache:
cached_raw: Final = await proxy_logging_obj.internal_usage_cache.dual_cache.async_get_cache(
key=key, parent_otel_span=parent_otel_span
)
if cached_raw is not None:
from_internal: Final = CacheCodec.deserialize(cached_raw, LiteLLM_TeamTableCachedObj)
if from_internal is not None:
return from_internal
decoded: Final = await user_api_key_cache.async_get_cache(
key=key,
parent_otel_span=parent_otel_span,
@ -2678,7 +2778,6 @@ async def get_team_object(
if not check_db_only:
cached_team_obj: Final = await _get_team_object_from_cache(
key=key,
proxy_logging_obj=proxy_logging_obj,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
)
@ -2841,7 +2940,6 @@ async def get_team_object_by_alias(
cached_team_obj: Final = await _get_team_object_from_cache(
key=cache_key,
proxy_logging_obj=proxy_logging_obj,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
)

View file

@ -11,6 +11,7 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import EMPTY_MAPPING
from litellm.integrations.otel.runtime import seed_request_identity
from litellm.litellm_core_utils.core_helpers import is_expected_client_error
from litellm.proxy._types import (
LitellmUserRoles,
ProxyErrorTypes,
@ -109,7 +110,12 @@ class UserAPIKeyAuthExceptionHandler:
request=request,
use_x_forwarded_for=general_settings.get("use_x_forwarded_for") is True,
)
verbose_proxy_logger.exception(
log_fn: Final = (
verbose_proxy_logger.error
if is_expected_client_error(e) and not litellm.log_client_error_tracebacks
else verbose_proxy_logger.exception
)
log_fn(
"litellm.proxy.proxy_server.user_api_key_auth(): Exception occured - %s\nRequester IP Address:%s",
e,
requester_ip,

View file

@ -87,7 +87,10 @@ from litellm.proxy.common_utils.http_parsing_utils import (
populate_request_with_path_params,
)
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
team_membership_auth_cache_key,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.utils import (
@ -1970,8 +1973,10 @@ async def _user_api_key_auth_builder(
# Check 3. Check if user is in their team budget
if not skip_budget_checks and valid_token.team_member_spend is not None:
if prisma_client is not None:
_cache_key: Final = f"{valid_token.team_id}_{valid_token.user_id}"
_user_id: Final = valid_token.user_id
_team_id: Final = valid_token.team_id
if prisma_client is not None and _user_id is not None and _team_id is not None:
_cache_key: Final = team_membership_auth_cache_key(team_id=_team_id, user_id=_user_id)
team_member_info = await user_api_key_cache.async_get_cache(
key=_cache_key,
@ -1979,25 +1984,21 @@ async def _user_api_key_auth_builder(
)
if team_member_info is None:
# read from DB
_user_id: Final = valid_token.user_id
_team_id: Final = valid_token.team_id
if _user_id is not None and _team_id is not None:
_db_member: Final = await TeamMembershipRepository(prisma_client).table.find_first(
where={
"user_id": _user_id,
"team_id": _team_id,
},
include={"litellm_budget_table": True},
_db_member: Final = await TeamMembershipRepository(prisma_client).table.find_first(
where={
"user_id": _user_id,
"team_id": _team_id,
},
include={"litellm_budget_table": True},
)
if _db_member is not None:
team_member_info = LiteLLM_TeamMembership(**_db_member.dict())
await user_api_key_cache.async_set_cache(
key=_cache_key,
value=team_member_info,
model_type=LiteLLM_TeamMembership,
ttl=5,
)
if _db_member is not None:
team_member_info = LiteLLM_TeamMembership(**_db_member.dict())
await user_api_key_cache.async_set_cache(
key=_cache_key,
value=team_member_info,
model_type=LiteLLM_TeamMembership,
ttl=5,
)
if team_member_info is not None and team_member_info.litellm_budget_table is not None:
team_member_budget: Final = team_member_info.litellm_budget_table.max_budget
@ -2013,11 +2014,16 @@ async def _user_api_key_auth_builder(
max_budget=team_member_budget,
)
if team_member_spend > team_member_budget:
_entity_id: Final = f"{valid_token.user_id}:{valid_token.team_id}"
raise litellm.BudgetExceededError(
current_cost=team_member_spend,
max_budget=team_member_budget,
message=(
f"Budget has been exceeded! TeamMember={_entity_id} "
f"Current cost: {team_member_spend}, Max budget: {team_member_budget}"
),
entity_type=Litellm_EntityType.TEAM_MEMBER.value,
entity_id=f"{valid_token.user_id}:{valid_token.team_id}",
entity_id=_entity_id,
)
# Check 3. If token is expired

View file

@ -33,7 +33,7 @@ from litellm.constants import (
UNSAFE_PROXY_RESPONSE_HEADERS,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket, is_expected_client_error
from litellm.litellm_core_utils.dd_tracing import NullTracer, tracer
from litellm.litellm_core_utils.get_supported_openai_params import (
get_supported_openai_params,
@ -1138,24 +1138,25 @@ async def open_sse_before_first_byte(
)
def _is_azure_model_router_request(model: str) -> bool:
def _is_azure_model_router_request(model: str, hidden_params: Mapping[str, object] | None = None) -> bool:
"""
Check if the requested model is an Azure Model Router.
Check if a request went down the Azure Model Router route.
Azure Model Router models follow the pattern:
- azure_ai/model_router/<deployment-name>
- azure_ai/model-router
- model_router/<deployment-name>
- model-router
``model`` here is what the *client* sent, a model group alias with no ``model_router/``
prefix, so matching on it alone only works when the operator happened to put "model-router"
in the alias. Where the response is in hand its stamp answers this outright, so callers
should pass ``hidden_params``.
Args:
model: The requested model name
hidden_params: ``_hidden_params`` from the response, when the caller has it
Returns:
bool: True if this is an Azure Model Router request
"""
model_lower: Final = model.lower()
return "model-router" in model_lower or "model_router" in model_lower
from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo
return AzureFoundryModelInfo.is_model_router_call(model=model, hidden_params=hidden_params)
def _override_openai_response_model(
@ -1223,7 +1224,7 @@ def _override_openai_response_model(
return
# Check if this is an Azure Model Router request - if so, preserve the actual model used
if _is_azure_model_router_request(requested_model):
if _is_azure_model_router_request(requested_model, hidden_params):
verbose_proxy_logger.debug(
"%s: Azure Model Router detected - preserving actual model used from response instead of overriding to router model.",
log_context,
@ -1379,7 +1380,12 @@ def _log_llm_api_exception(e: Exception) -> None:
"litellm.proxy.proxy_server._handle_llm_api_exception(): client disconnected, upstream LLM request cancelled"
)
return
verbose_proxy_logger.exception("litellm.proxy.proxy_server._handle_llm_api_exception(): Exception occured - %s", e)
log_fn: Final = (
verbose_proxy_logger.error
if is_expected_client_error(e) and not litellm.log_client_error_tracebacks
else verbose_proxy_logger.exception
)
log_fn("litellm.proxy.proxy_server._handle_llm_api_exception(): Exception occured - %s", e)
async def _cancel_llm_call_on_client_disconnect(

View file

@ -12,6 +12,7 @@ from litellm.proxy.common_utils.config_sync_pubsub import (
)
if TYPE_CHECKING:
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.caching.redis_cache import RedisCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -30,15 +31,24 @@ def auth_cache_invalidation_channel(redis_cache: "RedisCache") -> str:
@dataclass(frozen=True, slots=True)
class _CacheInvalidationMessage:
cache_key: str
new_value: float | None = None
ttl: float | None = None
def _cache_invalidation_message_json(cache_key: str) -> str:
return json.dumps(asdict(_CacheInvalidationMessage(cache_key=cache_key)))
def _cache_invalidation_message_json(cache_key: str, new_value: float | None = None, ttl: float | None = None) -> str:
message: Final = asdict(_CacheInvalidationMessage(cache_key=cache_key, new_value=new_value, ttl=ttl))
return json.dumps({field: value for field, value in message.items() if value is not None})
def _cache_key_from_message_data(data: object) -> str | None:
def _finite_number_or_none(value: object) -> float | None:
if isinstance(value, bool) or not isinstance(value, (int, float)):
return None
return float(value)
def _message_from_data(data: object) -> _CacheInvalidationMessage | None:
if isinstance(data, bytes):
data = data.decode("utf-8", errors="replace")
data = data.decode("utf-8", errors="replace") # rebind-ok: normalizing the wire payload to str
if not isinstance(data, str):
return None
try:
@ -48,14 +58,28 @@ def _cache_key_from_message_data(data: object) -> str | None:
if not isinstance(parsed, dict):
return None
cache_key: Final = parsed.get("cache_key")
return cache_key if isinstance(cache_key, str) else None
if not isinstance(cache_key, str):
return None
return _CacheInvalidationMessage(
cache_key=cache_key,
new_value=_finite_number_or_none(parsed.get("new_value")),
ttl=_finite_number_or_none(parsed.get("ttl")),
)
async def publish_auth_cache_invalidation(cache_key: str) -> None:
async def publish_auth_cache_invalidation(
cache_key: str, new_value: float | None = None, ttl: float | None = None
) -> None:
"""
Best-effort broadcast so every worker drops its local in-memory copy of a
mutated management object; without this, only the handling worker and Redis
are evicted and other workers keep serving the stale object until its TTL.
Passing ``new_value`` broadcasts a SET instead of a delete: every subscriber
(including the publishing worker's own, which receives its own message)
writes the value into its additional in-memory caches rather than deleting
the key. A spend reset uses this so the handler's self-delivered message
cannot erase the freshly-written post-reset counter or floor marker.
"""
redis_cache: Final = coordination_redis_cache()
if redis_cache is None:
@ -68,7 +92,10 @@ async def publish_auth_cache_invalidation(cache_key: str) -> None:
cache_key,
)
return
await client.publish(auth_cache_invalidation_channel(redis_cache), _cache_invalidation_message_json(cache_key))
await client.publish(
auth_cache_invalidation_channel(redis_cache),
_cache_invalidation_message_json(cache_key, new_value=new_value, ttl=ttl),
)
except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors
verbose_proxy_logger.warning("auth cache invalidation publish for %s failed: %s", cache_key, e)
@ -95,15 +122,17 @@ async def evict_and_broadcast(cache_keys: Sequence[str], user_api_key_cache: "Us
class AuthCacheInvalidationSubscriber:
__slots__ = ("_redis_cache", "_task", "_user_api_key_cache")
__slots__ = ("_additional_in_memory_caches", "_redis_cache", "_task", "_user_api_key_cache")
def __init__(
self,
redis_cache: "RedisCache",
user_api_key_cache: "UserApiKeyCache",
additional_in_memory_caches: Sequence["InMemoryCache"] = (),
) -> None:
self._redis_cache = redis_cache
self._user_api_key_cache = user_api_key_cache
self._additional_in_memory_caches = tuple(additional_in_memory_caches)
self._task: asyncio.Task[None] | None = None
def start(self) -> None:
@ -160,12 +189,18 @@ class AuthCacheInvalidationSubscriber:
def _apply_message(self, message: object) -> None:
data: Final = message.get("data") if isinstance(message, dict) else None
cache_key: Final = _cache_key_from_message_data(data)
if cache_key is None:
parsed: Final = _message_from_data(data)
if parsed is None:
return
if parsed.new_value is not None:
for additional_cache in self._additional_in_memory_caches:
additional_cache.set_cache(parsed.cache_key, parsed.new_value, ttl=parsed.ttl)
return
in_memory_cache: Final = self._user_api_key_cache.in_memory_cache
if in_memory_cache is not None:
in_memory_cache.delete_cache(cache_key)
in_memory_cache.delete_cache(parsed.cache_key)
for additional_cache in self._additional_in_memory_caches:
additional_cache.delete_cache(parsed.cache_key)
@staticmethod
async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None:

View file

@ -200,6 +200,21 @@ def end_user_restricted_registry_cache_key() -> str:
return "end_user_restricted_registry"
def team_membership_auth_cache_key(team_id: str, user_id: str) -> str:
"""Cache key one team member's ``LiteLLM_TeamMembership`` row is stored under for the admission check."""
return f"{team_id}_{user_id}"
def team_membership_reservation_cache_key(user_id: str, team_id: str) -> str:
"""Cache key the pre-call budget reservation stores the same ``LiteLLM_TeamMembership`` row under.
Deliberately not unified with ``team_membership_auth_cache_key``: the two readers wrote independent
keys before this file existed, so a fix that invalidates one must invalidate both explicitly rather
than assume a single write is visible to both.
"""
return f"team_membership:{user_id}:{team_id}"
def get_management_object_ttl(cache: DualCache) -> float:
"""
In-memory TTL for management-object cache writes (keys, teams, users, budgets, ...).

View file

@ -0,0 +1,30 @@
# Web search via Microsoft Foundry (Grounding with Bing Search / the built-in
# web_search tool), called through the Foundry Responses API.
#
# Configure the provider with env vars (setup and pricing are in the LiteLLM docs;
# the code lives in litellm/llms/azure/search/transformation.py):
# BING_GROUNDING_PROJECT_ENDPOINT (required) the Foundry project endpoint
# BING_GROUNDING_MODEL (required) a model deployment in that project
# BING_GROUNDING_CONNECTION_ID (optional) a Grounding with Bing connection id;
# without it the built-in web_search tool is used
# BING_GROUNDING_TOKEN (optional) an Entra bearer token; without it (and
# without api_key) azure-identity mints one
model_list:
- model_name: claude-sonnet
litellm_params:
model: bedrock/us.anthropic.claude-sonnet-5
aws_region_name: us-east-1
search_tools:
- search_tool_name: bing-grounding-search
litellm_params:
search_provider: bing_grounding
# Optional: an Azure API key instead of BING_GROUNDING_TOKEN / azure-identity
# api_key: os.environ/AZURE_AI_API_KEY
litellm_settings:
callbacks: ["websearch_interception"]
websearch_interception_params:
enabled_providers: ["bedrock"]
search_tool_name: bing-grounding-search

View file

@ -1,6 +1,7 @@
import asyncio
import io
import traceback
from collections.abc import Sequence
from typing import Final
import orjson
@ -33,10 +34,10 @@ async def uploadfile_to_bytesio(upload: UploadFile) -> io.BytesIO:
async def batch_to_bytesio(
uploads: list[UploadFile] | None,
uploads: Sequence[UploadFile] | None,
) -> list[io.BytesIO] | None:
"""
Convert a list of UploadFiles to a list of BytesIO buffers, or None.
Convert a sequence of UploadFiles to a list of BytesIO buffers, or None.
"""
if not uploads:
return None

View file

@ -210,7 +210,6 @@ async def _patch_team_caches_add_access_group(
for team_id in team_ids:
cached_team = await _get_team_object_from_cache(
key=f"team_id:{team_id}",
proxy_logging_obj=proxy_logging_obj,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
)
@ -240,7 +239,6 @@ async def _patch_team_caches_remove_access_group(
for team_id in team_ids:
cached_team = await _get_team_object_from_cache(
key=f"team_id:{team_id}",
proxy_logging_obj=proxy_logging_obj,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
)

View file

@ -16,7 +16,7 @@ import traceback
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Annotated, Final, NamedTuple, Protocol, TypedDict, TypeVar, cast
from typing import Annotated, Final, NamedTuple, NoReturn, Protocol, TypedDict, TypeVar, cast
import fastapi
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
@ -56,6 +56,7 @@ from litellm.proxy._types import (
PatchTeamRequest,
ProxyErrorTypes,
ProxyException,
ResetSpendRequest,
SpecialManagementEndpointEnums,
SpecialModelNames,
SpecialProxyStrings,
@ -84,6 +85,7 @@ from litellm.proxy.auth.auth_checks import (
get_team_membership,
get_team_object,
get_user_object,
invalidate_team_member_spend_state,
)
from litellm.proxy.auth.auth_utils import (
enforce_batch_enqueued_token_limit_is_admin_only,
@ -3392,7 +3394,7 @@ async def team_member_update(
Update team member budgets and team member role
"""
from litellm.proxy.proxy_server import premium_user, prisma_client
from litellm.proxy.proxy_server import premium_user, prisma_client, user_api_key_cache
if prisma_client is None:
raise HTTPException(status_code=500, detail={"error": "No db connected"})
@ -3491,6 +3493,12 @@ async def team_member_update(
budget_patch=budget_patch,
team_default_budget_id=team_default_budget_id,
)
if budget_patch:
await invalidate_team_member_spend_state(
user_id=received_user_id,
team_id=data.team_id,
user_api_key_cache=user_api_key_cache,
)
### update team member role
if data.role is not None:
@ -3527,6 +3535,125 @@ async def team_member_update(
)
def _check_not_resetting_own_spend(user_id: str, user_api_key_dict: UserAPIKeyAuth) -> None:
"""
_verify_team_access authorizes a team admin (or org admin) over their own
team, with no check that the target user_id differs from the caller. Left
unchecked, that admin could target their own LiteLLM_TeamMembership row and
repeatedly reset it to 0 right before it crosses their per-member cap,
consuming the shared team budget without the configured limit ever binding.
Only a proxy admin may reset an admin's own spend.
"""
if user_id == user_api_key_dict.user_id and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
_raise_reset_spend_error(status.HTTP_403_FORBIDDEN, "Cannot reset your own spend. Ask a proxy admin.")
def _raise_reset_spend_error(status_code: int, message: str) -> NoReturn:
detail: Final = {"error": message} # mutable-ok: HTTPException.detail takes a dict
raise HTTPException(status_code=status_code, detail=detail)
def _validate_team_member_reset_spend_value(
reset_to: object,
membership: LiteLLM_TeamMembership,
) -> float:
if not isinstance(reset_to, (int, float)):
_raise_reset_spend_error(status.HTTP_400_BAD_REQUEST, "reset_to must be a float")
reset_to_float: Final = float(reset_to)
if not math.isfinite(reset_to_float) or reset_to_float < 0:
_raise_reset_spend_error(status.HTTP_400_BAD_REQUEST, "reset_to must be a finite number >= 0")
current_spend: Final = membership.spend or 0.0
if reset_to_float > current_spend:
_raise_reset_spend_error(
status.HTTP_400_BAD_REQUEST,
f"reset_to ({reset_to_float}) must be <= current spend ({current_spend})",
)
max_budget: Final = membership.litellm_budget_table.max_budget if membership.litellm_budget_table else None
if max_budget is not None and reset_to_float > max_budget:
_raise_reset_spend_error(
status.HTTP_400_BAD_REQUEST,
f"reset_to ({reset_to_float}) must be <= budget ({max_budget})",
)
return reset_to_float
@router.post(
"/team/{team_id}/member/{user_id}/reset_spend",
tags=["team management"], # mutable-ok: FastAPI's `tags` param is typed as list[str], not Sequence
dependencies=(Depends(user_api_key_auth),),
)
@management_endpoint_wrapper
async def reset_team_member_spend_fn(
team_id: str,
user_id: str,
data: ResetSpendRequest,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
):
"""
Reset a team member's tracked spend against their per-member budget.
A member's spend is tracked separately from both their own personal
budget and the team's own budget (LiteLLM_TeamMembership.spend), so
neither /user/update nor /team/update can clear it: this is the only
endpoint that does. The cross-pod spend counter and cached membership
reads are invalidated so the reset takes effect on the member's next
request rather than waiting on the membership cache's TTL.
"""
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
if prisma_client is None:
_raise_reset_spend_error(status.HTTP_500_INTERNAL_SERVER_ERROR, "DB not connected. prisma_client is None")
team_obj: Final = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
check_db_only=True,
)
await _verify_team_access(team_obj=team_obj, user_api_key_dict=user_api_key_dict)
_check_not_resetting_own_spend(user_id=user_id, user_api_key_dict=user_api_key_dict)
membership_where: Final = { # mutable-ok: prisma client requires a plain dict where= argument
"user_id_team_id": {"user_id": user_id, "team_id": team_id} # mutable-ok: same prisma where= argument
}
_membership_row: Final = await _team_membership_db(prisma_client).find_unique(
where=membership_where,
include={"litellm_budget_table": True}, # mutable-ok: prisma client requires a plain dict include= argument
)
if _membership_row is None:
_raise_reset_spend_error(status.HTTP_404_NOT_FOUND, f"User {user_id} is not a member of team {team_id}.")
membership: Final = LiteLLM_TeamMembership.model_validate(_membership_row.model_dump())
current_spend: Final = membership.spend or 0.0
reset_to: Final = _validate_team_member_reset_spend_value(data.reset_to, membership)
await _team_membership_db(prisma_client).update(
where=membership_where,
data={"spend": reset_to}, # mutable-ok: prisma client requires a plain dict data= argument
)
await invalidate_team_member_spend_state(
user_id=user_id,
team_id=team_id,
user_api_key_cache=user_api_key_cache,
new_spend=reset_to,
)
return { # mutable-ok: matches this router's established untyped-response-dict convention
"team_id": team_id,
"user_id": user_id,
"spend": reset_to,
"previous_spend": current_spend,
"max_budget": membership.litellm_budget_table.max_budget if membership.litellm_budget_table else None,
}
def _create_results_from_response(
members: list[Member],
response: TeamAddMemberResponse,

View file

@ -2555,6 +2555,12 @@ async def _authoritative_floor_spend(
if db_spend is None:
return None
# a spend reset that committed during the DB read above wrote the post-reset
# floor to the marker; keep it over this read's now-stale pre-commit value
rechecked: Final = spend_counter_cache.in_memory_cache.get_cache(key=marker_key)
if rechecked is not None:
return float(rechecked)
spend_counter_cache.in_memory_cache.set_cache(
key=marker_key,
value=db_spend,
@ -6798,6 +6804,7 @@ class ProxyConfig:
subscriber: Final = AuthCacheInvalidationSubscriber(
redis_cache=redis_cache,
user_api_key_cache=user_api_key_cache,
additional_in_memory_caches=(spend_counter_cache.in_memory_cache,),
)
self.auth_cache_invalidation_subscriber = subscriber
subscriber.start()

View file

@ -25,7 +25,11 @@ from litellm.proxy._types import (
from litellm.proxy.auth.auth_utils import get_model_from_request
from litellm.proxy.auth.budget_throttle import should_throttle_budget_exceeded
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.user_api_key_cache import end_user_cache_key, tag_cache_key
from litellm.proxy.common_utils.user_api_key_cache import (
end_user_cache_key,
tag_cache_key,
team_membership_reservation_cache_key,
)
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.router import Router
@ -546,7 +550,9 @@ async def _get_team_member_budget_counter(
if team_object is None or team_object.team_id is None or user_object is None or valid_token.user_id is None:
return None
membership_cache_key: Final = f"team_membership:{valid_token.user_id}:{team_object.team_id}"
membership_cache_key: Final = team_membership_reservation_cache_key(
user_id=valid_token.user_id, team_id=team_object.team_id
)
cached_team_membership: Final = await user_api_key_cache.async_get_cache(key=membership_cache_key)
team_membership: LiteLLM_TeamMembership | None = None
if isinstance(cached_team_membership, LiteLLM_TeamMembership):

View file

@ -444,7 +444,9 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
or None
)
raw_model: Final = cast(str, kwargs.get("model") or "")
model_name: Final = reconstruct_model_name(raw_model, custom_llm_provider, metadata or {})
model_name: Final = (
standard_logging_payload.get("model") if standard_logging_payload is not None else None
) or reconstruct_model_name(raw_model, custom_llm_provider, metadata or {})
try:
payload: Final[SpendLogsPayload] = SpendLogsPayload(

View file

@ -91,7 +91,7 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.prometheus import PrometheusLogger
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert
from litellm.litellm_core_utils.core_helpers import coerce_token_limit
from litellm.litellm_core_utils.core_helpers import coerce_token_limit, is_expected_client_error
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
@ -2575,20 +2575,36 @@ class ProxyLogging:
api_key="",
)
# log the custom exception
await litellm_logging_obj.async_failure_handler(
exception=original_exception,
traceback_exception=traceback.format_exc(),
await self._dispatch_proxy_only_failure_handlers(
litellm_logging_obj=litellm_logging_obj,
original_exception=original_exception,
)
threading.Thread(
target=litellm_logging_obj.failure_handler,
args=(
original_exception,
traceback.format_exc(),
),
daemon=True,
).start()
@staticmethod
async def _dispatch_proxy_only_failure_handlers(
litellm_logging_obj: Logging,
original_exception: Exception | None,
) -> None:
"""Runs the async failure handler plus the threaded sync handler. Expected
client (4xx) errors skip traceback formatting unless
litellm.log_client_error_tracebacks is set."""
include_traceback: Final = litellm.log_client_error_tracebacks or not is_expected_client_error(
original_exception
)
traceback_str: Final = traceback.format_exc() if include_traceback else ""
await litellm_logging_obj.async_failure_handler(
exception=original_exception,
traceback_exception=traceback_str,
)
threading.Thread(
target=litellm_logging_obj.failure_handler,
args=(
original_exception,
traceback_str,
),
daemon=True,
).start()
async def post_call_success_hook(
self,

View file

@ -4,6 +4,7 @@ from typing import Any, Final
from fastapi import APIRouter, Depends, File, Form, Request, Response, UploadFile
from fastapi.responses import ORJSONResponse
from starlette.datastructures import UploadFile as StarletteUploadFile
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
@ -759,7 +760,14 @@ async def video_edit(
)
data: Final = await _read_request_body(request=request)
data["video_id"] = video_reference_to_id(data.pop("video", None))
uploaded_video: Final = data.pop("video", None)
if isinstance(uploaded_video, StarletteUploadFile):
video_files: Final = await batch_to_bytesio((uploaded_video,))
if video_files:
data["video"] = video_files[0]
data["video_id"] = ""
else:
data["video_id"] = video_reference_to_id(uploaded_video)
decoded: Final = decode_video_id_with_provider(data["video_id"])
provider_from_id: Final = decoded.get("custom_llm_provider")

View file

@ -277,6 +277,8 @@ def rerank(
if api_key is None:
raise ValueError("TogetherAI API key is required, please set 'TOGETHERAI_API_KEY' in your environment")
api_base = dynamic_api_base or optional_params.api_base or litellm.api_base or "https://api.together.ai/v1"
response = together_rerank.rerank(
model=model,
query=query,
@ -286,6 +288,7 @@ def rerank(
return_documents=return_documents,
max_chunks_per_doc=max_chunks_per_doc,
api_key=api_key,
api_base=api_base,
_is_async=_is_async,
)
elif _custom_llm_provider == litellm.LlmProviders.JINA_AI:

View file

@ -168,6 +168,11 @@ from litellm.router_utils.pre_call_checks.model_rate_limit_check import (
from litellm.router_utils.pre_call_checks.prompt_caching_deployment_check import (
PromptCachingDeploymentCheck,
)
from litellm.router_utils.reasoning_effort_capability import (
deployment_is_catalog_mapped,
intersect_supported_reasoning_efforts,
resolve_supported_reasoning_efforts,
)
from litellm.router_utils.router_callbacks.track_deployment_metrics import (
increment_deployment_failures_for_current_minute,
increment_deployment_successes_for_current_minute,
@ -8341,6 +8346,7 @@ class Router:
) = litellm.get_llm_provider(
model=deployment.litellm_params.model,
custom_llm_provider=deployment.litellm_params.get("custom_llm_provider", None),
api_base=deployment.litellm_params.api_base,
)
# done reading model["litellm_params"]
# Check if provider is supported: either in enum or JSON-configured
@ -9448,6 +9454,8 @@ class Router:
except Exception:
model_info = None
deployment_is_mapped = deployment_is_catalog_mapped(model_info, model_info_dict)
# get llm provider
litellm_model, llm_provider = "", ""
try:
@ -9490,6 +9498,7 @@ class Router:
"model_group": user_facing_model_group_name,
"providers": [llm_provider],
**model_info,
"supported_reasoning_efforts": None,
}
)
else:
@ -9567,6 +9576,11 @@ class Router:
if model_info.get("rpm", None) is not None and _deployment_rpm is None:
_deployment_rpm = model_info.get("rpm")
model_group_info.supported_reasoning_efforts = intersect_supported_reasoning_efforts(
model_group_info.supported_reasoning_efforts,
resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=deployment_is_mapped),
)
if _deployment_tpm is not None:
if total_tpm is None:
total_tpm = 0

View file

@ -19,7 +19,7 @@ import asyncio
import random
import re
from collections.abc import Iterator, Mapping, Sequence
from itertools import accumulate, islice
from itertools import accumulate, islice, takewhile
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast
@ -275,6 +275,7 @@ _DEFAULT_REMINDER_MARKERS: Final = ((_REMINDER_OPEN, _REMINDER_CLOSE),)
_TRUNCATION_MARKER: Final = "..."
_TRUNCATION_HEAD_FRACTION: Final = 0.3
_MIN_QUOTED_TURN_CHARS: Final = 120
_CJK_CHARACTER: Final = re.compile("[぀-ヿㇰ-ㇿ㐀-䶿一-鿿豈-﫿ヲ-ン\U00020000-\U0003ffff]")
@ -593,11 +594,40 @@ def _iter_context_turns_newest_first(
)
def _turns_within_budget(
turns: Sequence[tuple[str, str]],
budget_chars: int,
) -> tuple[tuple[str, str], ...]:
"""The newest-first turns that fit budget_chars, quoted whole wherever they fit.
Bounding the block rather than every turn in it is what lets an ordinary conversation reach the
classifier intact: a per-turn cap cuts a 785 character turn even when the whole block would have
been 353 characters, which is three orders of magnitude below anything the classifier call is
near. Once the budget does run out the older turns are dropped entire rather than shortened, so
at most one turn is ever cut and the rest read as themselves. A remainder too small to carry a
sentence buys less signal than the ellipses it would arrive wrapped in, so that turn is dropped.
The boundary turn is cut to leave room for the marker rather than to the remainder itself, so the
quoted block never exceeds budget_chars; the marker is part of what the budget buys, not an extra
charged on top of it.
"""
spent: Final = accumulate(len(text) for _, text in turns)
fitting: Final = tuple(takewhile(lambda pair: pair[1] <= budget_chars, zip(turns, spent)))
remaining: Final = budget_chars - (fitting[-1][1] if fitting else 0)
whole: Final = tuple(turn for turn, _ in fitting)
cut_to: Final = remaining - len(_TRUNCATION_MARKER)
if len(whole) == len(turns) or cut_to < _MIN_QUOTED_TURN_CHARS:
return whole
boundary_role, boundary_text = turns[len(whole)]
return (*whole, (boundary_role, _truncate(boundary_text, cut_to)))
def _extract_prior_turns(
messages: Sequence[Mapping[str, object]],
current_ask: str | None,
window_size: int,
per_turn_chars: int,
budget_chars: int,
per_turn_chars: int | None,
include_assistant: bool,
marker_pairs: tuple[tuple[str, str], ...] = _DEFAULT_REMINDER_MARKERS,
) -> tuple[tuple[str, str], ...]:
@ -612,19 +642,29 @@ def _extract_prior_turns(
window_size counts turns of every eligible role, so with assistant turns included it is the last N
of the conversation rather than the last N asks. A turn carrying only tool calls or thinking
blocks flattens to empty text and is skipped, so it never spends a slot.
Three bounds apply and the tightest wins: window_size caps how many turns, budget_chars caps the
block they form, and per_turn_chars optionally caps any single one of them before the block is
measured. They are separate because they answer separate questions, and only the block bound
tracks what the classifier call actually costs.
"""
if window_size <= 0 or not messages:
return ()
prior: Final = islice(
(
turn
for turn in _iter_context_turns_newest_first(messages, include_assistant, marker_pairs)
if turn[1] != current_ask
),
window_size,
prior: Final = tuple(
islice(
(
turn
for turn in _iter_context_turns_newest_first(messages, include_assistant, marker_pairs)
if turn[1] != current_ask
),
window_size,
)
)
return tuple((role, _truncate(text, per_turn_chars)) for role, text in reversed(tuple(prior)))
clamped: Final = (
prior if per_turn_chars is None else tuple((role, _truncate(text, per_turn_chars)) for role, text in prior)
)
return tuple(reversed(_turns_within_budget(clamped, budget_chars)))
def _decision_is_pinnable(decision: StandardLoggingRoutingDecision | None) -> bool:
@ -1363,6 +1403,7 @@ class ComplexityRouter(CustomLogger):
messages,
current_ask=prompt,
window_size=self.config.classifier_context_window_size,
budget_chars=self.config.classifier_context_budget_chars,
per_turn_chars=self.config.classifier_context_per_turn_chars,
include_assistant=include_assistant,
marker_pairs=self._reminder_markers,

View file

@ -49,7 +49,7 @@ TIER_SEVERITY_ORDER: Final[tuple[ComplexityTier, ...]] = (
DEFAULT_TIER_DISTANCE_PENALTY: Final[float] = 0.5
DEFAULT_CLASSIFIER_CONTEXT_WINDOW_SIZE: Final[int] = 3
DEFAULT_CLASSIFIER_CONTEXT_PER_TURN_CHARS: Final[int] = 200
DEFAULT_CLASSIFIER_CONTEXT_BUDGET_CHARS: Final[int] = 8000
class KeywordTierRule(BaseModel):
@ -645,12 +645,30 @@ class ComplexityRouterConfig(BaseModel):
"classifier_type is 'llm'."
),
)
classifier_context_per_turn_chars: int = Field(
default=DEFAULT_CLASSIFIER_CONTEXT_PER_TURN_CHARS,
classifier_context_budget_chars: int = Field(
default=DEFAULT_CLASSIFIER_CONTEXT_BUDGET_CHARS,
ge=0,
description=(
"Maximum characters of prior-turn text quoted to the LLM classifier, across the whole "
"context window, per classification call. Turns are taken newest first and quoted whole "
"while they fit, so a conversation small enough to quote entirely is never cut; once the "
"budget runs out the older turns are dropped whole and only the turn straddling the "
"boundary is truncated, into whatever space is left. The current ask and the caller's "
"system prompt sit outside this budget and are always sent in full, as does the numbering "
"each quoted turn carries. A budget under 120 leaves no room to quote a turn and "
"suppresses the block; set classifier_context_window_size to 0 to turn context off "
"deliberately. Only applies when classifier_type is 'llm'."
),
)
classifier_context_per_turn_chars: int | None = Field(
default=None,
gt=0,
description=(
"Maximum character length for each prior turn's text in the classifier context window. "
"Turns exceeding this are truncated. Only applies when classifier_type is 'llm'."
"Optional cap on each individual prior turn's text, applied before "
"classifier_context_budget_chars bounds the block. Unset by default, so one long turn may "
"spend the whole budget, which is usually what a follow-up needs; set it when no single "
"turn should dominate the context the classifier sees. A capped turn keeps its opening "
"and its ending with the middle elided. Only applies when classifier_type is 'llm'."
),
)
classifier_context_include_assistant_turns: bool = Field(
@ -662,9 +680,9 @@ class ComplexityRouterConfig(BaseModel):
"word 'yes'. When enabled, classifier_context_window_size counts the last N turns of the "
"conversation across both roles rather than the last N user turns, and assistant text is "
"sent to the classifier model, which may be a different deployment or provider than the "
"routed completion model. Assistant replies share classifier_context_per_turn_chars with "
"user turns, so raise it if replies are truncated before the part that carries the "
"difficulty. Off by default because enabling it shifts tier decisions, and therefore "
"routed completion model. Assistant replies spend classifier_context_budget_chars "
"alongside user turns, so raise it if the oldest turns stop being quoted once replies "
"join the window. Off by default because enabling it shifts tier decisions, and therefore "
"spend, for an already-deployed router. Only applies when classifier_type is 'llm'."
),
)

View file

@ -0,0 +1,146 @@
"""Resolve which reasoning_effort values a deployment, and by intersection a model group, accepts.
The model map's supports_*_reasoning_effort flags are the only signal, and each level's polarity
mirrors how a request path reads that same flag. medium and high are unconditional for a reasoning
model. minimal and low are opt-out: openai/chat/gpt_5_transformation.py refuses them only when the
map says false. xhigh and max are opt-in. none is opt-out everywhere except the azure gpt-5 family,
whose config raises UnsupportedParamsError without an explicit true.
xhigh is gated on the request path by the openai and azure gpt-5 configs. max is not gated there at
all: every entry carrying supports_max_reasoning_effort is Claude-family, and
anthropic/chat/transformation.py gates max on the output_config path while its reasoning_effort
path maps any level to a thinking budget. Making max opt-in is a deliberate trade, then, since an
explicit flag is the only signal that the tier is a real one rather than litellm rounding the level
to a budget, and a missing flag costs advisory metadata rather than a rejected request.
A deployment the map describes with no effort flags at all resolves to None rather than to the
opt-out defaults. 689 of the map's 854 reasoning entries carry no flag, and the o-series, xai and
bedrock nova entries among them take neither none nor minimal, so composing a set out of the
defaults alone would advertise levels those providers reject.
The advertisement order is the REASONING_EFFORT declaration order, which is presentation only. It
is not a strength scale and does not reconcile with bedrock's output_config ceiling order in
llms/bedrock/common_utils.py, which ranks max below xhigh while the thinking-budget constants rank
it above.
"""
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Final, get_args
import litellm
from litellm.types.llms.openai import REASONING_EFFORT
REASONING_EFFORT_ADVERTISEMENT_ORDER: Final = get_args(REASONING_EFFORT)
_EMPTY_ENTRY: Final[Mapping[str, object]] = MappingProxyType({})
_EFFORT_FLAGS: Final = (
("none", "supports_none_reasoning_effort"),
("minimal", "supports_minimal_reasoning_effort"),
("low", "supports_low_reasoning_effort"),
("xhigh", "supports_xhigh_reasoning_effort"),
("max", "supports_max_reasoning_effort"),
)
_OPT_OUT_EFFORTS: Final = ("minimal", "low")
_OPT_IN_EFFORTS: Final = ("xhigh", "max")
_UNCONDITIONAL_EFFORTS: Final = frozenset(("medium", "high"))
def _bare_model_entry(model_info: Mapping[str, object]) -> Mapping[str, object]:
"""The unprefixed twin of a provider-prefixed map entry, which is where the flags often live:
azure/gpt-5-mini carries none of them while gpt-5-mini carries all three. The request-path
gates resolve through the same twin (_supports_factory, #20885), so reading it here is what
keeps the advertisement and the gate on the same answer."""
key: Final = model_info.get("key")
provider: Final = model_info.get("litellm_provider")
if not isinstance(key, str) or not isinstance(provider, str) or not key.startswith(f"{provider}/"):
return _EMPTY_ENTRY
entry: Final[Mapping[str, object] | None] = litellm.model_cost.get(key.removeprefix(f"{provider}/"))
return entry if entry is not None else _EMPTY_ENTRY
def _declared_effort_flags(model_info: Mapping[str, object]) -> Mapping[str, object]:
bare: Final = _bare_model_entry(model_info)
return MappingProxyType(
{
effort: model_info.get(flag) if model_info.get(flag) is not None else bare.get(flag)
for effort, flag in _EFFORT_FLAGS
}
)
def _supports_none_reasoning_effort(model_info: Mapping[str, object], flag: object) -> bool:
"""Opt-in only where a request path refuses the level. AzureOpenAIGPT5Config raises
UnsupportedParamsError on reasoning_effort='none' without an explicit true, and it is selected
only for the gpt-5 family, so every other azure deployment keeps the opt-out default."""
if model_info.get("litellm_provider") != "azure":
return flag is not False
from litellm.llms.azure.chat.gpt_5_transformation import AzureOpenAIGPT5Config
key: Final = model_info.get("key")
if not isinstance(key, str) or not AzureOpenAIGPT5Config.is_model_gpt_5_model(key):
return flag is not False
return flag is True
def deployment_is_catalog_mapped(
resolved_model_info: Mapping[str, object] | None,
operator_model_info: Mapping[str, object],
) -> bool:
"""Whether the model map described this deployment, as opposed to the operator describing it.
Every deployment is registered in the cost map under its own id, so a mode the operator wrote
on an off-map deployment reads back here exactly like one the catalog supplied. Excluding it is
what stops such a deployment from claiming to be a known non-reasoning model and emptying the
levels its mapped siblings agree on.
"""
if resolved_model_info is None or resolved_model_info.get("mode") is None:
return False
return operator_model_info.get("mode") is None
def resolve_supported_reasoning_efforts(
model_info: Mapping[str, object],
*,
deployment_is_mapped: bool,
) -> tuple[str, ...] | None:
"""None = nothing is known about this deployment, so it must not narrow its group; () = a known
model that accepts no effort level, which correctly empties the group.
Telling those apart needs provenance the flattened ModelInfo does not carry. A deployment the
map does not describe arrives with supports_reasoning None, exactly like a mapped non-reasoning
model: 2273 of the map's 3165 entries omit the key rather than setting it false, so reading an
unset flag as () would let one custom deployment empty every level its mapped siblings agree
on. deployment_is_mapped is that provenance, and an operator who wants either answer for an
off-map deployment gets it by setting supports_reasoning explicitly.
"""
supports_reasoning: Final = model_info.get("supports_reasoning")
if supports_reasoning is not True:
return () if supports_reasoning is False or deployment_is_mapped else None
flags: Final = _declared_effort_flags(model_info)
if all(value is None for value in flags.values()):
return None
opt_out: Final = frozenset(effort for effort in _OPT_OUT_EFFORTS if flags[effort] is not False)
opt_in: Final = frozenset(effort for effort in _OPT_IN_EFFORTS if flags[effort] is True)
none_level: Final = (
frozenset(("none",)) if _supports_none_reasoning_effort(model_info, flags["none"]) else frozenset()
)
allowed: Final = opt_out | _UNCONDITIONAL_EFFORTS | opt_in | none_level
return tuple(effort for effort in REASONING_EFFORT_ADVERTISEMENT_ORDER if effort in allowed)
def intersect_supported_reasoning_efforts(
current: Sequence[str] | None,
resolved: Sequence[str] | None,
) -> tuple[str, ...] | None:
"""Deployments without metadata (None) never narrow the group; an effort survives only when
every deployment with metadata accepts it, so the group offers nothing routing could reject."""
if resolved is None:
return tuple(current) if current is not None else None
if current is None:
return tuple(resolved)
keep: Final = frozenset(current) & frozenset(resolved)
return tuple(effort for effort in REASONING_EFFORT_ADVERTISEMENT_ORDER if effort in keep)

View file

@ -685,6 +685,7 @@ ANTHROPIC_API_ONLY_HEADERS: Final = { # fails if calling anthropic on vertex ai
class AnthropicThinkingParam(TypedDict, total=False):
type: ReadOnly[Literal["enabled", "adaptive", "disabled"]]
budget_tokens: int
display: ReadOnly[Literal["summarized", "omitted"]]
class ANTHROPIC_HOSTED_TOOLS(str, Enum):

View file

@ -1,4 +1,5 @@
import json
from collections.abc import Sequence
from enum import Enum
from typing import TYPE_CHECKING, Any, Final, Literal
@ -396,7 +397,7 @@ class OutputConfigBlock(TypedDict, total=False):
class CommonRequestObject(TypedDict, total=False): # common request object across sync + async flows
additionalModelRequestFields: dict
additionalModelResponseFieldPaths: list[str]
additionalModelResponseFieldPaths: Sequence[str]
inferenceConfig: InferenceConfig
system: list[SystemContentBlock]
toolConfig: ToolConfigBlock

View file

@ -1840,7 +1840,7 @@ ResponsesAPIStreamingResponse = Annotated[
]
REASONING_EFFORT = Literal["none", "minimal", "low", "medium", "high", "xhigh"]
REASONING_EFFORT = Literal["none", "minimal", "low", "medium", "high", "xhigh", "max"]
class OpenAIRealtimeStreamSession(TypedDict, total=False):

View file

@ -637,6 +637,7 @@ class ModelGroupInfo(BaseModel):
supports_url_context: bool = Field(default=False)
supports_reasoning: bool = Field(default=False)
supports_function_calling: bool = Field(default=False)
supported_reasoning_efforts: tuple[str, ...] | None = Field(default=None)
supported_openai_params: list[str] | None = Field(default=[])
configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None

View file

@ -3855,6 +3855,7 @@ class SearchProviders(str, Enum):
TINYFISH = "tinyfish"
AGENTCORE = "agentcore"
NIMBLE = "nimble"
BING_GROUNDING = "bing_grounding"
# Create a set of all search provider values for quick lookup

View file

@ -8610,6 +8610,12 @@ class ProviderConfigManager:
)
return BedrockPassthroughConfig()
elif LlmProviders.BEDROCK_MANTLE == provider:
from litellm.llms.bedrock_mantle.passthrough.transformation import (
BedrockMantlePassthroughConfig,
)
return BedrockMantlePassthroughConfig()
elif LlmProviders.VLLM == provider or LlmProviders.HOSTED_VLLM == provider:
from litellm.llms.vllm.passthrough.transformation import (
VLLMPassthroughConfig,
@ -9096,6 +9102,7 @@ class ProviderConfigManager:
from litellm.llms.apiserpent.search.transformation import (
APISerpentSearchConfig,
)
from litellm.llms.azure.search.transformation import BingGroundingSearchConfig
from litellm.llms.bedrock.search.transformation import AgentCoreSearchConfig
from litellm.llms.brave.search.transformation import BraveSearchConfig
from litellm.llms.dataforseo.search.transformation import DataForSEOSearchConfig
@ -9137,6 +9144,7 @@ class ProviderConfigManager:
SearchProviders.TINYFISH: TinyfishSearchConfig,
SearchProviders.AGENTCORE: AgentCoreSearchConfig,
SearchProviders.NIMBLE: NimbleSearchConfig,
SearchProviders.BING_GROUNDING: BingGroundingSearchConfig,
}
config_class: Final = PROVIDER_TO_CONFIG_MAP.get(provider, None)
if config_class is None:

View file

@ -5,6 +5,8 @@ from collections.abc import Coroutine
from functools import partial
from typing import Final, Literal, overload
from httpx._types import FileContent
import litellm
from litellm.constants import DEFAULT_VIDEO_ENDPOINT_MODEL
from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT
@ -1344,6 +1346,8 @@ async def avideo_edit(
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
*,
video: FileContent | None = None,
**kwargs,
) -> VideoObject:
"""
@ -1359,6 +1363,7 @@ async def avideo_edit(
video_edit,
video_id=video_id,
prompt=prompt,
video=video,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
@ -1396,6 +1401,8 @@ def video_edit(
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
*,
video: FileContent | None = None,
**kwargs,
) -> VideoObject | Coroutine[object, object, VideoObject]:
"""
@ -1444,6 +1451,7 @@ def video_edit(
return base_llm_http_handler.video_edit_handler(
prompt=prompt,
video_id=video_id,
video_file=video,
video_provider_config=provider_config,
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,

View file

@ -17155,6 +17155,14 @@
"notes": "Web Search on Amazon Bedrock AgentCore, billed by AWS on the gateway"
}
},
"bing_grounding/search": {
"input_cost_per_query": 0.035,
"litellm_provider": "bing_grounding",
"mode": "search",
"metadata": {
"notes": "Grounding with Bing Search (G1 SKU): $35 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment."
}
},
"tinyfish/search": {
"input_cost_per_query": 0.0,
"litellm_provider": "tinyfish",
@ -49008,12 +49016,13 @@
"output_cost_per_token": 3.3e-05,
"output_cost_per_token_above_272k_tokens": 4.95e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1000000,
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"use_openai_responses_path": true,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
@ -49040,12 +49049,13 @@
"output_cost_per_token": 1.32e-05,
"output_cost_per_token_above_272k_tokens": 1.98e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1000000,
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"use_openai_responses_path": true,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
@ -49072,12 +49082,13 @@
"output_cost_per_token": 1.32e-06,
"output_cost_per_token_above_272k_tokens": 1.98e-06,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1000000,
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"use_openai_responses_path": true,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [

View file

@ -1,6 +1,6 @@
{
"ANN001": {
"limit": 3018
"limit": 3016
},
"ANN002": {
"limit": 71
@ -9,7 +9,7 @@
"limit": 827
},
"ANN201": {
"limit": 2016
"limit": 2015
},
"ANN202": {
"limit": 852
@ -57,7 +57,7 @@
"limit": 3
},
"BLE001": {
"limit": 2919
"limit": 2918
},
"C401": {
"limit": 8
@ -168,7 +168,7 @@
"limit": 3
},
"RET504": {
"limit": 176
"limit": 175
},
"RUF012": {
"limit": 240

View file

@ -0,0 +1,962 @@
#!/usr/bin/env python3
"""Ban row-rewriting DML from Prisma migrations.
Migrations run synchronously at proxy boot, before the process serves traffic, so
anything whose cost scales with existing table size turns into downtime. A single
`UPDATE` with no batching over a spend-log-sized table is minutes of unavailability
plus a doubled heap that plain autovacuum will not give back.
What is banned is the row-rewriting DML behind that, not everything whose cost
scales that way. A non-concurrent `CREATE INDEX`, an `ALTER COLUMN ... TYPE` that is
not binary coercible, a volatile `DEFAULT` on a new column, a `CREATE TABLE ... AS
SELECT` or `SELECT ... INTO` filling a new table from an existing one, the rename
that pairs with one of those to swap a table out, and a `REFRESH MATERIALIZED VIEW`
all read the whole table and all pass. That is deliberate: a rule wide enough to
reach them fires on most ordinary migrations, and a marker everyone adds by reflex
stops carrying information. The outage this was written for was a backfill.
Flagged, per statement, by its leading keyword:
UPDATE rewrites every matching row, and `WHERE` does not bound the scan
DELETE same scan, and the dead tuples outlive the migration
MERGE both of the above in one statement
INSERT only when its rows come from a query rather than a literal `VALUES`
list. The query counts wherever it sits, since Postgres takes it
parenthesised, and `TABLE t` is one as much as a `SELECT` is. An
insert bounded by a `VALUES` list passes, written bare or in
parentheses, and so do the scalar subqueries in that list and the
`RETURNING` and `ON CONFLICT` clauses written after it, none of which
supply the rows. A `VALUES` reached through a subquery or joined to a
query by a set operation bounds nothing
WITH a CTE-led statement containing any of the above. An `INSERT` is read
against the part of the statement holding it, so a writable CTE
bounded by its own `VALUES` list is not handed the query the statement
ends with as the rows it copies
Referential actions (`ON DELETE CASCADE`, `ON UPDATE CASCADE`) are schema, never a
statement's leading keyword, so they pass.
A statement wrapped in `EXPLAIN` is judged on the statement itself, because the
`ANALYZE` form runs it rather than only planning it, and a rewrite left under one
rewrites the table on the way to printing its timings. Explaining a rewrite without
`ANALYZE` is flagged too: nothing here needs the plan of a statement it is being
told not to run at boot, and a marker is a cheap answer if one ever does.
Statements inside dollar-quoted bodies are scanned too. `DO $$ ... $$` is this
repo's idiom for conditional DDL, so a body is where an `UPDATE` would otherwise
hide. A `CREATE FUNCTION` or `CREATE PROCEDURE` body is the exception, because
defining a routine only stores it: that body is read when the same migration names
the routine somewhere else, which is what defining a backfill and then running it
looks like, and left alone when nothing calls it. A routine whose name needed
quoting is read either way, since quoting is blanked at the call sites too and a
call written there could never be found. The SQL an `EXECUTE` runs is scanned the same way, since a rewrite reads the
same to Postgres whether it is spelled out or handed over as a string, and so is a
literal parked in a variable some `EXECUTE` in the same body then runs by name,
however it got there: an assignment with `:=`, the bare `=` PL/pgSQL takes as the
same operator, a query returning it through `INTO`, or a loop walking the query it
came out of. So is the body of a `DO` written in single quotes rather than dollar
quotes. A literal nothing runs is text, however much it reads like a statement, so
an error message naming a `DELETE` the application handles stays a message.
Each literal is read on its own, so a keyword built by concatenating fragments that
do not contain it (`'UPD' || 'ATE ...'`) is not caught. Every fragment is scanned,
so a concatenation is caught wherever the keyword survives whole in one of them,
which covers `'UPDATE ' || quote_ident(t)` and the rest of the readable shapes. The
gap needs a keyword deliberately split down the middle, and this check is a guard
against a rewrite reaching a boot unnoticed, not a defence against someone hiding
one on purpose.
Line numbers always count against the whole migration file, however deeply the
statement is nested, so a reported line points at the statement and the markers
below line up with the statements they exempt.
Add a column and let the application populate it, or run the rewrite as an opt-in
batched job outside boot. When a rewrite is genuinely bounded and must ship inside
the migration, put `-- data-migration-ok: <reason>` on the statement or on the line
above it, naming what bounds it. The reason is required. A marker sharing a line
with the statement it follows exempts that statement alone, so the next statement
down is still checked rather than picking the marker up as its own. A marker on an
`EXECUTE` or on the assignment feeding one covers the single-quoted SQL that
statement hands off, so it goes where the migration reads rather than inside the
string. A dollar-quoted payload is not a string to this check but a region read like
any other body, so a rewrite inside one takes its marker on the rewrite itself. That
placement is deliberate rather than an oversight: a marker covering a whole body
would let one written for a `DO` block silence a rewrite added to that block later.
`GRANDFATHERED` freezes the violations that predate this check. Prisma records a
checksum for every applied migration and this repo treats applied files as
immutable, so those two cannot take an inline marker. The set is closed; a new
migration belongs nowhere in it.
"""
from __future__ import annotations
import re
import sys
from collections.abc import Iterator, Mapping
from dataclasses import dataclass
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[2]
MIGRATIONS_DIR = REPO_ROOT / "litellm-proxy-extras" / "litellm_proxy_extras" / "migrations"
GRANDFATHERED = frozenset(
{
"20260817000000_shadow_eval_multi_key",
"20260818224500_add_shadow_eval_stopped_by",
}
)
MARKER = re.compile(r"--[ \t]*data-migration-ok:[ \t]*(\S.*?)[ \t]*$", re.MULTILINE)
DOLLAR_TAG = re.compile(r"\$(?:[A-Za-z_][A-Za-z0-9_]*)?\$")
FIRST_WORD = re.compile(r"[A-Za-z_][A-Za-z0-9_]*")
STATEMENT = re.compile(r"[^;]+")
RUN_BY_NAME = re.compile(r"\bEXECUTE\s+([A-Za-z_][A-Za-z0-9_]*)", re.IGNORECASE)
INTO_TARGETS = re.compile(
r"\bINTO\s+(?:STRICT\s+)?"
r"([A-Za-z_][A-Za-z0-9_]*(?:\s*,\s*[A-Za-z_][A-Za-z0-9_]*)*)",
re.IGNORECASE,
)
LOOP_TARGET = re.compile(r"\bFOR(?:EACH)?\s+([A-Za-z_][A-Za-z0-9_]*)\s+IN\b", re.IGNORECASE)
LOOP_HEADER = re.compile(r"\bFOR(?:EACH)?\b.*?\bLOOP\b", re.IGNORECASE | re.DOTALL)
WORD_OR_ASSIGN = re.compile(r"[A-Za-z_][A-Za-z0-9_]*|:=|(?<![<>!:=])=(?![=>])")
PRECEDING_WORD = re.compile(r"([A-Za-z_][A-Za-z0-9_]*)[^A-Za-z0-9_]*$")
QUALIFIER_GAP = re.compile(r"[\s.]*")
EXPLAIN_OPTIONS = re.compile(r"\bEXPLAIN\b(?:\s+(?:ANALYZE|ANALYSE|VERBOSE)\b)+", re.IGNORECASE)
DEFINES_A_ROUTINE = re.compile(
r"\bCREATE\b(?:\s+OR\s+REPLACE)?\s+(?:FUNCTION|PROCEDURE)\b", re.IGNORECASE
)
QUALIFIED_NAME = r"(?:\"[^\"]*\"|[A-Za-z_][A-Za-z0-9_$]*)"
ROUTINE_NAME = re.compile(rf"\s*(?:{QUALIFIED_NAME}\s*\.\s*)?({QUALIFIED_NAME})")
OPENS_A_CALL = re.compile(r"\s*\(")
NAMES_AN_INDEX = re.compile(r"\bCREATE\b.+\bINDEX\b", re.IGNORECASE | re.DOTALL)
INTRODUCES_A_RELATION = frozenset({"TABLE", "INTO", "REFERENCES", "EXISTS", "COPY"})
REWRITES_ROWS = frozenset({"UPDATE", "DELETE", "MERGE"})
JOINS_QUERIES = ("UNION", "INTERSECT", "EXCEPT")
SET_OPERATION = re.compile(rf"\b(?:{'|'.join(JOINS_QUERIES)})\b", re.IGNORECASE)
STATEMENT_KEYWORDS = REWRITES_ROWS | frozenset(
{
"INSERT",
"SELECT",
"WITH",
"ALTER",
"CREATE",
"DROP",
"TRUNCATE",
"COMMENT",
"GRANT",
"REVOKE",
"COPY",
"SET",
"PERFORM",
"RAISE",
"RETURN",
"EXECUTE",
"DO",
"CALL",
"REINDEX",
"REFRESH",
"VACUUM",
"ANALYZE",
}
)
GUARDS_A_CONDITION = frozenset({"IF", "ELSIF", "ELSEIF", "CASE", "WHEN", "WHILE", "EXIT", "ASSERT"})
OPENS_A_BLOCK = frozenset({"BEGIN", "THEN", "ELSE", "LOOP"})
NEVER_A_VARIABLE = frozenset({"INTO", "USING"})
BIND_VALUES = re.compile(r"\bUSING\b", re.IGNORECASE)
WRITES_ROWS = re.compile(r"\bINSERT\b", re.IGNORECASE)
GUIDANCE = """
Migrations apply at proxy boot, before it serves traffic, so a statement whose cost
scales with table size is downtime. Add the column and let the application backfill
it, or move the rewrite to a batched job outside boot.
If the rewrite is genuinely bounded and has to ship in the migration, mark the
statement with the bound spelled out:
-- data-migration-ok: <what bounds this>
UPDATE ...
"""
@dataclass(frozen=True, slots=True)
class Violation:
migration: str
line: int
keyword: str
def render(self) -> str:
location = f"{MIGRATIONS_DIR.relative_to(REPO_ROOT)}/{self.migration}/migration.sql"
return f"{location}:{self.line}: {self.keyword} rewrites existing rows at boot"
@dataclass(frozen=True, slots=True)
class Marker:
start: int
end: int
standalone: bool
@dataclass(frozen=True, slots=True)
class Markers:
sql: str
written: tuple[Marker, ...]
def exempt(self, start: int, end: int) -> bool:
"""Whether the statement spanning `start` to `end` carries a marker."""
return any(self.speaks_for(marker, start, end) for marker in self.written)
def speaks_for(self, marker: Marker, start: int, end: int) -> bool:
"""Whether a marker is written against this statement. One alone on its line speaks for
the statement below it, which is how a marker written above a rewrite exempts it, and one
sharing its line with code speaks for the statement it follows. Either is matched by where
it sits rather than by the line it lands on, so a second statement sharing that line does
not inherit the exemption. A marker inside a statement speaks for it whichever kind it is,
which is how one on the opening line of a long statement still covers the whole of it."""
if start <= marker.start < end:
return True
if marker.standalone:
return self.on_the_line_below(marker.end, start)
return self.only_separators(end, marker.start)
def on_the_line_below(self, start: int, end: int) -> bool:
"""Whether a marker on its own line is written directly above the statement, which means
one line break and nothing else that carries meaning. A blank line between the two leaves
the marker reading as a note about the file rather than a bound on what follows it."""
return self.only_separators(start, end) and self.sql[start:end].count("\n") == 1
def only_separators(self, start: int, end: int) -> bool:
"""Whether nothing but statement separators lie between two points, which is what makes a
marker and the statement it follows adjacent however they are laid out."""
return start <= end and not self.sql[start:end].strip(" \t\r\n;")
def blank(text: str) -> str:
return "".join(character if character == "\n" else " " for character in text)
def undouble(literal: str) -> str:
"""The SQL a single-quoted literal stands for, with each doubled quote read back as the one it
escapes. `mask` hands the literal on raw, `''` and all, so re-lexing it as SQL needs the escapes
resolved first: left doubled, the first quote of a pair opens an empty string and closes it on
the second, and a `--` or `/*` in what was a nested string is then bare and blanks the code
after it."""
return literal.replace("''", "'")
def defuse_escapes(literal: str) -> str:
"""The literal made safe to re-lex without moving anything: each doubled quote becomes a real
quote and a space, so a `--` or `/*` in a nested string stays inside its string the way
`undouble` achieves it, while the pair keeps its two characters. Every newline and every
character after a resolved escape then holds the offset it had in the document, so a rewrite
scanned out of the literal reports its true file line and lines up with the file's markers,
which `undouble` cannot promise because it shrinks the text as it collapses each pair."""
return literal.replace("''", "' ")
def mask(
sql: str,
) -> tuple[str, tuple[tuple[int, int], ...], tuple[tuple[int, int], ...], tuple[tuple[int, int], ...]]:
"""Blank comments and quoted text, keeping offsets, and locate the spans that can still
hold SQL: dollar-quoted bodies, and the single-quoted literals `EXECUTE` runs. Also locate
the double-quoted identifiers that open a call (`"backfill"(`), so a routine invoked through
one can be found by name even though the call is blanked here the way every other quoted run
of text is. Whether an identifier opens a call is read from the masked text rather than the
raw SQL, so a comment sitting between the name and its parenthesis, blanked to spaces here, is
skipped exactly as whitespace is. A double-quoted identifier that opens no call, a column,
index, or constraint name, is left out, so it never masquerades as a call to a like-named
routine, as is one whose parenthesis is a column list rather than an argument list, the table
of a `CREATE TABLE`, `INSERT INTO`, `REFERENCES`, `COPY`, or `CREATE INDEX`, which
`names_a_relation` reads from the word before the name."""
chunks: list[str] = []
bodies: list[tuple[int, int]] = []
literals: list[tuple[int, int]] = []
identifiers: list[tuple[int, int]] = []
index = 0
length = len(sql)
while index < length:
pair = sql[index : index + 2]
if pair == "--":
stop = sql.find("\n", index)
stop = length if stop == -1 else stop
chunks.append(blank(sql[index:stop]))
index = stop
continue
if pair == "/*":
stop = skip_block_comment(sql, index)
chunks.append(blank(sql[index:stop]))
index = stop
continue
character = sql[index]
if character in "'\"":
stop = skip_quoted(sql, index, character)
if character == "'":
closed = sql[stop - 1 : stop] == character
literals.append((index + 1, max(index + 1, stop - 1 if closed else stop)))
else:
identifiers.append((index, stop))
chunks.append(blank(sql[index:stop]))
index = stop
continue
if character == "$":
tag = DOLLAR_TAG.match(sql, index)
if tag is not None:
closing = sql.find(tag.group(), tag.end())
body_end = length if closing == -1 else closing
stop = length if closing == -1 else closing + len(tag.group())
bodies.append((tag.end(), body_end))
chunks.append(blank(sql[index:stop]))
index = stop
continue
chunks.append(character)
index += 1
masked = "".join(chunks)
calls = tuple(
(start, end)
for start, end in identifiers
if OPENS_A_CALL.match(masked, end) and not names_a_relation(masked[:start])
)
return masked, tuple(bodies), tuple(literals), calls
def skip_block_comment(sql: str, start: int) -> int:
depth = 1
index = start + 2
while index < len(sql) and depth > 0:
pair = sql[index : index + 2]
if pair == "/*":
depth += 1
index += 2
elif pair == "*/":
depth -= 1
index += 2
else:
index += 1
return index
def skip_quoted(sql: str, start: int, quote: str) -> int:
"""One quoted run, up to and including its closing quote. A doubled quote is an escaped
quote sitting inside the run rather than the end of it. Closing on the first and reopening
on the second would mask the same span, which is why this looked like it needed no special
case, but the run is also handed on whole as one literal, and splitting it there offers the
tail of a string to be read as SQL in its own right."""
index = start + 1
while True:
stop = sql.find(quote, index)
if stop == -1:
return len(sql)
if sql[stop + 1 : stop + 2] == quote:
index = stop + 2
continue
return stop + 1
def strip_parens(statement: str) -> str:
"""Blank parenthesised groups in place, so an `IF EXISTS (SELECT ...)` guard does not
stand in for the statement it guards."""
chunks: list[str] = []
depth = 0
for character in statement:
if character == "(":
depth += 1
chunks.append(" ")
elif character == ")":
depth = max(depth - 1, 0)
chunks.append(" ")
elif depth > 0 and character != "\n":
chunks.append(" ")
else:
chunks.append(character)
return "".join(chunks)
def strip_explain(statement: str) -> str:
"""Blank an `EXPLAIN` written with bare options, since the `ANALYZE` among them would
otherwise stand in for the keyword of the statement being explained. That statement is
the one worth reading: `EXPLAIN ANALYZE` runs it rather than only planning it, so a
rewrite underneath rewrites the table for real. The parenthesised option list needs
nothing here, already being blanked as a group."""
return EXPLAIN_OPTIONS.sub(lambda match: blank(match.group()), statement)
def leading_keyword(statement: str) -> re.Match[str] | None:
"""The statement's own keyword, looking past what wraps it: a parenthesised guard,
PL/pgSQL block syntax such as `BEGIN`, `IF ... THEN` and `END`, and an `EXPLAIN`.
Offsets survive both strips, so the match still points into `statement` itself."""
return next(
(
word
for word in FIRST_WORD.finditer(strip_explain(strip_parens(statement)))
if word.group().upper() in STATEMENT_KEYWORDS
),
None,
)
def offending_keyword(statement: str) -> str | None:
word = leading_keyword(statement)
if word is None:
return None
keyword = word.group().upper()
if keyword in REWRITES_ROWS:
return keyword
if keyword == "INSERT":
source = row_source_keyword(statement)
return None if source is None else f"INSERT ... {source}"
if keyword == "WITH":
nested = next((name for name in sorted(REWRITES_ROWS) if contains(statement, name)), None)
if nested is not None:
return f"WITH ... {nested}"
if contains(statement, "INSERT"):
source = insert_row_source(statement)
if source is not None:
return f"WITH ... INSERT ... {source}"
return None
def insert_row_source(statement: str) -> str | None:
"""Which keyword supplies the rows to an `INSERT` written somewhere inside a `WITH`
statement. Only the parts that hold that insert are read, because a writable CTE sits
beside the query the statement ends with and reading the whole thing hands the insert
the outer `SELECT` as its row source: `WITH c AS (INSERT ... VALUES (1) RETURNING "x")
SELECT * FROM c` adds one literal row and copies nothing. A CTE keeps its insert in a
parenthesised group, and the statement's own insert, if it is the one writing, runs from
the keyword to the end, found in the text outside every parenthesis so a group's insert
is not counted twice."""
inserts = [group for group in parenthesised_groups(statement) if contains(group, "INSERT")]
written = WRITES_ROWS.search(strip_parens(statement))
if written is not None:
inserts.append(statement[written.start() :])
sources = (row_source_keyword(insert) for insert in inserts)
return next((source for source in sources if source is not None), None)
def row_source_keyword(statement: str) -> str | None:
"""Which keyword supplies an `INSERT` its rows, or `None` when a literal `VALUES` list
does. A query outside every parenthesis is the row source outright. Failing that, a
set operation at that same level joins several terms, and the insert is a rewrite when
any one of them is a query, so each term is read on its own rather than the statement
read whole. Failing that, a `VALUES` outside every parenthesis is itself the row source,
so the scalar subqueries and helper CTEs nested within that list do not make the insert
a rewrite. Failing all three, the rows come from a parenthesised group, which Postgres
accepts and which reading only the unparenthesised text would let through:
`INSERT INTO "t" ("a") (SELECT ...)` copies a whole table. Each group at that level is
read on its own terms until one of them supplies the rows, since the ones before it are
the column list and the ones after it are the conflict target and the rest of the clauses
an insert is allowed to carry. A wrapped `VALUES` list is the row source as much as a
wrapped query is, so it ends the search rather than being skipped over: reading past it
reaches a `RETURNING (SELECT ...)` or a `DO UPDATE SET "a" = (SELECT ...)` written after
it and calls that scalar subquery the rows the insert copies. The group is read on its
own terms before it is allowed to end the search, because a `VALUES` list joined to a
query by a set operation inside the group supplies every row the query does, and
stopping on the word `VALUES` alone would pass the whole copy."""
outer = strip_parens(statement)
joined = row_source_in(outer)
if joined is not None:
return joined
if SET_OPERATION.search(outer):
sources = (row_source_keyword(term) for term in set_operation_terms(statement, outer))
return next((source for source in sources if source is not None), None)
if contains(outer, "VALUES"):
return None
groups = list(parenthesised_groups(statement))
if not groups:
return row_source_in(statement)
for group in groups:
source = row_source_keyword(group)
if source is not None:
return source
if contains(strip_parens(group), "VALUES"):
return None
return None
def set_operation_terms(statement: str, outer: str) -> Iterator[str]:
"""The terms a top-level set operation joins. The operators are read from the text outside
every parenthesis, which `strip_parens` blanks in place rather than removing, so their
offsets are offsets into the statement itself and each term comes back from the original
text with its own parentheses intact. Reading them at that level is what keeps a set
operation written inside a `VALUES` list from cutting the list in half. An `ALL` or a
`DISTINCT` stays at the head of the term that follows, where it names no row source and
so reads as nothing."""
edges = [0]
for operation in SET_OPERATION.finditer(outer):
edges += [operation.start(), operation.end()]
edges.append(len(statement))
for opens, closes in zip(edges[::2], edges[1::2]):
yield statement[opens:closes]
def parenthesised_groups(statement: str) -> Iterator[str]:
"""What each group of parentheses closed at the statement's outermost level holds, in the
order they are written. One of them is where an `INSERT` keeps a row source it has
wrapped, since Postgres takes `INSERT INTO "t" ("a") (SELECT ...)` and `... (VALUES (1))`
alike, and reading a group on its own terms is what stops a scalar subquery nested inside
a wrapped `VALUES` list standing in for the rows."""
depth = 0
opens = None
for index, character in enumerate(statement):
if character == "(":
if depth == 0:
opens = index
depth += 1
elif character == ")":
depth = max(depth - 1, 0)
if depth == 0 and opens is not None:
yield statement[opens + 1 : index]
def row_source_in(text: str) -> str | None:
return next((word for word in ("SELECT", "TABLE") if contains(text, word)), None)
def hands_off_sql(statement: str, executed: frozenset[str]) -> bool:
"""Whether a statement gives the server a string literal to run as SQL. `EXECUTE` runs one
outright, and so does `DO`, whose body is a string wherever it is not dollar-quoted. An
assignment parks one in a variable, which counts only when something further down runs
that variable by name, since a string the migration never executes is text."""
if leads_with(statement, "EXECUTE") or leads_with(statement, "DO"):
return True
return bool(assigned_names(statement) & executed)
def assigned_names(statement: str) -> frozenset[str]:
"""The candidate variable names a statement writes to. An assignment is read as every
word ahead of its operator, since a declaration carries its type and sometimes a leading
`DECLARE` alongside the name, and none of that is worth parsing when the only question
is which name is executed. A query assigns through the target list after its `INTO`
instead, and a loop through the variable it walks its query with, which is how a rewrite
reaches a variable with no operator appearing at all."""
names = {word.lower() for word in assignment_reach(statement)}
for targets in INTO_TARGETS.finditer(statement):
if names_a_table(statement[: targets.start()]):
continue
names.update(word.group().lower() for word in FIRST_WORD.finditer(targets.group(1)))
names.update(loop.group(1).lower() for loop in LOOP_TARGET.finditer(statement))
return frozenset(names)
def names_a_table(before: str) -> bool:
"""Whether the `INTO` this text runs up to introduces a table rather than a query's
target list. `INSERT INTO` is the one that does, and reading its table as somewhere a
string was parked would have an insert scanned for the SQL its own literals spell out.
An `INSERT` that really does assign reaches its `INTO` through a `RETURNING` list, so
the word immediately before is what separates the two."""
word = PRECEDING_WORD.search(before)
return word is not None and word.group(1).upper() == "INSERT"
def names_a_relation(before: str) -> bool:
"""Whether the parenthesised quoted identifier this text runs up to names a table with a
column list rather than opening a routine call. The two look alike, a name then a `(`, so
an uncalled routine sharing a name with a table would otherwise read as called. The word
immediately before tells most of them apart: `CREATE TABLE`, `INSERT INTO`, a foreign key's
`REFERENCES`, `CREATE TABLE IF NOT EXISTS`, and `COPY` each put a table there, and none can
precede a call. `ON` is the ambiguous one, since it introduces the table of a `CREATE INDEX`
but also a join condition that may itself be a call, so it counts only inside a statement
that creates an index, leaving `JOIN ... ON f()` and an index predicate's `WHERE f()` as
calls. A bare schema qualifier is read through: `INSERT INTO public."Foo"` parks the table's
introducing word a hop back behind `public.`, so any word ahead of the name that a dot follows,
touching or spaced as `public . "Foo"`, is the qualifier and the one before it decides. The
introducing word is settled before that, so a quoted schema, which blanks to spaces and leaves
`INTO` itself as the word ahead of the name however the dot is spaced, still reads as a relation,
while a genuine `SELECT public."f"()` reads through its qualifier to the `SELECT` and stays a call.
A word only introduces the name when nothing but whitespace and qualifier dots lies between them,
so a `(` in that gap keeps it from reaching across: a schema-qualified call inside a `CREATE INDEX`
expression, `ON "Foo" (public."f"(col))`, leaves `ON` behind the paren and the call stays a call."""
word = PRECEDING_WORD.search(before)
if word is None:
return False
gap = before[word.end(1) :]
if QUALIFIER_GAP.fullmatch(gap):
keyword = word.group(1).upper()
if keyword in INTRODUCES_A_RELATION:
return True
if keyword == "ON":
return NAMES_AN_INDEX.search(before[before.rfind(";") + 1 :]) is not None
if "." in gap:
return names_a_relation(before[: word.start(1)])
return False
def assignment_reach(statement: str) -> tuple[str, ...]:
"""The words the statement's assignment is reached through, empty where it holds none.
PL/pgSQL spells the operator `:=` and takes a bare `=` as the same thing, so both count,
the second only where none of the words reached so far `marks_a_comparison`. The search
stops at the first operator that reads as an assignment, because a statement holds one
at most and everything after it is the expression being assigned, where an `=` only ever
compares: that is what keeps `ok := stmt = '<sql>'` from reading as a write to `stmt`.
What comes before can still be a comparison the assignment sits behind, as in
`IF n = 1 THEN stmt = '<sql>'`, and a word opening a block ends what it is reached
through, since nothing ahead of the `THEN` describes what follows it."""
reached: list[str] = []
compares = False
for token in WORD_OR_ASSIGN.finditer(statement):
word = token.group().upper()
if word == ":=":
return tuple(reached)
if word == "=":
if not compares:
return tuple(reached)
continue
if word in OPENS_A_BLOCK:
reached.clear()
compares = False
continue
reached.append(word)
compares = compares or marks_a_comparison(word)
return ()
def marks_a_comparison(word: str) -> bool:
"""Whether reaching a bare `=` through this word means the operator tests a variable
rather than writing one. These are all that tell the two apart: an assignment is reached
with a name and perhaps a type, while a comparison is reached either through a statement
carrying its own keyword or through a word that guards a condition."""
return word in STATEMENT_KEYWORDS or word in GUARDS_A_CONDITION
def executed_names(masked: str) -> frozenset[str]:
"""The variables handed to an `EXECUTE` by name. Reading these off the masked text keeps
an `EXECUTE` written inside a comment or a string from counting. Masking blanks a literal
in place rather than removing it, so `EXECUTE '...'` leaves whatever follows the literal
looking like the name being run. Only `INTO` and `USING` can sit there, since the syntax
allows nothing else between an `EXECUTE` and the semicolon ending it, and neither is ever
a variable, so both are dropped rather than left to collide with a query reaching one."""
return frozenset(
match.group(1).lower()
for match in RUN_BY_NAME.finditer(masked)
if match.group(1).upper() not in NEVER_A_VARIABLE
)
def leads_with(statement: str, keyword: str) -> bool:
word = leading_keyword(statement)
return word is not None and word.group().upper() == keyword
def contains(statement: str, keyword: str) -> bool:
return re.search(rf"\b{keyword}\b", statement, re.IGNORECASE) is not None
def read_markers(sql: str) -> Markers:
return Markers(
sql,
tuple(
Marker(match.start(), match.end(), alone_on_its_line(sql, match.start()))
for match in MARKER.finditer(sql)
),
)
def alone_on_its_line(sql: str, start: int) -> bool:
return not sql[sql.rfind("\n", 0, start) + 1 : start].strip()
def scan(sql: str, migration: str, markers: Markers) -> Iterator[Violation]:
yield from scan_region(sql, sql, migration, markers, 0)
def scan_region(
document: str, region: str, migration: str, markers: Markers, offset: int
) -> Iterator[Violation]:
"""Violations in one region of `document`, whose text begins at `offset`. Positions are
always counted against the whole document, so a statement nested in a dollar-quoted body
reports its real file line and lines up with the markers read from that file. A single-quoted
literal that `DO` or `EXECUTE` runs as SQL has each doubled quote turned into a quote and a space
before it is scanned, so a `--` or `/*` in one of its nested strings blanks nothing and the
statement after it stays visible, and since that keeps every character on its offset, the
statement reports its true file line and lines up with the markers."""
masked, bodies, literals, identifiers = mask(region)
executed = executed_names(masked)
runnable = executed_literals(masked, literals, executed)
for match in STATEMENT.finditer(masked):
exempt = markers.exempt(offset + statement_start(match), offset + match.end())
for clause, base in clauses(match.group(), match.start()):
if hands_off_sql(clause, executed) and not exempt:
commands_end = base + bind_values_start(clause)
for start, end in literals:
if base <= start and end <= commands_end:
yield from scan_region(
document,
defuse_escapes(region[start:end]),
migration,
markers,
offset + start,
)
keyword = offending_keyword(clause)
if keyword is None or exempt:
continue
yield Violation(migration, line_of(document, offset + keyword_start(clause, base)), keyword)
for body in bodies:
if not runs_when_applied(masked, region, bodies, runnable, identifiers, body):
continue
start, end = body
yield from scan_region(document, region[start:end], migration, markers, offset + start)
def executed_literals(
masked: str, literals: tuple[tuple[int, int], ...], executed: frozenset[str]
) -> tuple[tuple[int, int], ...]:
"""The single-quoted literals a region runs as SQL, where a call to a routine the same
migration defines is as real as one written in the open. `DO '...'` runs its body and
`EXECUTE` runs the string it is handed, so a definition named inside one of those is called,
while a name in a message string or any literal nothing executes stays text. These are the
spans the direct scan already recurses into, read here so a call written in one is found when
the migration is searched for the routine's name."""
return tuple(
(start, end)
for match in STATEMENT.finditer(masked)
for clause, base in clauses(match.group(), match.start())
if hands_off_sql(clause, executed)
for start, end in literals
if base <= start and end <= base + bind_values_start(clause)
)
def runs_when_applied(
masked: str,
region: str,
bodies: tuple[tuple[int, int], ...],
runnable: tuple[tuple[int, int], ...],
identifiers: tuple[tuple[int, int], ...],
body: tuple[int, int],
) -> bool:
"""Whether a dollar-quoted body runs while the migration is being applied. A `DO` block runs
where it is written, and so does every other use of this quoting. A `CREATE FUNCTION` or a
`CREATE PROCEDURE` only stores its body, which runs when something calls the routine, so a
definition nothing calls rewrites no rows at boot and reporting it names a line that never
executes. Skipping every definition instead would let a migration define a backfill and then
run it unseen, which is the shape this check exists to catch, so the body is read whenever
the same migration names the routine anywhere outside the definition. The definition is
found in the masked text, where one written inside a comment has already been blanked, and
the name is read from the region at those same offsets, since masking blanks a quoted
identifier in place. A call written as a quoted identifier is blanked there too, and
`\"backfill\"()` is the same call as `backfill()` in Postgres, so the double-quoted call sites
are put back before the search and a routine invoked through one is found. A quoted name that
opens no call, a column or table sharing the routine's name, stays blanked and cannot be read
as a call it never makes. A definition whose
own name needs those quotes is read rather than trusted, since matching such a name once it is
put back in the open would be unreliable."""
start, end = body
opens = masked.rfind(";", 0, start) + 1
defined = DEFINES_A_ROUTINE.search(masked, opens, start)
if defined is None:
return True
named = ROUTINE_NAME.match(region, defined.end(), start)
if named is None or named.group(1).startswith('"'):
return True
restored = outside_definition(masked, region, bodies, runnable, identifiers, opens, end)
return contains(restored, re.escape(named.group(1)))
def outside_definition(
masked: str,
region: str,
bodies: tuple[tuple[int, int], ...],
runnable: tuple[tuple[int, int], ...],
identifiers: tuple[tuple[int, int], ...],
opens: int,
closes: int,
) -> str:
"""The migration's text with one routine definition blanked out and every runnable body put
back: the dollar-quoted bodies and the single-quoted literals `DO` and `EXECUTE` run as SQL.
Masking blanks all of them alike, and a `DO` block, dollar-quoted or single-quoted, is the
ordinary way a migration runs a routine it has just defined, so a call written inside one has
to stay readable. Each comes back with its comments blanked, since a name written in a comment
is documentation rather than a call, while its string literals stay readable because `EXECUTE`
runs one as SQL and the call can be written inside it. A single-quoted payload is undoubled as
it goes back, so a `--` or `/*` in one of its nested strings blanks nothing and the call after
it stays visible, and it is padded to the span it fills so the later offsets still land. The
double-quoted call sites come back verbatim, so a routine invoked as `\"backfill\"()` reads as
the call it is, while a like-named identifier that opens no call was never collected and stays
blanked. The definition is blanked after they are restored, which takes its own body and
any identifier standing inside it with it, so a routine that names itself recursively does not
thereby count as called."""
text = list(masked)
for start, end in bodies:
text[start:end] = without_comments(region[start:end])
for start, end in runnable:
text[start:end] = without_comments(undouble(region[start:end])).ljust(end - start)
for start, end in identifiers:
text[start:end] = region[start:end]
text[opens:closes] = blank(region[opens:closes])
return "".join(text)
def without_comments(sql: str) -> str:
"""The text with its comments blanked in place and everything else kept, read with the same
lexing as `mask` so a `--` inside a string literal blanks nothing. A dollar-quoted body
nested within is read the same way on its own, which keeps a stray quote inside it from
reaching past its closing tag."""
chunks: list[str] = []
index = 0
length = len(sql)
while index < length:
pair = sql[index : index + 2]
if pair == "--":
stop = sql.find("\n", index)
stop = length if stop == -1 else stop
chunks.append(blank(sql[index:stop]))
index = stop
continue
if pair == "/*":
stop = skip_block_comment(sql, index)
chunks.append(blank(sql[index:stop]))
index = stop
continue
character = sql[index]
if character in "'\"":
stop = skip_quoted(sql, index, character)
chunks.append(sql[index:stop])
index = stop
continue
if character == "$":
tag = DOLLAR_TAG.match(sql, index)
if tag is not None:
closing = sql.find(tag.group(), tag.end())
body_end = length if closing == -1 else closing
stop = length if closing == -1 else closing + len(tag.group())
chunks.append(sql[index : tag.end()])
chunks.append(without_comments(sql[tag.end() : body_end]))
chunks.append(sql[body_end:stop])
index = stop
continue
chunks.append(character)
index += 1
return "".join(chunks)
def clauses(statement: str, start: int) -> Iterator[tuple[str, int]]:
"""The statements written inside one semicolon-delimited run, each with where it begins. A
`FOR ... LOOP` header takes no semicolon of its own, so the first statement of the loop body
is written into the same run, and reading the pair as one statement lets the header's row
source stand in as the keyword for both. That hides the statement the loop repeats, which is
the shape a row-by-row backfill takes. Splitting after each header, nested ones included,
reads the header and the body as the separate statements Postgres runs them as."""
edges = (0, *(header.end() for header in LOOP_HEADER.finditer(statement)), len(statement))
for opens, closes in zip(edges, edges[1:]):
if opens < closes:
yield statement[opens:closes], start + opens
def bind_values_start(statement: str) -> int:
"""Where a statement stops handing commands to the server and starts listing bind values.
The expressions after `USING` are values substituted into the command, never commands in
their own right, so one that merely spells out a rewrite is not running it. Read off the
masked text, so a `USING` written inside the command string is not mistaken for this one,
and only once the parentheses have closed, so that the `USING` of a `JOIN` in a subquery
that helps build the command does not cut the command short and hide the rest of it."""
for keyword in BIND_VALUES.finditer(statement):
preceding = statement[: keyword.start()]
if preceding.count("(") == preceding.count(")"):
return keyword.start()
return len(statement)
def statement_start(statement: re.Match[str]) -> int:
"""Where the statement's own text begins, past the whitespace and blanked comments it picked
up from whatever sat between it and the statement before it, one of which can be a marker."""
text = statement.group()
return statement.start() + len(text) - len(text.lstrip())
def keyword_start(clause: str, base: int) -> int:
word = leading_keyword(clause)
return base + (0 if word is None else word.start())
def line_of(sql: str, offset: int) -> int:
return sql.count("\n", 0, offset) + 1
def scan_migration(directory: Path) -> tuple[Violation, ...]:
sql = (directory / "migration.sql").read_text(encoding="utf-8")
return tuple(scan(sql, directory.name, read_markers(sql)))
def stale_grandfathers(found: Mapping[str, tuple[Violation, ...]]) -> tuple[str, ...]:
clean = (name for name in GRANDFATHERED & found.keys() if not found[name])
missing = GRANDFATHERED - found.keys()
return tuple(sorted((*clean, *missing)))
def main() -> int:
if not MIGRATIONS_DIR.is_dir():
print(f"migrations directory not found: {MIGRATIONS_DIR}", file=sys.stderr)
return 2
directories = tuple(sorted(path for path in MIGRATIONS_DIR.iterdir() if (path / "migration.sql").is_file()))
found = {directory.name: scan_migration(directory) for directory in directories}
violations = tuple(
violation for name, results in found.items() if name not in GRANDFATHERED for violation in results
)
for violation in violations:
print(violation.render())
stale = stale_grandfathers(found)
for name in stale:
print(f"{name}: listed in GRANDFATHERED but no longer violates; remove it from the set")
if violations:
print(f"\n{len(violations)} data-rewriting statement(s) in migrations.")
print(GUIDANCE)
if violations or stale:
return 1
print(f"No data-rewriting statements in {len(directories)} migrations.")
return 0
if __name__ == "__main__":
raise SystemExit(main())

View file

@ -94,9 +94,9 @@ E2E_FIXTURE_MODE=record E2E_FIXTURE_DIR=/tmp/e2e-fixtures E2E_RESET_SPEND_LOGS=1
E2E_FIXTURE_MODE=replay E2E_FIXTURE_DIR=/tmp/e2e-fixtures E2E_RESET_SPEND_LOGS=1 uv run pytest tests/e2e/llm_translation/test_chat_completions_contract_e2e.py
```
Point the proxy at bogus provider credentials for the replay run and it still has to pass: that is the whole proof that nothing left the process. Bundles are never committed. `tests/e2e/.fixtures` is gitignored because a bundle holds verbatim provider response bodies and hard-fails after seven days, and publishing one for CI is LIT-5748
Point the proxy at bogus provider credentials for the replay run and it still has to pass: that is the whole proof that nothing left the process. Bundles are never committed. `tests/e2e/.fixtures` is gitignored because a bundle holds verbatim provider response bodies and hard-fails after seven days. CI records and replays this lane on a schedule in `.github/workflows/e2e_record_replay.yml`, publishing the bundle as a private `e2e-fixtures-bundle` artifact instead of committing it, selecting the tests with the `@pytest.mark.replayable` marker, and proving the bogus-credentials replay hermetic by counting provider egress with `.github/scripts/e2e_egress_sentinel.py`
Current limits: CI wiring is LIT-5748, Bedrock cannot be mounted (SigV4 signs the Host header, so a rewritten api_base fails signature verification), deployments baked into the proxy's config file cannot be edge-wired (only `/model/new` registrations can carry the edge api_base), and a file upload routed by `custom_llm_provider` through the proxy's `files_settings` block never passes a deployment at all, so the batches `model_param` and `provider_fallback` scenarios keep uploading live in every mode
Current limits: Bedrock cannot be mounted (SigV4 signs the Host header, so a rewritten api_base fails signature verification), deployments baked into the proxy's config file cannot be edge-wired (only `/model/new` registrations can carry the edge api_base), and a file upload routed by `custom_llm_provider` through the proxy's `files_settings` block never passes a deployment at all, so the batches `model_param` and `provider_fallback` scenarios keep uploading live in every mode
## Typing

View file

@ -61,11 +61,15 @@ E2E_FIXTURE_MODE=record E2E_FIXTURE_DIR=/tmp/e2e-fixtures uv run pytest tests/e2
E2E_FIXTURE_MODE=replay E2E_FIXTURE_DIR=/tmp/e2e-fixtures uv run pytest tests/e2e/quota_management/spend_tracking/test_provider_edge_spend_e2e.py -v
```
Bundles stay local. `tests/e2e/.fixtures` is gitignored because a bundle holds verbatim provider response bodies and expires seven days after it was recorded, so record the suite you want before you replay it and never commit the result; publishing bundles for CI is LIT-5748
Bundles stay local. `tests/e2e/.fixtures` is gitignored because a bundle holds verbatim provider response bodies and expires seven days after it was recorded, so record the suite you want before you replay it and never commit the result. CI keeps its bundle out of git too, as a private GitHub Actions artifact rather than a committed file, for the same reason
In CI the `.github/workflows/e2e_record_replay.yml` lane runs record and replay on a schedule. A Saturday cron records the `replayable` marker's tests against the real providers and publishes the bundle as a private `e2e-fixtures-bundle` artifact carrying a SHA-256 sidecar; weekday crons pull that artifact by its pinned digest, verify the checksum before extracting, and replay it with provider credentials deliberately set to bogus values, so a run that ever reached a real provider would fail instead of passing. An egress sentinel (`.github/scripts/e2e_egress_sentinel.py`) pins the provider hostnames to a local sink for the whole replay job and counts every connection that reaches them, and the job asserts that count is zero, so hermeticity is proven by measurement rather than by an absent bill. A red Saturday publishes no bundle, so the next weekday finds nothing fresh and fails loudly rather than replaying a week-old recording, and the seven-day freshness gate hard-fails any bundle that has drifted too far from the live providers. Run the lane on demand from the Actions tab with the `mode` input: `record` re-records and republishes, `replay` replays the current bundle. A test joins the lane by carrying `@pytest.mark.replayable` on top of its edge wiring, so add that marker only to a test whose provider traffic actually replays with zero egress
One sharp edge: a replayed response reuses the recorded provider response id, and that id is the primary key of `LiteLLM_SpendLogs`, so replaying against a database that still holds the record run's rows silently dedupes the spend writes and a spend assertion fails with zero rows. Run both commands above with `E2E_RESET_SPEND_LOGS=1` (and `DATABASE_URL` set in the pytest env) so each session truncates the spend log table after itself, or point replay at a fresh database
Replay answers any provider call that drifted from the recording with an HTTP 599 whose body names the computed and closest recorded keys, so the test fails loudly instead of silently going live, and a bundle older than seven days fails at collection time naming its age; either way the fix is to re-record. Only tests that register edge-wired deployments participate: everything else hits its provider live in every mode, so record exactly the suite you replay. If the proxy runs in a container, set `E2E_PROVIDER_EDGE_ADVERTISE_HOST` (e.g. `host.docker.internal`) so the api_base the proxy stores can reach the edge on the pytest host, and `E2E_PROVIDER_EDGE_BIND_HOST=0.0.0.0` so the edge accepts it. The suites wired to the edge today are `quota_management/spend_tracking/test_provider_edge_spend_e2e.py`, `llm_translation/test_chat_completions_contract_e2e.py`, the OpenAI registrations in `llm_translation/test_embeddings_endpoint_e2e.py`, the Anthropic tests in `llm_translation/test_messages_e2e.py`, streamed and not, and the OpenAI batch deployment behind `batches/`. A streamed response replays as the chunk sequence the provider sent rather than one buffered body. See `CLAUDE.md` in this directory for the bundle format, the edge design, and the current limits (Bedrock, CI wiring)
Another sharp edge, same root: record and replay derive every per-test token deterministically (the model name included, so a replay regenerates the exact requests the record run sent), which means an edge-wired deployment left in the database by an interrupted earlier run carries the same model name as the fresh one the current run registers. The proxy then holds two deployments under one model group and load-balances across both, and because the leftover's `api_base` points at the earlier run's edge process, which is gone, the calls that land on it fail with a connection error that reads like a transport bug rather than the stale row it is. Give each record or replay run a fresh database, or let a run finish so its own teardown deletes what it registered, and never reuse one long-lived proxy across back-to-back record/replay sessions. CI hands every job its own empty database and its own proxy, so it never sees this
Replay answers any provider call that drifted from the recording with an HTTP 599 whose body names the computed and closest recorded keys, so the test fails loudly instead of silently going live, and a bundle older than seven days fails at collection time naming its age; either way the fix is to re-record. Only tests that register edge-wired deployments participate: everything else hits its provider live in every mode, so record exactly the suite you replay. If the proxy runs in a container, set `E2E_PROVIDER_EDGE_ADVERTISE_HOST` (e.g. `host.docker.internal`) so the api_base the proxy stores can reach the edge on the pytest host, and `E2E_PROVIDER_EDGE_BIND_HOST=0.0.0.0` so the edge accepts it. The suites wired to the edge today are `quota_management/spend_tracking/test_provider_edge_spend_e2e.py`, `llm_translation/test_chat_completions_contract_e2e.py`, the OpenAI registrations in `llm_translation/test_embeddings_endpoint_e2e.py`, the Anthropic tests in `llm_translation/test_messages_e2e.py`, streamed and not, and the OpenAI batch deployment behind `batches/`. A streamed response replays as the chunk sequence the provider sent rather than one buffered body. See `CLAUDE.md` in this directory for the bundle format, the edge design, and the current limits (Bedrock). The scheduled CI record/replay lane is described above
Tests marked `@pytest.mark.e2e` hard-fail when no proxy answers `/health/liveliness`, so a run that goes red with `No live proxy` at setup means the proxy isn't up; they never skip for a missing proxy, so an absent proxy can't be mistaken for a pass

View file

@ -43,6 +43,11 @@ def pytest_configure(config: pytest.Config) -> None:
"markers",
"covers(cell_id, *, exercised_on=()): coverage-registry cell(s) this test covers",
)
config.addinivalue_line(
"markers",
"replayable: edge-wired test whose provider traffic replays from a fixture bundle, so it makes "
"zero provider calls in replay mode; the record/replay CI lane selects it with -m replayable",
)
config.addinivalue_line(
"markers",
"load: heavy throughput/load test; collected last so it never perturbs latency-sensitive suites",

View file

@ -0,0 +1,3 @@
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
store_model_in_db: true

View file

@ -13,7 +13,7 @@ from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody
from proxy_client import ProxyClient
from pydantic import BaseModel
pytestmark = pytest.mark.e2e
pytestmark = [pytest.mark.e2e, pytest.mark.replayable]
OPENAI_BACKEND = "openai/gpt-4o-mini"
CHAT_PATH = "/chat/completions"

View file

@ -40,6 +40,7 @@ def _openai_embeddings_params() -> LiteLLMParamsBody:
class TestEmbeddingsEndpoint:
@pytest.mark.replayable
@pytest.mark.covers("llm.embeddings.openai.basic.nonstream.works")
def test_embeddings_returns_vector(
self, endpoints_client: EndpointsClient, resources: ResourceManager
@ -109,6 +110,7 @@ class TestEmbeddingsEndpoint:
f"embedding vector is all zeros: {result.body[:300]}"
)
@pytest.mark.replayable
@pytest.mark.covers("llm.embeddings.openai.basic.nonstream.works")
def test_array_input_returns_vectors(
self, endpoints_client: EndpointsClient, resources: ResourceManager
@ -129,6 +131,7 @@ class TestEmbeddingsEndpoint:
parsed = EmbeddingsResult.model_validate_json(result.body)
assert len(parsed.data) == 3, f"expected 3 vectors: {result.body[:300]}"
@pytest.mark.replayable
@pytest.mark.covers("llm.embeddings.openai.input_validation.nonstream.works")
def test_missing_model_returns_client_error(
self, endpoints_client: EndpointsClient, resources: ResourceManager
@ -141,6 +144,7 @@ class TestEmbeddingsEndpoint:
)
assert_client_error(result, "embeddings missing model")
@pytest.mark.replayable
@pytest.mark.covers("llm.embeddings.openai.input_validation.nonstream.works")
def test_missing_input_returns_error(
self, endpoints_client: EndpointsClient, resources: ResourceManager

View file

@ -24,7 +24,7 @@ from models import (
)
from pydantic import BaseModel
pytestmark = pytest.mark.e2e
pytestmark = [pytest.mark.e2e, pytest.mark.replayable]
class _OptionalMessagesBody(BaseModel):
@ -171,7 +171,7 @@ class TestAnthropicMessages:
model=model,
max_tokens=64,
stream=True,
messages=[ChatMessage(role="user", content="Count from one to three.")],
messages=[ChatMessage(role="user", content="Count from 1 to 20, one number per line.")],
),
)
require_successful_call(result)

View file

@ -5,6 +5,7 @@
addopts = --strict-markers --strict-config --reruns 1 --only-rerun "kind='network'" --only-rerun "status_code=5[0-9][0-9]"
markers =
e2e: live test that requires a running proxy and real provider keys
replayable: edge-wired test whose provider traffic replays from a fixture bundle, so it makes zero provider calls in replay mode; the record/replay CI lane selects it with -m replayable
load: heavy throughput/load test; collected last so it never perturbs latency-sensitive suites
weekly: real-provider anomaly load test that spends real money; deselected unless E2E_WEEKLY_ANOMALY is set
managed_files: needs a proxy running with require_managed_files enabled; deselected unless E2E_MANAGED_FILES_STACK is set

View file

@ -18,7 +18,7 @@ from lifecycle import ResourceManager
from models import LiteLLMParamsBody
from spend_e2e_client import SpendClient, unique_marker, unwrap
pytestmark = pytest.mark.e2e
pytestmark = [pytest.mark.e2e, pytest.mark.replayable]
@pytest.mark.covers("quota_management.spend_tracking.chat_completions.logs_cost")

View file

@ -7,6 +7,7 @@ import pytest
import litellm
import asyncio
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
@pytest.fixture(scope="session")
@ -38,6 +39,8 @@ def setup_and_teardown():
yield
# Teardown code (executes after the yield point)
# LoggingWorker carries still-queued coroutines onto the next test's loop, where they'd log into that test's callbacks
asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue())
loop.close() # Close the loop created earlier
asyncio.set_event_loop(None) # Remove the reference to the loop

View file

@ -1,6 +1,9 @@
import os
import pytest
import asyncio
import subprocess
import sys
from pathlib import Path
from typing import Optional
from unittest.mock import AsyncMock, patch
@ -24,12 +27,20 @@ from mcp.types import Tool as MCPTool, CallToolResult, TextContent
class TestMCPLogger(CustomLogger):
def __init__(self):
self.standard_logging_payload = None
self.mcp_tool_call_payloads = []
super().__init__()
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
print("success event")
self.standard_logging_payload = kwargs.get("standard_logging_object", None)
print(f"Captured standard_logging_payload: {self.standard_logging_payload}")
payload = kwargs.get("standard_logging_object", None)
self.standard_logging_payload = payload
# Async success events from other calls (e.g. a mocked acompletion whose
# log task is delivered late) race with the MCP event for the single
# last-writer slot; keep MCP tool calls in their own list so assertions
# are order-independent.
if payload is not None and payload.get("call_type") == "call_mcp_tool":
self.mcp_tool_call_payloads.append(payload)
print(f"Captured standard_logging_payload: {payload}")
def _set_authorized_user(server_ids):
@ -138,7 +149,11 @@ async def test_mcp_cost_tracking():
# wait 1-2 seconds for logging to be processed
await asyncio.sleep(2)
logged_standard_logging_payload = test_logger.standard_logging_payload
logged_standard_logging_payload = (
test_logger.mcp_tool_call_payloads[-1]
if test_logger.mcp_tool_call_payloads
else None
)
print("logged_standard_logging_payload", logged_standard_logging_payload)
# Add assertions
@ -277,7 +292,11 @@ async def test_mcp_cost_tracking_per_tool():
# wait for logging to be processed
await asyncio.sleep(2)
logged_standard_logging_payload_1 = test_logger.standard_logging_payload
logged_standard_logging_payload_1 = (
test_logger.mcp_tool_call_payloads[-1]
if test_logger.mcp_tool_call_payloads
else None
)
print(
"logged_standard_logging_payload_1", logged_standard_logging_payload_1
)
@ -290,6 +309,7 @@ async def test_mcp_cost_tracking_per_tool():
# Reset logger for second test
test_logger.standard_logging_payload = None
test_logger.mcp_tool_call_payloads.clear()
# Test 2: Call cheap_tool - should cost 0.1
response2 = await mcp_server_tool_call(
@ -300,7 +320,11 @@ async def test_mcp_cost_tracking_per_tool():
# wait for logging to be processed
await asyncio.sleep(2)
logged_standard_logging_payload_2 = test_logger.standard_logging_payload
logged_standard_logging_payload_2 = (
test_logger.mcp_tool_call_payloads[-1]
if test_logger.mcp_tool_call_payloads
else None
)
print(
"logged_standard_logging_payload_2", logged_standard_logging_payload_2
)
@ -329,16 +353,7 @@ async def test_mcp_cost_tracking_per_tool():
assert mock_client.call_tool.call_count == 2
class MCPLoggerHook(CustomLogger):
def __init__(self):
self.standard_logging_payload = None
super().__init__()
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
print("success event")
self.standard_logging_payload = kwargs.get("standard_logging_object", None)
print(f"Captured standard_logging_payload: {self.standard_logging_payload}")
class MCPLoggerHook(TestMCPLogger):
async def async_post_mcp_tool_call_hook(
self, kwargs, response_obj: MCPPostCallResponseObject, start_time, end_time
) -> Optional[MCPPostCallResponseObject]:
@ -436,9 +451,55 @@ async def test_mcp_tool_call_hook():
await asyncio.sleep(2)
# check logged standard logging payload
logged_standard_logging_payload = test_logger.standard_logging_payload
logged_standard_logging_payload = (
test_logger.mcp_tool_call_payloads[-1]
if test_logger.mcp_tool_call_payloads
else None
)
print("logged_standard_logging_payload", logged_standard_logging_payload)
assert (
logged_standard_logging_payload is not None
), "Standard logging payload should not be None"
assert logged_standard_logging_payload["response_cost"] == 1.42
_QUEUED_LOGGING_OUTLIVES_TEST = '''
import time
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
ran_at = []
async def _record_run():
ran_at.append(time.monotonic())
async def test_1_leaves_logging_queued_behind_a_stopped_worker():
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(_record_run())
await GLOBAL_LOGGING_WORKER.stop()
assert ran_at == []
async def test_2_starts_after_the_previous_tests_logging_ran():
started_at = time.monotonic()
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(_record_run())
await GLOBAL_LOGGING_WORKER.flush()
assert [t < started_at for t in ran_at] == [True, False]
'''
def test_logging_queued_by_one_test_is_drained_before_the_next(tmp_path: Path):
"""Regression: a logging coroutine queued by one test must not run inside a later test (it would log into that
test's callbacks, which is how test_mcp_tool_call_hook captured a gpt-4o-mini payload under xdist)."""
(tmp_path / "conftest.py").write_text((Path(__file__).parent / "conftest.py").read_text())
(tmp_path / "pyproject.toml").write_text('[tool.pytest.ini_options]\nasyncio_mode = "auto"\n')
(tmp_path / "test_queued_logging.py").write_text(_QUEUED_LOGGING_OUTLIVES_TEST)
result = subprocess.run(
[sys.executable, "-m", "pytest", "-q", "-p", "no:cacheprovider", "test_queued_logging.py"],
cwd=tmp_path,
capture_output=True,
text=True,
timeout=120,
)
assert result.returncode == 0, result.stdout + result.stderr

View file

@ -0,0 +1,152 @@
import uuid
import pytest
from .actors import Actor
from .conftest import create_scratch_team
pytestmark = pytest.mark.asyncio(loop_scope="session")
_SEED_SPEND = 5.0
_RESET_TO = 2.0
# POST /team/{team_id}/member/{user_id}/reset_spend. The handler gate is
# _verify_team_access (proxy admin / team admin of this team / org admin of
# the team's org) — the same gate /team/member_update uses, so this mirrors
# that file's matrix exactly.
_MATRIX = [
("alpha/proxy_admin", Actor.PROXY_ADMIN, "alpha", 200),
("alpha/org_admin", Actor.ORG_ADMIN, "alpha", 200),
("alpha/team_admin", Actor.TEAM_ADMIN, "alpha", 200),
("alpha/internal_user", Actor.INTERNAL_USER, "alpha", 403),
("alpha/owner", Actor.OWNER, "alpha", 403),
("alpha/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "alpha", 403),
("alpha/cross_org_user", Actor.CROSS_ORG_USER, "alpha", 403),
("alpha/service_account", Actor.SERVICE_ACCOUNT, "alpha", 403),
("alpha/org_b_admin", Actor.ORG_B_ADMIN, "alpha", 403),
("beta/proxy_admin", Actor.PROXY_ADMIN, "beta", 200),
("beta/org_admin", Actor.ORG_ADMIN, "beta", 403),
("beta/team_admin", Actor.TEAM_ADMIN, "beta", 403),
("beta/org_b_admin", Actor.ORG_B_ADMIN, "beta", 200),
]
async def _seed_target(prisma, world, shape: str, team_id: str, member_id: str) -> None:
if shape == "alpha":
await create_scratch_team(
prisma,
team_id,
organization_id=world.org_a_id,
admin_user_ids=[world.keys[Actor.TEAM_ADMIN].user_id],
)
elif shape == "beta":
await create_scratch_team(prisma, team_id, organization_id=world.org_b_id)
else: # pragma: no cover - guard
pytest.fail(f"unknown shape={shape}")
await prisma.db.litellm_teammembership.create(
data={"user_id": member_id, "team_id": team_id, "spend": _SEED_SPEND}
)
@pytest.mark.parametrize(
"actor,shape,expected_status",
[(a, sh, s) for (_id, a, sh, s) in _MATRIX],
ids=[s[0] for s in _MATRIX],
)
async def test_team_member_reset_spend_authz_matrix(
actor: Actor,
shape: str,
expected_status: int,
proxy_client,
prisma,
scratch,
world,
):
member_id = scratch.tag("member")
await _seed_target(prisma, world, shape, scratch.prefix, member_id)
caller = world.keys[actor]
resp = await proxy_client.post(
f"/team/{scratch.prefix}/member/{member_id}/reset_spend",
headers={"Authorization": f"Bearer {caller.cleartext}"},
json={"reset_to": _RESET_TO},
)
assert (
resp.status_code == expected_status
), f"{actor.value} {shape}: {resp.status_code} {resp.text}"
row = await prisma.db.litellm_teammembership.find_unique(
where={"user_id_team_id": {"user_id": member_id, "team_id": scratch.prefix}}
)
assert row is not None
if expected_status == 200:
assert row.spend == _RESET_TO
else:
assert row.spend == _SEED_SPEND, "denied but spend reset"
async def test_team_member_reset_spend_missing_team_is_404(proxy_client, world):
resp = await proxy_client.post(
f"/team/behavior-pin-no-such-team/member/{uuid.uuid4().hex}/reset_spend",
headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"},
json={"reset_to": 0.0},
)
assert resp.status_code == 404, resp.text
async def test_team_member_reset_spend_missing_membership_is_404(
proxy_client, prisma, scratch, world
):
"""A well-formed team but a user_id with no LiteLLM_TeamMembership row is 404."""
await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id)
resp = await proxy_client.post(
f"/team/{scratch.prefix}/member/{uuid.uuid4().hex}/reset_spend",
headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"},
json={"reset_to": 0.0},
)
assert resp.status_code == 404, resp.text
async def test_team_member_reset_spend_above_current_spend_is_400(
proxy_client, prisma, scratch, world
):
member_id = scratch.tag("member")
await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id)
await prisma.db.litellm_teammembership.create(
data={"user_id": member_id, "team_id": scratch.prefix, "spend": 1.0}
)
resp = await proxy_client.post(
f"/team/{scratch.prefix}/member/{member_id}/reset_spend",
headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"},
json={"reset_to": 5.0},
)
assert resp.status_code == 400, resp.text
async def test_team_member_reset_spend_team_admin_cannot_reset_own_spend(
proxy_client, prisma, scratch, world
):
"""A team admin targeting their own LiteLLM_TeamMembership row is 403: unchecked, an
admin could repeatedly zero their own spend right before it crosses their per-member
cap, consuming the shared team budget without the configured limit ever binding."""
team_admin = world.keys[Actor.TEAM_ADMIN]
await create_scratch_team(
prisma,
scratch.prefix,
organization_id=world.org_a_id,
admin_user_ids=[team_admin.user_id],
)
await prisma.db.litellm_teammembership.create(
data={"user_id": team_admin.user_id, "team_id": scratch.prefix, "spend": _SEED_SPEND}
)
resp = await proxy_client.post(
f"/team/{scratch.prefix}/member/{team_admin.user_id}/reset_spend",
headers={"Authorization": f"Bearer {team_admin.cleartext}"},
json={"reset_to": 0.0},
)
assert resp.status_code == 403, resp.text
row = await prisma.db.litellm_teammembership.find_unique(
where={"user_id_team_id": {"user_id": team_admin.user_id, "team_id": scratch.prefix}}
)
assert row is not None and row.spend == _SEED_SPEND, "denied but spend reset"

View file

@ -0,0 +1,199 @@
"""
Tests for the Grounding with Bing Search (Microsoft Foundry) integration.
"""
import json
from unittest.mock import AsyncMock, Mock, patch
import pytest
import litellm
from tests.search_tests.base_search_unit_tests import BaseSearchTest
PROJECT_ENDPOINT = "https://acct.services.ai.azure.com/api/projects/proj"
_ANSWER_TEXT = (
"LiteLLM is an open source LLM gateway ([github.com](https://github.com/BerriAI/litellm))\n"
"The docs live on docs.litellm.ai ([docs.litellm.ai](https://docs.litellm.ai/))"
)
def _annotation(marker: str, url: str, title: str) -> dict:
start = _ANSWER_TEXT.index(marker)
return {
"type": "url_citation",
"url": url,
"title": title,
"start_index": start,
"end_index": start + len(marker),
}
MOCK_BING_GROUNDING_RESPONSE = {
"id": "resp_mock",
"object": "response",
"status": "completed",
"model": "gpt-4.1",
"output": [
{"type": "web_search_call", "status": "completed"},
{
"type": "message",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": _ANSWER_TEXT,
"annotations": [
_annotation(
"([github.com](https://github.com/BerriAI/litellm))",
"https://github.com/BerriAI/litellm",
"BerriAI/litellm - GitHub",
),
_annotation(
"([docs.litellm.ai](https://docs.litellm.ai/))",
"https://docs.litellm.ai/",
"LiteLLM Docs",
),
],
}
],
},
],
"usage": {"input_tokens": 100, "output_tokens": 50},
}
def _mock_response():
response = Mock()
response.status_code = 200
response.headers = {}
response.content = json.dumps(MOCK_BING_GROUNDING_RESPONSE).encode()
return response
@pytest.mark.skip(reason="Local only tested search providers")
class TestBingGroundingSearch(BaseSearchTest):
"""
E2E tests for Grounding with Bing Search that make real API calls.
Inherits from BaseSearchTest to run standard search tests.
"""
def get_search_provider(self) -> str:
return "bing_grounding"
class TestBingGroundingSearchTransformation:
"""
Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked.
Transformation details are unit-tested in tests/test_litellm/llms/azure/search/.
"""
@pytest.fixture(autouse=True)
def _server_env(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("BING_GROUNDING_PROJECT_ENDPOINT", PROJECT_ENDPOINT)
monkeypatch.setenv("BING_GROUNDING_MODEL", "gpt-4.1")
monkeypatch.setenv("BING_GROUNDING_TOKEN", "test-entra-token")
monkeypatch.delenv("BING_GROUNDING_CONNECTION_ID", raising=False)
def test_bing_grounding_search_request_and_response(self):
with patch( # test-quality-ok: litellm.search has no client injection seam
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post",
return_value=_mock_response(),
) as mock_post:
response = litellm.search(
query="what is litellm",
search_provider="bing_grounding",
max_results=5,
country="us",
)
assert mock_post.called
call_kwargs = mock_post.call_args.kwargs
assert call_kwargs["url"] == f"{PROJECT_ENDPOINT}/openai/v1/responses"
assert call_kwargs["headers"]["Authorization"] == "Bearer test-entra-token"
request_body = call_kwargs["json"]
assert request_body["model"] == "gpt-4.1"
assert request_body["input"] == "what is litellm"
assert request_body["tools"] == [
{"type": "web_search", "user_location": {"type": "approximate", "country": "US"}}
]
assert response.object == "search"
assert len(response.results) == 2
assert response.results[0].url == "https://github.com/BerriAI/litellm"
assert response.results[0].title == "BerriAI/litellm - GitHub"
assert response.results[0].snippet == "LiteLLM is an open source LLM gateway"
assert response.results[1].url == "https://docs.litellm.ai/"
assert response.results[1].snippet == "The docs live on docs.litellm.ai"
def test_connection_mode_sends_the_bing_grounding_tool(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv(
"BING_GROUNDING_CONNECTION_ID",
"/subscriptions/sub/resourceGroups/rg/providers/Microsoft.CognitiveServices"
"/accounts/acct/projects/proj/connections/bing-conn",
)
with patch( # test-quality-ok: litellm.search has no client injection seam
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post",
return_value=_mock_response(),
) as mock_post:
litellm.search(
query="what is litellm",
search_provider="bing_grounding",
max_results=3,
)
request_body = mock_post.call_args.kwargs["json"]
assert request_body["tools"] == [
{
"type": "bing_grounding",
"bing_grounding": {
"search_configurations": [
{
"project_connection_id": (
"/subscriptions/sub/resourceGroups/rg/providers/Microsoft.CognitiveServices"
"/accounts/acct/projects/proj/connections/bing-conn"
),
"count": 3,
}
]
},
}
]
@pytest.mark.asyncio
async def test_bing_grounding_asearch(self):
with patch( # test-quality-ok: litellm.asearch has no client injection seam
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=AsyncMock(return_value=_mock_response()),
) as mock_post:
response = await litellm.asearch(
query="what is litellm",
search_provider="bing_grounding",
)
assert mock_post.call_args.kwargs["json"]["tools"] == [{"type": "web_search"}]
assert len(response.results) == 2
def test_web_search_mode_is_not_billed_the_g1_price(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
with patch( # test-quality-ok: litellm.search has no client injection seam
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post",
return_value=_mock_response(),
):
response = litellm.search(query="pricing check", search_provider="bing_grounding")
assert response._hidden_params["response_cost"] == 0.0
def test_connection_mode_tracks_the_g1_cost(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("BING_GROUNDING_CONNECTION_ID", "conn-id")
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
with patch( # test-quality-ok: litellm.search has no client injection seam
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post",
return_value=_mock_response(),
):
response = litellm.search(query="pricing check", search_provider="bing_grounding")
assert response._hidden_params["response_cost"] == pytest.approx(0.035)

View file

@ -28,6 +28,14 @@ if TYPE_CHECKING:
from redis.asyncio.cluster import RedisCluster as _AsyncRedisClusterType
class _NodeClassWithPerConnectionRecovery:
def update_active_connections_for_reconnect(self) -> None: ...
class _NodeClassWithoutPerConnectionRecovery:
pass
class _FakeClusterNode:
def __init__(self, name: str, raises: Exception | None = None, response: object = None) -> None:
self.name = name
@ -47,7 +55,9 @@ class _FakeNodesManager:
def _build_cluster_instance() -> "_AsyncRedisClusterType":
cluster_cls = get_litellm_async_redis_cluster_class()
cluster_cls = get_litellm_async_redis_cluster_class(
cluster_node_class=_NodeClassWithoutPerConnectionRecovery
)
instance = cluster_cls.__new__(cluster_cls)
instance.RedisClusterRequestTTL = 1
instance.reinitialize_counter = 0
@ -58,6 +68,33 @@ def _build_cluster_instance() -> "_AsyncRedisClusterType":
return instance
def test_per_connection_recovery_redis_py_gets_the_unmodified_upstream_class() -> None:
"""Regression (redis-py 8.x): when upstream ClusterNode already recovers a node-level
connection error per-connection, the factory must NOT install the copied override,
whose node.disconnect() also kills connections other coroutines are mid-operation on."""
from redis.asyncio.cluster import RedisCluster
cluster_cls = get_litellm_async_redis_cluster_class(
cluster_node_class=_NodeClassWithPerConnectionRecovery
)
assert cluster_cls is RedisCluster
def test_pre_recovery_redis_py_still_gets_the_node_isolation_override() -> None:
"""Old redis-py (5.x) responds to a node-level error with a full-cluster aclose(),
so those versions must keep litellm's per-node isolation override."""
from redis.asyncio.cluster import RedisCluster
cluster_cls = get_litellm_async_redis_cluster_class(
cluster_node_class=_NodeClassWithoutPerConnectionRecovery
)
assert cluster_cls is not RedisCluster
assert issubclass(cluster_cls, RedisCluster)
assert "_execute_command" in cluster_cls.__dict__
@pytest.mark.asyncio
@pytest.mark.parametrize("error_cls", [RedisConnectionError, RedisTimeoutError])
async def test_node_level_error_resets_only_that_node_not_the_whole_client(error_cls: type[Exception]) -> None:

View file

@ -2,7 +2,7 @@ import datetime
import json
import os
import unittest
from typing import TYPE_CHECKING, List, Literal, Optional, Tuple
from typing import TYPE_CHECKING, Final, List, Literal, Optional, Tuple
from unittest.mock import ANY, MagicMock, Mock, patch
import httpx
@ -1585,10 +1585,16 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch):
assert result_dict["summary"] == "custom_summary"
print("✓ Dict input is passed through without modification")
# Test 5: None/unknown values return None
result_unknown = handler._map_reasoning_effort("unknown_value")
assert result_unknown is None
print("✓ Unknown reasoning_effort values return None")
# Test 5: every REASONING_EFFORT level reaches the provider, and anything else (a typo, an
# unshipped level, "default") is dropped so the request still succeeds at the provider default
from litellm.types.llms.openai import Reasoning
for effort in ("max", "xhigh", "none"):
result_passthrough = handler._map_reasoning_effort(effort)
assert result_passthrough == Reasoning(effort=effort)
for dropped in ("ultra", "hgih", "unknown_value", "", "default"):
assert handler._map_reasoning_effort(dropped) is None
print("✓ Enumerated levels pass through and unknown ones are dropped")
print(
"✓ All reasoning_effort behaviors work correctly with flag/env var control"
@ -2438,6 +2444,32 @@ def test_map_optional_params_preserves_reasoning_summary():
assert responses_api_request["reasoning"]["summary"] == "detailed"
@pytest.mark.parametrize("reasoning_effort", ["max", "high"])
def test_transform_request_bedrock_mantle_tools_keeps_reasoning_effort(monkeypatch, reasoning_effort):
"""Regression for reasoning_effort=max being dropped on the chat -> Responses bridge (issue #38084)."""
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
monkeypatch.setattr(litellm, "reasoning_auto_summary", False)
monkeypatch.delenv("LITELLM_REASONING_AUTO_SUMMARY", raising=False)
handler: Final = LiteLLMResponsesTransformationHandler()
result: Final = handler.transform_request(
model="openai.gpt-5.6-sol",
messages=[{"role": "user", "content": "Say pong"}],
optional_params={
"reasoning_effort": reasoning_effort,
"tools": [{"type": "function", "function": {"name": "get_weather", "parameters": {"type": "object"}}}],
},
litellm_params={"custom_llm_provider": "bedrock_mantle"},
headers={},
litellm_logging_obj=Mock(),
)
assert result["reasoning"] == {"effort": reasoning_effort}
def test_map_optional_params_tool_choice_chat_nested_to_responses_api():
"""Chat tool_choice must become Responses ToolChoiceFunction (top-level name)."""
from litellm.completion_extras.litellm_responses_transformation.transformation import (

View file

@ -432,3 +432,103 @@ class TestLangsmithRedactUserApiKeyInfo:
)
assert data["inputs"]["metadata"]["user_api_key_hash"] == "abc123"
class TestLangsmithRootRunIdConsistency:
"""Regression tests for LIT-5878 / #37269.
A request that carries a session/trace header (e.g. x-claude-code-session-id)
fans the header value out into litellm metadata as both trace_id and
session_id. LangSmith then rejected the whole ingest batch twice over:
a root run whose trace_id does not match the run id embedded in dotted_order
(400), and a run-body session_id that does not reference an existing tracer
session (404, or 422 for non-UUID values).
"""
def _prepare(self, request_metadata):
payload = {
"id": "slp-1",
"response": {"choices": []},
"metadata": {},
"startTime": 1.0,
"endTime": 2.0,
"request_tags": [],
"error_str": None,
"status": "success",
"response_cost": 0.0,
"prompt_tokens": 1,
"completion_tokens": 1,
"total_tokens": 2,
}
logger = LangsmithLogger(
langsmith_api_key="test-key",
langsmith_project="test-project",
)
return logger._prepare_log_data(
kwargs={
"litellm_params": {"metadata": request_metadata},
"standard_logging_object": payload,
},
response_obj=None,
start_time=1.0,
end_time=2.0,
credentials={
"LANGSMITH_API_KEY": "test-key",
"LANGSMITH_PROJECT": "test-project",
"LANGSMITH_BASE_URL": "https://api.smith.langchain.com",
},
)
def test_header_derived_ids_yield_self_consistent_root_run(self):
header_value = "ed29c3bb-44fa-4eec-9b7b-fecaa3e82d64"
data = self._prepare({"trace_id": header_value, "session_id": header_value})
assert data["trace_id"] == data["id"]
assert data["trace_id"] != header_value
assert data["dotted_order"].endswith(data["id"])
assert len(data["dotted_order"]) == 22 + len(data["id"])
assert "session_id" not in data
def test_distinct_session_id_is_still_forwarded(self):
data = self._prepare({"session_id": "11111111-2222-3333-4444-555555555555"})
assert data["session_id"] == "11111111-2222-3333-4444-555555555555"
def test_trace_id_only_root_run_is_overridden(self):
data = self._prepare({"trace_id": "ed29c3bb-44fa-4eec-9b7b-fecaa3e82d64"})
assert data["trace_id"] == data["id"]
assert data["trace_id"] != "ed29c3bb-44fa-4eec-9b7b-fecaa3e82d64"
assert data["dotted_order"].endswith(data["id"])
def test_root_run_without_caller_ids_is_self_consistent(self):
data = self._prepare({})
assert data["trace_id"] == data["id"]
assert data["dotted_order"].endswith(data["id"])
def test_child_run_keeps_caller_trace_id(self):
data = self._prepare(
{
"trace_id": "trace-1",
"parent_run_id": "parent-1",
"run_id": "child-1",
}
)
assert data["trace_id"] == "trace-1"
assert data["id"] == "child-1"
assert data["parent_run_id"] == "parent-1"
def test_caller_supplied_dotted_order_and_trace_id_are_untouched(self):
dotted = "20260820T000000000000Ztrace-1.20260820T000001000000Zrun-1"
data = self._prepare(
{
"trace_id": "trace-1",
"run_id": "run-1",
"dotted_order": dotted,
}
)
assert data["trace_id"] == "trace-1"
assert data["dotted_order"] == dotted

View file

@ -1647,15 +1647,27 @@ def _signature_for(signer_cls, url: str, method: str, body: bytes | None, header
return signer.signature(signer.string_to_sign(request, canonical_request), request)
def _as_s3_canonicalizes(url: str) -> str:
"""
The path S3 rebuilds from the wire path: percent-encode everything outside the unreserved
set, without normalizing or double-encoding. `=` becomes `%3D`, `%20` stays `%20`.
"""
from urllib.parse import quote, unquote, urlsplit, urlunsplit
split = urlsplit(url)
return urlunsplit(split._replace(path=quote(unquote(split.path), safe="/~")))
def _assert_signed_for_s3_canonicalization(url: str, method: str, body: bytes | None, headers: dict[str, str]) -> None:
"""
S3 rebuilds the canonical request from the wire path with single percent-encoding, which
botocore models as S3SigV4Auth; plain SigV4Auth double-encodes it (%2520 for a space) and S3
answers 403 SignatureDoesNotMatch. Assert we signed the path the way S3 reads it.
answers 403 SignatureDoesNotMatch. Assert we sent an already-encoded path and signed it the
way S3 reads it.
"""
from botocore.auth import S3SigV4Auth, SigV4Auth
assert "%20" in url
assert url == _as_s3_canonicalizes(url)
sent_signature = headers["Authorization"].split("Signature=")[1].strip()
assert sent_signature == _signature_for(S3SigV4Auth, url, method, body, headers)
assert sent_signature != _signature_for(SigV4Auth, url, method, body, headers)
@ -1744,3 +1756,103 @@ async def test_download_signs_object_key_with_space_the_way_s3_does():
body=None,
headers=call.kwargs["headers"],
)
_RESERVED_CHAR_KEYS = (
"2026-08-21/time-05-29-36_resp_bGl0ZWxsbTpjdXN0b20=.json",
"session=logs/2026-08-21/time-05-29-36_abc.json",
"a+b/2026-08-21/time-05-29-36_abc.json",
"a&b/2026-08-21/time-05-29-36_abc.json",
"a#b/2026-08-21/time-05-29-36_abc.json",
"a?b/2026-08-21/time-05-29-36_abc.json",
"a%b/2026-08-21/time-05-29-36_abc.json",
_KEY_WITH_SPACE,
)
def _element_for(s3_object_key: str):
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
return s3BatchLoggingElement(
s3_object_key=s3_object_key,
payload={"test": "sigv4"},
s3_object_download_filename="log.json",
)
def _expected_wire_url(s3_object_key: str) -> str:
"""The URL boto3 itself would put on the wire for this key."""
from urllib.parse import quote
return f"https://logs-bucket.s3.us-east-1.amazonaws.com/{quote(s3_object_key, safe='/')}"
@pytest.mark.parametrize("s3_object_key", _RESERVED_CHAR_KEYS)
@pytest.mark.asyncio
async def test_async_upload_percent_encodes_reserved_characters_in_object_key(s3_object_key):
from unittest.mock import AsyncMock, MagicMock
logger = _logger_for_signing()
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put.return_value = response
await logger.async_upload_data_to_s3(_element_for(s3_object_key))
call = logger.async_httpx_client.put.call_args
assert call[0][0] == _expected_wire_url(s3_object_key)
_assert_signed_for_s3_canonicalization(
url=call[0][0],
method="PUT",
body=call.kwargs["data"].encode("utf-8"),
headers=call.kwargs["headers"],
)
@pytest.mark.parametrize("s3_object_key", _RESERVED_CHAR_KEYS)
def test_sync_upload_percent_encodes_reserved_characters_in_object_key(s3_object_key):
from unittest.mock import MagicMock
logger = _logger_for_signing()
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
mock_sync_client = MagicMock()
mock_sync_client.put.return_value = response
with patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client):
logger.upload_data_to_s3(_element_for(s3_object_key))
call = mock_sync_client.put.call_args
assert call[0][0] == _expected_wire_url(s3_object_key)
_assert_signed_for_s3_canonicalization(
url=call[0][0],
method="PUT",
body=call.kwargs["data"].encode("utf-8"),
headers=call.kwargs["headers"],
)
@pytest.mark.parametrize("s3_object_key", _RESERVED_CHAR_KEYS)
@pytest.mark.asyncio
async def test_download_percent_encodes_reserved_characters_in_object_key(s3_object_key):
from unittest.mock import AsyncMock, MagicMock
logger = _logger_for_signing()
response = MagicMock()
response.status_code = 200
response.json = MagicMock(return_value={"downloaded": "data"})
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.get.return_value = response
assert await logger._download_object_from_s3(s3_object_key) == {"downloaded": "data"}
call = logger.async_httpx_client.get.call_args
assert call[0][0] == _expected_wire_url(s3_object_key)
_assert_signed_for_s3_canonicalization(
url=call[0][0],
method="GET",
body=None,
headers=call.kwargs["headers"],
)

View file

@ -478,10 +478,10 @@ def test_generic_cost_per_token_minimax_m3_above_512k_tokens(_local_model_cost_m
],
)
def test_generic_cost_per_token_bedrock_mantle_gpt56_long_context(_local_model_cost_map, model):
"""Bedrock GPT-5.6 supports a 1M context window, billed at the long-context rates above 272K."""
"""Bedrock GPT-5.6 enforces a 1,050,000-token context window, billed at the long-context rates above 272K."""
model_cost_map = litellm.model_cost[model]
assert model_cost_map["max_input_tokens"] == 1000000
assert model_cost_map["max_input_tokens"] == 1050000
cached_tokens = 100000
completion_tokens = 1000

View file

@ -19,9 +19,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
def test_get_format_from_file_id():
unified_file_id = (
"litellm_proxy:application/pdf;unified_id,cbbe3534-8bf8-4386-af00-f5f6b7e370bf"
)
unified_file_id = "litellm_proxy:application/pdf;unified_id,cbbe3534-8bf8-4386-af00-f5f6b7e370bf"
format = get_format_from_file_id(unified_file_id)
@ -48,9 +46,7 @@ def test_update_messages_with_model_file_ids():
model_file_id_mapping = {file_id: {"my_model_id": "provider_file_id"}}
updated_messages = update_messages_with_model_file_ids(
messages, model_id, model_file_id_mapping
)
updated_messages = update_messages_with_model_file_ids(messages, model_id, model_file_id_mapping)
assert updated_messages == [
{
@ -143,9 +139,7 @@ def test_add_system_prompt_to_messages_merge_with_first_system():
{"role": "system", "content": "Existing system prompt."},
{"role": "user", "content": "Hello"},
]
result = add_system_prompt_to_messages(
messages, "You are helpful.", merge_with_first_system=True
)
result = add_system_prompt_to_messages(messages, "You are helpful.", merge_with_first_system=True)
assert result == [
{"role": "system", "content": "You are helpful.\n\nExisting system prompt."},
{"role": "user", "content": "Hello"},
@ -155,9 +149,7 @@ def test_add_system_prompt_to_messages_merge_with_first_system():
def test_add_system_prompt_to_messages_merge_with_first_system_adds_new_when_no_system():
"""When merge_with_first_system=True but no system message, adds new one at start."""
messages = [{"role": "user", "content": "Hello"}]
result = add_system_prompt_to_messages(
messages, "You are helpful.", merge_with_first_system=True
)
result = add_system_prompt_to_messages(messages, "You are helpful.", merge_with_first_system=True)
assert result == [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "Hello"},
@ -492,14 +484,8 @@ def test_update_messages_with_model_file_ids_tolerates_non_dict_content_items():
messages_token_ids_batch = [{"role": "user", "content": [[15496, 995], [9906, 0]]}]
# Both should pass through unchanged without raising.
assert (
update_messages_with_model_file_ids(messages_token_ids, "model-A", {})
== messages_token_ids
)
assert (
update_messages_with_model_file_ids(messages_token_ids_batch, "model-A", {})
== messages_token_ids_batch
)
assert update_messages_with_model_file_ids(messages_token_ids, "model-A", {}) == messages_token_ids
assert update_messages_with_model_file_ids(messages_token_ids_batch, "model-A", {}) == messages_token_ids_batch
class TestExtractFileDataBareStr:
@ -645,9 +631,7 @@ class TestUnpackLegacyDefs:
definitions = {
f"L{i}": {
"type": "object",
"properties": {
f"x{j}": {"$ref": f"#/definitions/L{i + 1}"} for j in range(fanout)
},
"properties": {f"x{j}": {"$ref": f"#/definitions/L{i + 1}"} for j in range(fanout)},
}
for i in range(depth)
}
@ -712,9 +696,7 @@ class TestUnpackLegacyDefs:
schema = {
"type": "object",
"properties": {
f"r{i}": {"$ref": f"#/components/schemas/T{i}"} for i in range(50)
},
"properties": {f"r{i}": {"$ref": f"#/components/schemas/T{i}"} for i in range(50)},
"components": {
"schemas": {
f"T{i}": {
@ -739,9 +721,7 @@ class TestTextCompletionPromptToMessages:
text_completion_prompt_to_messages,
)
assert text_completion_prompt_to_messages("summarize this") == (
{"role": "user", "content": "summarize this"},
)
assert text_completion_prompt_to_messages("summarize this") == ({"role": "user", "content": "summarize this"},)
def test_list_of_strings_becomes_one_message_each(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
@ -970,3 +950,80 @@ class TestCustomToolFormatShapeConversion:
for weird in ({}, {"type": "grammar"}, {"type": "future_format", "x": 1}):
assert convert_custom_tool_format_to_chat_shape(dict(weird)) in (weird, {"type": "grammar", "grammar": {}})
assert convert_custom_tool_format_to_responses_shape(dict(weird)) == weird
# --- x-litellm-model upload-path decoding (litellm #29830) -------------------
def _xlitellm_encoded(raw_id: str, model: str) -> str:
from litellm.proxy.openai_files_endpoints.common_utils import (
encode_file_id_with_model,
)
return encode_file_id_with_model(raw_id, model)
def test_update_messages_with_model_file_ids_decodes_xlitellm_encoded_id():
"""x-litellm-model upload returns `file-<b64(litellm:<raw>;model,<m>)>`.
Without decoding, the encoded id leaks to upstream OpenAI and errors as
'Files [...] were not found'. Decode it back to raw provider id."""
raw_id = "file-ExTuCawUqxEMjVFK6xwR9B"
encoded_id = _xlitellm_encoded(raw_id, "gpt-5.1")
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "Summarize this."},
{"type": "file", "file": {"file_id": encoded_id}},
],
}
]
updated = update_messages_with_model_file_ids(messages, "model-A", {})
assert updated[0]["content"][1]["file"]["file_id"] == raw_id
def test_update_responses_input_with_model_file_ids_decodes_xlitellm_encoded_id():
"""Same bug on /v1/responses path. Without decoding the encoded id (>64
chars), OpenAI rejects with 'string too long. Expected ... maximum length
64'."""
from litellm.litellm_core_utils.prompt_templates.common_utils import (
update_responses_input_with_model_file_ids,
)
raw_id = "file-ExTuCawUqxEMjVFK6xwR9B"
encoded_id = _xlitellm_encoded(raw_id, "gpt-5.1")
input_items = [
{
"role": "user",
"content": [
{"type": "input_text", "text": "Summarize."},
{"type": "input_file", "file_id": encoded_id},
],
}
]
updated = update_responses_input_with_model_file_ids(input_items)
assert updated[0]["content"][1]["file_id"] == raw_id
def test_update_messages_xlitellm_decode_does_not_override_mapping():
"""If the call-site already resolved a provider id via the mapping, that
wins. The new decode fallback runs only when no mapping match."""
raw_id = "file-ExTuCawUqxEMjVFK6xwR9B"
encoded_id = _xlitellm_encoded(raw_id, "gpt-5.1")
mapping = {encoded_id: {"model-A": "provider-explicit-id"}}
messages = [
{
"role": "user",
"content": [
{"type": "file", "file": {"file_id": encoded_id}},
],
}
]
updated = update_messages_with_model_file_ids(messages, "model-A", mapping)
assert updated[0]["content"][0]["file"]["file_id"] == "provider-explicit-id"

View file

@ -1,6 +1,8 @@
import base64
import json
import logging
import os
from typing import Final
from unittest.mock import MagicMock, patch
import pytest
@ -3309,6 +3311,66 @@ def test_get_tool_calls_from_response_include_all_choices_reads_every_choice():
assert names == ["tool_alpha", "tool_beta"]
def test_get_tool_calls_from_response_silences_redacted_arguments(caplog):
from litellm.litellm_core_utils.prompt_templates.factory import (
get_tool_calls_from_response,
)
response: Final = {
"choices": [
{
"message": {
"tool_calls": [
{
"id": "call_1",
"function": {
"name": "Read",
"arguments": "redacted-by-litellm",
},
}
]
}
}
]
}
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
tool_calls: Final = get_tool_calls_from_response(response)
assert tool_calls == [{"id": "call_1", "name": "Read", "arguments": {}}]
assert "Failed to parse tool call arguments" not in caplog.text
def test_get_tool_calls_from_response_warns_for_malformed_arguments(caplog):
from litellm.litellm_core_utils.prompt_templates.factory import (
get_tool_calls_from_response,
)
response: Final = {
"choices": [
{
"message": {
"tool_calls": [
{
"id": "call_1",
"function": {
"name": "Read",
"arguments": "not-json",
},
}
]
}
}
]
}
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
tool_calls: Final = get_tool_calls_from_response(response)
assert tool_calls == [{"id": "call_1", "name": "Read", "arguments": {}}]
assert "Failed to parse tool call arguments" in caplog.text
def test_group_tool_exchanges_pairs_assistant_with_its_tool_rows():
from litellm.litellm_core_utils.prompt_templates.factory import group_tool_exchanges

View file

@ -255,3 +255,26 @@ class TestRedactNestedMatchAndRegexKeys:
def test_passes_through_none_and_str(self):
assert redact_nested_match_and_regex_keys(None) is None
assert redact_nested_match_and_regex_keys("plain") == "plain"
class TestIsExpectedClientError:
def test_status_ranges(self):
from litellm.litellm_core_utils.core_helpers import is_expected_client_error
class WithStatusCode(Exception):
def __init__(self, status_code):
self.status_code = status_code
class WithCode(Exception):
def __init__(self, code):
self.code = code
assert is_expected_client_error(WithStatusCode(400)) is True
assert is_expected_client_error(WithStatusCode(429)) is True
assert is_expected_client_error(WithStatusCode(499)) is True
assert is_expected_client_error(WithStatusCode(500)) is False
assert is_expected_client_error(WithStatusCode(399)) is False
assert is_expected_client_error(WithCode("403")) is True
assert is_expected_client_error(WithCode("invalid_request_error")) is False
assert is_expected_client_error(Exception("no status")) is False
assert is_expected_client_error(None) is False

Some files were not shown because too many files have changed in this diff Show more