mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
chore: merge main into litellm_batch_jsonl_line_item_callbacks
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
0ca85b47ca
268 changed files with 40139 additions and 1487 deletions
|
|
@ -12,6 +12,9 @@ parameters:
|
|||
migration_source_sha:
|
||||
type: string
|
||||
default: ""
|
||||
routing_parity_base:
|
||||
type: string
|
||||
default: ""
|
||||
orbs:
|
||||
codecov: codecov/codecov@4.0.1
|
||||
node: circleci/node@5.1.0 # Add this line to declare the node orb
|
||||
|
|
@ -176,6 +179,9 @@ commands:
|
|||
image:
|
||||
type: string
|
||||
default: postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26
|
||||
server_args:
|
||||
type: string
|
||||
default: ""
|
||||
steps:
|
||||
- run:
|
||||
name: Start PostgreSQL
|
||||
|
|
@ -186,7 +192,7 @@ commands:
|
|||
-e POSTGRES_PASSWORD=postgres \
|
||||
-e POSTGRES_DB=<< parameters.db_name >> \
|
||||
-p 5432:5432 \
|
||||
<< parameters.image >>
|
||||
<< parameters.image >> << parameters.server_args >>
|
||||
- wait_for_service:
|
||||
url: tcp://localhost:5432
|
||||
timeout: "60"
|
||||
|
|
@ -3108,6 +3114,10 @@ jobs:
|
|||
parameters:
|
||||
suite:
|
||||
type: string
|
||||
mode:
|
||||
type: enum
|
||||
enum: [standard, replica]
|
||||
default: standard
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
resource_class: large
|
||||
|
|
@ -3142,18 +3152,19 @@ jobs:
|
|||
command: cd ui/litellm-dashboard && NEXT_TELEMETRY_DISABLED=1 npm run build
|
||||
- start_postgres:
|
||||
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
|
||||
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
|
||||
- start_redis
|
||||
- run:
|
||||
name: Run owned integration contracts
|
||||
command: bash .circleci/scripts/run_integration.sh << parameters.suite >>
|
||||
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> << parameters.mode >>
|
||||
no_output_timeout: 15m
|
||||
- run:
|
||||
name: Stop owned database and Redis
|
||||
when: always
|
||||
command: |
|
||||
mkdir -p test-results/integration-<< parameters.suite >>
|
||||
docker logs postgres-db > test-results/integration-<< parameters.suite >>/postgres.log 2>&1 || true
|
||||
docker logs redis-cache > test-results/integration-<< parameters.suite >>/redis.log 2>&1 || true
|
||||
mkdir -p test-results/services-<< parameters.suite >>-<< parameters.mode >>
|
||||
docker logs postgres-db > test-results/services-<< parameters.suite >>-<< parameters.mode >>/postgres.log 2>&1 || true
|
||||
docker logs redis-cache > test-results/services-<< parameters.suite >>-<< parameters.mode >>/redis.log 2>&1 || true
|
||||
docker rm -f postgres-db redis-cache
|
||||
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
|
||||
- store_test_results:
|
||||
|
|
@ -3161,6 +3172,76 @@ jobs:
|
|||
- store_artifacts:
|
||||
path: test-results
|
||||
|
||||
routing_parity:
|
||||
parameters:
|
||||
suite:
|
||||
type: string
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
resource_class: large
|
||||
working_directory: ~/project
|
||||
steps:
|
||||
- setup_litellm_test_deps
|
||||
- run:
|
||||
name: Check out base product code
|
||||
environment:
|
||||
ROUTING_PARITY_BASE: << pipeline.parameters.routing_parity_base >>
|
||||
command: |
|
||||
[[ "$ROUTING_PARITY_BASE" =~ ^[0-9a-f]{40}$ ]] || exit 1
|
||||
git fetch --depth 1 origin "$ROUTING_PARITY_BASE"
|
||||
git rm -r -f --quiet litellm enterprise litellm-proxy-extras
|
||||
git checkout "$ROUTING_PARITY_BASE" -- litellm enterprise litellm-proxy-extras
|
||||
git reset --quiet
|
||||
test -f litellm/rust_bridge/_native.abi3.so
|
||||
- start_postgres:
|
||||
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
|
||||
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
|
||||
- start_redis
|
||||
- run:
|
||||
name: Run base side
|
||||
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> parity base
|
||||
no_output_timeout: 15m
|
||||
- run:
|
||||
name: Stop base database and Redis
|
||||
when: always
|
||||
command: |
|
||||
mkdir -p test-results/services-<< parameters.suite >>-parity-base
|
||||
docker logs postgres-db > test-results/services-<< parameters.suite >>-parity-base/postgres.log 2>&1 || true
|
||||
docker logs redis-cache > test-results/services-<< parameters.suite >>-parity-base/redis.log 2>&1 || true
|
||||
docker rm -f postgres-db redis-cache
|
||||
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
|
||||
- run:
|
||||
name: Check out head product code
|
||||
command: |
|
||||
git rm -r -f --quiet litellm enterprise litellm-proxy-extras
|
||||
git checkout "$CIRCLE_SHA1" -- litellm enterprise litellm-proxy-extras
|
||||
git reset --quiet
|
||||
test -f litellm/rust_bridge/_native.abi3.so
|
||||
- start_postgres:
|
||||
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
|
||||
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
|
||||
- start_redis
|
||||
- run:
|
||||
name: Run head side
|
||||
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> parity head
|
||||
no_output_timeout: 15m
|
||||
- run:
|
||||
name: Stop head database and Redis
|
||||
when: always
|
||||
command: |
|
||||
mkdir -p test-results/services-<< parameters.suite >>-parity-head
|
||||
docker logs postgres-db > test-results/services-<< parameters.suite >>-parity-head/postgres.log 2>&1 || true
|
||||
docker logs redis-cache > test-results/services-<< parameters.suite >>-parity-head/redis.log 2>&1 || true
|
||||
docker rm -f postgres-db redis-cache
|
||||
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
|
||||
- run:
|
||||
name: Compare routing parity
|
||||
command: PYTHONPATH="$PWD/tests" .venv/bin/python -m integration._support.routing check test-results/parity-<< parameters.suite >>/base test-results/parity-<< parameters.suite >>/head
|
||||
- store_test_results:
|
||||
path: test-results
|
||||
- store_artifacts:
|
||||
path: test-results
|
||||
|
||||
unit:
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
|
|
@ -3224,8 +3305,22 @@ workflows:
|
|||
branches:
|
||||
only: main
|
||||
jobs: *migration_jobs
|
||||
routing_parity:
|
||||
when:
|
||||
not:
|
||||
equal: ["", << pipeline.parameters.routing_parity_base >>]
|
||||
jobs:
|
||||
- routing_parity:
|
||||
name: routing-parity-<< matrix.suite >>
|
||||
matrix:
|
||||
parameters:
|
||||
suite: [management, accounting, database, providers, extensions, cost, mcp]
|
||||
integration:
|
||||
unless: << pipeline.parameters.run_migration_tests >>
|
||||
unless:
|
||||
or:
|
||||
- << pipeline.parameters.run_migration_tests >>
|
||||
- not:
|
||||
equal: ["", << pipeline.parameters.routing_parity_base >>]
|
||||
jobs:
|
||||
- integration_contracts:
|
||||
name: integration-<< matrix.suite >>
|
||||
|
|
@ -3237,8 +3332,23 @@ workflows:
|
|||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- integration_contracts:
|
||||
name: integration-<< matrix.suite >>-replica
|
||||
matrix:
|
||||
parameters:
|
||||
suite: [management, database]
|
||||
mode: [replica]
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
build_and_test:
|
||||
unless: << pipeline.parameters.run_migration_tests >>
|
||||
unless:
|
||||
or:
|
||||
- << pipeline.parameters.run_migration_tests >>
|
||||
- not:
|
||||
equal: ["", << pipeline.parameters.routing_parity_base >>]
|
||||
jobs:
|
||||
- using_litellm_on_windows:
|
||||
filters: &main_branches
|
||||
|
|
|
|||
52
.circleci/scripts/prepare_replica_roles.py
Normal file
52
.circleci/scripts/prepare_replica_roles.py
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
import psycopg
|
||||
|
||||
DATABASE_URL: Final = os.environ["DATABASE_URL"]
|
||||
|
||||
|
||||
def postgres_url() -> str:
|
||||
parsed: Final = urlsplit(DATABASE_URL)
|
||||
return urlunsplit(parsed._replace(path="/postgres"))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
with psycopg.connect(postgres_url(), autocommit=True) as admin:
|
||||
admin.execute("CREATE EXTENSION IF NOT EXISTS pg_stat_statements")
|
||||
admin.execute("CREATE ROLE litellm_writer LOGIN PASSWORD 'litellm-writer' NOSUPERUSER")
|
||||
admin.execute("CREATE ROLE litellm_reader LOGIN PASSWORD 'litellm-reader' NOSUPERUSER NOINHERIT")
|
||||
admin.execute("ALTER ROLE litellm_reader SET default_transaction_read_only = on")
|
||||
admin.execute("ALTER DATABASE circle_test OWNER TO litellm_writer")
|
||||
admin.execute("GRANT CONNECT ON DATABASE circle_test TO litellm_reader")
|
||||
with psycopg.connect(DATABASE_URL, autocommit=True) as admin:
|
||||
admin.execute("GRANT USAGE ON SCHEMA public TO litellm_reader")
|
||||
admin.execute(
|
||||
"ALTER DEFAULT PRIVILEGES FOR ROLE litellm_writer IN SCHEMA public GRANT SELECT ON TABLES TO litellm_reader"
|
||||
)
|
||||
admin.execute("GRANT SELECT ON ALL TABLES IN SCHEMA public TO litellm_reader")
|
||||
|
||||
parsed: Final = urlsplit(DATABASE_URL)
|
||||
reader_url: Final = urlunsplit(
|
||||
parsed._replace(netloc=f"litellm_reader:litellm-reader@{parsed.hostname}:{parsed.port}")
|
||||
)
|
||||
writer_url: Final = urlunsplit(
|
||||
parsed._replace(netloc=f"litellm_writer:litellm-writer@{parsed.hostname}:{parsed.port}")
|
||||
)
|
||||
with psycopg.connect(reader_url, autocommit=True) as reader:
|
||||
assert reader.execute("SHOW transaction_read_only").fetchone() == ("on",)
|
||||
try:
|
||||
reader.execute("CREATE TABLE integration_readonly_probe (id int)")
|
||||
except psycopg.errors.ReadOnlySqlTransaction:
|
||||
pass
|
||||
else:
|
||||
raise AssertionError("litellm_reader executed a write statement")
|
||||
with psycopg.connect(writer_url, autocommit=True) as writer:
|
||||
assert writer.execute("SELECT current_user").fetchone() == ("litellm_writer",)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -7,7 +7,15 @@ if [ "${GITHUB_ACTIONS:-}" = true ]; then
|
|||
fi
|
||||
|
||||
suite="${1:?integration suite required}"
|
||||
results="test-results/integration-${suite}"
|
||||
mode="${2:-standard}"
|
||||
side="${3:-}"
|
||||
if [ "$mode" = replica ]; then
|
||||
results="test-results/integration-${suite}-replica"
|
||||
elif [ "$mode" = parity ]; then
|
||||
results="test-results/parity-${suite}/${side:?parity side required}"
|
||||
else
|
||||
results="test-results/integration-${suite}"
|
||||
fi
|
||||
mkdir -p "$results"
|
||||
integration_identity="$(.venv/bin/python -c 'import uuid; print(uuid.uuid4().hex)')"
|
||||
upstream_pid=""
|
||||
|
|
@ -80,6 +88,18 @@ export INTEGRATION_ORDER_SEED="$INTEGRATION_SEED"
|
|||
|
||||
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma > "$results/prisma-generate.log" 2>&1
|
||||
|
||||
export INTEGRATION_PROXY_DATABASE_URL=""
|
||||
export INTEGRATION_PROXY_READ_REPLICA_URL=""
|
||||
export INTEGRATION_ROUTING=""
|
||||
if [ "$mode" = replica ] || [ "$mode" = parity ]; then
|
||||
.venv/bin/python .circleci/scripts/prepare_replica_roles.py > "$results/prepare-replica-roles.log" 2>&1
|
||||
export INTEGRATION_PROXY_DATABASE_URL="postgresql://litellm_writer:litellm-writer@127.0.0.1:5432/circle_test"
|
||||
export INTEGRATION_PROXY_READ_REPLICA_URL="postgresql://litellm_reader:litellm-reader@127.0.0.1:5432/circle_test"
|
||||
fi
|
||||
if [ "$mode" = parity ]; then
|
||||
export INTEGRATION_ROUTING=capture
|
||||
fi
|
||||
|
||||
sudo iptables -N integration_only
|
||||
guard_created=true
|
||||
sudo iptables -A integration_only -o lo -j ACCEPT
|
||||
|
|
@ -137,8 +157,12 @@ start_proxy() {
|
|||
else
|
||||
cost_map_env=("LITELLM_LOCAL_MODEL_COST_MAP=True")
|
||||
fi
|
||||
local -a database_env=("DATABASE_URL=${INTEGRATION_PROXY_DATABASE_URL:-$DATABASE_URL}")
|
||||
if [ -n "$INTEGRATION_PROXY_READ_REPLICA_URL" ]; then
|
||||
database_env+=("DATABASE_URL_READ_REPLICA=$INTEGRATION_PROXY_READ_REPLICA_URL")
|
||||
fi
|
||||
setsid env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" INTEGRATION_RUN_ID="$integration_identity" \
|
||||
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
|
||||
"${database_env[@]}" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
|
||||
INTEGRATION_UPSTREAM_URL="$INTEGRATION_UPSTREAM_URL" \
|
||||
LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" LITELLM_SALT_KEY="$LITELLM_SALT_KEY" LITELLM_UI_PATH="$LITELLM_UI_PATH" PROXY_BASE_URL="http://127.0.0.1:$port" \
|
||||
LITELLM_MODE=PRODUCTION STORE_MODEL_IN_DB=True "${cost_map_env[@]}" \
|
||||
|
|
@ -195,6 +219,9 @@ env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
|
|||
INTEGRATION_SEED="$INTEGRATION_SEED" \
|
||||
INTEGRATION_ORDER_SEED="$INTEGRATION_ORDER_SEED" \
|
||||
LITELLM_LOCAL_MODEL_COST_MAP=True AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
|
||||
INTEGRATION_PROXY_DATABASE_URL="$INTEGRATION_PROXY_DATABASE_URL" \
|
||||
INTEGRATION_PROXY_READ_REPLICA_URL="$INTEGRATION_PROXY_READ_REPLICA_URL" \
|
||||
INTEGRATION_ROUTING="$INTEGRATION_ROUTING" \
|
||||
.venv/bin/python tests/integration/run.py "$suite" --results "$results"
|
||||
|
||||
if [ "${INTEGRATION_COVERAGE:-0}" = 1 ]; then
|
||||
|
|
|
|||
13
.github/pull_request_template.md
vendored
13
.github/pull_request_template.md
vendored
|
|
@ -1,6 +1,8 @@
|
|||
<!-- The whole description's target audience is humans, not AI agents: write it in plain, simple,
|
||||
everyday engineering language, extremely parsable and readable at a glance. This goes double for
|
||||
the TLDR, User Flow, and Caveats sections -->
|
||||
the TLDR, User Flow, and Caveats sections
|
||||
Drop every section you have nothing to put in, heading included: a bare "## Relevant issues" or
|
||||
"## Affected release" with nothing under it must not appear in the final description -->
|
||||
|
||||
## TLDR
|
||||
|
||||
|
|
@ -21,6 +23,7 @@ How it solves it:
|
|||
<!-- Two ordered lists, Before and After, walking the same end user through the same task, written strictly from that user's seat
|
||||
Read the linked issue, ticket, or customer thread first so the flow reflects the real application and the routes its users actually hit; don't invent a generic scenario
|
||||
Lead each list with one plain sentence saying where the flow fails (Before) or succeeds (After), then number the steps
|
||||
Keep it tight: aim for 3 to 5 steps per list, one line each, roughly 20 words max, and never pad a shorter flow with filler steps to hit the count. Cover the one path the PR changes and fold variants (case, other field, second endpoint) into a clause on the step they belong to rather than their own steps. The example below is the target length
|
||||
Every step is something the user does or observes: the HTTP method and full URL they hit, what they sent, and what visibly came back (status code, error text, the shape of an ID). UI steps name the page URL and what is on screen
|
||||
No LiteLLM internals: never name functions, files, DB tables, config classes, hooks, callbacks, or code paths. "The upload hands back an ID that looks like OpenAI's own `file-abc123` instead of the scrambled one the gateway returned" is right, "no managed-file row was registered" is wrong
|
||||
Keep the two lists step-for-step identical until they diverge, so the changed step is obvious
|
||||
|
|
@ -45,15 +48,15 @@ After: the same request comes back with real token counts, so the dashboard show
|
|||
|
||||
## Relevant issues
|
||||
|
||||
<!-- e.g., "Fixes #000" -->
|
||||
<!-- e.g., "Fixes #000". Drop the section if there is none -->
|
||||
|
||||
## Affected release
|
||||
|
||||
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1", and add the `backport-stable` label so the fix is cherry-picked onto the rc line before the stable is tagged. Leave the section blank otherwise -->
|
||||
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1", and add the `backport-stable` label so the fix is cherry-picked onto the rc line before the stable is tagged. Drop the section otherwise -->
|
||||
|
||||
## Linear ticket
|
||||
|
||||
<!-- if you are an internal contributor, add "Resolves " followed by the Linear ticket e.g., "Resolves LIT-1234" to link the Linear ticket to the GitHub PR. If you don't have one, leave the section blank rather than guessing -->
|
||||
<!-- if you are an internal contributor, add "Resolves " followed by the Linear ticket e.g., "Resolves LIT-1234" to link the Linear ticket to the GitHub PR. If you don't have one, drop the section rather than guessing -->
|
||||
|
||||
## Pre-Submission checklist
|
||||
|
||||
|
|
@ -134,7 +137,7 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac
|
|||
human reader
|
||||
If you assumed something instead of testing it, e.g. "only reproduces with X on" or "no
|
||||
user-observable behavior difference", list it here too with what breaks if it is wrong
|
||||
Leave this section empty if there are none -->
|
||||
Drop this section if there are none -->
|
||||
|
||||
## QA runbook
|
||||
|
||||
|
|
|
|||
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -108,6 +108,7 @@ jobs:
|
|||
tests/test_litellm/completion_extras
|
||||
tests/test_litellm/containers
|
||||
tests/test_litellm/endpoints
|
||||
tests/test_litellm/files
|
||||
tests/test_litellm/images
|
||||
tests/test_litellm/interactions
|
||||
tests/test_litellm/messages
|
||||
|
|
|
|||
|
|
@ -33,11 +33,11 @@ End-to-end tests belong in `tests/e2e/` and must follow the harness conventions
|
|||
|
||||
When creating PRs, target the repository's current default branch for both internal and external / OSS contributions. Check it with `python3 scripts/default_branch.py --branch` instead of assuming a branch name or relying on cached `origin/HEAD`
|
||||
|
||||
When writing a PR body, treat the comments and imperative instructions inside .github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule
|
||||
When writing a PR body, treat the comments and imperative instructions inside .github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule. A section you have nothing to put in (Relevant issues, Affected release, Linear ticket, Caveats, QA runbook, and so on) is removed entirely, heading included, never left as an empty title
|
||||
|
||||
Same applies for filing bug reports and feature requests, with .github/ISSUE_TEMPLATE/bug_report.yml and .github/ISSUE_TEMPLATE/feature_request.yml, respectively
|
||||
|
||||
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just leave the section blank
|
||||
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just drop the section
|
||||
|
||||
Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, or the main use case runs through a headful agentic coding tool like Claude Code or Codex, drive that surface yourself and embed your own before and after screenshots of it in the PR (the Admin UI page, or what the coding tool shows), next to an ordered list of the URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, and what fields to fill out so a reviewer can reproduce it
|
||||
|
||||
|
|
|
|||
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -3414,6 +3414,7 @@ dependencies = [
|
|||
name = "litellm-types"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
]
|
||||
|
|
|
|||
274
litellm-rust/crates/core-utils/src/dot_notation_indexing.rs
Normal file
274
litellm-rust/crates/core-utils/src/dot_notation_indexing.rs
Normal file
|
|
@ -0,0 +1,274 @@
|
|||
use serde_json::Value;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
enum Segment {
|
||||
Field(String),
|
||||
Every,
|
||||
Index(usize),
|
||||
}
|
||||
|
||||
fn parse_segments(path: &str) -> Option<Vec<Segment>> {
|
||||
let mut segments = Vec::new();
|
||||
let mut rest = path;
|
||||
while !rest.is_empty() {
|
||||
if let Some(after_open) = rest.strip_prefix('[') {
|
||||
let (inside, after) = after_open.split_once(']')?;
|
||||
segments.push(match inside {
|
||||
"*" => Segment::Every,
|
||||
index => Segment::Index(index.trim().parse().ok()?),
|
||||
});
|
||||
rest = after.strip_prefix('.').unwrap_or(after);
|
||||
continue;
|
||||
}
|
||||
let end = rest.find(['.', '[']).unwrap_or(rest.len());
|
||||
let (field, after) = rest.split_at(end);
|
||||
if !field.is_empty() {
|
||||
segments.push(Segment::Field(field.to_string()));
|
||||
}
|
||||
rest = after.strip_prefix('.').unwrap_or(after);
|
||||
}
|
||||
Some(segments)
|
||||
}
|
||||
|
||||
fn without_path(value: Value, segments: &[Segment]) -> Value {
|
||||
let Some((segment, tail)) = segments.split_first() else {
|
||||
return value;
|
||||
};
|
||||
match (segment, value) {
|
||||
(Segment::Field(name), Value::Object(object)) => Value::Object(
|
||||
object
|
||||
.into_iter()
|
||||
.filter_map(|(key, item)| {
|
||||
if key != *name {
|
||||
return Some((key, item));
|
||||
}
|
||||
(!tail.is_empty()).then(|| (key, without_path(item, tail)))
|
||||
})
|
||||
.collect(),
|
||||
),
|
||||
(Segment::Every, Value::Array(items)) => Value::Array(
|
||||
items
|
||||
.into_iter()
|
||||
.map(|item| without_path(item, tail))
|
||||
.collect(),
|
||||
),
|
||||
(Segment::Index(index), Value::Array(items)) => Value::Array(
|
||||
items
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(position, item)| {
|
||||
if position == *index {
|
||||
without_path(item, tail)
|
||||
} else {
|
||||
item
|
||||
}
|
||||
})
|
||||
.collect(),
|
||||
),
|
||||
(_, value) => value,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn delete_nested_value(value: Value, path: &str) -> Value {
|
||||
match parse_segments(path) {
|
||||
Some(segments) => without_path(value, &segments),
|
||||
None => value,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[fixture]
|
||||
fn body() -> Value {
|
||||
json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
})
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::top_level_field("top", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}}
|
||||
}))]
|
||||
#[case::whole_object("meta", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::nested_field("meta.inner.drop", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::trailing_dot("meta.inner.drop.", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::leading_and_doubled_dots(".meta..inner.drop", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::field_in_every_element("tools[*].examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::whole_array_field_in_every_element("tools[*].arr", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"]},
|
||||
{"name": "t1", "examples": ["b"]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::field_in_indexed_element("tools[1].examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::padded_index("tools[ 1 ].examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::field_right_after_bracket("tools[0]examples", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::index_then_wildcard("tools[0].arr[*].f", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"k": 1}, {"k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::wildcard_then_index_only_where_it_exists("tools[*].arr[1].f", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
#[case::nested_wildcards("tools[*].arr[*].f", json!({
|
||||
"tools": [
|
||||
{"name": "t0", "examples": ["a"], "arr": [{"k": 1}, {"k": 2}]},
|
||||
{"name": "t1", "examples": ["b"], "arr": [{"k": 3}]}
|
||||
],
|
||||
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
|
||||
"top": 0.7
|
||||
}))]
|
||||
fn deletes_the_addressed_field(body: Value, #[case] path: &str, #[case] expected: Value) {
|
||||
assert_eq!(delete_nested_value(body, path), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty_path("")]
|
||||
#[case::missing_field("missing")]
|
||||
#[case::missing_parent("missing.field")]
|
||||
#[case::field_through_a_scalar("top.value")]
|
||||
#[case::field_on_an_array("tools.name")]
|
||||
#[case::index_on_an_object("meta[0].user")]
|
||||
#[case::wildcard_on_an_object("meta[*].user")]
|
||||
#[case::wildcard_over_scalars("tools[*].examples[*].name")]
|
||||
#[case::index_out_of_range("tools[5].name")]
|
||||
#[case::every_element_itself("tools[*]")]
|
||||
#[case::indexed_element_itself("tools[0]")]
|
||||
#[case::nested_element_itself("tools[*].arr[0]")]
|
||||
#[case::negative_index("tools[-1].name")]
|
||||
#[case::non_numeric_index("tools[x].name")]
|
||||
#[case::empty_index("tools[].name")]
|
||||
#[case::unclosed_bracket("top[0")]
|
||||
fn leaves_the_value_untouched(body: Value, #[case] path: &str) {
|
||||
assert_eq!(delete_nested_value(body.clone(), path), body);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::wildcards_indices_and_nesting(
|
||||
json!({"tools": [
|
||||
{"name": "t0", "configs": [{"id": "c0", "remove_me": 1, "keep": 1}, {"id": "c1", "remove_me": 2, "keep": 2}], "metadata": {"drop_this": 1, "preserve": 1}},
|
||||
{"name": "t1", "configs": [{"id": "c0", "remove_me": 3, "keep": 3}, {"id": "c1", "remove_me": 4, "keep": 4}], "metadata": {"drop_this": 2, "preserve": 2}},
|
||||
{"name": "t2", "configs": [{"id": "c0", "remove_me": 5, "keep": 5}], "metadata": {"drop_this": 3, "preserve": 3}}
|
||||
]}),
|
||||
&["tools[*].configs[1].remove_me", "tools[1].metadata.drop_this", "tools[*].configs[*].id"],
|
||||
json!({"tools": [
|
||||
{"name": "t0", "configs": [{"remove_me": 1, "keep": 1}, {"keep": 2}], "metadata": {"drop_this": 1, "preserve": 1}},
|
||||
{"name": "t1", "configs": [{"remove_me": 3, "keep": 3}, {"keep": 4}], "metadata": {"preserve": 2}},
|
||||
{"name": "t2", "configs": [{"remove_me": 5, "keep": 5}], "metadata": {"drop_this": 3, "preserve": 3}}
|
||||
]}),
|
||||
)]
|
||||
#[case::simple_and_wildcard_nesting(
|
||||
json!({
|
||||
"tools": [{"name": "t1", "simple_nested": {"remove": 1, "keep": 2}, "complex": [{"nested": {"remove": 3, "keep": 4}}]}],
|
||||
"top_level_remove": "should_go",
|
||||
"top_level_keep": "should_stay"
|
||||
}),
|
||||
&["tools[*].simple_nested.remove", "tools[*].complex[*].nested.remove"],
|
||||
json!({
|
||||
"tools": [{"name": "t1", "simple_nested": {"keep": 2}, "complex": [{"nested": {"keep": 4}}]}],
|
||||
"top_level_remove": "should_go",
|
||||
"top_level_keep": "should_stay"
|
||||
}),
|
||||
)]
|
||||
#[case::triple_nested_wildcards(
|
||||
json!({"tools": [{"name": "t1", "arr1": [
|
||||
{"arr2": [{"field": 1, "keep": 1}, {"field": 2, "keep": 2}]},
|
||||
{"arr2": [{"field": 3, "keep": 3}]}
|
||||
]}]}),
|
||||
&["tools[*].arr1[*].arr2[*].field"],
|
||||
json!({"tools": [{"name": "t1", "arr1": [
|
||||
{"arr2": [{"keep": 1}, {"keep": 2}]},
|
||||
{"arr2": [{"keep": 3}]}
|
||||
]}]}),
|
||||
)]
|
||||
fn applies_paths_in_sequence(
|
||||
#[case] value: Value,
|
||||
#[case] paths: &[&str],
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let deleted = paths
|
||||
.iter()
|
||||
.fold(value, |value, path| delete_nested_value(value, path));
|
||||
assert_eq!(deleted, expected);
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,93 @@
|
|||
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub fn get_provider_specific_headers(
|
||||
provider_specific_header: Option<&ProviderSpecificHeaders>,
|
||||
custom_llm_provider: &str,
|
||||
) -> Map<String, Value> {
|
||||
let entries: &[ProviderSpecificHeader] = match provider_specific_header {
|
||||
None => &[],
|
||||
Some(ProviderSpecificHeaders::One(entry)) => std::slice::from_ref(entry),
|
||||
Some(ProviderSpecificHeaders::Many(entries)) => entries,
|
||||
};
|
||||
entries
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
entry
|
||||
.custom_llm_provider
|
||||
.split(',')
|
||||
.any(|scoped| scoped.trim() == custom_llm_provider)
|
||||
})
|
||||
.flat_map(|entry| entry.extra_headers.clone())
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::single_entry_for_the_provider(
|
||||
json!({"custom_llm_provider": "anthropic", "extra_headers": {"Authorization": "Bearer t", "Custom-Header": "v"}}),
|
||||
json!({"Authorization": "Bearer t", "Custom-Header": "v"}),
|
||||
)]
|
||||
#[case::single_entry_for_another_provider(
|
||||
json!({"custom_llm_provider": "openai", "extra_headers": {"Authorization": "Bearer t"}}),
|
||||
json!({}),
|
||||
)]
|
||||
#[case::provider_in_a_comma_separated_scope(
|
||||
json!({"custom_llm_provider": "bedrock,anthropic,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}}),
|
||||
json!({"anthropic-beta": "context-1m-2025-08-07"}),
|
||||
)]
|
||||
#[case::provider_missing_from_a_comma_separated_scope(
|
||||
json!({"custom_llm_provider": "bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "test"}}),
|
||||
json!({}),
|
||||
)]
|
||||
#[case::scope_with_spaces(
|
||||
json!({"custom_llm_provider": "bedrock, anthropic , vertex_ai", "extra_headers": {"anthropic-beta": "test"}}),
|
||||
json!({"anthropic-beta": "test"}),
|
||||
)]
|
||||
#[case::scope_names_must_match_exactly(
|
||||
json!({"custom_llm_provider": "anthropic_text", "extra_headers": {"anthropic-beta": "test"}}),
|
||||
json!({}),
|
||||
)]
|
||||
#[case::entries_scope_independently(
|
||||
json!([
|
||||
{"custom_llm_provider": "anthropic,bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}},
|
||||
{"custom_llm_provider": "bedrock", "extra_headers": {"x-bedrock-only": "no"}},
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"authorization": "Bearer sk-ant-oat01-fake-token"}}
|
||||
]),
|
||||
json!({"anthropic-beta": "context-1m-2025-08-07", "authorization": "Bearer sk-ant-oat01-fake-token"}),
|
||||
)]
|
||||
#[case::later_entries_win(
|
||||
json!([
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "first"}},
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "second"}}
|
||||
]),
|
||||
json!({"x-scoped": "second"}),
|
||||
)]
|
||||
#[case::empty_list(json!([]), json!({}))]
|
||||
#[case::entry_without_scope(json!({"extra_headers": {"x-scoped": "yes"}}), json!({}))]
|
||||
#[case::entry_without_headers(json!({"custom_llm_provider": "anthropic"}), json!({}))]
|
||||
fn provider_specific_headers_match_the_scoped_provider(
|
||||
#[case] configured: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let configured: ProviderSpecificHeaders = serde_json::from_value(configured).unwrap();
|
||||
assert_eq!(
|
||||
Value::Object(get_provider_specific_headers(
|
||||
Some(&configured),
|
||||
"anthropic"
|
||||
)),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn no_configured_headers_match_nothing() {
|
||||
assert_eq!(get_provider_specific_headers(None, "anthropic"), Map::new());
|
||||
}
|
||||
}
|
||||
|
|
@ -1,7 +1,9 @@
|
|||
pub mod call_arguments;
|
||||
pub mod core_helpers;
|
||||
pub mod dot_notation_indexing;
|
||||
pub mod exception_mapping_utils;
|
||||
pub mod get_llm_provider_logic;
|
||||
pub mod get_provider_specific_headers;
|
||||
pub mod params;
|
||||
pub mod prompt_templates;
|
||||
pub mod secret_redaction;
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
use litellm_http::request::string_headers as shared_string_headers;
|
||||
pub(super) use litellm_http::request::{has_bearer_auth, has_header, truncate_error_body};
|
||||
pub(super) use litellm_http::request::truncate_error_body;
|
||||
use litellm_llms::{
|
||||
anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
|
||||
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
|
|
@ -18,8 +20,34 @@ pub enum Error {
|
|||
Transport(#[from] litellm_http::transport::Error),
|
||||
#[error(transparent)]
|
||||
Headers(#[from] litellm_http::request::HeaderError),
|
||||
#[error(transparent)]
|
||||
Secret(#[from] SecretError),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, thiserror::Error)]
|
||||
#[error(transparent)]
|
||||
pub struct SecretError(Arc<litellm_secrets::Error>);
|
||||
|
||||
impl SecretError {
|
||||
pub fn source_error(&self) -> &litellm_secrets::Error {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl From<litellm_secrets::Error> for Error {
|
||||
fn from(error: litellm_secrets::Error) -> Self {
|
||||
Self::Secret(SecretError(Arc::new(error)))
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialEq for SecretError {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
Arc::ptr_eq(&self.0, &other.0)
|
||||
}
|
||||
}
|
||||
|
||||
impl Eq for SecretError {}
|
||||
|
||||
impl From<LlmError> for Error {
|
||||
fn from(error: LlmError) -> Self {
|
||||
match error {
|
||||
|
|
|
|||
|
|
@ -12,6 +12,9 @@ mod common_utils;
|
|||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_secrets::source::EnvironmentSecrets;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine};
|
||||
use serde_json::Value;
|
||||
|
|
@ -31,9 +34,12 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesR
|
|||
api_base: request.api_base.map(Into::into),
|
||||
custom_llm_provider: request.custom_llm_provider.map(Into::into),
|
||||
extra_headers: request.extra_headers,
|
||||
provider_specific_header: request.provider_specific_header,
|
||||
timeout: request.timeout,
|
||||
shaping: request.shaping,
|
||||
};
|
||||
match litellm_host::run::run(messages_machine(), &LocalMessagesHost::new(call)).await? {
|
||||
let secrets = Arc::new(EnvironmentSecrets::python_compatible());
|
||||
match litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await? {
|
||||
MessagesOutput::Message(message) => Ok(*message),
|
||||
MessagesOutput::Streamed => Err(Error::Unsupported(
|
||||
"streamed responses need a streaming host",
|
||||
|
|
|
|||
|
|
@ -1,51 +1,102 @@
|
|||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_llms::base_llm::anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, MessagesAuthStrategy,
|
||||
use litellm_core_utils::{
|
||||
dot_notation_indexing::delete_nested_value,
|
||||
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
|
||||
get_provider_specific_headers::get_provider_specific_headers,
|
||||
settings::Lookup,
|
||||
};
|
||||
use litellm_llms::{
|
||||
anthropic::experimental_pass_through::messages::handler::shape_anthropic_messages_request,
|
||||
base_llm::anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, MessagesTransformContext,
|
||||
},
|
||||
};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers},
|
||||
common_utils::{messages_provider_config, string_headers},
|
||||
};
|
||||
use crate::messages::types::{MessagesRequest, ProviderMessagesRequest};
|
||||
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: MessagesRequest<'_>,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
|
||||
pub(super) struct ResolvedProvider<'a> {
|
||||
pub(super) model: &'a str,
|
||||
pub(super) provider: &'a str,
|
||||
pub(super) config: &'static dyn BaseAnthropicMessagesConfig,
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
) -> Result<ResolvedProvider<'a>, Error> {
|
||||
let CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
} = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
request
|
||||
.custom_llm_provider
|
||||
.map(|provider| CustomLlmProvider {
|
||||
model: request.model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for messages request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let model = provider_info.model.to_string();
|
||||
let provider = provider_info.custom_llm_provider;
|
||||
|
||||
let config = messages_provider_config(provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider.to_string()))?;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
Ok(ResolvedProvider {
|
||||
model,
|
||||
provider,
|
||||
config,
|
||||
})
|
||||
}
|
||||
|
||||
let headers =
|
||||
validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?;
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: MessagesRequest<'_>,
|
||||
resolved: ResolvedProvider<'_>,
|
||||
secrets: &dyn Lookup,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let ResolvedProvider {
|
||||
model,
|
||||
provider,
|
||||
config,
|
||||
} = resolved;
|
||||
let model = model.to_string();
|
||||
let env_lookup = |key: &str| secrets.get(key);
|
||||
|
||||
let typed_request: AnthropicMessagesRequest =
|
||||
serde_json::from_value(request.body).map_err(|err| {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
})?;
|
||||
let transformed = config.transform_anthropic_messages_request(AnthropicMessagesRequest {
|
||||
model: model.clone(),
|
||||
..typed_request
|
||||
})?;
|
||||
serde_json::from_value(request.body).map_err(invalid_request)?;
|
||||
let sanitized = shape_anthropic_messages_request(
|
||||
AnthropicMessagesRequest {
|
||||
model: model.clone(),
|
||||
..typed_request
|
||||
},
|
||||
request.shaping.reasoning_auto_summary,
|
||||
)?;
|
||||
let trimmed =
|
||||
without_additional_drop_params(sanitized, &request.shaping.additional_drop_params)?;
|
||||
let transformed = config.transform_anthropic_messages_request(
|
||||
trimmed,
|
||||
&MessagesTransformContext::new(request.shaping.capabilities, request.shaping.drop_params),
|
||||
)?;
|
||||
|
||||
let scoped = get_provider_specific_headers(request.provider_specific_header.as_ref(), provider);
|
||||
let forwarded = string_headers(Some(
|
||||
request
|
||||
.extra_headers
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(scoped)
|
||||
.collect(),
|
||||
))?;
|
||||
let authenticated = config.authenticate(forwarded, request.api_key, &env_lookup)?;
|
||||
let headers = config.request_headers(
|
||||
with_default_headers(authenticated, config.default_headers()),
|
||||
&transformed,
|
||||
);
|
||||
|
||||
let body = serde_json::to_value(transformed).map_err(|err| {
|
||||
Error::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
|
|
@ -65,33 +116,371 @@ pub(super) fn prepare_provider_request(
|
|||
})
|
||||
}
|
||||
|
||||
fn validate_environment(
|
||||
config: &dyn BaseAnthropicMessagesConfig,
|
||||
extra_headers: Option<Map<String, Value>>,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
let mut headers = string_headers(extra_headers)?;
|
||||
|
||||
let auth_strategy = config.auth_strategy();
|
||||
let already_authorized = has_header(&headers, auth_strategy.header_name())
|
||||
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers));
|
||||
if !already_authorized {
|
||||
let api_key = config.resolve_api_key(api_key, env_lookup)?;
|
||||
let auth_header = match auth_strategy {
|
||||
MessagesAuthStrategy::Bearer => {
|
||||
("authorization".to_string(), format!("Bearer {api_key}"))
|
||||
}
|
||||
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
|
||||
};
|
||||
headers.push(auth_header);
|
||||
}
|
||||
|
||||
for (name, value) in config.default_headers() {
|
||||
if !has_header(&headers, name) {
|
||||
headers.push((name.to_string(), value.to_string()));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(headers)
|
||||
fn invalid_request(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
}
|
||||
|
||||
fn without_additional_drop_params(
|
||||
request: AnthropicMessagesRequest,
|
||||
paths: &[String],
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
if paths.is_empty() {
|
||||
return Ok(request);
|
||||
}
|
||||
let Value::Object(fields) = serde_json::to_value(request).map_err(invalid_request)? else {
|
||||
return Err(Error::InvalidRequest(
|
||||
"Anthropic messages request did not serialize to an object".to_string(),
|
||||
));
|
||||
};
|
||||
let (required, optional): (Map<String, Value>, Map<String, Value>) = fields
|
||||
.into_iter()
|
||||
.partition(|(key, _)| matches!(key.as_str(), "model" | "messages"));
|
||||
let trimmed = paths.iter().fold(Value::Object(optional), |body, path| {
|
||||
delete_nested_value(body, path)
|
||||
});
|
||||
let merged: Map<String, Value> = required
|
||||
.into_iter()
|
||||
.chain(trimmed.as_object().cloned().unwrap_or_default())
|
||||
.collect();
|
||||
serde_json::from_value(Value::Object(merged)).map_err(invalid_request)
|
||||
}
|
||||
|
||||
fn with_default_headers(
|
||||
headers: Vec<(String, String)>,
|
||||
defaults: &[(&str, &str)],
|
||||
) -> Vec<(String, String)> {
|
||||
let missing: Vec<(String, String)> = defaults
|
||||
.iter()
|
||||
.filter(|(name, _)| {
|
||||
!headers
|
||||
.iter()
|
||||
.any(|(header, _)| header.eq_ignore_ascii_case(name))
|
||||
})
|
||||
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
|
||||
.collect();
|
||||
headers.into_iter().chain(missing).collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
use crate::messages::types::MessagesShaping;
|
||||
|
||||
#[fixture]
|
||||
fn shaping() -> MessagesShaping {
|
||||
MessagesShaping::default()
|
||||
}
|
||||
|
||||
fn prepare(request: MessagesRequest<'_>) -> Result<ProviderMessagesRequest, Error> {
|
||||
prepare_with_secrets(request, &|_: &str| None)
|
||||
}
|
||||
|
||||
fn prepare_with_secrets(
|
||||
request: MessagesRequest<'_>,
|
||||
secrets: &dyn Lookup,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let resolved = resolve_provider(request.model, request.custom_llm_provider)?;
|
||||
prepare_provider_request(request, resolved, secrets)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key(
|
||||
&[("ANTHROPIC_API_KEY", "sk-secret")],
|
||||
&[("x-api-key", "sk-secret")],
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
)]
|
||||
#[case::auth_token(
|
||||
&[("ANTHROPIC_AUTH_TOKEN", "token")],
|
||||
&[("authorization", "Bearer token")],
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
)]
|
||||
#[case::api_base(
|
||||
&[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_API_BASE", "https://gateway.test")],
|
||||
&[("x-api-key", "sk-secret")],
|
||||
"https://gateway.test/v1/messages"
|
||||
)]
|
||||
#[case::sdk_base_url(
|
||||
&[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_BASE_URL", "https://sdk.test")],
|
||||
&[("x-api-key", "sk-secret")],
|
||||
"https://sdk.test/v1/messages"
|
||||
)]
|
||||
fn credentials_and_base_come_from_the_resolved_secrets(
|
||||
shaping: MessagesShaping,
|
||||
#[case] secrets: &[(&str, &str)],
|
||||
#[case] expected_auth: &[(&str, &str)],
|
||||
#[case] expected_url: &str,
|
||||
) {
|
||||
let lookup = |name: &str| {
|
||||
secrets
|
||||
.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string())
|
||||
};
|
||||
let prepared = prepare_with_secrets(
|
||||
MessagesRequest {
|
||||
model: "claude-test",
|
||||
body: json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("anthropic"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: None,
|
||||
shaping,
|
||||
},
|
||||
&lookup,
|
||||
)
|
||||
.unwrap();
|
||||
let auth: Vec<(&str, &str)> = prepared
|
||||
.upstream_headers
|
||||
.iter()
|
||||
.filter(|(name, _)| matches!(name.as_str(), "x-api-key" | "authorization"))
|
||||
.map(|(name, value)| (name.as_str(), value.as_str()))
|
||||
.collect();
|
||||
assert_eq!(
|
||||
(auth.as_slice(), prepared.url.as_str()),
|
||||
(expected_auth, expected_url)
|
||||
);
|
||||
}
|
||||
|
||||
fn prepared_body(body: Value, shaping: MessagesShaping) -> Result<Value, Error> {
|
||||
prepare(MessagesRequest {
|
||||
model: "anthropic/claude-test",
|
||||
body,
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some("https://anthropic.test"),
|
||||
custom_llm_provider: Some("anthropic"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: None,
|
||||
shaping,
|
||||
})
|
||||
.map(|prepared| prepared.body)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::nothing_forwarded(
|
||||
&[],
|
||||
&[("x-version", "1"), ("content-type", "application/json")],
|
||||
&[("x-version", "1"), ("content-type", "application/json")],
|
||||
)]
|
||||
#[case::forwarded_header_wins_in_any_case(
|
||||
&[("X-Version", "custom"), ("x-api-key", "k")],
|
||||
&[("x-version", "1"), ("content-type", "application/json")],
|
||||
&[("X-Version", "custom"), ("x-api-key", "k"), ("content-type", "application/json")],
|
||||
)]
|
||||
#[case::no_defaults(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])]
|
||||
fn default_headers_fill_only_missing_names(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] defaults: &[(&str, &str)],
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
let owned = |headers: &[(&str, &str)]| -> Vec<(String, String)> {
|
||||
headers
|
||||
.iter()
|
||||
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
|
||||
.collect()
|
||||
};
|
||||
assert_eq!(
|
||||
with_default_headers(owned(forwarded), defaults),
|
||||
owned(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::top_level_and_nested_paths(
|
||||
json!({
|
||||
"max_tokens": 1024,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2048},
|
||||
"context_management": {"edits": [{"type": "clear_thinking_20251015"}]},
|
||||
"metadata": {"user_id": "u1"},
|
||||
"tools": [{"name": "lookup", "input_schema": {"type": "object"}, "input_examples": [{"q": "x"}]}]
|
||||
}),
|
||||
&["thinking", "context_management", "tools[*].input_examples"],
|
||||
json!({
|
||||
"max_tokens": 1024,
|
||||
"metadata": {"user_id": "u1"},
|
||||
"tools": [{"name": "lookup", "input_schema": {"type": "object"}}]
|
||||
}),
|
||||
)]
|
||||
#[case::no_paths(
|
||||
json!({"max_tokens": 16, "safeguards": [{"type": "dangerous_tool_use"}]}),
|
||||
&[],
|
||||
json!({"max_tokens": 16, "safeguards": [{"type": "dangerous_tool_use"}]}),
|
||||
)]
|
||||
#[case::model_and_messages_are_never_dropped(
|
||||
json!({"max_tokens": 16}),
|
||||
&["model", "messages", "messages[0].content"],
|
||||
json!({"max_tokens": 16}),
|
||||
)]
|
||||
fn prepared_body_drops_configured_paths(
|
||||
shaping: MessagesShaping,
|
||||
#[case] fields: Value,
|
||||
#[case] additional_drop_params: &[&str],
|
||||
#[case] expected_fields: Value,
|
||||
) {
|
||||
let with_messages = |fields: Value| -> Value {
|
||||
let Value::Object(fields) = fields else {
|
||||
unreachable!()
|
||||
};
|
||||
Value::Object(
|
||||
[
|
||||
("model".to_string(), json!("claude-test")),
|
||||
(
|
||||
"messages".to_string(),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
),
|
||||
]
|
||||
.into_iter()
|
||||
.chain(fields)
|
||||
.collect(),
|
||||
)
|
||||
};
|
||||
let shaping = MessagesShaping {
|
||||
additional_drop_params: additional_drop_params
|
||||
.iter()
|
||||
.map(ToString::to_string)
|
||||
.collect(),
|
||||
..shaping
|
||||
};
|
||||
assert_eq!(
|
||||
prepared_body(with_messages(fields), shaping),
|
||||
Ok(with_messages(expected_fields))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::model_prefix_picks_the_provider(
|
||||
"azure_ai/claude-test",
|
||||
None,
|
||||
&[("x-priority", "extra"), ("x-scoped", "azure_ai")]
|
||||
)]
|
||||
#[case::explicit_provider(
|
||||
"claude-test",
|
||||
Some("anthropic"),
|
||||
&[("x-priority", "scoped"), ("x-scoped", "anthropic")]
|
||||
)]
|
||||
#[case::provider_prefix_on_an_anthropic_model(
|
||||
"anthropic/claude-test",
|
||||
None,
|
||||
&[("x-priority", "scoped"), ("x-scoped", "anthropic")]
|
||||
)]
|
||||
fn provider_specific_headers_follow_the_resolved_provider(
|
||||
shaping: MessagesShaping,
|
||||
#[case] model: &str,
|
||||
#[case] custom_llm_provider: Option<&str>,
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
let configured: ProviderSpecificHeaders = serde_json::from_value(json!([
|
||||
{"custom_llm_provider": "azure_ai", "extra_headers": {"x-scoped": "azure_ai"}},
|
||||
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "anthropic", "x-priority": "scoped"}}
|
||||
]))
|
||||
.unwrap();
|
||||
let prepared = prepare(MessagesRequest {
|
||||
model,
|
||||
body: json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some("https://resource.services.ai.azure.com"),
|
||||
custom_llm_provider,
|
||||
extra_headers: Some(serde_json::from_value(json!({"x-priority": "extra"})).unwrap()),
|
||||
provider_specific_header: Some(configured),
|
||||
timeout: None,
|
||||
shaping,
|
||||
})
|
||||
.unwrap();
|
||||
let caller_headers: Vec<(&str, &str)> = prepared
|
||||
.upstream_headers
|
||||
.iter()
|
||||
.filter(|(name, _)| matches!(name.as_str(), "x-priority" | "x-scoped"))
|
||||
.map(|(name, value)| (name.as_str(), value.as_str()))
|
||||
.collect();
|
||||
assert_eq!(caller_headers, expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn prepared_body_carries_the_provider_stripped_model(shaping: MessagesShaping) {
|
||||
assert_eq!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "anthropic/claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Ok(json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn dropped_thinking_display_is_not_restored_by_auto_summary(shaping: MessagesShaping) {
|
||||
let shaping = MessagesShaping {
|
||||
reasoning_auto_summary: true,
|
||||
additional_drop_params: vec!["thinking.display".to_string()],
|
||||
..shaping
|
||||
};
|
||||
assert_eq!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 4096,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2048}
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Ok(json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 4096,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2048}
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn dropping_an_invalid_metadata_user_id_does_not_skip_its_validation(shaping: MessagesShaping) {
|
||||
let shaping = MessagesShaping {
|
||||
additional_drop_params: vec!["metadata.user_id".to_string()],
|
||||
..shaping
|
||||
};
|
||||
assert!(matches!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16,
|
||||
"metadata": {"user_id": 123}
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Err(Error::InvalidRequest(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn prepared_body_rejects_invalid_metadata_before_the_call(shaping: MessagesShaping) {
|
||||
assert_eq!(
|
||||
prepared_body(
|
||||
json!({
|
||||
"model": "claude-test",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16,
|
||||
"metadata": {"user_id": 123}
|
||||
}),
|
||||
shaping,
|
||||
),
|
||||
Err(Error::InvalidRequest(
|
||||
"metadata.user_id must be a string, got 123".to_string()
|
||||
))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,7 @@
|
|||
use std::{sync::Mutex, time::Duration};
|
||||
use std::{
|
||||
sync::{Arc, Mutex},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_auth::SecretValue;
|
||||
|
|
@ -9,15 +12,19 @@ use litellm_host::{
|
|||
machine::{HostChannel, MachineFault, RouteMachine},
|
||||
route::Route,
|
||||
};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse,
|
||||
utils::ProviderSpecificHeaders,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::messages_provider_config,
|
||||
handler::{decode_response, network, provider_error, send},
|
||||
prepare::prepare_provider_request,
|
||||
types::MessagesRequest,
|
||||
prepare::{prepare_provider_request, resolve_provider},
|
||||
types::{MessagesRequest, MessagesShaping},
|
||||
};
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
|
||||
|
|
@ -38,7 +45,9 @@ pub struct MessagesCall {
|
|||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub provider_specific_header: Option<ProviderSpecificHeaders>,
|
||||
pub timeout: Option<Duration>,
|
||||
pub shaping: MessagesShaping,
|
||||
}
|
||||
|
||||
impl MessagesCall {
|
||||
|
|
@ -120,22 +129,33 @@ impl Host<Messages> for LocalMessagesHost {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn messages_machine() -> MessagesMachine {
|
||||
RouteMachine::new(|host| Box::pin(execute(host)))
|
||||
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
|
||||
RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
|
||||
}
|
||||
|
||||
async fn execute(host: MessagesHost) -> Result<MessagesOutput, Error> {
|
||||
async fn execute(
|
||||
host: MessagesHost,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?;
|
||||
let stream = call.streams();
|
||||
let request = prepare_provider_request(MessagesRequest {
|
||||
model: &call.model,
|
||||
body: Value::Object(call.body.clone()),
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers.clone(),
|
||||
timeout: call.timeout,
|
||||
})?;
|
||||
let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?;
|
||||
let secrets = secrets.resolve(resolved.config.secret_names()).await?;
|
||||
let request = prepare_provider_request(
|
||||
MessagesRequest {
|
||||
model: &call.model,
|
||||
body: Value::Object(call.body.clone()),
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers.clone(),
|
||||
provider_specific_header: call.provider_specific_header.clone(),
|
||||
timeout: call.timeout,
|
||||
shaping: call.shaping.clone(),
|
||||
},
|
||||
resolved,
|
||||
secrets.as_ref(),
|
||||
)?;
|
||||
if stream && request.provider != ANTHROPIC_MESSAGES_PROVIDER {
|
||||
return Err(Error::Unsupported("streaming messages for this provider"));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,8 @@
|
|||
use std::time::Duration;
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use futures_util::future::BoxFuture;
|
||||
use litellm_http::request::{has_bearer_auth, has_header};
|
||||
use litellm_secrets::{SecretValue, source::SecretSource};
|
||||
use serde_json::{Map, Value, json};
|
||||
use tokio::{
|
||||
io::{AsyncReadExt, AsyncWriteExt},
|
||||
|
|
@ -8,12 +11,132 @@ use tokio::{
|
|||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::{
|
||||
has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body,
|
||||
},
|
||||
common_utils::{messages_provider_config, string_headers, truncate_error_body},
|
||||
messages,
|
||||
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
|
||||
};
|
||||
use crate::messages::types::MessagesRequest;
|
||||
use crate::messages::types::{MessagesRequest, MessagesShaping};
|
||||
|
||||
struct RecordingSecrets {
|
||||
values: Vec<(&'static str, String)>,
|
||||
fails: bool,
|
||||
requested: std::sync::Mutex<Vec<String>>,
|
||||
}
|
||||
|
||||
impl RecordingSecrets {
|
||||
fn new(values: Vec<(&'static str, String)>, fails: bool) -> Self {
|
||||
Self {
|
||||
values,
|
||||
fails,
|
||||
requested: std::sync::Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl SecretSource for RecordingSecrets {
|
||||
fn get_secret_str<'a>(
|
||||
&'a self,
|
||||
name: &'a str,
|
||||
) -> BoxFuture<'a, Result<Option<SecretValue>, litellm_secrets::Error>> {
|
||||
Box::pin(async move {
|
||||
self.requested.lock().unwrap().push(name.to_string());
|
||||
if self.fails {
|
||||
return Err(litellm_secrets::Error::ManagedSecretMissing);
|
||||
}
|
||||
Ok(self
|
||||
.values
|
||||
.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| SecretValue::new(value.clone())))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn secrets_call() -> MessagesCall {
|
||||
let Value::Object(body) = json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}) else {
|
||||
unreachable!("literal object")
|
||||
};
|
||||
MessagesCall {
|
||||
model: "claude-sonnet-4-5".into(),
|
||||
body,
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("anthropic".into()),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn route_reads_the_provider_credential_and_base_from_the_secret_source() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
let addr = listener.local_addr().expect("addr");
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
||||
let request = read_http_request(&mut socket).await;
|
||||
let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}"#;
|
||||
socket
|
||||
.write_all(write_response(response_body).as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
request
|
||||
});
|
||||
let secrets = Arc::new(RecordingSecrets::new(
|
||||
vec![
|
||||
("ANTHROPIC_API_KEY", "sk-from-manager".to_string()),
|
||||
("ANTHROPIC_BASE_URL", format!("http://{addr}")),
|
||||
],
|
||||
false,
|
||||
));
|
||||
|
||||
let output = litellm_host::run::run(
|
||||
messages_machine(secrets.clone()),
|
||||
&LocalMessagesHost::new(secrets_call()),
|
||||
)
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
||||
assert!(matches!(output, MessagesOutput::Message(_)));
|
||||
let request = server.await.expect("server task completes");
|
||||
assert!(
|
||||
request
|
||||
.to_ascii_lowercase()
|
||||
.contains("x-api-key: sk-from-manager"),
|
||||
"{request}"
|
||||
);
|
||||
let requested = secrets.requested.lock().unwrap().clone();
|
||||
assert_eq!(
|
||||
requested,
|
||||
messages_provider_config("anthropic")
|
||||
.unwrap()
|
||||
.secret_names()
|
||||
.iter()
|
||||
.map(ToString::to_string)
|
||||
.collect::<Vec<_>>()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn route_surfaces_a_secret_manager_failure_before_the_call() {
|
||||
let Err(error) = litellm_host::run::run(
|
||||
messages_machine(Arc::new(RecordingSecrets::new(Vec::new(), true))),
|
||||
&LocalMessagesHost::new(secrets_call()),
|
||||
)
|
||||
.await
|
||||
else {
|
||||
panic!("a secret manager failure fails the call");
|
||||
};
|
||||
assert!(
|
||||
matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)),
|
||||
"{error:?}"
|
||||
);
|
||||
}
|
||||
|
||||
async fn read_http_request(socket: &mut TcpStream) -> String {
|
||||
let mut request = Vec::new();
|
||||
|
|
@ -159,7 +282,9 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through()
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
|
@ -215,7 +340,9 @@ async fn messages_round_trip_builds_native_anthropic_request() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("anthropic"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
|
@ -268,7 +395,9 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: Some(headers),
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
|
@ -322,7 +451,9 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: Some(headers),
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("entra id request succeeds without api key");
|
||||
|
|
@ -346,7 +477,9 @@ async fn messages_requires_auth_when_no_key_and_no_header() {
|
|||
api_base: Some("http://127.0.0.1:1"),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_millis(50)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect_err("missing auth errors");
|
||||
|
|
@ -384,7 +517,9 @@ async fn messages_ignores_malformed_authorization_and_uses_api_key() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: Some(headers),
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect("falls back to api key");
|
||||
|
|
@ -425,7 +560,9 @@ async fn messages_maps_provider_error_status_to_http_error() {
|
|||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: Some("azure_ai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect_err("provider error propagates");
|
||||
|
|
@ -445,7 +582,9 @@ async fn messages_rejects_unsupported_provider() {
|
|||
api_base: Some("http://127.0.0.1:1"),
|
||||
custom_llm_provider: Some("openai"),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
timeout: Some(Duration::from_millis(50)),
|
||||
shaping: MessagesShaping::default(),
|
||||
})
|
||||
.await
|
||||
.expect_err("unsupported provider errors");
|
||||
|
|
|
|||
|
|
@ -1,8 +1,25 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig;
|
||||
use litellm_llms::{
|
||||
anthropic::common_utils::AnthropicModelCapabilities,
|
||||
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
|
||||
};
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct MessagesShaping {
|
||||
#[serde(default)]
|
||||
pub capabilities: AnthropicModelCapabilities,
|
||||
#[serde(default)]
|
||||
pub drop_params: bool,
|
||||
#[serde(default)]
|
||||
pub reasoning_auto_summary: bool,
|
||||
#[serde(default)]
|
||||
pub additional_drop_params: Vec<String>,
|
||||
}
|
||||
|
||||
pub struct MessagesRequest<'a> {
|
||||
pub model: &'a str,
|
||||
pub body: Value,
|
||||
|
|
@ -10,7 +27,9 @@ pub struct MessagesRequest<'a> {
|
|||
pub api_base: Option<&'a str>,
|
||||
pub custom_llm_provider: Option<&'a str>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub provider_specific_header: Option<ProviderSpecificHeaders>,
|
||||
pub timeout: Option<Duration>,
|
||||
pub shaping: MessagesShaping,
|
||||
}
|
||||
|
||||
pub struct ProviderMessagesRequest {
|
||||
|
|
@ -22,3 +41,86 @@ pub struct ProviderMessagesRequest {
|
|||
pub upstream_headers: Vec<(String, String)>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::anthropic::common_utils::SupportedEffortTiers;
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::nothing_projected(json!({}), MessagesShaping::default())]
|
||||
#[case::only_drop_params(
|
||||
json!({"drop_params": true}),
|
||||
MessagesShaping { drop_params: true, ..MessagesShaping::default() },
|
||||
)]
|
||||
#[case::only_reasoning_auto_summary(
|
||||
json!({"reasoning_auto_summary": true}),
|
||||
MessagesShaping { reasoning_auto_summary: true, ..MessagesShaping::default() },
|
||||
)]
|
||||
#[case::only_additional_drop_params(
|
||||
json!({"additional_drop_params": ["tools[*].input_examples"]}),
|
||||
MessagesShaping {
|
||||
additional_drop_params: vec!["tools[*].input_examples".to_string()],
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
)]
|
||||
#[case::partial_capabilities(
|
||||
json!({"capabilities": {"supports_reasoning": true}}),
|
||||
MessagesShaping {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
)]
|
||||
#[case::everything_the_python_host_projects(
|
||||
json!({
|
||||
"capabilities": {
|
||||
"supports_reasoning": true,
|
||||
"supports_adaptive_thinking": true,
|
||||
"thinking_always_on": false,
|
||||
"supports_legacy_thinking": false,
|
||||
"supports_output_config": true,
|
||||
"supports_sampling_params": false,
|
||||
"supports_speed": true,
|
||||
"effort_tiers": {"minimal": false, "low": true, "medium": true, "high": true, "xhigh": true, "max": false}
|
||||
},
|
||||
"drop_params": true,
|
||||
"reasoning_auto_summary": true,
|
||||
"additional_drop_params": ["metadata.user_id", "thinking"]
|
||||
}),
|
||||
MessagesShaping {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
thinking_always_on: false,
|
||||
supports_legacy_thinking: false,
|
||||
supports_output_config: true,
|
||||
supports_sampling_params: false,
|
||||
supports_speed: true,
|
||||
effort_tiers: SupportedEffortTiers {
|
||||
minimal: false,
|
||||
low: true,
|
||||
medium: true,
|
||||
high: true,
|
||||
xhigh: true,
|
||||
max: false,
|
||||
},
|
||||
},
|
||||
drop_params: true,
|
||||
reasoning_auto_summary: true,
|
||||
additional_drop_params: vec!["metadata.user_id".to_string(), "thinking".to_string()],
|
||||
},
|
||||
)]
|
||||
fn shaping_deserializes_with_defaults_for_absent_fields(
|
||||
#[case] projected: Value,
|
||||
#[case] expected: MessagesShaping,
|
||||
) {
|
||||
let shaping: MessagesShaping = serde_json::from_value(projected).unwrap();
|
||||
assert_eq!(shaping, expected);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
1568
litellm-rust/crates/llms/src/anthropic/common_utils.rs
Normal file
1568
litellm-rust/crates/llms/src/anthropic/common_utils.rs
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -0,0 +1,270 @@
|
|||
use litellm_types::llms::anthropic_messages::anthropic_request::{
|
||||
AnthropicMessage, AnthropicMessagesRequest,
|
||||
};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::{
|
||||
anthropic::common_utils::{
|
||||
flatten_unencrypted_web_search_results, sanitize_tool_use_ids, strip_empty_content_blocks,
|
||||
strip_provider_specific_fields,
|
||||
},
|
||||
base_llm::chat::transformation::Error,
|
||||
};
|
||||
|
||||
pub fn shape_anthropic_messages_request(
|
||||
request: AnthropicMessagesRequest,
|
||||
reasoning_auto_summary: bool,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
Ok(AnthropicMessagesRequest {
|
||||
messages: sanitize_anthropic_messages(request.messages),
|
||||
metadata: request
|
||||
.metadata
|
||||
.as_ref()
|
||||
.map(validate_anthropic_api_metadata)
|
||||
.transpose()?,
|
||||
thinking: with_reasoning_auto_summary(request.thinking, reasoning_auto_summary),
|
||||
..request
|
||||
})
|
||||
}
|
||||
|
||||
fn sanitize_anthropic_messages(messages: Vec<AnthropicMessage>) -> Vec<AnthropicMessage> {
|
||||
strip_provider_specific_fields(flatten_unencrypted_web_search_results(
|
||||
sanitize_tool_use_ids(strip_empty_content_blocks(messages)),
|
||||
))
|
||||
}
|
||||
|
||||
fn validate_anthropic_api_metadata(metadata: &Value) -> Result<Value, Error> {
|
||||
let Value::Object(fields) = metadata else {
|
||||
return Err(Error::InvalidRequest(format!(
|
||||
"metadata must be an object, got {metadata}"
|
||||
)));
|
||||
};
|
||||
match fields.get("user_id") {
|
||||
None | Some(Value::Null) => Ok(json!({})),
|
||||
Some(Value::String(user_id)) => Ok(json!({"user_id": user_id})),
|
||||
Some(other) => Err(Error::InvalidRequest(format!(
|
||||
"metadata.user_id must be a string, got {other}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn with_reasoning_auto_summary(thinking: Option<Value>, enabled: bool) -> Option<Value> {
|
||||
let Some(Value::Object(thinking)) = thinking else {
|
||||
return thinking;
|
||||
};
|
||||
if !enabled || thinking.get("type").and_then(Value::as_str) == Some("disabled") {
|
||||
return Some(Value::Object(thinking));
|
||||
}
|
||||
Some(Value::Object(
|
||||
thinking
|
||||
.into_iter()
|
||||
.filter(|(key, _)| key != "display")
|
||||
.chain([("display".to_string(), json!("summarized"))])
|
||||
.collect(),
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn messages(value: Value) -> Vec<AnthropicMessage> {
|
||||
serde_json::from_value(value).unwrap()
|
||||
}
|
||||
|
||||
fn request(body: Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(body).unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty_text_next_to_a_tool_use(
|
||||
json!([{"role": "assistant", "content": [
|
||||
{"type": "text", "text": " "},
|
||||
{"type": "tool_use", "id": "t", "name": "B", "input": {}}
|
||||
]}]),
|
||||
json!([{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "t", "name": "B", "input": {}}
|
||||
]}]),
|
||||
)]
|
||||
#[case::cross_provider_tool_ids(
|
||||
json!([
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]}
|
||||
]),
|
||||
json!([
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]}
|
||||
]),
|
||||
)]
|
||||
#[case::replayed_unencrypted_web_search_results(
|
||||
json!([
|
||||
{"role": "user", "content": "latest litellm version?"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "latest litellm version"}},
|
||||
{"type": "web_search_tool_result", "tool_use_id": "srvtoolu_1", "content": [{
|
||||
"type": "web_search_result",
|
||||
"url": "https://github.com/BerriAI/litellm/releases",
|
||||
"title": "Releases",
|
||||
"page_age": null,
|
||||
"encrypted_content": "",
|
||||
"snippet": "Latest release v1.95.0"
|
||||
}]}
|
||||
]},
|
||||
{"role": "user", "content": "which version?"}
|
||||
]),
|
||||
json!([
|
||||
{"role": "user", "content": "latest litellm version?"},
|
||||
{"role": "assistant", "content": [{
|
||||
"type": "text",
|
||||
"text": "Web search results for 'latest litellm version':\n\nTitle: Releases\nURL: https://github.com/BerriAI/litellm/releases\nSnippet: Latest release v1.95.0"
|
||||
}]},
|
||||
{"role": "user", "content": "which version?"}
|
||||
]),
|
||||
)]
|
||||
#[case::replayed_provider_specific_fields(
|
||||
json!([
|
||||
{"role": "assistant", "content": [{
|
||||
"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"},
|
||||
"provider_specific_fields": {"signature": "sig_abc"}
|
||||
}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]}
|
||||
]),
|
||||
json!([
|
||||
{"role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"}}]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]}
|
||||
]),
|
||||
)]
|
||||
#[case::ids_are_normalized_before_web_search_results_flatten(
|
||||
json!([
|
||||
{"role": "user", "content": "run it"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "thinking", "thinking": "", "signature": "sig"},
|
||||
{"type": "text", "text": ""},
|
||||
{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}, "provider_specific_fields": {"x": 1}},
|
||||
{"type": "server_tool_use", "id": "srv.1", "name": "web_search", "input": {"query": "q"}, "provider_specific_fields": {"x": 2}},
|
||||
{"type": "web_search_tool_result", "tool_use_id": "srv.1", "provider_specific_fields": {"x": 3}, "content": [
|
||||
{"type": "web_search_result", "url": "u", "title": "", "encrypted_content": "", "provider_specific_fields": {"x": 4}}
|
||||
]}
|
||||
]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": " "}]}
|
||||
]),
|
||||
json!([
|
||||
{"role": "user", "content": "run it"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}},
|
||||
{"type": "server_tool_use", "id": "srv_1", "name": "web_search", "input": {"query": "q"}},
|
||||
{"type": "text", "text": "Web search results:\n\nURL: u"}
|
||||
]},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]}
|
||||
]),
|
||||
)]
|
||||
fn sanitize_anthropic_messages_cleans_replayed_history(
|
||||
#[case] history: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
assert_eq!(
|
||||
serde_json::to_value(sanitize_anthropic_messages(messages(history))).unwrap(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::keeps_only_user_id(json!({"user_id": "u-1", "trace_id": "internal"}), Ok(json!({"user_id": "u-1"})))]
|
||||
#[case::null_user_id(json!({"user_id": null, "trace_id": "internal"}), Ok(json!({})))]
|
||||
#[case::no_user_id(json!({"trace_id": "internal"}), Ok(json!({})))]
|
||||
#[case::empty(json!({}), Ok(json!({})))]
|
||||
#[case::numeric_user_id(
|
||||
json!({"user_id": 123}),
|
||||
Err(Error::InvalidRequest("metadata.user_id must be a string, got 123".to_string())),
|
||||
)]
|
||||
#[case::boolean_user_id(
|
||||
json!({"user_id": true}),
|
||||
Err(Error::InvalidRequest("metadata.user_id must be a string, got true".to_string())),
|
||||
)]
|
||||
#[case::not_an_object(
|
||||
json!(["u-1"]),
|
||||
Err(Error::InvalidRequest(r#"metadata must be an object, got ["u-1"]"#.to_string())),
|
||||
)]
|
||||
fn validate_anthropic_api_metadata_passes_only_a_string_user_id(
|
||||
#[case] metadata: Value,
|
||||
#[case] expected: Result<Value, Error>,
|
||||
) {
|
||||
assert_eq!(validate_anthropic_api_metadata(&metadata), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::adaptive(
|
||||
Some(json!({"type": "adaptive", "budget_tokens": 5000})),
|
||||
true,
|
||||
Some(json!({"type": "adaptive", "budget_tokens": 5000, "display": "summarized"})),
|
||||
)]
|
||||
#[case::enabled(
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000})),
|
||||
true,
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})),
|
||||
)]
|
||||
#[case::no_type(Some(json!({})), true, Some(json!({"display": "summarized"})))]
|
||||
#[case::display_omitted_is_overridden(
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "omitted"})),
|
||||
true,
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})),
|
||||
)]
|
||||
#[case::display_summarized_is_kept(
|
||||
Some(json!({"type": "enabled", "display": "summarized"})),
|
||||
true,
|
||||
Some(json!({"type": "enabled", "display": "summarized"})),
|
||||
)]
|
||||
#[case::disabled_thinking(Some(json!({"type": "disabled"})), true, Some(json!({"type": "disabled"})))]
|
||||
#[case::flag_off(
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000})),
|
||||
false,
|
||||
Some(json!({"type": "enabled", "budget_tokens": 10000})),
|
||||
)]
|
||||
#[case::flag_off_keeps_callers_display(
|
||||
Some(json!({"type": "enabled", "display": "omitted"})),
|
||||
false,
|
||||
Some(json!({"type": "enabled", "display": "omitted"})),
|
||||
)]
|
||||
#[case::no_thinking(None, true, None)]
|
||||
#[case::non_object_thinking(Some(json!("enabled")), true, Some(json!("enabled")))]
|
||||
fn reasoning_auto_summary_marks_active_thinking_as_summarized(
|
||||
#[case] thinking: Option<Value>,
|
||||
#[case] enabled: bool,
|
||||
#[case] expected: Option<Value>,
|
||||
) {
|
||||
assert_eq!(with_reasoning_auto_summary(thinking, enabled), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn shaping_cleans_messages_metadata_and_thinking() {
|
||||
let sanitized = shape_anthropic_messages_request(
|
||||
request(json!({
|
||||
"model": "m",
|
||||
"messages": [{"role": "assistant", "content": [
|
||||
{"type": "text", "text": ""},
|
||||
{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}
|
||||
]}],
|
||||
"metadata": {"user_id": "u", "trace_id": "t"},
|
||||
"thinking": {"type": "enabled", "budget_tokens": 1024},
|
||||
"safeguards": [{"type": "dangerous_tool_use"}]
|
||||
})),
|
||||
true,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(sanitized).unwrap(),
|
||||
json!({
|
||||
"model": "m",
|
||||
"messages": [{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}
|
||||
]}],
|
||||
"metadata": {"user_id": "u"},
|
||||
"thinking": {"type": "enabled", "budget_tokens": 1024, "display": "summarized"},
|
||||
"safeguards": [{"type": "dangerous_tool_use"}]
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,643 @@
|
|||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
anthropic::{
|
||||
ANTHROPIC_OAUTH_TOKEN_PREFIX,
|
||||
common_utils::{
|
||||
ANTHROPIC_OAUTH_BETA_HEADER, beta, has_advisor_tool, is_anthropic_oauth_key,
|
||||
is_tool_search_used, join_beta_values, requires_native_compaction_beta,
|
||||
split_beta_values,
|
||||
},
|
||||
},
|
||||
base_llm::anthropic_messages::transformation::Headers,
|
||||
};
|
||||
|
||||
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
|
||||
const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
|
||||
const BETA_HEADER: &str = "anthropic-beta";
|
||||
const AUTHORIZATION: &str = "authorization";
|
||||
const API_KEY_HEADER: &str = "x-api-key";
|
||||
const DIRECT_BROWSER_ACCESS_HEADER: &str = "anthropic-dangerous-direct-browser-access";
|
||||
|
||||
fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
|
||||
headers
|
||||
.iter()
|
||||
.find(|(header, _)| header.eq_ignore_ascii_case(name))
|
||||
.map(|(_, value)| value.as_str())
|
||||
}
|
||||
|
||||
fn without(headers: Headers, names: &[&str]) -> Headers {
|
||||
headers
|
||||
.into_iter()
|
||||
.filter(|(header, _)| !names.iter().any(|name| header.eq_ignore_ascii_case(name)))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn existing_betas(headers: &[(String, String)]) -> impl Iterator<Item = String> + '_ {
|
||||
headers
|
||||
.iter()
|
||||
.filter(|(header, _)| header.eq_ignore_ascii_case(BETA_HEADER))
|
||||
.flat_map(|(_, value)| split_beta_values(Some(value)))
|
||||
}
|
||||
|
||||
fn with_oauth_bearer(headers: Headers, bearer: String) -> Headers {
|
||||
let beta =
|
||||
join_beta_values(existing_betas(&headers).chain([ANTHROPIC_OAUTH_BETA_HEADER.to_string()]));
|
||||
without(headers, &[API_KEY_HEADER, AUTHORIZATION, BETA_HEADER])
|
||||
.into_iter()
|
||||
.chain([
|
||||
(AUTHORIZATION.to_string(), bearer),
|
||||
(BETA_HEADER.to_string(), beta),
|
||||
(DIRECT_BROWSER_ACCESS_HEADER.to_string(), "true".to_string()),
|
||||
])
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn non_empty(value: Option<&str>) -> Option<&str> {
|
||||
value.map(str::trim).filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
pub fn authenticate(
|
||||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Headers, litellm_auth::Error> {
|
||||
if let Some(forwarded) = header_value(&headers, AUTHORIZATION)
|
||||
&& forwarded
|
||||
.strip_prefix("Bearer ")
|
||||
.is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX))
|
||||
{
|
||||
let bearer = forwarded.to_string();
|
||||
return Ok(with_oauth_bearer(headers, bearer));
|
||||
}
|
||||
if let Some(key) = api_key.filter(|key| key.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) {
|
||||
return Ok(with_oauth_bearer(headers, format!("Bearer {key}")));
|
||||
}
|
||||
if header_value(&headers, API_KEY_HEADER).is_some()
|
||||
|| header_value(&headers, AUTHORIZATION).is_some()
|
||||
{
|
||||
return Ok(headers);
|
||||
}
|
||||
let resolved_key = non_empty(api_key)
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty()));
|
||||
let auth = match resolved_key {
|
||||
Some(key) if is_anthropic_oauth_key(&key) => {
|
||||
(AUTHORIZATION.to_string(), format!("Bearer {key}"))
|
||||
}
|
||||
Some(key) => (API_KEY_HEADER.to_string(), key),
|
||||
None => match env_lookup(ANTHROPIC_AUTH_TOKEN_ENV).filter(|value| !value.trim().is_empty())
|
||||
{
|
||||
Some(token) => (AUTHORIZATION.to_string(), format!("Bearer {token}")),
|
||||
None => {
|
||||
return Err(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: ANTHROPIC_API_KEY_ENV,
|
||||
});
|
||||
}
|
||||
},
|
||||
};
|
||||
Ok(headers.into_iter().chain([auth]).collect())
|
||||
}
|
||||
|
||||
fn context_management_betas(
|
||||
context_management: Option<&Value>,
|
||||
) -> impl Iterator<Item = &'static str> {
|
||||
let edits = context_management
|
||||
.and_then(|value| value.get("edits"))
|
||||
.and_then(Value::as_array)
|
||||
.map(Vec::as_slice)
|
||||
.unwrap_or(&[]);
|
||||
let (compact, other) = edits.iter().fold((false, false), |(compact, other), edit| {
|
||||
match edit.get("type").and_then(Value::as_str) {
|
||||
Some("compact_20260112") => (true, other),
|
||||
_ => (compact, true),
|
||||
}
|
||||
});
|
||||
compact
|
||||
.then_some(beta::COMPACT_2026_01_12)
|
||||
.into_iter()
|
||||
.chain(other.then_some(beta::CONTEXT_MANAGEMENT_2025_06_27))
|
||||
}
|
||||
|
||||
fn uses_structured_output(request: &AnthropicMessagesRequest) -> bool {
|
||||
request.output_format.is_some()
|
||||
|| request
|
||||
.output_config
|
||||
.as_ref()
|
||||
.and_then(|config| config.get("format"))
|
||||
.is_some_and(|format| !format.is_null())
|
||||
}
|
||||
|
||||
fn messages_carry_output_config(request: &AnthropicMessagesRequest) -> bool {
|
||||
request
|
||||
.messages
|
||||
.iter()
|
||||
.any(|message| message.extra.contains_key("output_config"))
|
||||
}
|
||||
|
||||
pub fn feature_betas(request: &AnthropicMessagesRequest) -> Vec<&'static str> {
|
||||
let tools = request.tools.as_deref();
|
||||
[
|
||||
requires_native_compaction_beta(request.compaction.as_ref(), &request.messages)
|
||||
.then_some(beta::COMPACT_2026_09_04),
|
||||
uses_structured_output(request).then_some(beta::STRUCTURED_OUTPUT),
|
||||
(request.speed.as_deref() == Some("fast")).then_some(beta::FAST_MODE_2026_02_01),
|
||||
messages_carry_output_config(request).then_some(beta::PER_TURN_CONTROL_2026_07_01),
|
||||
has_advisor_tool(tools).then_some(beta::ADVISOR_TOOL_2026_03_01),
|
||||
is_tool_search_used(tools).then_some(beta::ADVANCED_TOOL_USE_2025_11_20),
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(context_management_betas(
|
||||
request.context_management.as_ref(),
|
||||
))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn with_feature_betas(headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
|
||||
let existing = existing_betas(&headers).collect::<Vec<_>>();
|
||||
let features = feature_betas(request);
|
||||
if existing.is_empty() && features.is_empty() {
|
||||
return headers;
|
||||
}
|
||||
let merged = join_beta_values(
|
||||
existing
|
||||
.into_iter()
|
||||
.chain(features.into_iter().map(str::to_string)),
|
||||
);
|
||||
without(headers, &[BETA_HEADER])
|
||||
.into_iter()
|
||||
.chain([(BETA_HEADER.to_string(), merged)])
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
const OAUTH_TOKEN: &str = "sk-ant-oat01-token";
|
||||
const OAUTH_BEARER: &str = "Bearer sk-ant-oat01-token";
|
||||
const REGULAR_KEY: &str = "sk-ant-api03-regular";
|
||||
const BROWSER_ACCESS: (&str, &str) = ("anthropic-dangerous-direct-browser-access", "true");
|
||||
|
||||
type Env = &'static [(&'static str, &'static str)];
|
||||
|
||||
fn request(fields: Value) -> AnthropicMessagesRequest {
|
||||
let mut body =
|
||||
json!({"model": "claude", "messages": [{"role": "user", "content": "Hello"}]});
|
||||
body.as_object_mut()
|
||||
.unwrap()
|
||||
.extend(fields.as_object().unwrap().clone());
|
||||
serde_json::from_value(body).unwrap()
|
||||
}
|
||||
|
||||
fn headers(pairs: &[(&str, &str)]) -> Headers {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn betas(values: &[&str]) -> String {
|
||||
values.join(",")
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn no_env() -> Env {
|
||||
&[]
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn full_env() -> Env {
|
||||
&[
|
||||
("ANTHROPIC_API_KEY", "sk-env"),
|
||||
("ANTHROPIC_AUTH_TOKEN", "env-token"),
|
||||
]
|
||||
}
|
||||
|
||||
fn authenticate_with(
|
||||
forwarded: &[(&str, &str)],
|
||||
api_key: Option<&str>,
|
||||
env: Env,
|
||||
) -> Result<Headers, litellm_auth::Error> {
|
||||
let lookup = |name: &str| {
|
||||
env.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string())
|
||||
};
|
||||
authenticate(headers(forwarded), api_key, &lookup)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::forwarded_bearer_drops_forwarded_and_deployment_keys(
|
||||
&[("X-Api-Key", REGULAR_KEY), ("Authorization", OAUTH_BEARER)],
|
||||
Some(REGULAR_KEY),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::forwarded_bearer_in_uppercase_authorization_header(
|
||||
&[("AUTHORIZATION", OAUTH_BEARER)],
|
||||
None,
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::forwarded_bearer_keeps_unrelated_headers_in_place(
|
||||
&[("anthropic-version", "2023-06-01"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
OAUTH_BEARER,
|
||||
&[("anthropic-version", "2023-06-01")],
|
||||
)]
|
||||
#[case::forwarded_bearer_wins_over_an_oauth_api_key(
|
||||
&[("authorization", OAUTH_BEARER)],
|
||||
Some("sk-ant-oat01-deployment"),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::api_key_authenticates_as_a_bearer(&[], Some(OAUTH_TOKEN), OAUTH_BEARER, &[])]
|
||||
#[case::api_key_removes_a_forwarded_x_api_key(
|
||||
&[("x-api-key", OAUTH_TOKEN)],
|
||||
Some(OAUTH_TOKEN),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::api_key_replaces_a_forwarded_non_oauth_bearer(
|
||||
&[("Authorization", "Bearer some-proxy-token")],
|
||||
Some(OAUTH_TOKEN),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
fn oauth_token_is_the_whole_credential(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] expected_bearer: &str,
|
||||
#[case] kept: &[(&str, &str)],
|
||||
full_env: Env,
|
||||
) {
|
||||
let expected = kept
|
||||
.iter()
|
||||
.copied()
|
||||
.chain([
|
||||
("authorization", expected_bearer),
|
||||
("anthropic-beta", ANTHROPIC_OAUTH_BETA_HEADER),
|
||||
BROWSER_ACCESS,
|
||||
])
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, full_env).unwrap(),
|
||||
headers(&expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::forwarded_bearer_merges_a_differently_cased_beta_header(
|
||||
&[("Anthropic-Beta", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
#[case::forwarded_bearer_dedupes_an_existing_oauth_beta(
|
||||
&[("anthropic-beta", "web-search-2025-03-05, oauth-2025-04-20"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
#[case::api_key_merges_the_existing_beta_header(
|
||||
&[("anthropic-beta", " web-search-2025-03-05 ,")],
|
||||
Some(OAUTH_TOKEN),
|
||||
)]
|
||||
#[case::forwarded_bearer_unions_every_beta_header_casing(
|
||||
&[("anthropic-beta", "oauth-2025-04-20"), ("ANTHROPIC-BETA", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
fn oauth_beta_merges_into_existing_betas(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
no_env: Env,
|
||||
) {
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, no_env).unwrap(),
|
||||
headers(&[
|
||||
("authorization", OAUTH_BEARER),
|
||||
(
|
||||
"anthropic-beta",
|
||||
&betas(&[ANTHROPIC_OAUTH_BETA_HEADER, "web-search-2025-03-05"])
|
||||
),
|
||||
BROWSER_ACCESS,
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::x_api_key_over_the_deployment_key(&[("x-api-key", "caller-key")], Some("sk-other"))]
|
||||
#[case::uppercase_x_api_key(&[("X-API-KEY", "caller-key")], None)]
|
||||
#[case::non_oauth_bearer(&[("Authorization", "Bearer some-proxy-token")], None)]
|
||||
#[case::non_oauth_bearer_over_a_regular_api_key(
|
||||
&[("authorization", "Bearer sk-ant-api03-forwarded")],
|
||||
Some(REGULAR_KEY),
|
||||
)]
|
||||
#[case::oauth_token_without_the_bearer_scheme(&[("authorization", OAUTH_TOKEN)], None)]
|
||||
#[case::oauth_token_behind_a_lowercase_bearer_scheme(
|
||||
&[("authorization", "bearer sk-ant-oat01-token")],
|
||||
None,
|
||||
)]
|
||||
fn forwarded_auth_header_is_kept_untouched(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
full_env: Env,
|
||||
) {
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, full_env).unwrap(),
|
||||
headers(forwarded)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key_param(Some("sk-param"), &[], ("x-api-key", "sk-param"))]
|
||||
#[case::api_key_param_over_env_key_and_auth_token(
|
||||
Some("sk-param"),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("x-api-key", "sk-param"),
|
||||
)]
|
||||
#[case::env_key_without_a_param(None, &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))]
|
||||
#[case::env_key_when_the_param_is_empty(Some(""), &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))]
|
||||
#[case::env_key_when_the_param_is_whitespace(
|
||||
Some(" "),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env")],
|
||||
("x-api-key", "sk-env"),
|
||||
)]
|
||||
#[case::env_key_over_auth_token(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("x-api-key", "sk-env"),
|
||||
)]
|
||||
#[case::auth_token_as_a_bearer(
|
||||
None,
|
||||
&[("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("authorization", "Bearer env-token"),
|
||||
)]
|
||||
#[case::auth_token_when_the_env_key_is_whitespace(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", " \t"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("authorization", "Bearer env-token"),
|
||||
)]
|
||||
#[case::oauth_env_key_as_a_plain_bearer(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", "sk-ant-oat01-env")],
|
||||
("authorization", "Bearer sk-ant-oat01-env"),
|
||||
)]
|
||||
fn credential_is_resolved_after_the_existing_headers(
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] env: Env,
|
||||
#[case] expected: (&str, &str),
|
||||
) {
|
||||
let forwarded = [("anthropic-beta", "web-search-2025-03-05")];
|
||||
assert_eq!(
|
||||
authenticate_with(&forwarded, api_key, env).unwrap(),
|
||||
headers(&[forwarded[0], expected])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_credentials(&[], None, &[])]
|
||||
#[case::empty_api_key(&[], Some(""), &[])]
|
||||
#[case::whitespace_only_env_values(
|
||||
&[],
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", " "), ("ANTHROPIC_AUTH_TOKEN", " \t")],
|
||||
)]
|
||||
#[case::unrelated_forwarded_headers(&[("anthropic-beta", "web-search-2025-03-05")], None, &[])]
|
||||
fn missing_credentials_are_an_auth_error(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] env: Env,
|
||||
) {
|
||||
assert!(matches!(
|
||||
authenticate_with(forwarded, api_key, env),
|
||||
Err(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: "ANTHROPIC_API_KEY",
|
||||
})
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_features(json!({}), &[])]
|
||||
#[case::output_format(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])]
|
||||
#[case::null_output_format(json!({"output_format": null}), &[])]
|
||||
#[case::output_config_format(
|
||||
json!({"output_config": {"format": {"type": "json_schema"}, "effort": "xhigh"}}),
|
||||
&[beta::STRUCTURED_OUTPUT]
|
||||
)]
|
||||
#[case::null_output_config_format(json!({"output_config": {"format": null}}), &[])]
|
||||
#[case::top_level_output_config_without_format(json!({"output_config": {"effort": "high"}}), &[])]
|
||||
#[case::fast_speed(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])]
|
||||
#[case::standard_speed(json!({"speed": "standard"}), &[])]
|
||||
#[case::compaction_param(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])]
|
||||
#[case::empty_compaction_param(json!({"compaction": {}}), &[beta::COMPACT_2026_09_04])]
|
||||
#[case::signed_compaction_block_in_history(
|
||||
json!({"messages": [
|
||||
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": "sig"}]},
|
||||
{"role": "user", "content": "Continue"},
|
||||
]}),
|
||||
&[beta::COMPACT_2026_09_04]
|
||||
)]
|
||||
#[case::unsigned_compaction_block_in_history(
|
||||
json!({"messages": [
|
||||
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": ""}]},
|
||||
{"role": "user", "content": "Continue"},
|
||||
]}),
|
||||
&[]
|
||||
)]
|
||||
#[case::advisor_tool(
|
||||
json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": "claude-opus-4-6"}]}),
|
||||
&[beta::ADVISOR_TOOL_2026_03_01]
|
||||
)]
|
||||
#[case::no_tools(json!({"tools": []}), &[])]
|
||||
#[case::regex_tool_search(
|
||||
json!({"tools": [{"type": "tool_search_tool_regex_20251119"}]}),
|
||||
&[beta::ADVANCED_TOOL_USE_2025_11_20]
|
||||
)]
|
||||
#[case::bm25_tool_search(
|
||||
json!({"tools": [{"type": "tool_search_tool_bm25_20251119"}]}),
|
||||
&[beta::ADVANCED_TOOL_USE_2025_11_20]
|
||||
)]
|
||||
#[case::unrelated_server_tool(json!({"tools": [{"type": "web_search_20250305", "name": "web_search"}]}), &[])]
|
||||
#[case::only_compact_edits(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}),
|
||||
&[beta::COMPACT_2026_01_12]
|
||||
)]
|
||||
#[case::only_other_edits(
|
||||
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919", "keep": {"type": "tool_uses", "value": 3}}]}}),
|
||||
&[beta::CONTEXT_MANAGEMENT_2025_06_27]
|
||||
)]
|
||||
#[case::compact_and_other_edits(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_tool_uses_20250919"}]}}),
|
||||
&[beta::COMPACT_2026_01_12, beta::CONTEXT_MANAGEMENT_2025_06_27]
|
||||
)]
|
||||
#[case::edit_without_a_type(json!({"context_management": {"edits": [{}]}}), &[beta::CONTEXT_MANAGEMENT_2025_06_27])]
|
||||
#[case::empty_edits(json!({"context_management": {"edits": []}}), &[])]
|
||||
#[case::context_management_without_edits(json!({"context_management": {}}), &[])]
|
||||
#[case::per_message_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
&[beta::PER_TURN_CONTROL_2026_07_01]
|
||||
)]
|
||||
#[case::per_message_null_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": null}]}),
|
||||
&[beta::PER_TURN_CONTROL_2026_07_01]
|
||||
)]
|
||||
fn feature_betas_follow_the_request(#[case] fields: Value, #[case] expected: &[&str]) {
|
||||
assert_eq!(feature_betas(&request(fields)), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_betas(&[("x-api-key", "k"), ("anthropic-version", "2023-06-01")], json!({}))]
|
||||
#[case::blank_beta_header(&[("Anthropic-Beta", " , "), ("x-api-key", "k")], json!({}))]
|
||||
fn headers_without_any_beta_value_are_untouched(
|
||||
#[case] input: &[(&str, &str)],
|
||||
#[case] fields: Value,
|
||||
) {
|
||||
assert_eq!(
|
||||
with_feature_betas(headers(input), &request(fields)),
|
||||
headers(input)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::feature_beta_is_appended(
|
||||
&[("x-api-key", "k")],
|
||||
json!({"speed": "fast"}),
|
||||
&[("x-api-key", "k"), ("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
)]
|
||||
#[case::existing_betas_are_normalized_without_features(
|
||||
&[("Anthropic-Beta", "web-search-2025-03-05, interleaved-thinking-2025-05-14 ,web-search-2025-03-05"), ("x-api-key", "k")],
|
||||
json!({}),
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "interleaved-thinking-2025-05-14,web-search-2025-03-05")],
|
||||
)]
|
||||
#[case::existing_advisor_beta_is_kept_without_an_advisor_tool(
|
||||
&[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)],
|
||||
json!({"tools": []}),
|
||||
&[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)],
|
||||
)]
|
||||
#[case::feature_already_sent_is_not_duplicated(
|
||||
&[("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
json!({"speed": "fast"}),
|
||||
&[("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
)]
|
||||
fn feature_betas_merge_into_the_headers(
|
||||
#[case] input: &[(&str, &str)],
|
||||
#[case] fields: Value,
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
assert_eq!(
|
||||
with_feature_betas(headers(input), &request(fields)),
|
||||
headers(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn differently_cased_beta_header_is_replaced_by_one_sorted_header() {
|
||||
let merged = with_feature_betas(
|
||||
headers(&[("Anthropic-Beta", "interleaved-thinking-2025-05-14")]),
|
||||
&request(
|
||||
json!({"messages": [{"role": "system", "content": "env", "output_config": {"effort": "low"}}]}),
|
||||
),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::PER_TURN_CONTROL_2026_07_01
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_beta_header_casing_is_unioned_into_one_header() {
|
||||
let merged = with_feature_betas(
|
||||
headers(&[
|
||||
("anthropic-beta", "interleaved-thinking-2025-05-14"),
|
||||
("Anthropic-Beta", "web-search-2025-03-05"),
|
||||
]),
|
||||
&request(json!({"speed": "fast"})),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
beta::FAST_MODE_2026_02_01,
|
||||
"interleaved-thinking-2025-05-14",
|
||||
"web-search-2025-03-05"
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_client_betas_survive_alongside_the_added_one() {
|
||||
let client_betas = [
|
||||
"claude-code-20250219",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
"effort-2025-11-24",
|
||||
];
|
||||
let merged = with_feature_betas(
|
||||
headers(&[("anthropic-beta", &betas(&client_betas))]),
|
||||
&request(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
"claude-code-20250219",
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
"effort-2025-11-24",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_feature_merges_with_the_oauth_beta_sorted_and_last() {
|
||||
let oauth_headers = authenticate_with(&[], Some(OAUTH_TOKEN), &[]).unwrap();
|
||||
let all_features = request(json!({
|
||||
"compaction": {"enabled": true},
|
||||
"output_format": {"type": "json_schema"},
|
||||
"speed": "fast",
|
||||
"tools": [{"type": "advisor_20260301"}, {"type": "tool_search_tool_bm25_20251119"}],
|
||||
"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_thinking_20251015"}]},
|
||||
"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}],
|
||||
}));
|
||||
assert_eq!(
|
||||
with_feature_betas(oauth_headers, &all_features),
|
||||
headers(&[
|
||||
("authorization", OAUTH_BEARER),
|
||||
BROWSER_ACCESS,
|
||||
(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
beta::ADVANCED_TOOL_USE_2025_11_20,
|
||||
beta::ADVISOR_TOOL_2026_03_01,
|
||||
beta::COMPACT_2026_01_12,
|
||||
beta::COMPACT_2026_09_04,
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
beta::FAST_MODE_2026_02_01,
|
||||
ANTHROPIC_OAUTH_BETA_HEADER,
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
beta::STRUCTURED_OUTPUT,
|
||||
])
|
||||
),
|
||||
])
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,2 +1,5 @@
|
|||
pub mod handler;
|
||||
pub mod headers;
|
||||
pub mod streaming_iterator;
|
||||
pub mod thinking;
|
||||
pub mod transformation;
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,9 +1,28 @@
|
|||
use crate::base_llm::{
|
||||
anthropic_messages::transformation::BaseAnthropicMessagesConfig, chat::transformation::Error,
|
||||
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::{
|
||||
headers::{authenticate, with_feature_betas},
|
||||
thinking::{ThinkingBudgets, ThinkingContext, translate_thinking},
|
||||
};
|
||||
use crate::{
|
||||
anthropic::common_utils::{
|
||||
AnthropicModelCapabilities, has_advisor_tool, strip_advisor_blocks,
|
||||
strip_encrypted_reasoning_blocks,
|
||||
},
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, Headers, MessagesTransformContext,
|
||||
},
|
||||
chat::transformation::Error,
|
||||
},
|
||||
};
|
||||
|
||||
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
|
||||
const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
|
||||
const ANTHROPIC_API_BASE_ENV: &str = "ANTHROPIC_API_BASE";
|
||||
const ANTHROPIC_BASE_URL_ENV: &str = "ANTHROPIC_BASE_URL";
|
||||
const DEFAULT_ANTHROPIC_API_BASE: &str = "https://api.anthropic.com";
|
||||
const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
|
||||
|
||||
|
|
@ -11,6 +30,26 @@ pub struct AnthropicMessagesConfig;
|
|||
|
||||
pub const ANTHROPIC_MESSAGES_CONFIG: AnthropicMessagesConfig = AnthropicMessagesConfig;
|
||||
|
||||
impl MessagesTransformContext {
|
||||
pub fn new(capabilities: AnthropicModelCapabilities, drop_params: bool) -> Self {
|
||||
Self::with_lookup(capabilities, drop_params, &ProcessEnvironment)
|
||||
}
|
||||
|
||||
pub fn with_lookup(
|
||||
capabilities: AnthropicModelCapabilities,
|
||||
drop_params: bool,
|
||||
env: &impl Lookup,
|
||||
) -> Self {
|
||||
Self {
|
||||
thinking: ThinkingContext {
|
||||
capabilities,
|
||||
budgets: ThinkingBudgets::from_lookup(env),
|
||||
},
|
||||
drop_params,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
|
|
@ -21,6 +60,35 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
|||
Ok(complete_anthropic_url(api_base, env_lookup))
|
||||
}
|
||||
|
||||
fn transform_anthropic_messages_request(
|
||||
&self,
|
||||
request: AnthropicMessagesRequest,
|
||||
context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
if request.max_tokens.is_none() {
|
||||
return Err(Error::InvalidRequest(
|
||||
"max_tokens is required for Anthropic /v1/messages API".to_string(),
|
||||
));
|
||||
}
|
||||
let request = drop_unsupported_params(request, context)?;
|
||||
let request = translate_thinking(request, &context.thinking)?;
|
||||
let context_management = request
|
||||
.context_management
|
||||
.as_ref()
|
||||
.and_then(map_openai_context_management_to_anthropic)
|
||||
.or_else(|| request.context_management.clone());
|
||||
let messages = if has_advisor_tool(request.tools.as_deref()) {
|
||||
request.messages
|
||||
} else {
|
||||
strip_advisor_blocks(request.messages)
|
||||
};
|
||||
Ok(AnthropicMessagesRequest {
|
||||
messages: strip_encrypted_reasoning_blocks(messages),
|
||||
context_management,
|
||||
..request
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
|
|
@ -28,6 +96,113 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
|||
) -> Result<String, Error> {
|
||||
resolve_anthropic_api_key(api_key, env_lookup).map_err(Error::from)
|
||||
}
|
||||
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[
|
||||
ANTHROPIC_API_KEY_ENV,
|
||||
ANTHROPIC_AUTH_TOKEN_ENV,
|
||||
ANTHROPIC_API_BASE_ENV,
|
||||
ANTHROPIC_BASE_URL_ENV,
|
||||
]
|
||||
}
|
||||
|
||||
fn authenticate(
|
||||
&self,
|
||||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Headers, Error> {
|
||||
authenticate(headers, api_key, env_lookup).map_err(Error::from)
|
||||
}
|
||||
|
||||
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
|
||||
with_feature_betas(headers, request)
|
||||
}
|
||||
}
|
||||
|
||||
fn unsupported_param(model: &str, param: &str, value: &str, hint: &str) -> Error {
|
||||
Error::InvalidRequest(format!(
|
||||
"{model} does not support {param}={value}. {hint}To drop unsupported params, set `litellm.drop_params = True`."
|
||||
))
|
||||
}
|
||||
|
||||
fn drop_unsupported_params(
|
||||
request: AnthropicMessagesRequest,
|
||||
context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
let capabilities = &context.thinking.capabilities;
|
||||
let model = request.model.clone();
|
||||
let reject = |param: &str, value: String, hint: &str| -> Result<(), Error> {
|
||||
if context.drop_params {
|
||||
return Ok(());
|
||||
}
|
||||
Err(unsupported_param(&model, param, &value, hint))
|
||||
};
|
||||
let speed = match request.speed.as_deref() {
|
||||
Some(speed) if !capabilities.supports_speed => {
|
||||
reject("speed", format!("'{speed}'"), "")?;
|
||||
None
|
||||
}
|
||||
_ => request.speed.clone(),
|
||||
};
|
||||
if capabilities.supports_sampling_params {
|
||||
return Ok(AnthropicMessagesRequest { speed, ..request });
|
||||
}
|
||||
let temperature = match request.temperature {
|
||||
Some(temperature) if temperature != 1.0 => {
|
||||
reject(
|
||||
"temperature",
|
||||
json!(temperature).to_string(),
|
||||
"Only temperature=1 is supported. ",
|
||||
)?;
|
||||
None
|
||||
}
|
||||
temperature => temperature,
|
||||
};
|
||||
if let Some(top_p) = request.top_p {
|
||||
reject("top_p", json!(top_p).to_string(), "")?;
|
||||
}
|
||||
if let Some(top_k) = request.top_k {
|
||||
reject("top_k", json!(top_k).to_string(), "")?;
|
||||
}
|
||||
Ok(AnthropicMessagesRequest {
|
||||
speed,
|
||||
temperature,
|
||||
top_p: None,
|
||||
top_k: None,
|
||||
..request
|
||||
})
|
||||
}
|
||||
|
||||
pub fn map_openai_context_management_to_anthropic(context_management: &Value) -> Option<Value> {
|
||||
match context_management {
|
||||
Value::Object(edits) if edits.contains_key("edits") => Some(context_management.clone()),
|
||||
Value::Array(entries) => {
|
||||
let edits: Vec<Value> = entries
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.filter(|entry| entry.get("type").and_then(Value::as_str) == Some("compaction"))
|
||||
.map(|entry| {
|
||||
let trigger = entry.get("compact_threshold").and_then(Value::as_f64).map(
|
||||
|threshold| json!({"type": "input_tokens", "value": threshold as i64}),
|
||||
);
|
||||
let passthrough = entry
|
||||
.iter()
|
||||
.filter(|(key, _)| !matches!(key.as_str(), "type" | "compact_threshold"))
|
||||
.map(|(key, value)| (key.clone(), value.clone()));
|
||||
Value::Object(
|
||||
[("type".to_string(), json!("compact_20260112"))]
|
||||
.into_iter()
|
||||
.chain(trigger.map(|trigger| ("trigger".to_string(), trigger)))
|
||||
.chain(passthrough)
|
||||
.collect::<Map<String, Value>>(),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
(!edits.is_empty()).then(|| json!({"edits": edits}))
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn non_empty(value: Option<&str>) -> Option<&str> {
|
||||
|
|
@ -64,70 +239,619 @@ pub fn resolve_anthropic_api_base(
|
|||
api_base: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String {
|
||||
let env = |name: &str| env_lookup(name).filter(|value| !value.trim().is_empty());
|
||||
non_empty(api_base)
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(ANTHROPIC_API_BASE_ENV).filter(|value| !value.trim().is_empty()))
|
||||
.or_else(|| env(ANTHROPIC_API_BASE_ENV))
|
||||
.or_else(|| env(ANTHROPIC_BASE_URL_ENV))
|
||||
.unwrap_or_else(|| DEFAULT_ANTHROPIC_API_BASE.to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::process::Command;
|
||||
|
||||
use rstest::{fixture, rstest};
|
||||
|
||||
use super::*;
|
||||
use crate::anthropic::common_utils::{ENCRYPTED_REASONING_SIGNATURE_PREFIX, beta};
|
||||
|
||||
#[test]
|
||||
fn url_defaults_to_public_anthropic_endpoint() {
|
||||
type Env = &'static [(&'static str, &'static str)];
|
||||
|
||||
const BOTH_BASE_ENVS: Env = &[
|
||||
(ANTHROPIC_API_BASE_ENV, "https://api-base.example.com"),
|
||||
(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com"),
|
||||
];
|
||||
const API_KEY_ENV: Env = &[(ANTHROPIC_API_KEY_ENV, "sk-env")];
|
||||
const MISSING_API_KEY: &str =
|
||||
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable";
|
||||
const LOW_BUDGET_ENV: &str = "DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET";
|
||||
const PROCESS_ENV_PROBE: &str = "LITELLM_MESSAGES_TRANSFORM_CONTEXT_PROBE";
|
||||
|
||||
fn merged(base: Value, fields: Value) -> Value {
|
||||
Value::Object(
|
||||
base.as_object()
|
||||
.unwrap()
|
||||
.clone()
|
||||
.into_iter()
|
||||
.chain(fields.as_object().unwrap().clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
fn body(fields: Value) -> Value {
|
||||
merged(
|
||||
json!({
|
||||
"model": "claude",
|
||||
"max_tokens": 1024,
|
||||
"messages": [{"role": "user", "content": "Hello"}]
|
||||
}),
|
||||
fields,
|
||||
)
|
||||
}
|
||||
|
||||
fn request(fields: Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(body(fields)).unwrap()
|
||||
}
|
||||
|
||||
fn no_env(_: &str) -> Option<String> {
|
||||
None
|
||||
}
|
||||
|
||||
fn env(vars: Env) -> impl Fn(&str) -> Option<String> {
|
||||
move |name| {
|
||||
vars.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
fn headers(pairs: &[(&str, &str)]) -> Headers {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn transform(
|
||||
fields: Value,
|
||||
capabilities: AnthropicModelCapabilities,
|
||||
drop_params: bool,
|
||||
) -> Result<Value, Error> {
|
||||
ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(
|
||||
request(fields),
|
||||
&MessagesTransformContext::with_lookup(capabilities, drop_params, &no_env),
|
||||
)
|
||||
.map(|transformed| serde_json::to_value(transformed).unwrap())
|
||||
}
|
||||
|
||||
fn invalid(message: &str) -> Result<Value, Error> {
|
||||
Err(Error::InvalidRequest(message.to_string()))
|
||||
}
|
||||
|
||||
fn advisor_history() -> Value {
|
||||
json!([
|
||||
{"role": "user", "content": "Build a worker pool."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "text", "text": "Let me consult the advisor."},
|
||||
{"type": "server_tool_use", "id": "srvtoolu_abc123", "name": "advisor", "input": {}},
|
||||
{"type": "advisor_tool_result", "tool_use_id": "srvtoolu_abc123", "content": {"type": "advisor_result", "text": "Use channels."}},
|
||||
{"type": "text", "text": "Here is the implementation."}
|
||||
]}
|
||||
])
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn unmapped() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities::default()
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn sampling_removed() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities {
|
||||
supports_sampling_params: false,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn fast_mode() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities {
|
||||
supports_speed: true,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::alone(json!({"max_tokens": null}))]
|
||||
#[case::ahead_of_the_param_gate(json!({"max_tokens": null, "speed": "fast"}))]
|
||||
fn missing_max_tokens_is_rejected(#[case] fields: Value, unmapped: AnthropicModelCapabilities) {
|
||||
assert_eq!(
|
||||
complete_anthropic_url(None, &|_| None),
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
transform(fields, unmapped, false),
|
||||
invalid("max_tokens is required for Anthropic /v1/messages API")
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::sampling_params_on_a_sampling_model(
|
||||
unmapped(),
|
||||
false,
|
||||
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40})
|
||||
)]
|
||||
#[case::sampling_params_on_a_sampling_model_under_drop_params(
|
||||
unmapped(),
|
||||
true,
|
||||
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40})
|
||||
)]
|
||||
#[case::unit_temperature_on_a_sampling_removed_model(
|
||||
sampling_removed(),
|
||||
false,
|
||||
json!({"temperature": 1.0})
|
||||
)]
|
||||
#[case::unit_temperature_on_a_sampling_removed_model_under_drop_params(
|
||||
sampling_removed(),
|
||||
true,
|
||||
json!({"temperature": 1.0})
|
||||
)]
|
||||
#[case::speed_on_a_fast_mode_model(fast_mode(), false, json!({"speed": "fast"}))]
|
||||
#[case::speed_on_a_fast_mode_model_under_drop_params(fast_mode(), true, json!({"speed": "fast"}))]
|
||||
#[case::native_context_management_edits(unmapped(), false, json!({"context_management": {"edits": [{
|
||||
"type": "clear_tool_uses_20250919",
|
||||
"trigger": {"type": "input_tokens", "value": 30000},
|
||||
"keep": {"type": "tool_uses", "value": 3},
|
||||
"clear_at_least": {"type": "input_tokens", "value": 5000},
|
||||
"exclude_tools": ["web_search"],
|
||||
"clear_tool_inputs": false
|
||||
}]}}))]
|
||||
#[case::first_party_billing_header_system_block(unmapped(), false, json!({"system": [
|
||||
{"type": "text", "text": "x-anthropic-billing-header: cc_version=1"},
|
||||
{"type": "text", "text": "real system prompt"}
|
||||
]}))]
|
||||
#[case::anthropic_signed_reasoning_history(unmapped(), false, json!({"messages": [
|
||||
{"role": "user", "content": "Solve it."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "thinking", "thinking": "plan", "signature": "EqQBCkYIAxgCIkA_anthropic_signed"},
|
||||
{"type": "redacted_thinking", "data": "EmwKAhgBEgy_anthropic_minted"},
|
||||
{"type": "text", "text": "The answer."}
|
||||
]}
|
||||
]}))]
|
||||
#[case::advisor_history_alongside_the_advisor_tool(unmapped(), false, json!({
|
||||
"messages": advisor_history(),
|
||||
"tools": [{"type": "advisor_20260301", "name": "advisor"}]
|
||||
}))]
|
||||
fn request_is_forwarded_unchanged(
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] drop_params: bool,
|
||||
#[case] fields: Value,
|
||||
) {
|
||||
assert_eq!(
|
||||
transform(fields.clone(), capabilities, drop_params),
|
||||
Ok(body(fields))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::temperature(sampling_removed(), json!({"temperature": 0.3}), json!({}))]
|
||||
#[case::top_p(sampling_removed(), json!({"top_p": 0.9}), json!({}))]
|
||||
#[case::top_k(sampling_removed(), json!({"top_k": 40}), json!({}))]
|
||||
#[case::every_sampling_param_keeping_the_rest(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40, "stream": true}),
|
||||
json!({"stream": true})
|
||||
)]
|
||||
#[case::speed_on_a_sampling_model(
|
||||
unmapped(),
|
||||
json!({"speed": "fast", "temperature": 0.5}),
|
||||
json!({"temperature": 0.5})
|
||||
)]
|
||||
#[case::speed_on_a_sampling_removed_model(
|
||||
sampling_removed(),
|
||||
json!({"speed": "fast", "temperature": 1.0}),
|
||||
json!({"temperature": 1.0})
|
||||
)]
|
||||
fn removed_params_are_dropped_under_drop_params(
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
assert_eq!(transform(fields, capabilities, true), Ok(body(expected)));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::temperature(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.3}),
|
||||
"claude does not support temperature=0.3. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::temperature_just_below_one(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.99}),
|
||||
"claude does not support temperature=0.99. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::whole_number_temperature_keeps_its_decimal(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 2.0}),
|
||||
"claude does not support temperature=2.0. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_p(
|
||||
sampling_removed(),
|
||||
json!({"top_p": 0.9}),
|
||||
"claude does not support top_p=0.9. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_k(
|
||||
sampling_removed(),
|
||||
json!({"top_k": 5}),
|
||||
"claude does not support top_k=5. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_k_next_to_unit_temperature(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 1.0, "top_k": 5}),
|
||||
"claude does not support top_k=5. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::temperature_ahead_of_top_k(
|
||||
sampling_removed(),
|
||||
json!({"temperature": 0.5, "top_k": 5}),
|
||||
"claude does not support temperature=0.5. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::top_p_ahead_of_top_k(
|
||||
sampling_removed(),
|
||||
json!({"top_p": 0.9, "top_k": 5}),
|
||||
"claude does not support top_p=0.9. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::speed(
|
||||
unmapped(),
|
||||
json!({"speed": "fast"}),
|
||||
"claude does not support speed='fast'. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
#[case::speed_ahead_of_sampling_params(
|
||||
sampling_removed(),
|
||||
json!({"speed": "fast", "temperature": 0.5}),
|
||||
"claude does not support speed='fast'. To drop unsupported params, set `litellm.drop_params = True`."
|
||||
)]
|
||||
fn removed_params_are_rejected_without_drop_params(
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] message: &str,
|
||||
) {
|
||||
assert_eq!(transform(fields, capabilities, false), invalid(message));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::compaction_threshold(
|
||||
json!([{"type": "compaction", "compact_threshold": 200000}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 200000}}]}))
|
||||
)]
|
||||
#[case::other_keys_pass_through(
|
||||
json!([{"type": "compaction", "compact_threshold": 150000, "instructions": "Focus on preserving code snippets"}]),
|
||||
Some(json!({"edits": [{
|
||||
"type": "compact_20260112",
|
||||
"trigger": {"type": "input_tokens", "value": 150000},
|
||||
"instructions": "Focus on preserving code snippets"
|
||||
}]}))
|
||||
)]
|
||||
#[case::float_threshold_is_truncated(
|
||||
json!([{"type": "compaction", "compact_threshold": 150000.9}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}))
|
||||
)]
|
||||
#[case::compaction_without_threshold(
|
||||
json!([{"type": "compaction"}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112"}]}))
|
||||
)]
|
||||
#[case::non_numeric_threshold_is_dropped(
|
||||
json!([{"type": "compaction", "compact_threshold": "150000"}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112"}]}))
|
||||
)]
|
||||
#[case::non_object_entries_are_skipped(
|
||||
json!([42, "compaction", null, [], {"type": "compaction", "compact_threshold": 1000}]),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}}]}))
|
||||
)]
|
||||
#[case::only_compaction_entries_are_mapped_in_order(
|
||||
json!([
|
||||
{"type": "compaction", "compact_threshold": 1000},
|
||||
{"type": "other", "compact_threshold": 5},
|
||||
{"type": "compaction", "instructions": "second"}
|
||||
]),
|
||||
Some(json!({"edits": [
|
||||
{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}},
|
||||
{"type": "compact_20260112", "instructions": "second"}
|
||||
]}))
|
||||
)]
|
||||
#[case::list_without_compaction(json!([{"type": "other"}]), None)]
|
||||
#[case::empty_list(json!([]), None)]
|
||||
#[case::anthropic_edits_pass_through(
|
||||
json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}))
|
||||
)]
|
||||
#[case::object_without_edits(json!({"type": "compaction"}), None)]
|
||||
#[case::scalar(json!("compaction"), None)]
|
||||
fn openai_context_management_maps_to_anthropic_edits(
|
||||
#[case] context_management: Value,
|
||||
#[case] expected: Option<Value>,
|
||||
) {
|
||||
assert_eq!(
|
||||
map_openai_context_management_to_anthropic(&context_management),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::openai_list_is_mapped(
|
||||
json!([{"type": "compaction", "compact_threshold": 200000}]),
|
||||
json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 200000}}]})
|
||||
)]
|
||||
#[case::unmappable_list_is_kept(json!([{"type": "other"}]), json!([{"type": "other"}]))]
|
||||
#[case::unmappable_object_is_kept(json!({"type": "other"}), json!({"type": "other"}))]
|
||||
fn context_management_reaches_the_wire(
|
||||
#[case] context_management: Value,
|
||||
#[case] expected: Value,
|
||||
unmapped: AnthropicModelCapabilities,
|
||||
) {
|
||||
assert_eq!(
|
||||
transform(
|
||||
json!({"context_management": context_management}),
|
||||
unmapped,
|
||||
false
|
||||
),
|
||||
Ok(body(json!({"context_management": expected})))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::without_tools(json!({}))]
|
||||
#[case::with_only_other_tools(json!({"tools": [{"name": "get_weather", "input_schema": {"type": "object"}}]}))]
|
||||
fn advisor_history_is_stripped_without_the_advisor_tool(
|
||||
#[case] tools: Value,
|
||||
unmapped: AnthropicModelCapabilities,
|
||||
) {
|
||||
let stripped = json!([
|
||||
{"role": "user", "content": "Build a worker pool."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "text", "text": "Let me consult the advisor."},
|
||||
{"type": "text", "text": "Here is the implementation."}
|
||||
]}
|
||||
]);
|
||||
assert_eq!(
|
||||
transform(
|
||||
merged(tools.clone(), json!({"messages": advisor_history()})),
|
||||
unmapped,
|
||||
false
|
||||
),
|
||||
Ok(body(merged(tools, json!({"messages": stripped}))))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn bridge_minted_reasoning_is_stripped_from_the_wire(unmapped: AnthropicModelCapabilities) {
|
||||
let messages = json!([
|
||||
{"role": "user", "content": "Solve it."},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "thinking", "thinking": "plan", "signature": format!("{ENCRYPTED_REASONING_SIGNATURE_PREFIX}gAAAA_1")},
|
||||
{"type": "redacted_thinking", "data": format!("{ENCRYPTED_REASONING_SIGNATURE_PREFIX}gAAAA_2")},
|
||||
{"type": "text", "text": "The answer."}
|
||||
]},
|
||||
{"role": "user", "content": "And the next one?"}
|
||||
]);
|
||||
assert_eq!(
|
||||
transform(json!({"messages": messages}), unmapped, false),
|
||||
Ok(body(json!({"messages": [
|
||||
{"role": "user", "content": "Solve it."},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "The answer."}]},
|
||||
{"role": "user", "content": "And the next one?"}
|
||||
]})))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_appends_messages_suffix_to_custom_base() {
|
||||
fn thinking_is_translated_with_the_context_budgets() {
|
||||
let context = MessagesTransformContext::with_lookup(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..Default::default()
|
||||
},
|
||||
false,
|
||||
&env(&[(LOW_BUDGET_ENV, "2000")]),
|
||||
);
|
||||
let transformed = ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(
|
||||
request(json!({"max_tokens": 4096, "reasoning_effort": "low"})),
|
||||
&context,
|
||||
)
|
||||
.map(|transformed| serde_json::to_value(transformed).unwrap());
|
||||
assert_eq!(
|
||||
complete_anthropic_url(Some("https://proxy.internal"), &|_| None),
|
||||
"https://proxy.internal/v1/messages"
|
||||
transformed,
|
||||
Ok(body(json!({
|
||||
"max_tokens": 4096,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2000}
|
||||
})))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_leaves_complete_messages_endpoint_untouched() {
|
||||
fn new_reads_thinking_budgets_from_the_process_environment() {
|
||||
if std::env::var_os(PROCESS_ENV_PROBE).is_some() {
|
||||
assert_eq!(
|
||||
MessagesTransformContext::new(sampling_removed(), true),
|
||||
MessagesTransformContext {
|
||||
thinking: ThinkingContext {
|
||||
capabilities: sampling_removed(),
|
||||
budgets: ThinkingBudgets {
|
||||
low: 2000,
|
||||
..ThinkingBudgets::default()
|
||||
},
|
||||
},
|
||||
drop_params: true,
|
||||
}
|
||||
);
|
||||
return;
|
||||
}
|
||||
let (_, test_path) = concat!(
|
||||
module_path!(),
|
||||
"::new_reads_thinking_budgets_from_the_process_environment"
|
||||
)
|
||||
.split_once("::")
|
||||
.unwrap();
|
||||
let other_tiers = ["MINIMAL", "MEDIUM", "HIGH", "XHIGH", "MAX"]
|
||||
.map(|tier| format!("DEFAULT_REASONING_EFFORT_{tier}_THINKING_BUDGET"));
|
||||
let output = other_tiers
|
||||
.iter()
|
||||
.fold(
|
||||
Command::new(std::env::current_exe().unwrap()),
|
||||
|mut command, name| {
|
||||
command.env_remove(name);
|
||||
command
|
||||
},
|
||||
)
|
||||
.args([test_path, "--exact"])
|
||||
.env(PROCESS_ENV_PROBE, "1")
|
||||
.env(LOW_BUDGET_ENV, "2000")
|
||||
.output()
|
||||
.unwrap();
|
||||
let stdout = String::from_utf8_lossy(&output.stdout);
|
||||
assert!(
|
||||
output.status.success() && stdout.contains("1 passed"),
|
||||
"{stdout}{}",
|
||||
String::from_utf8_lossy(&output.stderr)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::public_endpoint_by_default(None, &[], "https://api.anthropic.com")]
|
||||
#[case::explicit_api_base_beats_env(
|
||||
Some("https://explicit.example.com"),
|
||||
BOTH_BASE_ENVS,
|
||||
"https://explicit.example.com"
|
||||
)]
|
||||
#[case::explicit_api_base_is_trimmed(
|
||||
Some(" https://explicit.example.com "),
|
||||
&[],
|
||||
"https://explicit.example.com"
|
||||
)]
|
||||
#[case::blank_api_base_falls_back_to_env(
|
||||
Some(" "),
|
||||
BOTH_BASE_ENVS,
|
||||
"https://api-base.example.com"
|
||||
)]
|
||||
#[case::api_base_env_beats_base_url_env(None, BOTH_BASE_ENVS, "https://api-base.example.com")]
|
||||
#[case::base_url_env_without_api_base_env(
|
||||
None,
|
||||
&[(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
|
||||
"https://base-url.example.com"
|
||||
)]
|
||||
#[case::blank_api_base_env_falls_back_to_base_url_env(
|
||||
None,
|
||||
&[(ANTHROPIC_API_BASE_ENV, " \t "), (ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
|
||||
"https://base-url.example.com"
|
||||
)]
|
||||
#[case::blank_envs_fall_back_to_public_endpoint(
|
||||
None,
|
||||
&[(ANTHROPIC_API_BASE_ENV, ""), (ANTHROPIC_BASE_URL_ENV, " ")],
|
||||
"https://api.anthropic.com"
|
||||
)]
|
||||
fn api_base_resolution(
|
||||
#[case] api_base: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(resolve_anthropic_api_base(api_base, &env(vars)), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::public_endpoint(None, &[], "https://api.anthropic.com/v1/messages")]
|
||||
#[case::base_url_env(
|
||||
None,
|
||||
&[(ANTHROPIC_BASE_URL_ENV, "https://custom.example.com")],
|
||||
"https://custom.example.com/v1/messages"
|
||||
)]
|
||||
#[case::custom_base(Some("https://proxy.internal"), &[], "https://proxy.internal/v1/messages")]
|
||||
#[case::trailing_slash(Some("https://proxy.internal/"), &[], "https://proxy.internal/v1/messages")]
|
||||
#[case::complete_endpoint(
|
||||
Some("https://proxy.internal/v1/messages"),
|
||||
&[],
|
||||
"https://proxy.internal/v1/messages"
|
||||
)]
|
||||
#[case::complete_endpoint_with_trailing_slash(
|
||||
Some("https://proxy.internal/v1/messages/"),
|
||||
&[],
|
||||
"https://proxy.internal/v1/messages"
|
||||
)]
|
||||
fn complete_url_ends_in_the_messages_path(
|
||||
#[case] api_base: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(
|
||||
complete_anthropic_url(Some("https://proxy.internal/v1/messages"), &|_| None),
|
||||
"https://proxy.internal/v1/messages"
|
||||
ANTHROPIC_MESSAGES_CONFIG.get_complete_url(api_base, "claude", &env(vars)),
|
||||
Ok(expected.to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::param_beats_env(Some("sk-param"), API_KEY_ENV, Ok("sk-param"))]
|
||||
#[case::param_is_trimmed(Some(" sk-param "), &[], Ok("sk-param"))]
|
||||
#[case::blank_param_falls_back_to_env(Some(" "), API_KEY_ENV, Ok("sk-env"))]
|
||||
#[case::env_without_param(None, API_KEY_ENV, Ok("sk-env"))]
|
||||
#[case::blank_env_is_missing(None, &[(ANTHROPIC_API_KEY_ENV, " ")], Err(MISSING_API_KEY))]
|
||||
#[case::nothing_is_missing(None, &[], Err(MISSING_API_KEY))]
|
||||
fn api_key_resolution(
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected: Result<&str, &str>,
|
||||
) {
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(api_key, &env(vars)).map_err(|error| error.to_string()),
|
||||
expected.map(str::to_string).map_err(str::to_string)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_falls_back_to_env_base() {
|
||||
let with_env = |key: &str| {
|
||||
(key == ANTHROPIC_API_BASE_ENV).then(|| "https://env.anthropic".to_string())
|
||||
};
|
||||
fn config_reports_a_missing_key_as_an_auth_error() {
|
||||
assert_eq!(
|
||||
complete_anthropic_url(Some(" "), &with_env),
|
||||
"https://env.anthropic/v1/messages"
|
||||
ANTHROPIC_MESSAGES_CONFIG.resolve_api_key(None, &no_env),
|
||||
Err(Error::Auth(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: ANTHROPIC_API_KEY_ENV,
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn api_key_prefers_param_then_env_then_errors() {
|
||||
fn config_authenticates_with_the_anthropic_auth_token() {
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(Some("sk-param"), &|_| None).unwrap(),
|
||||
"sk-param"
|
||||
ANTHROPIC_MESSAGES_CONFIG.authenticate(
|
||||
vec![],
|
||||
None,
|
||||
&env(&[("ANTHROPIC_AUTH_TOKEN", "auth-token")])
|
||||
),
|
||||
Ok(headers(&[("authorization", "Bearer auth-token")]))
|
||||
);
|
||||
let with_env = |key: &str| (key == ANTHROPIC_API_KEY_ENV).then(|| "sk-env".to_string());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn config_requests_the_betas_the_request_features_need() {
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(Some(" "), &with_env).unwrap(),
|
||||
"sk-env"
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(None, &|_| None)
|
||||
.expect_err("missing key")
|
||||
.to_string(),
|
||||
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
|
||||
ANTHROPIC_MESSAGES_CONFIG.request_headers(
|
||||
headers(&[("x-api-key", "sk")]),
|
||||
&request(json!({"speed": "fast"}))
|
||||
),
|
||||
headers(&[
|
||||
("x-api-key", "sk"),
|
||||
("anthropic-beta", beta::FAST_MODE_2026_02_01)
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::absent(None, None)]
|
||||
#[case::blank(Some(" \t "), None)]
|
||||
#[case::padded(Some(" value "), Some("value"))]
|
||||
fn non_empty_trims_and_drops_blank_values(
|
||||
#[case] value: Option<&str>,
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
assert_eq!(non_empty(value), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_strategy_and_default_headers_match_anthropic() {
|
||||
assert_eq!(
|
||||
|
|
@ -142,4 +866,26 @@ mod tests {
|
|||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn secret_names_cover_every_credential_and_base_lookup() {
|
||||
let requested = std::cell::RefCell::new(Vec::<String>::new());
|
||||
let record = |name: &str| -> Option<String> {
|
||||
requested.borrow_mut().push(name.to_string());
|
||||
None
|
||||
};
|
||||
let _ = ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record);
|
||||
let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
|
||||
let requested = requested.into_inner();
|
||||
assert!(!requested.is_empty());
|
||||
let undeclared: Vec<&String> = requested
|
||||
.iter()
|
||||
.filter(|name| {
|
||||
!ANTHROPIC_MESSAGES_CONFIG
|
||||
.secret_names()
|
||||
.contains(&name.as_str())
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(undeclared, Vec::<&String>::new());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
pub mod batches;
|
||||
pub mod chat;
|
||||
pub mod common_utils;
|
||||
pub mod count_tokens;
|
||||
pub mod experimental_pass_through;
|
||||
|
||||
|
|
|
|||
|
|
@ -4,14 +4,15 @@ use litellm_types::llms::anthropic_messages::{
|
|||
},
|
||||
anthropic_response::AnthropicMessagesResponse,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::{
|
||||
anthropic::experimental_pass_through::messages::transformation::{
|
||||
ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty,
|
||||
},
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::{BaseAnthropicMessagesConfig, MessagesAuthStrategy},
|
||||
anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, Headers, MessagesAuthStrategy, MessagesTransformContext,
|
||||
},
|
||||
chat::transformation::Error,
|
||||
},
|
||||
};
|
||||
|
|
@ -21,7 +22,6 @@ const AZURE_API_BASE_ENV: &str = "AZURE_API_BASE";
|
|||
const ANTHROPIC_PATH_SEGMENT: &str = "/anthropic";
|
||||
const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
|
||||
const SYSTEM_ROLE: &str = "system";
|
||||
const TEXT_BLOCK_TYPE: &str = "text";
|
||||
|
||||
pub struct AzureAnthropicMessagesConfig {
|
||||
anthropic: AnthropicMessagesConfig,
|
||||
|
|
@ -45,6 +45,7 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
fn transform_anthropic_messages_request(
|
||||
&self,
|
||||
request: AnthropicMessagesRequest,
|
||||
context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
let mut request = fold_system_role_messages(request);
|
||||
if let Some(system) = request.system.as_mut() {
|
||||
|
|
@ -54,7 +55,8 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
.messages
|
||||
.iter_mut()
|
||||
.for_each(strip_scope_from_message);
|
||||
self.anthropic.transform_anthropic_messages_request(request)
|
||||
self.anthropic
|
||||
.transform_anthropic_messages_request(request, context)
|
||||
}
|
||||
|
||||
fn transform_anthropic_messages_response(
|
||||
|
|
@ -74,6 +76,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
resolve_azure_api_key(api_key, env_lookup)
|
||||
}
|
||||
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[AZURE_API_KEY_ENV, AZURE_API_BASE_ENV]
|
||||
}
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
self.anthropic.auth_strategy()
|
||||
}
|
||||
|
|
@ -85,6 +91,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
|
||||
self.anthropic.default_headers()
|
||||
}
|
||||
|
||||
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
|
||||
self.anthropic.request_headers(headers, request)
|
||||
}
|
||||
}
|
||||
|
||||
pub fn resolve_azure_api_key(
|
||||
|
|
@ -143,17 +153,7 @@ fn strip_scope_from_message(message: &mut AnthropicMessage) {
|
|||
}
|
||||
|
||||
fn text_content_block(text: String) -> ContentBlock {
|
||||
let extra = Map::from_iter([
|
||||
(
|
||||
"type".to_string(),
|
||||
Value::String(TEXT_BLOCK_TYPE.to_string()),
|
||||
),
|
||||
("text".to_string(), Value::String(text)),
|
||||
]);
|
||||
ContentBlock {
|
||||
cache_control: None,
|
||||
extra,
|
||||
}
|
||||
ContentBlock::text(text)
|
||||
}
|
||||
|
||||
fn content_into_blocks(content: MessageContent) -> Vec<ContentBlock> {
|
||||
|
|
@ -202,6 +202,7 @@ mod tests {
|
|||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
use crate::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
|
||||
fn request_from(value: serde_json::Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(value).expect("valid request")
|
||||
|
|
@ -346,7 +347,7 @@ mod tests {
|
|||
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms"),
|
||||
);
|
||||
|
||||
|
|
@ -373,10 +374,13 @@ mod tests {
|
|||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}));
|
||||
let once = AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms");
|
||||
let twice = AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(once.clone())
|
||||
.transform_anthropic_messages_request(
|
||||
once.clone(),
|
||||
&MessagesTransformContext::default(),
|
||||
)
|
||||
.expect("request transforms");
|
||||
assert_eq!(once, twice);
|
||||
assert_eq!(to_value(once)["system"], json!("plain string system"));
|
||||
|
|
@ -408,9 +412,21 @@ mod tests {
|
|||
"inference_geo": "us",
|
||||
"litellm_metadata": {"trace": "abc"}
|
||||
});
|
||||
let context = MessagesTransformContext::with_lookup(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
supports_legacy_thinking: true,
|
||||
supports_output_config: true,
|
||||
supports_speed: true,
|
||||
..Default::default()
|
||||
},
|
||||
false,
|
||||
&|_: &str| None,
|
||||
);
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request_from(body.clone()))
|
||||
.transform_anthropic_messages_request(request_from(body.clone()), &context)
|
||||
.expect("request transforms"),
|
||||
);
|
||||
assert_eq!(transformed, body);
|
||||
|
|
@ -430,7 +446,7 @@ mod tests {
|
|||
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms"),
|
||||
);
|
||||
|
||||
|
|
@ -460,7 +476,7 @@ mod tests {
|
|||
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request)
|
||||
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
|
||||
.expect("request transforms"),
|
||||
);
|
||||
|
||||
|
|
@ -485,9 +501,21 @@ mod tests {
|
|||
{"role": "assistant", "content": "hello"}
|
||||
]
|
||||
});
|
||||
let context = MessagesTransformContext::with_lookup(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
supports_legacy_thinking: true,
|
||||
supports_output_config: true,
|
||||
supports_speed: true,
|
||||
..Default::default()
|
||||
},
|
||||
false,
|
||||
&|_: &str| None,
|
||||
);
|
||||
let transformed = to_value(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.transform_anthropic_messages_request(request_from(body.clone()))
|
||||
.transform_anthropic_messages_request(request_from(body.clone()), &context)
|
||||
.expect("request transforms"),
|
||||
);
|
||||
assert_eq!(transformed, body);
|
||||
|
|
@ -500,6 +528,57 @@ mod tests {
|
|||
assert!(err.is_data());
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::compact_context_management_edit(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}),
|
||||
&[],
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "compact-2026-01-12")]
|
||||
)]
|
||||
#[case::forwarded_beta_merged_with_structured_output(
|
||||
json!({"output_config": {"format": {"type": "json_schema"}}}),
|
||||
&[("anthropic-beta", "web-search-2025-03-05")],
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "structured-outputs-2025-11-13,web-search-2025-03-05")]
|
||||
)]
|
||||
#[case::no_feature_needs_a_beta(json!({}), &[], &[("x-api-key", "k")])]
|
||||
fn request_headers_carry_the_anthropic_feature_betas(
|
||||
#[case] fields: serde_json::Value,
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
let pairs = |pairs: &[(&str, &str)]| -> Vec<(String, String)> {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
};
|
||||
let serde_json::Value::Object(fields) = fields else {
|
||||
panic!("case fields are an object")
|
||||
};
|
||||
let request = request_from(serde_json::Value::Object(
|
||||
[
|
||||
("model".to_string(), json!("claude-sonnet")),
|
||||
("max_tokens".to_string(), json!(16)),
|
||||
(
|
||||
"messages".to_string(),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
),
|
||||
]
|
||||
.into_iter()
|
||||
.chain(fields)
|
||||
.collect(),
|
||||
));
|
||||
assert_eq!(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG.request_headers(
|
||||
pairs(&[("x-api-key", "k")])
|
||||
.into_iter()
|
||||
.chain(pairs(forwarded))
|
||||
.collect(),
|
||||
&request
|
||||
),
|
||||
pairs(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transform_response_passes_through() {
|
||||
let response: AnthropicMessagesResponse = serde_json::from_value(json!({
|
||||
|
|
@ -521,4 +600,26 @@ mod tests {
|
|||
assert_eq!(value["stop_sequence"], json!(null));
|
||||
assert_eq!(value["content"][0]["text"], json!("hello"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn secret_names_cover_every_credential_and_base_lookup() {
|
||||
let requested = std::cell::RefCell::new(Vec::<String>::new());
|
||||
let record = |name: &str| -> Option<String> {
|
||||
requested.borrow_mut().push(name.to_string());
|
||||
None
|
||||
};
|
||||
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record);
|
||||
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
|
||||
let requested = requested.into_inner();
|
||||
assert!(!requested.is_empty());
|
||||
let undeclared: Vec<&String> = requested
|
||||
.iter()
|
||||
.filter(|name| {
|
||||
!AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.secret_names()
|
||||
.contains(&name.as_str())
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(undeclared, Vec::<&String>::new());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,14 @@
|
|||
use litellm_http::request::{has_bearer_auth, has_header};
|
||||
use litellm_types::llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
};
|
||||
|
||||
use crate::base_llm::chat::transformation::Error;
|
||||
use crate::{
|
||||
anthropic::experimental_pass_through::messages::thinking::ThinkingContext,
|
||||
base_llm::chat::transformation::Error,
|
||||
};
|
||||
|
||||
pub type Headers = Vec<(String, String)>;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum MessagesAuthStrategy {
|
||||
|
|
@ -19,6 +25,12 @@ impl MessagesAuthStrategy {
|
|||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
|
||||
pub struct MessagesTransformContext {
|
||||
pub thinking: ThinkingContext,
|
||||
pub drop_params: bool,
|
||||
}
|
||||
|
||||
pub trait BaseAnthropicMessagesConfig: Sync {
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
|
|
@ -30,6 +42,7 @@ pub trait BaseAnthropicMessagesConfig: Sync {
|
|||
fn transform_anthropic_messages_request(
|
||||
&self,
|
||||
request: AnthropicMessagesRequest,
|
||||
_context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
Ok(request)
|
||||
}
|
||||
|
|
@ -48,6 +61,8 @@ pub trait BaseAnthropicMessagesConfig: Sync {
|
|||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error>;
|
||||
|
||||
fn secret_names(&self) -> &'static [&'static str];
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
MessagesAuthStrategy::Header("x-api-key")
|
||||
}
|
||||
|
|
@ -56,10 +71,225 @@ pub trait BaseAnthropicMessagesConfig: Sync {
|
|||
false
|
||||
}
|
||||
|
||||
fn authenticate(
|
||||
&self,
|
||||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Headers, Error> {
|
||||
let strategy = self.auth_strategy();
|
||||
if has_header(&headers, strategy.header_name())
|
||||
|| (self.accepts_bearer_auth() && has_bearer_auth(&headers))
|
||||
{
|
||||
return Ok(headers);
|
||||
}
|
||||
let api_key = self.resolve_api_key(api_key, env_lookup)?;
|
||||
let auth_header = match strategy {
|
||||
MessagesAuthStrategy::Bearer => {
|
||||
("authorization".to_string(), format!("Bearer {api_key}"))
|
||||
}
|
||||
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
|
||||
};
|
||||
Ok(headers.into_iter().chain([auth_header]).collect())
|
||||
}
|
||||
|
||||
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
|
||||
&[
|
||||
("anthropic-version", "2023-06-01"),
|
||||
("content-type", "application/json"),
|
||||
]
|
||||
}
|
||||
|
||||
fn request_headers(&self, headers: Headers, _request: &AnthropicMessagesRequest) -> Headers {
|
||||
headers
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
const X_API_KEY: MessagesAuthStrategy = MessagesAuthStrategy::Header("x-api-key");
|
||||
|
||||
struct StubConfig {
|
||||
strategy: MessagesAuthStrategy,
|
||||
accepts_bearer: bool,
|
||||
}
|
||||
|
||||
impl BaseAnthropicMessagesConfig for StubConfig {
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[]
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
_api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
Ok(String::new())
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
api_key
|
||||
.map(str::to_string)
|
||||
.ok_or(Error::MissingField("api_key"))
|
||||
}
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
self.strategy
|
||||
}
|
||||
|
||||
fn accepts_bearer_auth(&self) -> bool {
|
||||
self.accepts_bearer
|
||||
}
|
||||
}
|
||||
|
||||
struct DefaultsConfig;
|
||||
|
||||
impl BaseAnthropicMessagesConfig for DefaultsConfig {
|
||||
fn secret_names(&self) -> &'static [&'static str] {
|
||||
&[]
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
_api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
Ok(String::new())
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
api_key
|
||||
.map(str::to_string)
|
||||
.ok_or(Error::MissingField("api_key"))
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_config_adds_its_key_next_to_a_forwarded_bearer() {
|
||||
assert_eq!(
|
||||
DefaultsConfig.authenticate(
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
Some("sk"),
|
||||
&|_| None
|
||||
),
|
||||
Ok(headers(&[
|
||||
("authorization", "Bearer forwarded"),
|
||||
("x-api-key", "sk")
|
||||
]))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_request_headers_are_the_given_headers() {
|
||||
let request: AnthropicMessagesRequest = serde_json::from_value(serde_json::json!({
|
||||
"model": "claude",
|
||||
"max_tokens": 16,
|
||||
"speed": "fast",
|
||||
"messages": [{"role": "user", "content": "hi"}]
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
DefaultsConfig.request_headers(headers(&[("x-api-key", "sk")]), &request),
|
||||
headers(&[("x-api-key", "sk")])
|
||||
);
|
||||
}
|
||||
|
||||
fn headers(pairs: &[(&str, &str)]) -> Headers {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::own_header_is_kept(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("x-api-key", "forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("x-api-key", "forwarded")]))
|
||||
)]
|
||||
#[case::own_header_in_any_casing_is_kept(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("X-Api-Key", "forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("X-Api-Key", "forwarded")]))
|
||||
)]
|
||||
#[case::accepted_bearer_is_kept(
|
||||
X_API_KEY,
|
||||
true,
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("authorization", "Bearer forwarded")]))
|
||||
)]
|
||||
#[case::bearer_the_provider_does_not_accept_gets_the_key_too(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("authorization", "Bearer forwarded"), ("x-api-key", "sk")]))
|
||||
)]
|
||||
#[case::blank_bearer_gets_the_key(
|
||||
X_API_KEY,
|
||||
true,
|
||||
headers(&[("authorization", "Bearer ")]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("authorization", "Bearer "), ("x-api-key", "sk")]))
|
||||
)]
|
||||
#[case::key_goes_in_the_provider_header(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[("content-type", "application/json")]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("content-type", "application/json"), ("x-api-key", "sk")]))
|
||||
)]
|
||||
#[case::key_goes_in_a_bearer(
|
||||
MessagesAuthStrategy::Bearer,
|
||||
false,
|
||||
headers(&[]),
|
||||
Some("sk"),
|
||||
Ok(headers(&[("authorization", "Bearer sk")]))
|
||||
)]
|
||||
#[case::bearer_strategy_keeps_a_forwarded_authorization(
|
||||
MessagesAuthStrategy::Bearer,
|
||||
false,
|
||||
headers(&[("authorization", "Bearer forwarded")]),
|
||||
None,
|
||||
Ok(headers(&[("authorization", "Bearer forwarded")]))
|
||||
)]
|
||||
#[case::missing_key_is_an_error(
|
||||
X_API_KEY,
|
||||
false,
|
||||
headers(&[]),
|
||||
None,
|
||||
Err(Error::MissingField("api_key"))
|
||||
)]
|
||||
fn default_authenticate_applies_the_key_unless_a_credential_is_forwarded(
|
||||
#[case] strategy: MessagesAuthStrategy,
|
||||
#[case] accepts_bearer: bool,
|
||||
#[case] forwarded: Headers,
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] expected: Result<Headers, Error>,
|
||||
) {
|
||||
let config = StubConfig {
|
||||
strategy,
|
||||
accepts_bearer,
|
||||
};
|
||||
assert_eq!(config.authenticate(forwarded, api_key, &|_| None), expected);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -245,6 +245,8 @@ pub struct ModelInfo {
|
|||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_272k_tokens_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_32k_tokens: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_batches: Option<f64>,
|
||||
/// Flex service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -283,6 +285,8 @@ pub struct ModelInfo {
|
|||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_272k_tokens_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_32k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_512k_tokens: Option<f64>,
|
||||
|
|
@ -377,6 +381,8 @@ pub struct ModelInfo {
|
|||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_272k_tokens_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_32k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_512k_tokens: Option<f64>,
|
||||
|
|
@ -498,6 +504,8 @@ pub struct ModelInfo {
|
|||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_272k_tokens_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_32k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_512k_tokens: Option<f64>,
|
||||
|
|
|
|||
|
|
@ -2,9 +2,11 @@ use bytes::Bytes;
|
|||
use litellm_core::messages::{
|
||||
Error,
|
||||
route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput},
|
||||
types::MessagesShaping,
|
||||
};
|
||||
use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use pyo3::{
|
||||
exceptions::{PyException, PyValueError},
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
|
|
@ -18,9 +20,10 @@ use crate::{
|
|||
marshal::{optional_timeout, python_timeout_seconds},
|
||||
};
|
||||
|
||||
/// The Anthropic Messages body fields a caller may pass besides `model` and `messages`,
|
||||
/// as `AnthropicMessagesRequestOptionalParams` declares them.
|
||||
const BODY_FIELDS: [&str; 20] = [
|
||||
const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.messages.route_host";
|
||||
const REQUEST_ERROR_MARKER: &str = "messages_request_error";
|
||||
|
||||
const BODY_FIELDS: [&str; 22] = [
|
||||
"max_tokens",
|
||||
"metadata",
|
||||
"stop_sequences",
|
||||
|
|
@ -35,14 +38,46 @@ const BODY_FIELDS: [&str; 20] = [
|
|||
"top_p",
|
||||
"mcp_servers",
|
||||
"context_management",
|
||||
"compaction",
|
||||
"container",
|
||||
"output_format",
|
||||
"speed",
|
||||
"output_config",
|
||||
"cache_control",
|
||||
"reasoning_effort",
|
||||
"safeguards",
|
||||
];
|
||||
|
||||
fn merge_headers(
|
||||
forwarded: Option<Map<String, Value>>,
|
||||
extra_headers: Option<Map<String, Value>>,
|
||||
) -> Option<Map<String, Value>> {
|
||||
let merged: Map<String, Value> = forwarded
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(extra_headers.into_iter().flatten())
|
||||
.collect();
|
||||
(!merged.is_empty()).then_some(merged)
|
||||
}
|
||||
|
||||
fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
|
||||
match error {
|
||||
Error::Transport(TransportError::Http { status, body }) => {
|
||||
let error = RustUpstreamError::new_err((status, body));
|
||||
error
|
||||
.value(py)
|
||||
.setattr("headers", Vec::<(String, String)>::new())?;
|
||||
Ok(error)
|
||||
}
|
||||
Error::InvalidRequest(message) => {
|
||||
let error = PyValueError::new_err(message);
|
||||
error.value(py).setattr(REQUEST_ERROR_MARKER, true)?;
|
||||
Ok(error)
|
||||
}
|
||||
other => Ok(messages_error_to_pyerr(other)),
|
||||
}
|
||||
}
|
||||
|
||||
/// The Python side of the Messages route: projects the prepared arguments and builds the
|
||||
/// public response, chunks and exceptions.
|
||||
pub(super) struct MessagesRouteHost {
|
||||
|
|
@ -84,19 +119,65 @@ impl MessagesRouteHost {
|
|||
.map(|value| python_timeout_seconds(py, value.unbind()))
|
||||
.transpose()?
|
||||
.flatten();
|
||||
let custom_llm_provider = string("custom_llm_provider")?;
|
||||
let shaping = self.shaping(py, &model, custom_llm_provider.as_deref(), arguments)?;
|
||||
Ok(MessagesCall {
|
||||
model,
|
||||
body,
|
||||
api_key: string("api_key")?,
|
||||
api_base: string("api_base")?,
|
||||
custom_llm_provider: string("custom_llm_provider")?,
|
||||
extra_headers: argument("extra_headers")?
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()?,
|
||||
extra_headers: self.merged_headers(py, arguments)?,
|
||||
provider_specific_header: self.provider_specific_header(py, arguments)?,
|
||||
custom_llm_provider,
|
||||
timeout: optional_timeout(timeout),
|
||||
shaping,
|
||||
})
|
||||
}
|
||||
|
||||
fn merged_headers(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) -> PyResult<Option<Map<String, Value>>> {
|
||||
let request = self.request.bind(py);
|
||||
let mapping = |name: &str| -> PyResult<Option<Map<String, Value>>> {
|
||||
lookup(arguments, request, name)?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()
|
||||
};
|
||||
Ok(merge_headers(
|
||||
mapping("headers")?,
|
||||
mapping("extra_headers")?,
|
||||
))
|
||||
}
|
||||
|
||||
fn provider_specific_header(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) -> PyResult<Option<ProviderSpecificHeaders>> {
|
||||
lookup(arguments, self.request.bind(py), "provider_specific_header")?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn shaping(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) -> PyResult<MessagesShaping> {
|
||||
let projected = py.import(ROUTE_HOST_MODULE)?.getattr("shaping")?.call1((
|
||||
model,
|
||||
custom_llm_provider,
|
||||
arguments,
|
||||
))?;
|
||||
from_py(&projected)
|
||||
}
|
||||
|
||||
fn provider(&self, py: Python<'_>) -> String {
|
||||
self.request
|
||||
.bind(py)
|
||||
|
|
@ -112,7 +193,7 @@ impl MessagesRouteHost {
|
|||
return error;
|
||||
}
|
||||
let mapped = py
|
||||
.import("litellm.rust_bridge.messages.route_host")
|
||||
.import(ROUTE_HOST_MODULE)
|
||||
.and_then(|module| module.getattr("map_failure"))
|
||||
.and_then(|map| map.call1((error.value(py), self.request.bind(py), self.provider(py))))
|
||||
.and_then(|mapped| {
|
||||
|
|
@ -148,7 +229,7 @@ impl RouteHost for MessagesRouteHost {
|
|||
fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult<Py<PyAny>> {
|
||||
match response {
|
||||
MessagesOutput::Message(message) => py
|
||||
.import("litellm.rust_bridge.messages.route_host")?
|
||||
.import(ROUTE_HOST_MODULE)?
|
||||
.getattr("response")?
|
||||
.call1((to_py(py, message.as_ref())?,))
|
||||
.map(Bound::unbind),
|
||||
|
|
@ -161,17 +242,12 @@ impl RouteHost for MessagesRouteHost {
|
|||
}
|
||||
|
||||
fn classify(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
|
||||
let native = match error {
|
||||
Error::Transport(TransportError::Http { status, body }) => {
|
||||
let error = RustUpstreamError::new_err((status, body));
|
||||
error
|
||||
.value(py)
|
||||
.setattr("headers", Vec::<(String, String)>::new())?;
|
||||
error
|
||||
}
|
||||
other => messages_error_to_pyerr(other),
|
||||
};
|
||||
Ok(self.map_failure(py, native))
|
||||
if let Error::Secret(source) = &error
|
||||
&& let Some(original) = crate::secrets::python_error(py, source.source_error())
|
||||
{
|
||||
return Ok(original);
|
||||
}
|
||||
Ok(self.map_failure(py, native_error(py, error)?))
|
||||
}
|
||||
|
||||
fn host_error(error: &PyErr) -> Error {
|
||||
|
|
@ -184,3 +260,62 @@ impl RouteHost for MessagesRouteHost {
|
|||
visit.call(&self.request)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn map(value: Value) -> Map<String, Value> {
|
||||
serde_json::from_value(value).unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::extra_over_forwarded(
|
||||
Some(json!({"X-Priority": "forwarded", "X-Forwarded-Only": "keep"})),
|
||||
Some(json!({"X-Priority": "extra", "X-Extra-Only": "also-keep"})),
|
||||
Some(json!({"X-Priority": "extra", "X-Forwarded-Only": "keep", "X-Extra-Only": "also-keep"})),
|
||||
)]
|
||||
#[case::only_forwarded(Some(json!({"X-Forwarded": "yes"})), None, Some(json!({"X-Forwarded": "yes"})))]
|
||||
#[case::only_extra_headers(
|
||||
None,
|
||||
Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})),
|
||||
Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})),
|
||||
)]
|
||||
#[case::nothing(None, Some(json!({})), None)]
|
||||
fn headers_merge_forwarded_then_extra(
|
||||
#[case] forwarded: Option<Value>,
|
||||
#[case] extra_headers: Option<Value>,
|
||||
#[case] expected: Option<Value>,
|
||||
) {
|
||||
assert_eq!(
|
||||
merge_headers(forwarded.map(map), extra_headers.map(map)),
|
||||
expected.map(map)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::rejected_request(Error::InvalidRequest("does not support top_k=5".into()), true)]
|
||||
#[case::unresolvable_provider(Error::InvalidProvider("openai".into()), false)]
|
||||
#[case::upstream_failure(
|
||||
Error::Transport(TransportError::Http { status: 400, body: "bad".into() }),
|
||||
false,
|
||||
)]
|
||||
fn only_request_rejections_carry_the_request_error_marker(
|
||||
#[case] error: Error,
|
||||
#[case] marked: bool,
|
||||
) {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let native = native_error(py, error).unwrap();
|
||||
let marker = native
|
||||
.value(py)
|
||||
.getattr_opt(REQUEST_ERROR_MARKER)
|
||||
.unwrap()
|
||||
.map(|value| value.extract::<bool>().unwrap());
|
||||
assert_eq!(marker.unwrap_or(false), marked);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -39,11 +39,12 @@ fn run_messages(
|
|||
"the Rust Messages route does not serve this provider",
|
||||
));
|
||||
}
|
||||
let secrets = crate::secrets::source(py)?;
|
||||
run_legacy_call(
|
||||
py,
|
||||
SURFACE,
|
||||
PublicCall::capture(&request, &args, &kwargs)?,
|
||||
crate::logger::LoggedMachine::new(messages_machine()),
|
||||
crate::logger::LoggedMachine::new(messages_machine(secrets)),
|
||||
MessagesRouteHost::new(request.unbind()),
|
||||
asynchronous,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8,3 +8,6 @@ repository.workspace = true
|
|||
[dependencies]
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -17,12 +17,48 @@ pub enum MessageContent {
|
|||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ContentBlock {
|
||||
#[serde(rename = "type", default, skip_serializing_if = "Option::is_none")]
|
||||
pub block_type: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub text: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub thinking: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub signature: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub data: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub name: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tool_use_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub content: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub provider_specific_fields: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_control: Option<CacheControl>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
impl ContentBlock {
|
||||
pub fn text(text: impl Into<String>) -> Self {
|
||||
Self {
|
||||
block_type: Some("text".to_string()),
|
||||
text: Some(text.into()),
|
||||
..Self::default()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_type(&self, block_type: &str) -> bool {
|
||||
self.block_type.as_deref() == Some(block_type)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct CacheControl {
|
||||
#[serde(rename = "type", skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -85,6 +121,126 @@ pub struct AnthropicMessagesRequest {
|
|||
pub speed: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub inference_geo: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub reasoning_effort: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub compaction: Option<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
impl AnthropicMessage {
|
||||
pub fn blocks(&self) -> &[ContentBlock] {
|
||||
match &self.content {
|
||||
MessageContent::Blocks(blocks) => blocks,
|
||||
MessageContent::Text(_) => &[],
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_blocks(self, blocks: Vec<ContentBlock>) -> Self {
|
||||
Self {
|
||||
content: MessageContent::Blocks(blocks),
|
||||
..self
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn round_trip<T: serde::de::DeserializeOwned + Serialize>(value: &Value) -> Value {
|
||||
let parsed: T = serde_json::from_value(value.clone()).unwrap();
|
||||
serde_json::to_value(parsed).unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::text(json!({"type": "text", "text": "hi"}))]
|
||||
#[case::text_with_citations_and_cache_control(json!({
|
||||
"type": "text",
|
||||
"text": "hi",
|
||||
"citations": [{"type": "char_location", "cited_text": "x"}],
|
||||
"cache_control": {"type": "ephemeral", "ttl": "1h", "scope": "global", "future": 1}
|
||||
}))]
|
||||
#[case::image(json!({"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "AA=="}}))]
|
||||
#[case::thinking(json!({"type": "thinking", "thinking": "hmm", "signature": "sig"}))]
|
||||
#[case::redacted_thinking(json!({"type": "redacted_thinking", "data": "opaque"}))]
|
||||
#[case::tool_use(json!({"type": "tool_use", "id": "toolu_1", "name": "f", "input": {"q": [1, null]}}))]
|
||||
#[case::tool_result_with_text(json!({"type": "tool_result", "tool_use_id": "toolu_1", "content": "ok", "is_error": false}))]
|
||||
#[case::tool_result_with_blocks(json!({"type": "tool_result", "tool_use_id": "toolu_1", "content": [{"type": "text", "text": "ok"}]}))]
|
||||
#[case::web_search_result_with_nulls(json!({
|
||||
"type": "web_search_tool_result",
|
||||
"tool_use_id": "srvtoolu_1",
|
||||
"content": [{"type": "web_search_result", "url": "u", "page_age": null, "encrypted_content": ""}]
|
||||
}))]
|
||||
#[case::provider_specific_fields(json!({"type": "tool_use", "id": "t", "name": "f", "input": {}, "provider_specific_fields": {"x": 1}}))]
|
||||
#[case::untyped(json!({"unknown": {"nested": true}}))]
|
||||
fn content_block_round_trips_unchanged(#[case] block: Value) {
|
||||
assert_eq!(round_trip::<ContentBlock>(&block), block);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn text_constructor_serializes_as_a_text_block() {
|
||||
assert_eq!(
|
||||
serde_json::to_value(ContentBlock::text("hello")).unwrap(),
|
||||
json!({"type": "text", "text": "hello"})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::same_type(json!({"type": "tool_use"}), "tool_use", true)]
|
||||
#[case::other_type(json!({"type": "tool_result"}), "tool_use", false)]
|
||||
#[case::prefix_of_type(json!({"type": "tool_use"}), "tool", false)]
|
||||
#[case::no_type(json!({"text": "x"}), "text", false)]
|
||||
fn is_type_matches_the_exact_block_type(
|
||||
#[case] block: Value,
|
||||
#[case] block_type: &str,
|
||||
#[case] expected: bool,
|
||||
) {
|
||||
let block: ContentBlock = serde_json::from_value(block).unwrap();
|
||||
assert_eq!(block.is_type(block_type), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::string_content(json!({"role": "user", "content": "hi"}), vec![])]
|
||||
#[case::block_content(
|
||||
json!({"role": "user", "content": [{"type": "text", "text": "a"}, {"type": "text", "text": "b"}]}),
|
||||
vec![ContentBlock::text("a"), ContentBlock::text("b")],
|
||||
)]
|
||||
fn message_blocks_list_only_block_content(
|
||||
#[case] message: Value,
|
||||
#[case] expected: Vec<ContentBlock>,
|
||||
) {
|
||||
let message: AnthropicMessage = serde_json::from_value(message).unwrap();
|
||||
assert_eq!(message.blocks(), expected.as_slice());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::replaces_string_content(json!({"role": "assistant", "content": "old", "name": "kept"}))]
|
||||
#[case::replaces_block_content(json!({"role": "assistant", "content": [{"type": "text", "text": "old"}], "name": "kept"}))]
|
||||
fn with_blocks_replaces_content_and_keeps_the_rest(#[case] message: Value) {
|
||||
let message: AnthropicMessage = serde_json::from_value(message).unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(message.with_blocks(vec![ContentBlock::text("new")])).unwrap(),
|
||||
json!({"role": "assistant", "content": [{"type": "text", "text": "new"}], "name": "kept"})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::minimal(json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]}))]
|
||||
#[case::reasoning_effort_compaction_and_unknown_fields(json!({
|
||||
"model": "m",
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
|
||||
"max_tokens": 8,
|
||||
"reasoning_effort": "high",
|
||||
"compaction": {"type": "auto"},
|
||||
"safeguards": [{"type": "dangerous_tool_use", "classifier_context": {"v": 1}}],
|
||||
"metadata": {"user_id": "u"}
|
||||
}))]
|
||||
fn request_round_trips_unchanged(#[case] request: Value) {
|
||||
assert_eq!(round_trip::<AnthropicMessagesRequest>(&request), request);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,8 +9,6 @@ pub struct AnthropicMessagesResponse {
|
|||
pub role: String,
|
||||
pub model: String,
|
||||
pub content: Vec<Value>,
|
||||
// Anthropic always includes stop_reason / stop_sequence, null until the turn
|
||||
// ends; serialize them even when None so callers see the same shape as Python.
|
||||
pub stop_reason: Option<String>,
|
||||
pub stop_sequence: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -20,3 +18,61 @@ pub struct AnthropicMessagesResponse {
|
|||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn response(
|
||||
stop_reason: Option<&str>,
|
||||
stop_sequence: Option<&str>,
|
||||
usage: Option<Value>,
|
||||
container: Option<Value>,
|
||||
) -> AnthropicMessagesResponse {
|
||||
AnthropicMessagesResponse {
|
||||
id: "msg_1".to_string(),
|
||||
message_type: "message".to_string(),
|
||||
role: "assistant".to_string(),
|
||||
model: "claude".to_string(),
|
||||
content: vec![],
|
||||
stop_reason: stop_reason.map(str::to_string),
|
||||
stop_sequence: stop_sequence.map(str::to_string),
|
||||
usage,
|
||||
container,
|
||||
extra: Map::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::turn_in_progress(None, None, json!(null), json!(null))]
|
||||
#[case::ended_on_end_turn(Some("end_turn"), None, json!("end_turn"), json!(null))]
|
||||
#[case::ended_on_stop_sequence(Some("stop_sequence"), Some("###"), json!("stop_sequence"), json!("###"))]
|
||||
fn stop_fields_are_always_serialized(
|
||||
#[case] stop_reason: Option<&str>,
|
||||
#[case] stop_sequence: Option<&str>,
|
||||
#[case] expected_reason: Value,
|
||||
#[case] expected_sequence: Value,
|
||||
) {
|
||||
let body: Value = serde_json::to_value(response(stop_reason, stop_sequence, None, None))
|
||||
.expect("serializable");
|
||||
assert_eq!(body.get("stop_reason"), Some(&expected_reason));
|
||||
assert_eq!(body.get("stop_sequence"), Some(&expected_sequence));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::absent(None, None)]
|
||||
#[case::present(Some(json!({"input_tokens": 1})), Some(json!({"id": "c_1"})))]
|
||||
fn usage_and_container_are_omitted_only_when_none(
|
||||
#[case] usage: Option<Value>,
|
||||
#[case] container: Option<Value>,
|
||||
) {
|
||||
let body: Value =
|
||||
serde_json::to_value(response(None, None, usage.clone(), container.clone()))
|
||||
.expect("serializable");
|
||||
assert_eq!(body.get("usage").cloned(), usage);
|
||||
assert_eq!(body.get("container").cloned(), container);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,6 +3,21 @@ use serde_json::{Map, Value};
|
|||
|
||||
use crate::llms::openai::{ChatCompletionThinkingBlock, ChatCompletionToolCallChunk};
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ProviderSpecificHeader {
|
||||
#[serde(default)]
|
||||
pub custom_llm_provider: String,
|
||||
#[serde(default)]
|
||||
pub extra_headers: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum ProviderSpecificHeaders {
|
||||
One(ProviderSpecificHeader),
|
||||
Many(Vec<ProviderSpecificHeader>),
|
||||
}
|
||||
|
||||
/// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python
|
||||
/// path reports so cost tracking sees the same numbers on either path.
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
|
|
|
|||
|
|
@ -11,7 +11,10 @@ from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_
|
|||
from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRUsageInfo
|
||||
from litellm.llms.bedrock.batches.transformation import titan_embedding_usage_from_batch_output
|
||||
from litellm.llms.vertex_ai.batches.transformation import vertex_prompt_tokens_details
|
||||
from litellm.llms.vertex_ai.batches.transformation import (
|
||||
is_native_vertex_batch_output_row,
|
||||
native_vertex_batch_row_stats,
|
||||
)
|
||||
from litellm.types.llms.openai import Batch
|
||||
from litellm.types.utils import ModelInfo, Usage
|
||||
from litellm.utils import token_counter
|
||||
|
|
@ -31,6 +34,20 @@ class BatchCostUsageResult:
|
|||
|
||||
|
||||
_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"})
|
||||
|
||||
|
||||
def _uses_native_vertex_output(
|
||||
custom_llm_provider: str,
|
||||
model_name: str | None,
|
||||
first_row: Mapping[str, object] | None,
|
||||
) -> bool:
|
||||
if custom_llm_provider != "vertex_ai":
|
||||
return False
|
||||
if model_name and getattr(litellm, "disable_vertex_batch_output_transformation", False):
|
||||
return True
|
||||
return first_row is not None and is_native_vertex_batch_output_row(first_row)
|
||||
|
||||
|
||||
_TERMINAL_BATCH_STATUSES: Final = _COMPLETED_BATCH_STATUSES | frozenset({"failed", "cancelled", "expired"})
|
||||
|
||||
|
||||
|
|
@ -66,12 +83,9 @@ async def calculate_batch_cost_and_usage(
|
|||
deployment-specific pricing (e.g. input_cost_per_token_batches)
|
||||
is used instead of the global cost map.
|
||||
"""
|
||||
if (
|
||||
custom_llm_provider == "vertex_ai"
|
||||
and model_name
|
||||
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
|
||||
):
|
||||
return calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
|
||||
first_row: Final = file_content_dictionary[0] if file_content_dictionary else None
|
||||
if _uses_native_vertex_output(custom_llm_provider, model_name, first_row):
|
||||
return calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name, model_info=model_info)
|
||||
|
||||
return _aggregate_batch_cost_usage_models(
|
||||
entries=file_content_dictionary,
|
||||
|
|
@ -126,11 +140,11 @@ async def _handle_completed_batch(
|
|||
)
|
||||
|
||||
output_file_result: Final = (
|
||||
calculate_vertex_ai_batch_cost_and_usage(_get_file_content_as_dictionary(file_content), model_name)
|
||||
if (
|
||||
custom_llm_provider == "vertex_ai"
|
||||
and model_name
|
||||
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
|
||||
calculate_vertex_ai_batch_cost_and_usage(
|
||||
_iter_batch_output_entries(file_content), model_name, model_info=model_info
|
||||
)
|
||||
if _uses_native_vertex_output(
|
||||
custom_llm_provider, model_name, next(_iter_batch_output_entries(file_content), None)
|
||||
)
|
||||
else _aggregate_batch_cost_usage_models(
|
||||
entries=_iter_batch_output_entries(file_content),
|
||||
|
|
@ -332,69 +346,36 @@ def _aggregate_batch_cost_usage_models(
|
|||
|
||||
|
||||
def calculate_vertex_ai_batch_cost_and_usage(
|
||||
vertex_ai_batch_responses: list[dict],
|
||||
vertex_ai_batch_responses: Iterable[dict],
|
||||
model_name: str | None = None,
|
||||
model_info: ModelInfo | None = None,
|
||||
) -> BatchCostUsageResult:
|
||||
"""
|
||||
Calculate both cost and usage from raw Vertex AI batch responses.
|
||||
|
||||
Used only when ``litellm.disable_vertex_batch_output_transformation = True``.
|
||||
In that case the GCS predictions.jsonl is returned as-is, with each line in
|
||||
the native Vertex format:
|
||||
|
||||
{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}}}
|
||||
|
||||
usageMetadata contains promptTokenCount, candidatesTokenCount, totalTokenCount.
|
||||
|
||||
A row with no ``response`` is counted as failed - the same signal already
|
||||
used to skip it from cost/usage aggregation, since Vertex batch prediction
|
||||
output doesn't establish a distinct error shape in this (non-default) path.
|
||||
Cost and usage of a native Vertex predictions.jsonl, one
|
||||
`{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}, "modelVersion": ...}}`
|
||||
generateContent row or `{"request": ..., "response": {"embedding": {...}, "usageMetadata": {...}}}`
|
||||
embedding row per line. `model_name` (the deployment model) prices every row, else each row's own
|
||||
`modelVersion` does; a row without a usable response counts as failed.
|
||||
"""
|
||||
from litellm.cost_calculator import batch_cost_calculator
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
|
||||
total_prompt_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
|
||||
total_completion_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
|
||||
total_tokens = 0
|
||||
prompt_tokens = 0
|
||||
completion_tokens = 0
|
||||
successful_requests = 0 # rebind-ok: loop accumulator, matches total_cost/total_tokens above
|
||||
failed_requests = 0 # rebind-ok: loop accumulator, matches total_cost/total_tokens above
|
||||
actual_model_name: Final = model_name or "gemini-2.0-flash-001"
|
||||
|
||||
for response in vertex_ai_batch_responses:
|
||||
response_body = response.get("response")
|
||||
if response_body is None:
|
||||
failed_requests += 1
|
||||
continue
|
||||
successful_requests += 1
|
||||
|
||||
usage_metadata = response_body.get("usageMetadata", {})
|
||||
_prompt = usage_metadata.get("promptTokenCount", 0) or 0
|
||||
_completion = usage_metadata.get("candidatesTokenCount", 0) or 0
|
||||
_total = usage_metadata.get("totalTokenCount", 0) or (_prompt + _completion)
|
||||
|
||||
line_usage = Usage(
|
||||
prompt_tokens=_prompt,
|
||||
completion_tokens=_completion,
|
||||
total_tokens=_total,
|
||||
prompt_tokens_details=vertex_prompt_tokens_details(usage_metadata),
|
||||
row_stats: Final = tuple(
|
||||
native_vertex_batch_row_stats(
|
||||
row,
|
||||
model_name,
|
||||
model_info=model_info,
|
||||
calculate_usage=VertexGeminiConfig._calculate_usage,
|
||||
cost_calculator=batch_cost_calculator,
|
||||
)
|
||||
|
||||
try:
|
||||
p_cost, c_cost = batch_cost_calculator(
|
||||
usage=line_usage,
|
||||
model=actual_model_name,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
total_prompt_cost += p_cost
|
||||
total_completion_cost += c_cost
|
||||
except Exception as e:
|
||||
verbose_logger.debug("vertex_ai batch cost calculation error for line: %s", str(e))
|
||||
|
||||
prompt_tokens += _prompt
|
||||
completion_tokens += _completion
|
||||
total_tokens += _total
|
||||
|
||||
for row in vertex_ai_batch_responses
|
||||
)
|
||||
priced: Final = tuple(stats for stats in row_stats if stats is not None)
|
||||
total_prompt_cost: Final = sum(stats.prompt_cost for stats in priced)
|
||||
total_completion_cost: Final = sum(stats.completion_cost for stats in priced)
|
||||
prompt_tokens: Final = sum(stats.usage.prompt_tokens for stats in priced)
|
||||
completion_tokens: Final = sum(stats.usage.completion_tokens for stats in priced)
|
||||
total_tokens: Final = sum(stats.total_tokens for stats in priced)
|
||||
total_cost: Final = total_prompt_cost + total_completion_cost
|
||||
verbose_logger.info(
|
||||
"vertex_ai batch cost: cost=%s, prompt=%d, completion=%d, total=%d, successful=%d, failed=%d",
|
||||
|
|
@ -402,8 +383,8 @@ def calculate_vertex_ai_batch_cost_and_usage(
|
|||
prompt_tokens,
|
||||
completion_tokens,
|
||||
total_tokens,
|
||||
successful_requests,
|
||||
failed_requests,
|
||||
len(priced),
|
||||
len(row_stats) - len(priced),
|
||||
)
|
||||
|
||||
return BatchCostUsageResult(
|
||||
|
|
@ -413,9 +394,13 @@ def calculate_vertex_ai_batch_cost_and_usage(
|
|||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
),
|
||||
models=[actual_model_name],
|
||||
successful_requests=successful_requests,
|
||||
failed_requests=failed_requests,
|
||||
models=(
|
||||
[model_name]
|
||||
if model_name
|
||||
else list(dict.fromkeys(stats.model for stats in priced if stats.model is not None))
|
||||
),
|
||||
successful_requests=len(priced),
|
||||
failed_requests=len(row_stats) - len(priced),
|
||||
prompt_cost=total_prompt_cost,
|
||||
completion_cost=total_completion_cost,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -429,6 +429,7 @@ class LLMCachingHandler:
|
|||
kwargs=kwargs,
|
||||
cached_result=cached_result,
|
||||
is_async=False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
if not _should_defer_streaming_cache_hit_callbacks(cached_result=cached_result):
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
|
|||
DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5))
|
||||
DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10))
|
||||
DEFAULT_S3_BATCH_SIZE: Final = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512))
|
||||
DEFAULT_S3_MAX_CONCURRENT_UPLOADS: Final = int(os.getenv("DEFAULT_S3_MAX_CONCURRENT_UPLOADS", "16"))
|
||||
# https://docs.aws.amazon.com/AmazonS3/latest/userguide/object-keys.html
|
||||
MAX_S3_OBJECT_KEY_BYTES: Final = 1024
|
||||
S3_BOUNDED_OBJECT_KEY_HEAD_BYTES: Final = 64
|
||||
|
|
@ -1518,6 +1519,7 @@ PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES: Final = int(
|
|||
CLOUDZERO_EXPORT_INTERVAL_MINUTES: Final = int(os.getenv("CLOUDZERO_EXPORT_INTERVAL_MINUTES", 60))
|
||||
MCP_TOOL_NAME_PREFIX: Final = "mcp_tool"
|
||||
MAXIMUM_TRACEBACK_LINES_TO_LOG: Final = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100))
|
||||
PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: Final = 4096
|
||||
|
||||
# Headers to control callbacks
|
||||
X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks"
|
||||
|
|
@ -1770,6 +1772,8 @@ RESPONSES_SESSION_LOOKUP_MAX_ATTEMPTS: Final = max(1, int(os.getenv("RESPONSES_S
|
|||
RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL: Final = float(os.getenv("RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL", "0.2"))
|
||||
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000))
|
||||
PROXY_DB_LOOKUP_MAX_CONCURRENCY: Final = max(1, int(os.getenv("PROXY_DB_LOOKUP_MAX_CONCURRENCY", "25")))
|
||||
PROXY_DB_LOOKUP_DEADLINE_SECONDS: Final = max(0.1, float(os.getenv("PROXY_DB_LOOKUP_DEADLINE_SECONDS", "10")))
|
||||
PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS: Final = max(0.0, float(os.getenv("PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS", "30")))
|
||||
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute
|
||||
PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597))
|
||||
RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500")))
|
||||
|
|
|
|||
|
|
@ -2017,8 +2017,8 @@ def _deployment_model_info(
|
|||
return cast(ModelInfo, registered_deployment_info) # cast-ok: router registers deployment prices under its id
|
||||
if litellm_logging_obj is None:
|
||||
return None
|
||||
litellm_params: Final = getattr(litellm_logging_obj, "litellm_params", None)
|
||||
if litellm_params is None:
|
||||
litellm_params: Final = litellm_logging_obj.litellm_params
|
||||
if not litellm_params:
|
||||
return None
|
||||
return next(
|
||||
(
|
||||
|
|
@ -2036,7 +2036,9 @@ def _ocr_model_info(
|
|||
router_model_id: str | None,
|
||||
) -> OCRPricing | None:
|
||||
deployment_info: Final = _deployment_model_info(litellm_logging_obj, custom_pricing, router_model_id)
|
||||
litellm_params: Final = getattr(litellm_logging_obj, "litellm_params", None) if custom_pricing else None
|
||||
litellm_params: Final = (
|
||||
litellm_logging_obj.litellm_params if custom_pricing and litellm_logging_obj is not None else None
|
||||
)
|
||||
if litellm_params is None:
|
||||
return deployment_info
|
||||
return _layered_ocr_pricing(litellm_params, deployment_info)
|
||||
|
|
|
|||
|
|
@ -129,7 +129,7 @@ async def list_tools_with_pagination(
|
|||
)
|
||||
tools.extend(result.tools)
|
||||
|
||||
next_cursor = getattr(result, "next_cursor", None)
|
||||
next_cursor = result.next_cursor
|
||||
if not isinstance(next_cursor, str) or not next_cursor:
|
||||
return tools
|
||||
if next_cursor in seen_cursors:
|
||||
|
|
|
|||
|
|
@ -176,6 +176,22 @@ def create_file(
|
|||
if logging_obj is None:
|
||||
raise ValueError("logging_obj is required")
|
||||
client: Final = kwargs.get("client")
|
||||
if litellm_params_dict.get("passthrough") is True and (
|
||||
custom_llm_provider != "vertex_ai" or purpose != "batch"
|
||||
):
|
||||
raise litellm.exceptions.BadRequestError(
|
||||
message=(
|
||||
"`passthrough=True` uploads the file bytes unchanged for a native Vertex AI batch, so it needs "
|
||||
f"custom_llm_provider='vertex_ai' and purpose='batch', got '{custom_llm_provider}' and '{purpose}'."
|
||||
),
|
||||
model="n/a",
|
||||
llm_provider=custom_llm_provider or "n/a",
|
||||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="passthrough needs a vertex_ai batch",
|
||||
request=httpx.Request(method="create_file", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
|
||||
|
|
|
|||
|
|
@ -112,7 +112,7 @@ class ArizeLogger(OpenTelemetry):
|
|||
if value is None or value in ("", "None"):
|
||||
return None
|
||||
try:
|
||||
rate = float(value)
|
||||
rate: Final = float(value)
|
||||
except (TypeError, ValueError):
|
||||
verbose_logger.warning(
|
||||
"ArizeLogger: %s value %r is not a number; exporting the request",
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ from __future__ import annotations
|
|||
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -35,17 +35,6 @@ else:
|
|||
AsyncIOScheduler = Any
|
||||
|
||||
|
||||
class _PodLockManager(Protocol):
|
||||
"""The subset of PodLockManager this logger drives to serialize the export across pods."""
|
||||
|
||||
@property
|
||||
def redis_cache(self) -> object: ...
|
||||
|
||||
async def acquire_lock(self, cronjob_id: str) -> bool | None: ...
|
||||
|
||||
async def release_lock(self, cronjob_id: str) -> None: ...
|
||||
|
||||
|
||||
def _parse_metrics_marker(
|
||||
marker: object | None,
|
||||
) -> datetime | None:
|
||||
|
|
@ -237,13 +226,10 @@ class MavvrikFocusLogger(FocusLogger):
|
|||
"""Scheduler entry point — uses Mavvrik-specific pod-lock key."""
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj # noqa: PLC0415
|
||||
|
||||
pod_lock_manager: _PodLockManager | None = None
|
||||
if proxy_logging_obj is not None:
|
||||
writer: Final[object] = getattr(proxy_logging_obj, "db_spend_update_writer", None)
|
||||
if writer is not None:
|
||||
pod_lock_manager = getattr(writer, "pod_lock_manager", None)
|
||||
|
||||
if pod_lock_manager and pod_lock_manager.redis_cache:
|
||||
pod_lock_manager: Final = (
|
||||
proxy_logging_obj.db_spend_update_writer.pod_lock_manager if proxy_logging_obj is not None else None
|
||||
)
|
||||
if pod_lock_manager is not None and pod_lock_manager.redis_cache:
|
||||
acquired: Final = await pod_lock_manager.acquire_lock(cronjob_id=MAVVRIK_FOCUS_EXPORT_JOB_NAME)
|
||||
if not acquired:
|
||||
verbose_proxy_logger.debug("Mavvrik FOCUS export: unable to acquire pod lock")
|
||||
|
|
|
|||
|
|
@ -62,10 +62,10 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy.proxy_server import UserAPIKeyAuth as _UserAPIKeyAuth
|
||||
|
||||
Span = _Span | Any
|
||||
Tracer = _Tracer | Any
|
||||
Context = _Context | Any
|
||||
SpanExporter = _SpanExporter | Any
|
||||
UserAPIKeyAuth = _UserAPIKeyAuth | Any
|
||||
Tracer = _Tracer
|
||||
Context = _Context
|
||||
SpanExporter = _SpanExporter
|
||||
UserAPIKeyAuth = _UserAPIKeyAuth
|
||||
ManagementEndpointLoggingPayload = _ManagementEndpointLoggingPayload | Any
|
||||
else:
|
||||
Span = Any
|
||||
|
|
@ -2730,7 +2730,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
self.handle_callback_failure(callback_name=self.callback_name or "opentelemetry")
|
||||
verbose_logger.exception("OpenTelemetry logging error in set_attributes %s", str(e))
|
||||
|
||||
def _cast_as_primitive_value_type(self, value) -> str | bool | int | float:
|
||||
def _cast_as_primitive_value_type(self, value: object) -> str | bool | int | float:
|
||||
"""
|
||||
Casts the value to a primitive OTEL type if it is not already a primitive type.
|
||||
|
||||
|
|
|
|||
|
|
@ -2778,6 +2778,13 @@ class PrometheusLogger(CustomLogger):
|
|||
- increment deployment failure responses metric
|
||||
- increment deployment total requests metric
|
||||
|
||||
Both counters also carry a model_group label. When a deployment was
|
||||
actually selected, model_group is the router-resolved value and is
|
||||
trusted as-is. On a pre-routing reject (no deployment selected), it
|
||||
is caller-supplied via litellm_params.metadata and is bounded with
|
||||
_bounded_requested_model_label the same way requested_model is, so an
|
||||
unrecognized value cannot mint unbounded label series.
|
||||
|
||||
Args:
|
||||
request_kwargs: dict
|
||||
|
||||
|
|
@ -2844,6 +2851,7 @@ class PrometheusLogger(CustomLogger):
|
|||
label_api_base = api_base
|
||||
label_api_provider = llm_provider
|
||||
label_requested_model = model_group or litellm_model_name
|
||||
label_model_group = model_group
|
||||
else:
|
||||
label_litellm_model_name = ""
|
||||
label_model_id = ""
|
||||
|
|
@ -2852,6 +2860,7 @@ class PrometheusLogger(CustomLogger):
|
|||
label_requested_model = (
|
||||
_bounded_requested_model_label(litellm_model_name or model_group, router_originated=True) or ""
|
||||
)
|
||||
label_model_group = _bounded_requested_model_label(model_group, router_originated=True)
|
||||
|
||||
enum_values: Final = UserAPIKeyLabelValues(
|
||||
litellm_model_name=label_litellm_model_name,
|
||||
|
|
@ -2861,6 +2870,7 @@ class PrometheusLogger(CustomLogger):
|
|||
exception_status=exception_status,
|
||||
exception_class=(self._get_exception_class_name(exception) if exception else None),
|
||||
requested_model=label_requested_model,
|
||||
model_group=label_model_group,
|
||||
hashed_api_key=hashed_api_key,
|
||||
api_key_alias=api_key_alias,
|
||||
user_email=user_email,
|
||||
|
|
@ -2912,9 +2922,21 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id: str | None,
|
||||
api_base: str | None,
|
||||
llm_provider: str | None,
|
||||
model_group: str | None,
|
||||
):
|
||||
"""
|
||||
Set the deployment TPM and RPM limits metrics
|
||||
|
||||
Args:
|
||||
model_info: the deployment's static model_info config (id, tpm, rpm, etc.)
|
||||
litellm_params: the deployment's litellm_params, as a tpm/rpm fallback source
|
||||
litellm_model_name: the resolved deployment model name
|
||||
model_id: the deployment's model_id
|
||||
api_base: the deployment's api_base
|
||||
llm_provider: the deployment's custom_llm_provider
|
||||
model_group: the router-resolved model_group the deployment belongs to,
|
||||
from the caller's already-resolved enum_values.model_group (trusted,
|
||||
not caller-supplied at this call site)
|
||||
"""
|
||||
tpm: Final = model_info.get("tpm") or litellm_params.get("tpm")
|
||||
rpm: Final = model_info.get("rpm") or litellm_params.get("rpm")
|
||||
|
|
@ -2927,6 +2949,7 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id=model_id,
|
||||
api_base=api_base,
|
||||
api_provider=llm_provider,
|
||||
model_group=model_group,
|
||||
),
|
||||
)
|
||||
self.litellm_deployment_tpm_limit.labels(**_labels).set(tpm)
|
||||
|
|
@ -2939,6 +2962,7 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id=model_id,
|
||||
api_base=api_base,
|
||||
api_provider=llm_provider,
|
||||
model_group=model_group,
|
||||
),
|
||||
)
|
||||
self.litellm_deployment_rpm_limit.labels(**_labels).set(rpm)
|
||||
|
|
@ -3058,6 +3082,7 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id=model_id,
|
||||
api_base=api_base,
|
||||
llm_provider=llm_provider,
|
||||
model_group=enum_values.model_group,
|
||||
)
|
||||
|
||||
remaining_requests: int | None = None
|
||||
|
|
|
|||
|
|
@ -20,7 +20,8 @@ from litellm.constants import (
|
|||
)
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
_S3_LOG_PROMPTS_ONLY: Final = TypeAdapter(bool)
|
||||
_S3_BOOL: Final = TypeAdapter(bool)
|
||||
_UPLOAD_BOUND: Final = TypeAdapter(int)
|
||||
|
||||
|
||||
def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] | None = None) -> bool:
|
||||
|
|
@ -29,12 +30,42 @@ def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] |
|
|||
if raw is None or raw == "":
|
||||
return False
|
||||
try:
|
||||
return _S3_LOG_PROMPTS_ONLY.validate_python(raw.strip() if isinstance(raw, str) else raw)
|
||||
return _S3_BOOL.validate_python(raw.strip() if isinstance(raw, str) else raw)
|
||||
except ValidationError:
|
||||
verbose_logger.warning("s3 logging: s3_log_prompts_only=%r is not a boolean, logging prompts only", raw)
|
||||
return True
|
||||
|
||||
|
||||
def resolve_s3_max_concurrent_uploads(configured: object, fallback: int) -> int:
|
||||
if configured is None or configured == "":
|
||||
return fallback
|
||||
try:
|
||||
bound: Final = _UPLOAD_BOUND.validate_python(configured.strip() if isinstance(configured, str) else configured)
|
||||
except ValidationError:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_max_concurrent_uploads=%r is not an integer, using %s", configured, fallback
|
||||
)
|
||||
return fallback
|
||||
if bound < 1:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_max_concurrent_uploads=%r must be at least 1, using %s", configured, fallback
|
||||
)
|
||||
return fallback
|
||||
return bound
|
||||
|
||||
|
||||
def resolve_s3_batch_file_upload(configured: object) -> bool:
|
||||
if configured is None or configured == "":
|
||||
return False
|
||||
try:
|
||||
return _S3_BOOL.validate_python(configured.strip() if isinstance(configured, str) else configured)
|
||||
except ValidationError:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_batch_file_upload=%r is not a boolean, keeping per-request objects", configured
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def prompts_only_payload(payload: StandardLoggingPayload) -> StandardLoggingPayload:
|
||||
return {**payload, "response": None}
|
||||
|
||||
|
|
|
|||
|
|
@ -3,26 +3,33 @@ s3 Bucket Logging Integration
|
|||
|
||||
async_log_success_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
async_log_failure_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
NOTE 1: S3 does not provide a BATCH PUT API endpoint, so we create tasks to upload each element individually
|
||||
NOTE 1: S3 does not provide a BATCH PUT API endpoint; by default each element is uploaded concurrently (bounded by s3_max_concurrent_uploads), or with s3_batch_file_upload the whole flush is written as one .jsonl file
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
from urllib.parse import quote
|
||||
from uuid import uuid4
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_SECONDS
|
||||
from litellm.constants import (
|
||||
DEFAULT_S3_BATCH_SIZE,
|
||||
DEFAULT_S3_FLUSH_INTERVAL_SECONDS,
|
||||
DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
|
||||
)
|
||||
from litellm.integrations.s3 import (
|
||||
get_s3_object_download_filename,
|
||||
get_s3_object_key,
|
||||
prompts_only_payload,
|
||||
resolve_s3_batch_file_upload,
|
||||
resolve_s3_log_prompts_only,
|
||||
resolve_s3_max_concurrent_uploads,
|
||||
resolve_sse_params,
|
||||
)
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
|
|
@ -43,7 +50,20 @@ if TYPE_CHECKING:
|
|||
from botocore.credentials import Credentials
|
||||
|
||||
|
||||
def _s3_key_parent(s3_object_key: str) -> str:
|
||||
return s3_object_key.rsplit("/", 1)[0] if "/" in s3_object_key else ""
|
||||
|
||||
|
||||
class S3BatchUploadError(Exception):
|
||||
def __init__(self, failed: int, total: int) -> None:
|
||||
self.failed = failed
|
||||
self.total = total
|
||||
super().__init__(f"{failed} of {total} S3 uploads failed; events kept in queue for the next flush")
|
||||
|
||||
|
||||
class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
||||
preserve_events_added_during_flush = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
s3_bucket_name: str | None = None,
|
||||
|
|
@ -71,6 +91,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_server_side_encryption: str | None = None,
|
||||
s3_sse_kms_key_id: str | None = None,
|
||||
s3_log_prompts_only: bool | None = None,
|
||||
s3_max_concurrent_uploads: int = DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
|
||||
s3_batch_file_upload: bool = False,
|
||||
s3_callback_params_override: dict | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
|
|
@ -112,7 +134,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_server_side_encryption=s3_server_side_encryption,
|
||||
s3_sse_kms_key_id=s3_sse_kms_key_id,
|
||||
s3_log_prompts_only=s3_log_prompts_only,
|
||||
s3_max_concurrent_uploads=s3_max_concurrent_uploads,
|
||||
s3_batch_file_upload=s3_batch_file_upload,
|
||||
)
|
||||
self._upload_semaphore = asyncio.Semaphore(self.s3_max_concurrent_uploads)
|
||||
verbose_logger.debug("s3 logger using endpoint url %s", s3_endpoint_url)
|
||||
|
||||
# IMPORTANT
|
||||
|
|
@ -168,6 +193,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_server_side_encryption: str | None = None,
|
||||
s3_sse_kms_key_id: str | None = None,
|
||||
s3_log_prompts_only: bool | None = None,
|
||||
s3_max_concurrent_uploads: int = DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
|
||||
s3_batch_file_upload: bool = False,
|
||||
params_source: dict | None = None,
|
||||
):
|
||||
"""
|
||||
|
|
@ -226,6 +253,16 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
params.get("s3_sse_kms_key_id") or s3_sse_kms_key_id,
|
||||
)
|
||||
|
||||
configured_bound: Final = params.get("s3_max_concurrent_uploads")
|
||||
self.s3_max_concurrent_uploads = resolve_s3_max_concurrent_uploads(
|
||||
s3_max_concurrent_uploads if configured_bound is None or configured_bound == "" else configured_bound,
|
||||
DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
|
||||
)
|
||||
|
||||
self.s3_batch_file_upload = s3_batch_file_upload or resolve_s3_batch_file_upload(
|
||||
params.get("s3_batch_file_upload")
|
||||
)
|
||||
|
||||
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.
|
||||
|
|
@ -347,7 +384,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
verbose_logger.exception("s3 Layer Error - %s", e)
|
||||
self.handle_callback_failure(callback_name="S3Logger")
|
||||
|
||||
async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement):
|
||||
async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement) -> bool:
|
||||
try:
|
||||
import base64
|
||||
import hashlib
|
||||
|
|
@ -364,7 +401,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
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)
|
||||
json_string: Final = (
|
||||
batch_logging_element.body
|
||||
if batch_logging_element.body is not None
|
||||
else safe_dumps(batch_logging_element.payload)
|
||||
)
|
||||
|
||||
# Calculate SHA256 hash of the content
|
||||
content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
|
||||
|
|
@ -374,7 +415,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
# Prepare the request
|
||||
headers: Final = {
|
||||
"Content-Type": "application/json",
|
||||
"Content-Type": batch_logging_element.content_type,
|
||||
"Content-MD5": content_md5,
|
||||
"x-amz-content-sha256": content_hash,
|
||||
"Content-Language": "en",
|
||||
|
|
@ -421,27 +462,72 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
except Exception as e:
|
||||
verbose_logger.exception("Error uploading to s3: %s", e)
|
||||
self.handle_callback_failure(callback_name="S3Logger")
|
||||
return False
|
||||
return True
|
||||
|
||||
async def async_send_batch(self):
|
||||
async def async_send_batch(self) -> None:
|
||||
"""
|
||||
Sends runs from self.log_queue.
|
||||
|
||||
Sends runs from self.log_queue
|
||||
|
||||
Returns: None
|
||||
|
||||
Raises: Does not raise an exception, will only verbose_logger.exception()
|
||||
Raises S3BatchUploadError when any upload failed; CustomBatchLogger.flush_queue
|
||||
keeps the surviving queue entries for the next flush.
|
||||
"""
|
||||
verbose_logger.debug("s3_v2 logger - sending batch of %s", len(self.log_queue))
|
||||
if not self.log_queue:
|
||||
batch: Final = tuple(self.log_queue)
|
||||
if not batch:
|
||||
return
|
||||
verbose_logger.debug("s3_v2 logger - sending batch of %s", len(batch))
|
||||
|
||||
#########################################################
|
||||
# Flush the log queue to s3
|
||||
# the log queue can be bounded by DEFAULT_S3_BATCH_SIZE
|
||||
# see custom_batch_logger.py which triggers the flush
|
||||
#########################################################
|
||||
for payload in self.log_queue:
|
||||
asyncio.create_task(self.async_upload_data_to_s3(payload))
|
||||
uploads: Final = self._batch_file_elements(batch) if self._batch_file_mode_active() else batch
|
||||
results: Final = await asyncio.gather(*(self._upload_bounded(element) for element in uploads))
|
||||
failed: Final = tuple(element for element, ok in zip(uploads, results, strict=True) if not ok)
|
||||
if not failed:
|
||||
return
|
||||
self.log_queue = [*failed, *self.log_queue[len(batch) :]]
|
||||
raise S3BatchUploadError(failed=len(failed), total=len(uploads))
|
||||
|
||||
def _batch_file_mode_active(self) -> bool:
|
||||
if not self.s3_batch_file_upload:
|
||||
return False
|
||||
if litellm.cold_storage_custom_logger == "s3_v2":
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_batch_file_upload is ignored because s3_v2 is the cold storage logger; "
|
||||
"per-request objects are required for spend log lookups"
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
async def _upload_bounded(self, element: s3BatchLoggingElement) -> bool:
|
||||
async with self._upload_semaphore:
|
||||
return await self.async_upload_data_to_s3(element)
|
||||
|
||||
def _batch_file_elements(self, batch: tuple[s3BatchLoggingElement, ...]) -> tuple[s3BatchLoggingElement, ...]:
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
groups: Final = {
|
||||
parent: tuple(
|
||||
element for element in batch if element.body is None and _s3_key_parent(element.s3_object_key) == parent
|
||||
)
|
||||
for parent in sorted({_s3_key_parent(element.s3_object_key) for element in batch if element.body is None})
|
||||
}
|
||||
return tuple(element for element in batch if element.body is not None) + tuple(
|
||||
self._build_batch_file_element(elements, parent, now) for parent, elements in groups.items()
|
||||
)
|
||||
|
||||
def _build_batch_file_element(
|
||||
self, elements: tuple[s3BatchLoggingElement, ...], parent: str, now: datetime
|
||||
) -> s3BatchLoggingElement:
|
||||
batch_name: Final = f"batch_{now.strftime('%H-%M-%S')}_{uuid4().hex}"
|
||||
return s3BatchLoggingElement(
|
||||
payload={},
|
||||
body="\n".join(safe_dumps(element.payload) for element in elements),
|
||||
content_type="application/x-ndjson",
|
||||
s3_object_key=f"{parent}/{batch_name}.jsonl" if parent else f"{batch_name}.jsonl",
|
||||
s3_object_download_filename=f"{batch_name}.jsonl",
|
||||
)
|
||||
|
||||
def create_s3_batch_logging_element(
|
||||
self,
|
||||
|
|
@ -521,7 +607,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
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)
|
||||
json_string: Final = (
|
||||
batch_logging_element.body
|
||||
if batch_logging_element.body is not None
|
||||
else safe_dumps(batch_logging_element.payload)
|
||||
)
|
||||
|
||||
# Calculate SHA256 hash of the content
|
||||
content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
|
||||
|
|
@ -531,7 +621,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
# Prepare the request
|
||||
headers: Final = {
|
||||
"Content-Type": "application/json",
|
||||
"Content-Type": batch_logging_element.content_type,
|
||||
"Content-MD5": content_md5,
|
||||
"x-amz-content-sha256": content_hash,
|
||||
"Content-Language": "en",
|
||||
|
|
|
|||
|
|
@ -1849,7 +1849,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
for tool_call in tool_calls:
|
||||
# Handle both Anthropic-style input and OpenAI-style function.arguments
|
||||
query = None
|
||||
tool_args: dict | None = None # mutable-ok: the tool call's own arguments dict
|
||||
tool_args: dict[str, object] | None = None # mutable-ok: the tool call's own arguments dict
|
||||
if "input" in tool_call and isinstance(tool_call["input"], dict):
|
||||
tool_args = tool_call["input"]
|
||||
query = tool_args.get("query")
|
||||
|
|
|
|||
|
|
@ -365,7 +365,7 @@ def _budget_reservation_on_auth_object(user_api_key_auth: object) -> object:
|
|||
return getattr(user_api_key_auth, "budget_reservation", None)
|
||||
|
||||
|
||||
def budget_reservation_from_metadata(metadata: Mapping[str, object]) -> dict | None:
|
||||
def budget_reservation_from_metadata(metadata: Mapping[str, object]) -> dict[str, object] | None:
|
||||
stamped: Final = metadata.get("user_api_key_budget_reservation")
|
||||
if isinstance(stamped, dict):
|
||||
return stamped
|
||||
|
|
@ -776,3 +776,30 @@ def is_batch_line_item_event(kwargs: object) -> bool:
|
|||
return False
|
||||
typed_params: Final = cast(Mapping[str, object], litellm_params) # cast-ok: same narrowing limit
|
||||
return bool(typed_params.get("batch_parent_id"))
|
||||
|
||||
|
||||
_HIDDEN_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
_PROVIDER_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str])
|
||||
|
||||
|
||||
def set_provider_response_headers_in_hidden_params(
|
||||
response: _CarriesHiddenParams, headers: httpx.Headers | Mapping[str, str]
|
||||
) -> None:
|
||||
hidden_params: Final = response._hidden_params # pyright: ignore[reportPrivateUsage] # no public accessor
|
||||
existing_additional_headers: Final[object] = hidden_params.get("additional_headers")
|
||||
raw_headers: Final[dict[str, str]] = dict(headers) # mutable-ok: stored as the plain-dict hidden param
|
||||
additional_headers: Final[dict[str, object]] = { # mutable-ok: assigned into the plain-dict hidden params
|
||||
**process_response_headers(raw_headers),
|
||||
**(existing_additional_headers if isinstance(existing_additional_headers, Mapping) else _NO_HEADERS),
|
||||
}
|
||||
hidden_params["headers"] = raw_headers
|
||||
hidden_params["additional_headers"] = additional_headers
|
||||
|
||||
|
||||
def get_provider_response_headers_from_hidden_params(response: object) -> Mapping[str, str] | None:
|
||||
hidden_params: Final[object] = getattr(response, "_hidden_params", None)
|
||||
try:
|
||||
validated: Final = _HIDDEN_PARAMS_ADAPTER.validate_python(hidden_params)
|
||||
return _PROVIDER_HEADERS_ADAPTER.validate_python(validated.get("headers"))
|
||||
except ValidationError:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -60,13 +60,13 @@ class _HasProxyErrorType(Protocol):
|
|||
|
||||
|
||||
_MESSAGE_PATTERNS: Final[tuple[tuple[re.Pattern[str], str], ...]] = (
|
||||
(re.compile(r"upstream passthrough request failed", re.IGNORECASE), UPSTREAM_PASSTHROUGH),
|
||||
(
|
||||
re.compile(r"budget has been exceeded|max budget|exceeded.*budget|crossed budget", re.IGNORECASE),
|
||||
re.compile(r"budget has been exceeded|max budget|crossed budget", re.IGNORECASE),
|
||||
BUDGET_EXCEEDED,
|
||||
),
|
||||
(re.compile(r"no healthy deployments?|no deployments available", re.IGNORECASE), NO_HEALTHY_DEPLOYMENTS),
|
||||
(re.compile(r"not allowed to access model due to tags configuration", re.IGNORECASE), MODEL_ACCESS_DENIED),
|
||||
(re.compile(r"upstream passthrough request failed", re.IGNORECASE), UPSTREAM_PASSTHROUGH),
|
||||
(re.compile(r"is not supported for provider|not implemented", re.IGNORECASE), UNSUPPORTED_OPERATION),
|
||||
(
|
||||
re.compile(r"context window|context length|(prompt|input) is too long|tokens? ?> ?\d+ ?maximum", re.IGNORECASE),
|
||||
|
|
@ -155,6 +155,14 @@ _CLASS_CODE_TABLE: Final[tuple[tuple[tuple[type[BaseException], ...], str], ...]
|
|||
)
|
||||
|
||||
|
||||
def _exceeded_before_budget(message: str) -> bool:
|
||||
"""Linear-time equivalent of ``re.search(r"exceeded.*budget", message, re.IGNORECASE)``."""
|
||||
return any(
|
||||
(start := line.find("exceeded")) != -1 and line.find("budget", start + len("exceeded")) != -1
|
||||
for line in message.lower().split("\n")
|
||||
)
|
||||
|
||||
|
||||
def _classify_by_message(message: str, patterns: tuple[tuple[re.Pattern[str], str], ...]) -> str | None:
|
||||
return next((code for pattern, code in patterns if pattern.search(message)), None)
|
||||
|
||||
|
|
@ -183,7 +191,9 @@ def normalize_error(exc: Exception | None, status_code: str, message: str) -> st
|
|||
by_proxy_type: Final = _PROXY_ERROR_TYPE_MAP.get(proxy_type) if isinstance(proxy_type, str) else None
|
||||
if by_proxy_type is not None:
|
||||
return by_proxy_type
|
||||
by_message: Final = _classify_by_message(message, _MESSAGE_PATTERNS)
|
||||
by_message: Final = (
|
||||
BUDGET_EXCEEDED if _exceeded_before_budget(message) else _classify_by_message(message, _MESSAGE_PATTERNS)
|
||||
)
|
||||
if by_message is not None:
|
||||
return by_message
|
||||
by_class: Final = _classify_by_class(exc)
|
||||
|
|
|
|||
|
|
@ -72,6 +72,7 @@ from litellm.litellm_core_utils.classifier_logging import (
|
|||
is_classifier_call,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_provider_response_headers_from_hidden_params,
|
||||
is_expected_client_error,
|
||||
reconstruct_model_name,
|
||||
set_response_cost_in_hidden_params,
|
||||
|
|
@ -1567,14 +1568,14 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
attr = "debug"
|
||||
|
||||
if json_logs:
|
||||
callattr = getattr(verbose_logger, attr)
|
||||
callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug
|
||||
callattr(
|
||||
"RAW RESPONSE:\n{}\n\n".format(
|
||||
self.model_call_details.get("original_response", self.model_call_details)
|
||||
),
|
||||
)
|
||||
else:
|
||||
callattr = getattr(verbose_logger, attr)
|
||||
callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug
|
||||
callattr(
|
||||
"RAW RESPONSE:\n{}\n\n".format(
|
||||
self.model_call_details.get("original_response", self.model_call_details)
|
||||
|
|
@ -2354,6 +2355,15 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
return logging_result
|
||||
|
||||
def _surface_response_headers_from_result(self, logging_result: object) -> None:
|
||||
existing: Final[object] = self.model_call_details.get("response_headers")
|
||||
if existing is not None:
|
||||
return
|
||||
headers: Final = get_provider_response_headers_from_hidden_params(logging_result)
|
||||
if headers is None:
|
||||
return
|
||||
self.model_call_details["response_headers"] = headers
|
||||
|
||||
def _merge_hidden_params_from_response_into_metadata(self, logging_result: object) -> None:
|
||||
"""
|
||||
Copy response._hidden_params into litellm_params.metadata['hidden_params'].
|
||||
|
|
@ -2387,6 +2397,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
build_logging_payload: bool = True,
|
||||
):
|
||||
"""Resolve hidden params, compute response cost, and emit the standard logging payload."""
|
||||
self._surface_response_headers_from_result(logging_result)
|
||||
hidden_params: Final = getattr(logging_result, "_hidden_params", {})
|
||||
if hidden_params:
|
||||
if self.model_call_details.get("litellm_params") is not None:
|
||||
|
|
@ -2789,6 +2800,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if complete_streaming_response is not None:
|
||||
verbose_logger.debug("Logging Details LiteLLM-Success Call streaming complete")
|
||||
self.model_call_details["complete_streaming_response"] = complete_streaming_response
|
||||
self._surface_response_headers_from_result(complete_streaming_response)
|
||||
self.model_call_details["response_cost"] = self._response_cost_calculator(
|
||||
result=complete_streaming_response
|
||||
)
|
||||
|
|
@ -3315,6 +3327,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
print_verbose("Async success callbacks: Got a complete streaming response")
|
||||
|
||||
self.model_call_details["async_complete_streaming_response"] = complete_streaming_response
|
||||
self._surface_response_headers_from_result(complete_streaming_response)
|
||||
|
||||
try:
|
||||
if self.model_call_details.get("cache_hit", False) is True:
|
||||
|
|
@ -5195,7 +5208,7 @@ def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list[Custom
|
|||
for callback in _in_memory_loggers:
|
||||
if (
|
||||
isinstance(callback, OpenTelemetryV2)
|
||||
and getattr(callback, "callback_name", None) == callback_name
|
||||
and callback.callback_name == callback_name
|
||||
and (serves_a_destination or not _exports_nowhere(callback.config))
|
||||
):
|
||||
return callback
|
||||
|
|
@ -5886,7 +5899,7 @@ class StandardLoggingPayloadSetup:
|
|||
base_model: str | None,
|
||||
custom_pricing: bool | None,
|
||||
custom_llm_provider: str | None,
|
||||
init_response_obj: Any | BaseModel | dict,
|
||||
init_response_obj: object,
|
||||
api_base: str | None = None,
|
||||
) -> StandardLoggingModelInformation:
|
||||
model_cost_name: Final = _select_model_name_for_cost_calc(
|
||||
|
|
@ -5919,9 +5932,7 @@ class StandardLoggingPayloadSetup:
|
|||
return model_cost_information
|
||||
|
||||
@staticmethod
|
||||
def get_final_response_obj(
|
||||
response_obj: dict, init_response_obj: Any | BaseModel | dict, kwargs: dict
|
||||
) -> dict | str | list | None:
|
||||
def get_final_response_obj(response_obj: dict, init_response_obj: object, kwargs: dict) -> dict | str | list | None:
|
||||
"""
|
||||
Get final response object after redacting the message input/output from logging
|
||||
"""
|
||||
|
|
@ -6364,16 +6375,19 @@ def _get_status_fields(
|
|||
|
||||
|
||||
def _extract_response_obj_and_hidden_params(
|
||||
init_response_obj: Any | BaseModel | dict,
|
||||
init_response_obj: object,
|
||||
original_exception: Exception | None,
|
||||
) -> tuple[dict, dict | None]:
|
||||
"""Extract response_obj and hidden_params from init_response_obj."""
|
||||
hidden_params: dict | None = None
|
||||
hidden_params: dict | None = (
|
||||
getattr(init_response_obj, "_hidden_params", None)
|
||||
if isinstance(init_response_obj, BaseModel | HttpxBinaryResponseContent)
|
||||
else None
|
||||
)
|
||||
if init_response_obj is None:
|
||||
response_obj = {}
|
||||
elif isinstance(init_response_obj, BaseModel):
|
||||
response_obj = init_response_obj.model_dump()
|
||||
hidden_params = getattr(init_response_obj, "_hidden_params", None)
|
||||
elif isinstance(init_response_obj, dict):
|
||||
response_obj = init_response_obj
|
||||
elif isinstance(init_response_obj, HttpxBinaryResponseContent):
|
||||
|
|
@ -6669,7 +6683,7 @@ def get_standard_logging_object_payload(
|
|||
cost_breakdown=request_cost_breakdown,
|
||||
autorouter_savings=autorouter_savings,
|
||||
autorouter_savings_estimate=(
|
||||
{
|
||||
{ # mutable-ok: spend-log JSON serialization requires plain mappings
|
||||
"version": 3,
|
||||
"status": "unknown",
|
||||
"reason": "pending_projection",
|
||||
|
|
|
|||
|
|
@ -5,7 +5,12 @@ from typing import Final
|
|||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.types.utils import StandardLoggingZeroCostDiagnostic, Usage
|
||||
from litellm.types.utils import (
|
||||
CompletionTokensDetailsWrapper,
|
||||
PromptTokensDetailsWrapper,
|
||||
StandardLoggingZeroCostDiagnostic,
|
||||
Usage,
|
||||
)
|
||||
|
||||
ZERO_COST_COUNTER_NAME: Final = "litellm_zero_cost_requests_total"
|
||||
|
||||
|
|
@ -18,8 +23,8 @@ _NESTED_PRICING: Final = TypeAdapter(Mapping[str, object] | tuple[object, ...])
|
|||
_MAX_PRICING_DEPTH: Final = 4
|
||||
|
||||
|
||||
def _audio_tokens(details: object) -> int:
|
||||
audio_tokens: Final = getattr(details, "audio_tokens", None)
|
||||
def _audio_tokens(details: PromptTokensDetailsWrapper | CompletionTokensDetailsWrapper | None) -> int:
|
||||
audio_tokens: Final = details.audio_tokens if details is not None else None
|
||||
return audio_tokens if isinstance(audio_tokens, int) and audio_tokens > 0 else 0
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2003,11 +2003,11 @@ def strip_encrypted_reasoning_from_messages(messages: object) -> None:
|
|||
"""
|
||||
if not isinstance(messages, list):
|
||||
return
|
||||
for content in _anthropic_content_lists(cast(list[object], messages)): # cast-ok: untyped client json
|
||||
for content in anthropic_content_lists(cast(list[object], messages)): # cast-ok: untyped client json
|
||||
_strip_encrypted_reasoning_from_blocks(content)
|
||||
|
||||
|
||||
def _anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
|
||||
def anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
|
||||
return (
|
||||
cast(list[object], content) # cast-ok: narrowed by isinstance
|
||||
for message in messages
|
||||
|
|
|
|||
|
|
@ -446,7 +446,7 @@ def _render_chat_template(env, chat_template: str, bos_token: str, eos_token: st
|
|||
|
||||
|
||||
async def _afetch_and_extract_template(
|
||||
model: str, chat_template: Any | None, get_config_fn, get_template_fn
|
||||
model: str, chat_template: str | None, get_config_fn, get_template_fn
|
||||
) -> tuple[str, str, str]:
|
||||
"""
|
||||
Async version: Fetch template and tokens from HuggingFace.
|
||||
|
|
@ -500,7 +500,7 @@ async def _afetch_and_extract_template(
|
|||
|
||||
|
||||
def _fetch_and_extract_template(
|
||||
model: str, chat_template: Any | None, get_config_fn, get_template_fn
|
||||
model: str, chat_template: str | None, get_config_fn, get_template_fn
|
||||
) -> tuple[str, str, str]:
|
||||
"""
|
||||
Sync version: Fetch template and tokens from HuggingFace.
|
||||
|
|
|
|||
|
|
@ -1329,7 +1329,7 @@ class CustomStreamWrapper:
|
|||
"is_finished": chunk_finish_reason is not None,
|
||||
"finish_reason": chunk_finish_reason,
|
||||
"original_chunk": cached_chunk,
|
||||
"tool_calls": (getattr(cached_choice.delta, "tool_calls", None) if cached_choice is not None else None),
|
||||
"tool_calls": cached_choice.delta.tool_calls if cached_choice is not None else None,
|
||||
}
|
||||
|
||||
completion_obj["content"] = response_obj["text"]
|
||||
|
|
|
|||
|
|
@ -48,7 +48,7 @@ def _registry_api_key(agent_litellm_params: Mapping[str, object]) -> str | None:
|
|||
return configured_api_key if isinstance(configured_api_key, str) else None
|
||||
|
||||
|
||||
def _registry_headers(agent_litellm_params: Mapping[str, object]) -> dict[str, Any] | None:
|
||||
def _registry_headers(agent_litellm_params: Mapping[str, object]) -> dict[str, object] | None:
|
||||
stored_headers: Final = agent_litellm_params.get("headers")
|
||||
if not isinstance(stored_headers, Mapping):
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -685,7 +685,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
def _hoisted_top_level_system_message(self, data: dict) -> AllMessageValues | None:
|
||||
def _hoisted_top_level_system_message(self, data: Mapping[str, object]) -> AllMessageValues | None:
|
||||
"""Return the system message produced by translating the top-level prompt."""
|
||||
system: Final = data.get("system")
|
||||
if not system:
|
||||
|
|
|
|||
|
|
@ -11,7 +11,6 @@ from typing import (
|
|||
Final,
|
||||
Literal,
|
||||
Protocol,
|
||||
cast, # noqa: TID251 # rebuilt message_delta dict spans the ContentBlockDelta/MessageBlockDelta union
|
||||
get_args,
|
||||
)
|
||||
|
||||
|
|
@ -27,6 +26,7 @@ from litellm.types.llms.anthropic import (
|
|||
ContentBlockDelta,
|
||||
ContextManagementResponse,
|
||||
MessageBlockDelta,
|
||||
MessageDelta,
|
||||
StreamingContentBlockDeltaType,
|
||||
UsageDelta,
|
||||
UsageIteration,
|
||||
|
|
@ -1028,26 +1028,22 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
self,
|
||||
processed_chunk: ContentBlockDelta | MessageBlockDelta,
|
||||
) -> ContentBlockDelta | MessageBlockDelta:
|
||||
if processed_chunk.get("type") != "message_delta" or not self._refusal_text:
|
||||
if processed_chunk["type"] != "message_delta" or not self._refusal_text:
|
||||
return processed_chunk
|
||||
delta: Final = cast(Mapping[str, object], processed_chunk["delta"]) # cast-ok: keys checked before use
|
||||
delta: Final = processed_chunk["delta"]
|
||||
if delta.get("stop_reason") == "max_tokens":
|
||||
return processed_chunk
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
|
||||
refusal_stop_details,
|
||||
)
|
||||
|
||||
return cast( # cast-ok: rebuilt dict matches the message_delta TypedDict shape for this branch
|
||||
ContentBlockDelta | MessageBlockDelta,
|
||||
{ # mutable-ok: fresh translation payload; never mutated after construction
|
||||
**processed_chunk,
|
||||
"delta": { # mutable-ok: fresh message_delta payload; never mutated after construction
|
||||
**delta,
|
||||
"stop_reason": "refusal",
|
||||
"stop_details": refusal_stop_details(self._refusal_text),
|
||||
},
|
||||
},
|
||||
)
|
||||
refusal_delta: Final[MessageDelta] = {
|
||||
**delta,
|
||||
"stop_reason": "refusal",
|
||||
"stop_details": refusal_stop_details(self._refusal_text),
|
||||
}
|
||||
refusal_chunk: Final[MessageBlockDelta] = {**processed_chunk, "delta": refusal_delta}
|
||||
return refusal_chunk
|
||||
|
||||
@staticmethod
|
||||
def _delta_has_content(processed_chunk: Mapping[str, object]) -> bool:
|
||||
|
|
|
|||
|
|
@ -37,7 +37,7 @@ def _mapping_field(container: object, key: str) -> object | None:
|
|||
"""One key of a raw provider payload, or None when the payload is not a mapping."""
|
||||
if not isinstance(container, Mapping):
|
||||
return None
|
||||
return cast(Mapping[str, object], container).get(key) # cast-ok: raw payload, callers re-check every value
|
||||
return container.get(key)
|
||||
|
||||
|
||||
def _mapping_str_field(container: object, key: str) -> str | None:
|
||||
|
|
|
|||
|
|
@ -169,7 +169,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
cls,
|
||||
summary: Iterable[object],
|
||||
encrypted_content: object,
|
||||
) -> dict[str, Any] | None: # mutable-ok: API message payload
|
||||
) -> dict[str, object] | None: # mutable-ok: API message payload
|
||||
"""The one Anthropic block for a Responses reasoning item.
|
||||
|
||||
The item's encrypted reasoning rides the block's opaque field (`signature`, or
|
||||
|
|
@ -198,7 +198,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
@classmethod
|
||||
def _assistant_group_to_input_items(
|
||||
cls, group: tuple[Mapping[str, object], ...]
|
||||
) -> tuple[dict[str, Any], ...]: # mutable-ok: API message payload
|
||||
) -> tuple[dict[str, object], ...]: # mutable-ok: API message payload
|
||||
first: Final = group[0]
|
||||
btype: Final = first.get("type")
|
||||
if btype in ("thinking", "redacted_thinking"):
|
||||
|
|
|
|||
|
|
@ -1218,6 +1218,8 @@ class BedrockModelInfo(BaseLLMModelInfo):
|
|||
alt_model: Final = BedrockModelInfo.get_non_litellm_routing_model_name(model=model)
|
||||
if base_model in litellm.bedrock_converse_models or alt_model in litellm.bedrock_converse_models:
|
||||
return "converse"
|
||||
if _OPENAI_FAMILY_MODEL_RE.search(base_model):
|
||||
return "converse"
|
||||
return "invoke"
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from typing import (
|
|||
Literal,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
Protocol,
|
||||
TypedDict,
|
||||
TypeVar,
|
||||
Union,
|
||||
|
|
@ -24,6 +25,7 @@ import httpx
|
|||
from httpx import USE_CLIENT_DEFAULT
|
||||
from httpx._types import FileContent
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
import litellm.litellm_core_utils
|
||||
|
|
@ -43,6 +45,7 @@ from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
|||
SUBTITLE_RESPONSE_FORMATS,
|
||||
synthesize_subtitle_document,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
|
||||
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
|
||||
from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields
|
||||
from litellm.litellm_core_utils.realtime_errors import (
|
||||
|
|
@ -206,6 +209,7 @@ if TYPE_CHECKING:
|
|||
FakeAnthropicMessagesStreamIterator,
|
||||
)
|
||||
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.llms.openai_evals import (
|
||||
CancelEvalResponse,
|
||||
CancelRunResponse,
|
||||
|
|
@ -221,6 +225,21 @@ if TYPE_CHECKING:
|
|||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class _RealtimeClientWebSocket(Protocol):
|
||||
async def send_text(self, data: str) -> None: ...
|
||||
|
||||
async def close(self, code: int = ..., reason: str | None = ...) -> None: ...
|
||||
|
||||
|
||||
class _ResponsesClientWebSocket(Protocol):
|
||||
async def send_text(self, data: str) -> None: ...
|
||||
|
||||
async def receive_text(self) -> str: ...
|
||||
|
||||
async def close(self, code: int = ..., reason: str | None = ...) -> None: ...
|
||||
|
||||
|
||||
_ResponseT = TypeVar("_ResponseT")
|
||||
|
||||
|
||||
|
|
@ -237,6 +256,17 @@ class _MediaUploadKwargs(TypedDict, total=False):
|
|||
timeout: float | httpx.Timeout
|
||||
|
||||
|
||||
class _SignedBodyKwargs(TypedDict, total=False):
|
||||
data: ReadOnly[bytes]
|
||||
json: ReadOnly[dict[str, object]]
|
||||
|
||||
|
||||
def _signed_body_kwargs(*, signed_body: bytes | None, data: dict[str, object]) -> _SignedBodyKwargs:
|
||||
if signed_body is not None:
|
||||
return {"data": signed_body}
|
||||
return {"json": data}
|
||||
|
||||
|
||||
def _google_genai_streaming_hidden_params(
|
||||
*,
|
||||
api_base: str,
|
||||
|
|
@ -318,7 +348,9 @@ def _mask_presigned_request_headers(transformed_request: bytes | str | dict) ->
|
|||
}
|
||||
|
||||
|
||||
def _aws_signing_overrides(optional_params: Mapping[str, Any], litellm_params: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
def _aws_signing_overrides(
|
||||
optional_params: Mapping[str, object], litellm_params: Mapping[str, object]
|
||||
) -> Mapping[str, object]:
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: litellm_params[key]
|
||||
|
|
@ -1430,6 +1462,7 @@ class BaseLLMHTTPHandler:
|
|||
transformed: Final = provider_config.transform_audio_transcription_response(
|
||||
raw_response=response,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(transformed, response.headers)
|
||||
if not provider_config.supports_subtitle_synthesis:
|
||||
return transformed
|
||||
requested_format: Final = optional_params.get("response_format")
|
||||
|
|
@ -2739,7 +2772,7 @@ class BaseLLMHTTPHandler:
|
|||
stream=stream,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
|
||||
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -2926,7 +2959,7 @@ class BaseLLMHTTPHandler:
|
|||
stream=stream,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
|
||||
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -4540,7 +4573,7 @@ class BaseLLMHTTPHandler:
|
|||
api_key=litellm_params.api_key,
|
||||
model=model,
|
||||
)
|
||||
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
|
||||
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -4634,7 +4667,7 @@ class BaseLLMHTTPHandler:
|
|||
api_key=litellm_params.api_key,
|
||||
model=model,
|
||||
)
|
||||
body_kwargs: Final[dict[str, Any]] = {"data": signed_body} if signed_body is not None else {"json": data}
|
||||
body_kwargs: Final = _signed_body_kwargs(signed_body=signed_body, data=data)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -6186,6 +6219,7 @@ class BaseLLMHTTPHandler:
|
|||
"BasePassthroughConfig",
|
||||
"BaseContainerConfig",
|
||||
BaseEvalsAPIConfig,
|
||||
BaseRealtimeHTTPConfig,
|
||||
],
|
||||
):
|
||||
received_status_code: Final = (
|
||||
|
|
@ -6300,7 +6334,7 @@ class BaseLLMHTTPHandler:
|
|||
async def async_realtime(
|
||||
self,
|
||||
model: str,
|
||||
websocket: Any,
|
||||
websocket: _RealtimeClientWebSocket,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
provider_config: BaseRealtimeConfig,
|
||||
headers: dict,
|
||||
|
|
@ -6308,7 +6342,7 @@ class BaseLLMHTTPHandler:
|
|||
api_key: str | None = None,
|
||||
client: Any | None = None,
|
||||
timeout: float | None = None,
|
||||
user_api_key_dict: Any | None = None,
|
||||
user_api_key_dict: object | None = None,
|
||||
litellm_metadata: dict[str, object] | None = None,
|
||||
query_params: RealtimeQueryParams | None = None,
|
||||
):
|
||||
|
|
@ -6483,7 +6517,7 @@ class BaseLLMHTTPHandler:
|
|||
request_data: dict[str, object],
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
timeout: float | httpx.Timeout,
|
||||
provider_config: Any | None = None,
|
||||
provider_config: BaseRealtimeHTTPConfig | None = None,
|
||||
model: str | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
client: HTTPHandler | AsyncHTTPHandler | None = None,
|
||||
|
|
@ -6555,7 +6589,7 @@ class BaseLLMHTTPHandler:
|
|||
sdp_body: bytes,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
timeout: float | httpx.Timeout,
|
||||
provider_config: Any | None = None,
|
||||
provider_config: BaseRealtimeHTTPConfig | None = None,
|
||||
model: str | None = None,
|
||||
session_config: dict[str, object] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
|
|
@ -6633,13 +6667,13 @@ class BaseLLMHTTPHandler:
|
|||
async def async_responses_websocket(
|
||||
self,
|
||||
model: str,
|
||||
websocket: Any,
|
||||
websocket: _ResponsesClientWebSocket,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
responses_api_provider_config: BaseResponsesAPIConfig | None,
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
timeout: float | None = None,
|
||||
user_api_key_dict: Any | None = None,
|
||||
user_api_key_dict: "UserAPIKeyAuth | None" = None,
|
||||
litellm_metadata: dict[str, object] | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
first_message: str | None = None,
|
||||
|
|
@ -6928,11 +6962,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=image_edit_provider_config,
|
||||
)
|
||||
|
||||
return image_edit_provider_config.transform_image_edit_response(
|
||||
image_edit_response: Final = image_edit_provider_config.transform_image_edit_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
|
||||
return image_edit_response
|
||||
|
||||
async def async_image_edit_handler(
|
||||
self,
|
||||
|
|
@ -7027,11 +7063,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=image_edit_provider_config,
|
||||
)
|
||||
|
||||
return image_edit_provider_config.transform_image_edit_response(
|
||||
image_edit_response: Final = image_edit_provider_config.transform_image_edit_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
|
||||
return image_edit_response
|
||||
|
||||
def image_generation_handler(
|
||||
self,
|
||||
|
|
@ -7154,6 +7192,7 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
encoding=None,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(model_response, response.headers)
|
||||
|
||||
return model_response
|
||||
|
||||
|
|
@ -7261,6 +7300,7 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
encoding=None,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(model_response, response.headers)
|
||||
|
||||
return model_response
|
||||
|
||||
|
|
@ -7850,7 +7890,7 @@ class BaseLLMHTTPHandler:
|
|||
def video_create_character_handler(
|
||||
self,
|
||||
name: str,
|
||||
video: Any,
|
||||
video: FileTypes,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
|
|
@ -7934,7 +7974,7 @@ class BaseLLMHTTPHandler:
|
|||
async def async_video_create_character_handler(
|
||||
self,
|
||||
name: str,
|
||||
video: Any,
|
||||
video: FileTypes,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
|
|
@ -12045,11 +12085,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=text_to_speech_provider_config,
|
||||
)
|
||||
|
||||
return text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.headers)
|
||||
return speech_response
|
||||
|
||||
async def async_text_to_speech_handler(
|
||||
self,
|
||||
|
|
@ -12144,11 +12186,13 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=text_to_speech_provider_config,
|
||||
)
|
||||
|
||||
return text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.headers)
|
||||
return speech_response
|
||||
|
||||
#########################################################
|
||||
########## SKILLS API HANDLERS ##########################
|
||||
|
|
|
|||
|
|
@ -59,6 +59,7 @@ def resolve_fireworks_api_key(api_key: str | None) -> str | None:
|
|||
|
||||
|
||||
AZURE_FOUNDRY_FIREWORKS_MODEL_ID_PREFIX: Final = "FW-"
|
||||
FIREROUTER: Final = "firerouter"
|
||||
|
||||
|
||||
def resolve_fireworks_resource_name(model: str) -> str:
|
||||
|
|
@ -67,7 +68,7 @@ def resolve_fireworks_resource_name(model: str) -> str:
|
|||
return stripped
|
||||
if stripped.startswith(("routers/", "models/")):
|
||||
return f"accounts/fireworks/{stripped}"
|
||||
if stripped.endswith("-fast"):
|
||||
if stripped.endswith("-fast") or stripped == FIREROUTER or stripped.startswith(f"{FIREROUTER}/"):
|
||||
return f"accounts/fireworks/routers/{stripped}"
|
||||
return f"accounts/fireworks/models/{stripped}"
|
||||
|
||||
|
|
|
|||
|
|
@ -59,6 +59,13 @@ def get_base_model_for_pricing(model_name: str) -> str:
|
|||
def _resolve_model_info(model: str) -> ModelInfo:
|
||||
try:
|
||||
return get_model_info(model=model, custom_llm_provider="fireworks_ai")
|
||||
except Exception:
|
||||
return _resolve_routed_model_info(model)
|
||||
|
||||
|
||||
def _resolve_routed_model_info(model: str) -> ModelInfo:
|
||||
try:
|
||||
return get_model_info(model=model.removeprefix("fireworks_ai/"))
|
||||
except Exception:
|
||||
base_model: Final = get_base_model_for_pricing(model_name=model)
|
||||
return get_model_info(model=base_model, custom_llm_provider="fireworks_ai")
|
||||
|
|
@ -81,7 +88,7 @@ def cost_per_token(model: str, usage: Usage, current_time: datetime | None = Non
|
|||
return generic_cost_per_token(
|
||||
model=model,
|
||||
usage=usage,
|
||||
custom_llm_provider="fireworks_ai",
|
||||
custom_llm_provider=model_info["litellm_provider"],
|
||||
model_info=model_info,
|
||||
current_time=current_time,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from collections.abc import AsyncIterator, Iterator
|
|||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from httpx._models import Headers, Response
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -43,6 +44,37 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class _OllamaGenerateReasoning(BaseModel):
|
||||
"""The two `/api/generate` fields a reply's reasoning can arrive in."""
|
||||
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
# Absent and explicitly null are distinct here: Ollama omits `response` where it sends
|
||||
# no text, and sends null where the reply carries none, which stay "" and None downstream.
|
||||
response: str | None = ""
|
||||
thinking: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_response(cls, response_json: object) -> "_OllamaGenerateReasoning":
|
||||
try:
|
||||
return cls.model_validate(response_json)
|
||||
except ValidationError:
|
||||
return cls()
|
||||
|
||||
def split(self) -> tuple[str | None, str | None]:
|
||||
"""Reasoning reaches `/api/generate` either in the top-level `thinking` field or
|
||||
inline in `<think>` tags, never both. The field wins, matching `ollama_chat`."""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
_parse_content_for_reasoning,
|
||||
)
|
||||
|
||||
if self.thinking:
|
||||
return self.thinking, self.response
|
||||
if self.response is None:
|
||||
return None, None
|
||||
return _parse_content_for_reasoning(self.response)
|
||||
|
||||
|
||||
class OllamaConfig(BaseConfig):
|
||||
"""
|
||||
Reference: https://github.com/ollama/ollama/blob/main/docs/api.md#parameters
|
||||
|
|
@ -255,20 +287,17 @@ class OllamaConfig(BaseConfig):
|
|||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
_parse_content_for_reasoning,
|
||||
)
|
||||
|
||||
response_json: Final = raw_response.json()
|
||||
## RESPONSE OBJECT
|
||||
model_response.choices[0].finish_reason = "stop"
|
||||
if request_data.get("format", "") == "json":
|
||||
# Check if response field exists and is not empty before parsing JSON
|
||||
response_text = response_json.get("response", "")
|
||||
thinking: Final = _OllamaGenerateReasoning.from_response(response_json).thinking or None
|
||||
|
||||
if not response_text or not response_text.strip():
|
||||
# Handle empty response gracefully - set empty content
|
||||
message = litellm.Message(content="")
|
||||
message = litellm.Message(content="", reasoning_content=thinking)
|
||||
model_response.choices[0].message = message
|
||||
model_response.choices[0].finish_reason = "stop"
|
||||
else:
|
||||
|
|
@ -285,6 +314,7 @@ class OllamaConfig(BaseConfig):
|
|||
function_call: Final = response_content
|
||||
message = litellm.Message(
|
||||
content=None,
|
||||
reasoning_content=thinking,
|
||||
tool_calls=[
|
||||
{
|
||||
"id": f"call_{uuid.uuid4()}",
|
||||
|
|
@ -302,27 +332,18 @@ class OllamaConfig(BaseConfig):
|
|||
# Handle as regular JSON (new behavior)
|
||||
message = litellm.Message(
|
||||
content=json.dumps(response_content),
|
||||
reasoning_content=thinking,
|
||||
)
|
||||
model_response.choices[0].message = message
|
||||
model_response.choices[0].finish_reason = "stop"
|
||||
except json.JSONDecodeError:
|
||||
# If JSON parsing fails, treat as regular text response
|
||||
## output parse reasoning content from response_text
|
||||
reasoning_content: str | None = None
|
||||
content: str | None = None
|
||||
if response_text is not None:
|
||||
reasoning_content, content = _parse_content_for_reasoning(response_text)
|
||||
reasoning_content, content = _OllamaGenerateReasoning.from_response(response_json).split()
|
||||
message = litellm.Message(content=content, reasoning_content=reasoning_content)
|
||||
model_response.choices[0].message = message
|
||||
model_response.choices[0].finish_reason = "stop"
|
||||
else:
|
||||
response_text = response_json.get("response", "")
|
||||
content = None
|
||||
reasoning_content = None
|
||||
if response_text is not None and isinstance(response_text, str):
|
||||
reasoning_content, content = _parse_content_for_reasoning(response_text)
|
||||
else:
|
||||
content = response_text
|
||||
reasoning_content, content = _OllamaGenerateReasoning.from_response(response_json).split()
|
||||
model_response.choices[0].message.content = content
|
||||
model_response.choices[0].message.reasoning_content = reasoning_content
|
||||
model_response.created = int(time.time())
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from litellm import LlmProviders
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import DEFAULT_MAX_RETRIES
|
||||
from litellm.files.types import FileContentStreamingResult
|
||||
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.logging_utils import speech_request_body, track_llm_api_timing
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
|
|
@ -1404,7 +1405,6 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
organization: str | None = None,
|
||||
headers: dict | None = None,
|
||||
):
|
||||
response = None
|
||||
try:
|
||||
openai_aclient: Final = self._get_openai_client(
|
||||
is_async=True,
|
||||
|
|
@ -1428,8 +1428,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
)
|
||||
|
||||
request_data: Final = {**data, "extra_headers": headers} if headers else data
|
||||
response = await openai_aclient.images.generate(**request_data, timeout=timeout)
|
||||
stringified_response: Final = response.model_dump()
|
||||
raw_response: Final = await openai_aclient.images.with_raw_response.generate(
|
||||
**request_data, timeout=timeout
|
||||
)
|
||||
stringified_response: Final = raw_response.parse().model_dump()
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
input=prompt,
|
||||
|
|
@ -1437,11 +1439,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
additional_args={"complete_input_dict": data},
|
||||
original_response=stringified_response,
|
||||
)
|
||||
return convert_to_model_response_object(
|
||||
image_response: Final[ImageResponse] = convert_to_model_response_object(
|
||||
response_object=stringified_response,
|
||||
model_response_object=model_response,
|
||||
response_type="image_generation",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_response, raw_response.headers)
|
||||
return image_response
|
||||
except Exception as e:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -1512,9 +1516,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
|
||||
## COMPLETION CALL
|
||||
request_data: Final = {**data, "extra_headers": headers} if headers else data
|
||||
_response: Final = openai_client.images.generate(**request_data, timeout=timeout)
|
||||
raw_response: Final = openai_client.images.with_raw_response.generate(**request_data, timeout=timeout)
|
||||
|
||||
response: Final = _response.model_dump()
|
||||
response: Final = raw_response.parse().model_dump()
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
input=prompt,
|
||||
|
|
@ -1522,11 +1526,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
additional_args={"complete_input_dict": data},
|
||||
original_response=response,
|
||||
)
|
||||
return convert_to_model_response_object(
|
||||
image_response: Final[ImageResponse] = convert_to_model_response_object(
|
||||
response_object=response,
|
||||
model_response_object=model_response,
|
||||
response_type="image_generation",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(image_response, raw_response.headers)
|
||||
return image_response
|
||||
except OpenAIError as e:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -1609,7 +1615,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
input=input,
|
||||
**optional_params,
|
||||
)
|
||||
return HttpxBinaryResponseContent(response=response.response)
|
||||
speech_response: Final = HttpxBinaryResponseContent(response=response.response)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.response.headers)
|
||||
return speech_response
|
||||
|
||||
async def async_audio_speech(
|
||||
self,
|
||||
|
|
@ -1655,8 +1663,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
input=input,
|
||||
**optional_params,
|
||||
)
|
||||
|
||||
return HttpxBinaryResponseContent(response=response.response)
|
||||
speech_response: Final = HttpxBinaryResponseContent(response=response.response)
|
||||
set_provider_response_headers_in_hidden_params(speech_response, response.response.headers)
|
||||
return speech_response
|
||||
|
||||
|
||||
class OpenAIFilesAPI(BaseLLM):
|
||||
|
|
|
|||
|
|
@ -994,7 +994,7 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
|
||||
def _spread_text_rewrite_over_stream_events(
|
||||
self,
|
||||
stream_events: Sequence[Any],
|
||||
stream_events: Sequence[object],
|
||||
rewritten_text: str,
|
||||
guardrail_name: str,
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -4,11 +4,10 @@ import httpx
|
|||
from openai import AsyncOpenAI, OpenAI
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from aiohttp import ClientSession
|
||||
from litellm.litellm_core_utils.audio_utils.utils import get_audio_file_name
|
||||
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.audio_transcription.transformation import (
|
||||
BaseAudioTranscriptionConfig,
|
||||
|
|
@ -31,11 +30,6 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
data: dict,
|
||||
timeout: float | httpx.Timeout,
|
||||
):
|
||||
"""
|
||||
Helper to:
|
||||
- call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True
|
||||
- call openai_aclient.audio.transcriptions.create by default
|
||||
"""
|
||||
try:
|
||||
raw_response = await openai_aclient.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
|
||||
headers: Final = dict(raw_response.headers)
|
||||
|
|
@ -51,20 +45,11 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
data: dict,
|
||||
timeout: float | httpx.Timeout,
|
||||
):
|
||||
"""
|
||||
Helper to:
|
||||
- call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True
|
||||
- call openai_aclient.audio.transcriptions.create by default
|
||||
"""
|
||||
try:
|
||||
if litellm.return_response_headers is True:
|
||||
raw_response = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
|
||||
headers: Final = dict(raw_response.headers)
|
||||
response = raw_response.parse()
|
||||
return headers, response
|
||||
else:
|
||||
response = openai_client.audio.transcriptions.create(**data, timeout=timeout)
|
||||
return None, response
|
||||
raw_response: Final = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
|
||||
headers: Final = dict(raw_response.headers)
|
||||
response: Final = raw_response.parse()
|
||||
return headers, response
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
|
@ -133,11 +118,12 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
"complete_input_dict": data,
|
||||
},
|
||||
)
|
||||
_, response = self.make_sync_openai_audio_transcriptions_request(
|
||||
headers, response = self.make_sync_openai_audio_transcriptions_request(
|
||||
openai_client=openai_client,
|
||||
data=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
logging_obj.model_call_details["response_headers"] = headers
|
||||
|
||||
if isinstance(response, BaseModel):
|
||||
stringified_response = response.model_dump()
|
||||
|
|
@ -158,6 +144,7 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
hidden_params=hidden_params,
|
||||
response_type="audio_transcription",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(final_response, headers)
|
||||
return final_response
|
||||
|
||||
async def async_audio_transcriptions(
|
||||
|
|
@ -217,12 +204,14 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
|
|||
actual_model: Final = data.get("model", "whisper-1")
|
||||
hidden_params: Final = {"model": actual_model, "custom_llm_provider": "openai"}
|
||||
|
||||
return convert_to_model_response_object(
|
||||
final_response: Final[TranscriptionResponse] = convert_to_model_response_object(
|
||||
response_object=stringified_response,
|
||||
model_response_object=model_response,
|
||||
hidden_params=hidden_params,
|
||||
response_type="audio_transcription",
|
||||
)
|
||||
set_provider_response_headers_in_hidden_params(final_response, headers)
|
||||
return final_response
|
||||
except Exception as e:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ Vercel AI Gateway is OpenAI-compatible and supports embeddings via the /v1/embed
|
|||
Docs: https://vercel.com/docs/ai-gateway/openai-compat/embeddings
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
|
|
@ -161,12 +162,14 @@ class VercelAIGatewayEmbeddingConfig(BaseEmbeddingConfig):
|
|||
optional_params[param] = value
|
||||
return optional_params
|
||||
|
||||
def get_error_class(self, error_message: str, status_code: int, headers: Any) -> BaseLLMException:
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers
|
||||
) -> BaseLLMException:
|
||||
"""
|
||||
Get the error class for Vercel AI Gateway errors.
|
||||
"""
|
||||
return VercelAIGatewayException(
|
||||
message=error_message,
|
||||
status_code=status_code,
|
||||
headers=headers,
|
||||
headers=headers if isinstance(headers, httpx.Headers) else httpx.Headers(headers),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,7 +1,11 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Any, Final
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Final, Protocol
|
||||
from urllib.parse import unquote
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.llms.vertex_ai.common_utils import (
|
||||
VertexAIError,
|
||||
|
|
@ -9,35 +13,128 @@ from litellm.llms.vertex_ai.common_utils import (
|
|||
)
|
||||
from litellm.types.llms.openai import BatchJobStatus, CreateBatchRequest
|
||||
from litellm.types.llms.vertex_ai import *
|
||||
from litellm.types.utils import LiteLLMBatch, PromptTokensDetailsWrapper
|
||||
from litellm.types.llms.vertex_ai import GenerateContentResponseBody
|
||||
from litellm.types.utils import LiteLLMBatch, ModelInfo, Usage
|
||||
|
||||
_NATIVE_VERTEX_RESPONSE: Final = TypeAdapter(GenerateContentResponseBody)
|
||||
|
||||
|
||||
def vertex_prompt_tokens_details(
|
||||
usage_metadata: Mapping[str, object],
|
||||
) -> PromptTokensDetailsWrapper | None:
|
||||
raw_details: Final = usage_metadata.get("promptTokensDetails")
|
||||
if not isinstance(raw_details, list):
|
||||
return None
|
||||
def _int_field(mapping: Mapping[str, object], key: str) -> int:
|
||||
value: Final = mapping.get(key)
|
||||
if isinstance(value, int):
|
||||
return value
|
||||
return int(value) if isinstance(value, str) and value.isdigit() else 0
|
||||
|
||||
def _normalize(detail: object) -> tuple[str, int] | None:
|
||||
if not isinstance(detail, Mapping):
|
||||
|
||||
def vertex_embedding_prompt_token_count(vertex_response: Mapping[str, object]) -> int:
|
||||
"""
|
||||
Prompt tokens billed for one Vertex Gemini Embedding batch row.
|
||||
|
||||
Live rows report usage under `usageMetadata`; the documented `tokenCount` is kept as
|
||||
a fallback.
|
||||
"""
|
||||
usage_metadata: Final = vertex_response.get("usageMetadata")
|
||||
if isinstance(usage_metadata, Mapping):
|
||||
return _int_field(usage_metadata, "promptTokenCount")
|
||||
return _int_field(vertex_response, "tokenCount")
|
||||
|
||||
|
||||
def is_vertex_embedding_batch_output_response(response_body: Mapping[str, object]) -> bool:
|
||||
return isinstance(response_body.get("embedding"), dict)
|
||||
|
||||
|
||||
def is_native_vertex_batch_output_row(row: Mapping[str, object]) -> bool:
|
||||
return isinstance(row.get("request"), dict)
|
||||
|
||||
|
||||
class NativeVertexBatchCostCalculator(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
usage: Usage,
|
||||
model: str,
|
||||
custom_llm_provider: str | None = None,
|
||||
model_info: ModelInfo | None = None,
|
||||
) -> tuple[float, float]: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class NativeVertexBatchRowStats:
|
||||
usage: Usage
|
||||
total_tokens: int
|
||||
model: str | None
|
||||
prompt_cost: float
|
||||
completion_cost: float
|
||||
|
||||
|
||||
def _native_vertex_row_usage(
|
||||
response_body: Mapping[str, object],
|
||||
calculate_usage: Callable[[GenerateContentResponseBody], Usage],
|
||||
) -> Usage | None:
|
||||
if "usageMetadata" not in response_body:
|
||||
if not is_vertex_embedding_batch_output_response(response_body):
|
||||
return None
|
||||
modality: Final = detail.get("modality")
|
||||
token_count: Final = detail.get("tokenCount")
|
||||
if not isinstance(modality, str) or not isinstance(token_count, int):
|
||||
return None
|
||||
return modality.upper(), token_count
|
||||
|
||||
parsed_details: Final = tuple(_normalize(detail) for detail in raw_details)
|
||||
normalized: Final = tuple(detail for detail in parsed_details if detail is not None)
|
||||
if len(normalized) != len(parsed_details):
|
||||
prompt_tokens: Final = vertex_embedding_prompt_token_count(response_body)
|
||||
return Usage(prompt_tokens=prompt_tokens, completion_tokens=0, total_tokens=prompt_tokens)
|
||||
try:
|
||||
completion_response: Final = _NATIVE_VERTEX_RESPONSE.validate_python(response_body)
|
||||
except ValidationError as e:
|
||||
verbose_logger.debug("vertex_ai batch row response is not a GenerateContentResponse: %s", str(e))
|
||||
return None
|
||||
return calculate_usage(completion_response)
|
||||
|
||||
return PromptTokensDetailsWrapper(
|
||||
text_tokens=sum(token_count for modality, token_count in normalized if modality in ("TEXT", "DOCUMENT")),
|
||||
audio_tokens=sum(token_count for modality, token_count in normalized if modality == "AUDIO"),
|
||||
image_tokens=sum(token_count for modality, token_count in normalized if modality == "IMAGE"),
|
||||
video_tokens=sum(token_count for modality, token_count in normalized if modality == "VIDEO"),
|
||||
|
||||
def native_vertex_batch_row_stats(
|
||||
row: Mapping[str, object],
|
||||
model_name: str | None,
|
||||
*,
|
||||
model_info: ModelInfo | None,
|
||||
calculate_usage: Callable[[GenerateContentResponseBody], Usage],
|
||||
cost_calculator: NativeVertexBatchCostCalculator,
|
||||
) -> NativeVertexBatchRowStats | None:
|
||||
"""
|
||||
Usage and cost of one native Vertex predictions.jsonl row, a
|
||||
`{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}, "modelVersion": ...}}`
|
||||
generateContent object or a `{"request": ..., "response": {"embedding": {...}, "usageMetadata": {...}}}`
|
||||
embedding object (an embedding row without `usageMetadata` is billed from its documented `tokenCount`).
|
||||
`model_name` (the deployment model) prices the row unless it is a wildcard, else its own `modelVersion`
|
||||
does, else the wildcard name so explicit deployment prices still apply; a row without a response, a
|
||||
generateContent row without `response.usageMetadata`, and a row whose response fails validation are
|
||||
None (failed).
|
||||
"""
|
||||
response_body: Final = row.get("response")
|
||||
if not isinstance(response_body, dict):
|
||||
return None
|
||||
usage: Final = _native_vertex_row_usage(response_body, calculate_usage)
|
||||
if usage is None:
|
||||
return None
|
||||
total_tokens: Final = usage.total_tokens or (usage.prompt_tokens + usage.completion_tokens)
|
||||
model_version: Final = response_body.get("modelVersion")
|
||||
deployment_model: Final = model_name if model_name and "*" not in model_name else None
|
||||
model: Final = deployment_model or (model_version if isinstance(model_version, str) else model_name)
|
||||
if model is None:
|
||||
verbose_logger.warning(
|
||||
"vertex_ai batch output row could not be costed, so it is billed at $0 and the rest of the batch "
|
||||
"is still billed: the row has no modelVersion and the batch has no deployment model"
|
||||
)
|
||||
return NativeVertexBatchRowStats(
|
||||
usage=usage, total_tokens=total_tokens, model=None, prompt_cost=0.0, completion_cost=0.0
|
||||
)
|
||||
try:
|
||||
prompt_cost, completion_cost = cost_calculator(
|
||||
usage=usage, model=model, custom_llm_provider="vertex_ai", model_info=model_info
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # one unpriceable row must not abort the batch's cost accounting
|
||||
verbose_logger.warning(
|
||||
"vertex_ai batch output row could not be costed, so it is billed at $0 and the rest of the batch "
|
||||
"is still billed. model=%s error=%s",
|
||||
model,
|
||||
str(e),
|
||||
)
|
||||
return NativeVertexBatchRowStats(
|
||||
usage=usage, total_tokens=total_tokens, model=model, prompt_cost=0.0, completion_cost=0.0
|
||||
)
|
||||
return NativeVertexBatchRowStats(
|
||||
usage=usage, total_tokens=total_tokens, model=model, prompt_cost=prompt_cost, completion_cost=completion_cost
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from collections.abc import AsyncGenerator, Callable, Iterable, Iterator, Mappin
|
|||
from contextlib import aclosing
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, TypedDict
|
||||
from typing import IO, Any, Final, TypedDict
|
||||
from urllib.parse import quote, unquote
|
||||
|
||||
import httpx
|
||||
|
|
@ -41,6 +41,7 @@ from litellm.llms.base_llm.files.transformation import (
|
|||
BaseFileUploadStream,
|
||||
LiteLLMLoggingObj,
|
||||
)
|
||||
from litellm.llms.vertex_ai.batches.transformation import vertex_embedding_prompt_token_count
|
||||
from litellm.llms.vertex_ai.common_utils import (
|
||||
_convert_vertex_datetime_to_openai_datetime,
|
||||
get_vertex_ai_fine_tuned_endpoint_id,
|
||||
|
|
@ -56,6 +57,7 @@ from litellm.types.files import StreamingMediaUploadConfig
|
|||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
CreateFileRequest,
|
||||
FileContent,
|
||||
FileTypes,
|
||||
HttpxBinaryResponseContent,
|
||||
OpenAICreateFileRequestOptionalParams,
|
||||
|
|
@ -87,6 +89,8 @@ _EMBED_REQUEST_FIELD_BY_GEMINI_PARAM: Final = (
|
|||
_VERTEX_BATCH_FANNED_OUT_KEY_PATTERN: Final = re.compile(r"(?P<custom_id>[^#]*)#(?P<index>\d+)/(?P<total>\d+)")
|
||||
_JSONL_NEWLINE: Final = b"\n"
|
||||
_BATCH_OUTPUT_FIRST_ROW_PEEK_LIMIT_BYTES: Final = 32 * 1024 * 1024
|
||||
_PASSTHROUGH_MANAGED_GCS_PREFIX: Final = f"{VERTEX_AI_MANAGED_GCS_PREFIX}passthrough/"
|
||||
_RAW_UPLOAD_CHUNK_BYTES: Final = 1024 * 1024
|
||||
|
||||
|
||||
class _GcsObjectMetadataJson(TypedDict, total=False):
|
||||
|
|
@ -418,19 +422,6 @@ def _split_vertex_batch_key(vertex_output_row: Mapping[str, object]) -> tuple[st
|
|||
return unquote(match["custom_id"]), int(match["index"]), int(match["total"])
|
||||
|
||||
|
||||
def _embedding_prompt_token_count(vertex_response: _VertexEmbeddingResponse) -> int:
|
||||
"""
|
||||
Prompt tokens billed for one Vertex Gemini Embedding batch row.
|
||||
|
||||
Live rows report usage under `usageMetadata`; the documented `tokenCount` is kept as
|
||||
a fallback.
|
||||
"""
|
||||
usage_metadata = vertex_response.get("usageMetadata")
|
||||
if isinstance(usage_metadata, Mapping):
|
||||
return int(usage_metadata.get("promptTokenCount") or 0)
|
||||
return int(vertex_response.get("tokenCount") or 0)
|
||||
|
||||
|
||||
def _vertex_embeddings_rows_to_openai_batch_output_row(
|
||||
custom_id: str,
|
||||
vertex_output_rows: tuple[_VertexEmbeddingBatchRow, ...],
|
||||
|
|
@ -471,7 +462,7 @@ def _vertex_embeddings_rows_to_openai_batch_output_row(
|
|||
)
|
||||
|
||||
responses = tuple(row["response"] for row in vertex_output_rows)
|
||||
token_count = sum(_embedding_prompt_token_count(response) for response in responses)
|
||||
token_count = sum(vertex_embedding_prompt_token_count(response) for response in responses)
|
||||
body = EmbeddingResponse(
|
||||
model=model or "",
|
||||
data=[
|
||||
|
|
@ -528,6 +519,16 @@ def _model_from_managed_gcs_url(url: str) -> str | None:
|
|||
return match.group(1) if match else None
|
||||
|
||||
|
||||
def is_passthrough_managed_gcs_url(url: str) -> bool:
|
||||
decoded_url: Final = unquote(url)
|
||||
managed_prefix_start: Final = decoded_url.find(VERTEX_AI_MANAGED_GCS_PREFIX)
|
||||
return managed_prefix_start >= 0 and decoded_url.startswith(_PASSTHROUGH_MANAGED_GCS_PREFIX, managed_prefix_start)
|
||||
|
||||
|
||||
def is_passthrough_batch_upload(create_file_data: Mapping[str, object], litellm_params: Mapping[str, object]) -> bool:
|
||||
return create_file_data.get("purpose") == "batch" and litellm_params.get("passthrough") is True
|
||||
|
||||
|
||||
def _is_embeddings_batch_entry(openai_entry: Mapping[str, object]) -> bool:
|
||||
"""
|
||||
Whether an OpenAI batch JSONL line targets the embeddings endpoint.
|
||||
|
|
@ -791,6 +792,58 @@ class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
|
|||
return self._iter_vertex_jsonl_chunks()
|
||||
|
||||
|
||||
def _read_chunk_as_bytes(handle: IO[bytes]) -> bytes:
|
||||
chunk: Final[bytes | str] = handle.read(_RAW_UPLOAD_CHUNK_BYTES)
|
||||
return chunk.encode("utf-8") if isinstance(chunk, str) else bytes(chunk)
|
||||
|
||||
|
||||
def _iter_raw_file_chunks(file_content: FileTypes) -> Iterator[bytes]:
|
||||
content: Final[FileContent | str] = file_content[1] if isinstance(file_content, tuple) else file_content
|
||||
if isinstance(content, (bytes, bytearray)):
|
||||
yield from (
|
||||
bytes(content[offset : offset + _RAW_UPLOAD_CHUNK_BYTES])
|
||||
for offset in range(0, len(content), _RAW_UPLOAD_CHUNK_BYTES)
|
||||
)
|
||||
return
|
||||
if isinstance(content, str):
|
||||
yield content.encode("utf-8")
|
||||
return
|
||||
if isinstance(content, PathLike):
|
||||
with open(str(content), "rb") as handle:
|
||||
yield from iter(lambda: handle.read(_RAW_UPLOAD_CHUNK_BYTES), b"")
|
||||
return
|
||||
if not hasattr(content, "read"):
|
||||
raise ValueError("Unsupported file content type")
|
||||
seek: Final = getattr(content, "seek", None)
|
||||
if seek is None:
|
||||
raise ValueError(
|
||||
"Batch upload file handle must be seekable; got a non-seekable "
|
||||
"stream. Pass bytes, a path, or a seekable handle."
|
||||
)
|
||||
seek(0)
|
||||
yield from iter(lambda: _read_chunk_as_bytes(content), b"")
|
||||
|
||||
|
||||
class _RawFileUploadStream(BaseFileUploadStream):
|
||||
def __init__(self, file_content: FileTypes) -> None:
|
||||
self._file_content = file_content
|
||||
|
||||
def iter_bytes(self) -> Iterator[bytes]:
|
||||
return _iter_raw_file_chunks(self._file_content)
|
||||
|
||||
|
||||
def _managed_batch_object_name(raw_model: str, *, passthrough: bool) -> str:
|
||||
endpoint_id: Final = get_vertex_ai_fine_tuned_endpoint_id(raw_model)
|
||||
model_path: Final = (
|
||||
f"endpoints/{endpoint_id}"
|
||||
if endpoint_id is not None
|
||||
else (raw_model if "publishers/google/models" in raw_model else f"publishers/google/models/{raw_model}")
|
||||
)
|
||||
safe_model_path: Final = sanitize_cloud_object_path(model_path, fallback="model")
|
||||
prefix: Final = _PASSTHROUGH_MANAGED_GCS_PREFIX if passthrough else VERTEX_AI_MANAGED_GCS_PREFIX
|
||||
return f"{prefix}{safe_model_path}/{uuid.uuid4()}"
|
||||
|
||||
|
||||
class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
||||
"""
|
||||
Config for VertexAI Files
|
||||
|
|
@ -848,23 +901,34 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
if deployment_model
|
||||
else openai_jsonl_content[0].get("body", {}).get("model", "")
|
||||
)
|
||||
endpoint_id: Final = get_vertex_ai_fine_tuned_endpoint_id(raw_model)
|
||||
model_path: Final = (
|
||||
f"endpoints/{endpoint_id}"
|
||||
if endpoint_id is not None
|
||||
else (raw_model if "publishers/google/models" in raw_model else f"publishers/google/models/{raw_model}")
|
||||
)
|
||||
safe_model_path: Final = sanitize_cloud_object_path(model_path, fallback="model")
|
||||
object_name: Final = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
|
||||
return object_name
|
||||
return _managed_batch_object_name(raw_model, passthrough=False)
|
||||
|
||||
def get_object_name(self, file_data: FileTypes, purpose: str, deployment_model: str | None = None) -> str:
|
||||
def _get_passthrough_gcs_object_name(self, deployment_model: str | None) -> str:
|
||||
if not deployment_model:
|
||||
raise VertexAIError(
|
||||
status_code=400,
|
||||
message=(
|
||||
"Native Vertex batch passthrough uploads need the deployment model to name the GCS object, "
|
||||
"since native rows carry no model: pass `target_model_names` (proxy) or `model` (SDK)."
|
||||
),
|
||||
)
|
||||
return _managed_batch_object_name(deployment_model.removeprefix("vertex_ai/"), passthrough=True)
|
||||
|
||||
def get_object_name(
|
||||
self,
|
||||
file_data: FileTypes,
|
||||
purpose: str,
|
||||
deployment_model: str | None = None,
|
||||
passthrough: bool = False,
|
||||
) -> str:
|
||||
"""
|
||||
Get the object name for the request.
|
||||
|
||||
Reads only the first JSONL entry (streamed) for batch files, so a large
|
||||
upload is never materialized just to derive the GCS object name.
|
||||
"""
|
||||
if purpose == "batch" and passthrough:
|
||||
return self._get_passthrough_gcs_object_name(deployment_model)
|
||||
if purpose == "batch":
|
||||
## 1. If jsonl, derive the object name from the deployment model (or the first entry's)
|
||||
first_entry: Final = next(_iter_openai_jsonl_entries(file_data), None)
|
||||
|
|
@ -922,6 +986,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
file_data,
|
||||
purpose,
|
||||
deployment_model=configured_model if isinstance(configured_model, str) else None,
|
||||
passthrough=is_passthrough_batch_upload(data, litellm_params),
|
||||
)
|
||||
if object_prefix:
|
||||
object_name = f"{object_prefix}/{object_name}"
|
||||
|
|
@ -984,6 +1049,14 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
if file_data is None:
|
||||
raise ValueError("file is required")
|
||||
|
||||
if is_passthrough_batch_upload(create_file_data, litellm_params):
|
||||
return {
|
||||
"streaming_media_upload": StreamingMediaUploadConfig(
|
||||
body_stream=_RawFileUploadStream(file_data),
|
||||
content_type="application/json",
|
||||
)
|
||||
}
|
||||
|
||||
_, content_type = extract_file_metadata(file_data)
|
||||
if FilesAPIUtils.is_batch_jsonl_request(
|
||||
create_file_data=create_file_data,
|
||||
|
|
@ -1164,6 +1237,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
# transformation, e.g. if they consume raw `predictions.jsonl` directly.
|
||||
if getattr(litellm, "disable_vertex_batch_output_transformation", False):
|
||||
return HttpxBinaryResponseContent(response=raw_response)
|
||||
if is_passthrough_managed_gcs_url(str(raw_response.request.url)):
|
||||
return HttpxBinaryResponseContent(response=raw_response)
|
||||
|
||||
# Try to transform batch output if it's a JSONL file
|
||||
content: Final = raw_response.content
|
||||
|
|
@ -1209,7 +1284,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
Everything else is passed through unchanged, including a row that fails to
|
||||
transform mid-stream.
|
||||
"""
|
||||
if litellm.disable_vertex_batch_output_transformation:
|
||||
if litellm.disable_vertex_batch_output_transformation or is_passthrough_managed_gcs_url(request_url):
|
||||
return FileContentStreamingResult(stream_iterator=stream_iterator, headers=headers)
|
||||
|
||||
first_line, buffered = await _peek_first_jsonl_line(
|
||||
|
|
|
|||
|
|
@ -286,7 +286,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
"""
|
||||
Check if the model is Gemini 3 or newer.
|
||||
"""
|
||||
model_name = model.split("/")[-1].lower()
|
||||
model_name: Final = model.split("/")[-1].lower()
|
||||
is_vertex_fine_tuned_model: Final = model_name.isdigit() or (
|
||||
model.startswith("gemini/") and not model_name.startswith("gemini-")
|
||||
)
|
||||
|
|
|
|||
|
|
@ -57,6 +57,7 @@ from litellm.utils import (
|
|||
# Logging is imported lazily when needed to avoid loading litellm_logging at import time
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.router import Router
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
from litellm.constants import (
|
||||
|
|
@ -351,7 +352,7 @@ class LiteLLM:
|
|||
|
||||
|
||||
class Chat:
|
||||
def __init__(self, params, router_obj: Any | None):
|
||||
def __init__(self, params, router_obj: "Router | None"):
|
||||
self.params = params
|
||||
if self.params.get("acompletion", False) is True:
|
||||
self.params.pop("acompletion")
|
||||
|
|
@ -361,7 +362,7 @@ class Chat:
|
|||
|
||||
|
||||
class Completions:
|
||||
def __init__(self, params, router_obj: Any | None):
|
||||
def __init__(self, params, router_obj: "Router | None"):
|
||||
self.params = params
|
||||
self.router_obj = router_obj
|
||||
|
||||
|
|
@ -377,7 +378,7 @@ class Completions:
|
|||
|
||||
|
||||
class AsyncCompletions:
|
||||
def __init__(self, params, router_obj: Any | None):
|
||||
def __init__(self, params, router_obj: "Router | None"):
|
||||
self.params = params
|
||||
self.router_obj = router_obj
|
||||
|
||||
|
|
@ -1182,7 +1183,7 @@ def _is_claude_tool_target(custom_llm_provider: str | None, model: str) -> bool:
|
|||
return False
|
||||
|
||||
|
||||
def _without_anthropic_only_tool_keys(tool: dict) -> dict:
|
||||
def _without_anthropic_only_tool_keys(tool: dict[str, object]) -> dict[str, object]:
|
||||
kept: Final = {key: value for key, value in tool.items() if key not in _ANTHROPIC_ONLY_TOOL_KEYS}
|
||||
function: Final = tool.get("function")
|
||||
if not isinstance(function, dict):
|
||||
|
|
@ -1193,7 +1194,7 @@ def _without_anthropic_only_tool_keys(tool: dict) -> dict:
|
|||
}
|
||||
|
||||
|
||||
def _drop_anthropic_only_tool_keys(tools: list[dict] | None) -> list[dict] | None:
|
||||
def _drop_anthropic_only_tool_keys(tools: list[dict[str, object]] | None) -> list[dict[str, object]] | None:
|
||||
if tools is None:
|
||||
return None
|
||||
return [_without_anthropic_only_tool_keys(tool) if isinstance(tool, dict) else tool for tool in tools]
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -540,7 +540,9 @@ def llm_passthrough_route(
|
|||
)
|
||||
|
||||
## IS STREAMING REQUEST
|
||||
_streaming_request_data: dict = data if isinstance(data, dict) else (json if isinstance(json, dict) else {})
|
||||
_streaming_request_data: Final[dict[str, object]] = (
|
||||
data if isinstance(data, dict) else (json if isinstance(json, dict) else {})
|
||||
)
|
||||
is_streaming_request: Final = provider_config.is_streaming_request(
|
||||
endpoint=endpoint,
|
||||
request_data=_streaming_request_data,
|
||||
|
|
|
|||
|
|
@ -86,8 +86,8 @@ async def oauth_authorization_uses_gateway_credential(request: Request) -> bool:
|
|||
|
||||
|
||||
async def _opaque_bearer_is_gateway_credential(token: str) -> bool:
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
|
||||
is_envelope, # noqa: PLC0415 # envelope imports bridge types
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # envelope imports bridge types
|
||||
is_envelope,
|
||||
is_refresh_envelope,
|
||||
)
|
||||
from litellm.proxy._types import hash_token # noqa: PLC0415 # proxy import cycle
|
||||
|
|
|
|||
|
|
@ -2634,7 +2634,9 @@ def _build_aggregate_protected_resource_response(request: Request) -> dict:
|
|||
}
|
||||
|
||||
|
||||
def _build_aggregate_authorization_server_response(request: Request, token_exchange_available: bool) -> dict:
|
||||
def _build_aggregate_authorization_server_response(
|
||||
request: Request, token_exchange_available: bool
|
||||
) -> dict[str, object]:
|
||||
"""RFC 8414 metadata for the gateway as the aggregate authorization server.
|
||||
|
||||
The issuer is ``{base}/mcp`` and must stay equal to the value the
|
||||
|
|
|
|||
|
|
@ -3490,7 +3490,7 @@ class MCPServerManager:
|
|||
passthrough_server_ids: Final = [
|
||||
server.server_id
|
||||
for server in self.get_registry().values()
|
||||
if getattr(server, "auth_type", None) == MCPAuth.true_passthrough
|
||||
if server.auth_type == MCPAuth.true_passthrough
|
||||
]
|
||||
combined_servers.update(passthrough_server_ids)
|
||||
|
||||
|
|
|
|||
|
|
@ -111,6 +111,13 @@ _MCP_DESTINATIONS_SCOPE_KEY: Final = "litellm_otel_request_destinations"
|
|||
_MCP_PROTOCOL_VERSION_HEADER: Final = b"mcp-protocol-version"
|
||||
|
||||
|
||||
def reject_disallowed_mcp_origin(request: StarletteRequest) -> None:
|
||||
from litellm.proxy.proxy_server import origins # noqa: PLC0415 # proxy imports this module during startup
|
||||
|
||||
if "*" not in origins and any(origin not in origins for origin in request.headers.getlist("origin")):
|
||||
raise HTTPException(status_code=403, detail="Invalid Origin header")
|
||||
|
||||
|
||||
def unsupported_protocol_version(scope: Scope) -> str | None:
|
||||
"""Return the unsupported ``MCP-Protocol-Version`` header value, if any.
|
||||
|
||||
|
|
@ -1931,6 +1938,7 @@ if MCP_AVAILABLE:
|
|||
async def handle_streamable_http_mcp(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
"""Handle MCP requests through StreamableHTTP."""
|
||||
try:
|
||||
reject_disallowed_mcp_origin(StarletteRequest(scope))
|
||||
bad_version: Final = unsupported_protocol_version(scope)
|
||||
if bad_version is not None:
|
||||
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
|
||||
|
|
@ -2275,6 +2283,7 @@ if MCP_AVAILABLE:
|
|||
async def handle_sse_mcp(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
"""Handle MCP requests through SSE."""
|
||||
try:
|
||||
reject_disallowed_mcp_origin(StarletteRequest(scope))
|
||||
bad_version: Final = unsupported_protocol_version(scope)
|
||||
if bad_version is not None:
|
||||
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
|
||||
|
|
|
|||
|
|
@ -128,14 +128,15 @@ def with_mcp_proxy_identity(tool: Tool, server_id: str) -> Tool:
|
|||
|
||||
|
||||
def _mcp_proxy_identity(tool: Tool) -> MCPProxyToolIdentity:
|
||||
identity: Final = (tool.meta or {}).get(_MCP_PROXY_IDENTITY_META_KEY) # mutable-ok: absent metadata default
|
||||
identity: Final = None if tool.meta is None else tool.meta.get(_MCP_PROXY_IDENTITY_META_KEY)
|
||||
if not isinstance(identity, Mapping):
|
||||
raise TypeError("MCP proxy tool identity is missing")
|
||||
server_id: Final = identity.get("server_id")
|
||||
tool_name: Final = identity.get("tool_name")
|
||||
if not isinstance(server_id, str) or not isinstance(tool_name, str):
|
||||
raise TypeError("MCP proxy tool identity is invalid")
|
||||
return {"server_id": server_id, "tool_name": tool_name} # mutable-ok: TypedDict identity payload
|
||||
resolved: Final[MCPProxyToolIdentity] = {"server_id": server_id, "tool_name": tool_name}
|
||||
return resolved
|
||||
|
||||
|
||||
def mcp_proxy_tool_id(tool: Tool) -> str:
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -307,6 +307,7 @@ class KeyManagementRoutes(str, enum.Enum):
|
|||
# team usage routes
|
||||
TEAM_DAILY_ACTIVITY = "/team/daily/activity"
|
||||
TEAM_DAILY_ACTIVITY_AGGREGATED = "/team/daily/activity/aggregated"
|
||||
TEAM_DAILY_ACTIVITY_AGGREGATED_SEARCH = "/team/daily/activity/aggregated/search"
|
||||
|
||||
# team spend-log viewing
|
||||
SPEND_LOGS = "/spend/logs"
|
||||
|
|
@ -673,6 +674,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
KeyManagementRoutes.TEAM_KEY_BULK_UPDATE.value,
|
||||
KeyManagementRoutes.TEAM_DAILY_ACTIVITY.value,
|
||||
KeyManagementRoutes.TEAM_DAILY_ACTIVITY_AGGREGATED.value,
|
||||
KeyManagementRoutes.TEAM_DAILY_ACTIVITY_AGGREGATED_SEARCH.value,
|
||||
KeyManagementRoutes.SPEND_LOGS.value,
|
||||
KeyManagementRoutes.SPEND_LOGS_V2.value,
|
||||
KeyManagementRoutes.KEY_RESET_SPEND.value,
|
||||
|
|
@ -699,6 +701,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/user/list",
|
||||
"/user/daily/activity",
|
||||
"/user/daily/activity/aggregated",
|
||||
"/user/daily/activity/aggregated/search",
|
||||
# team
|
||||
"/team/new",
|
||||
"/team/update",
|
||||
|
|
@ -716,6 +719,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/team/permissions_bulk_update",
|
||||
"/team/daily/activity",
|
||||
"/team/daily/activity/aggregated",
|
||||
"/team/daily/activity/aggregated/search",
|
||||
"/team/spend/by_user",
|
||||
# gateway request counts (SGR); deployment-wide, admin-only
|
||||
"/gateway/daily/activity",
|
||||
|
|
@ -886,6 +890,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/team/permissions_update",
|
||||
"/team/daily/activity",
|
||||
"/team/daily/activity/aggregated",
|
||||
"/team/daily/activity/aggregated/search",
|
||||
"/team/spend/by_user",
|
||||
"/team/{team_id}/members/me",
|
||||
# POST/GET the team's logging callbacks, and DELETE one of them. Every
|
||||
|
|
@ -901,6 +906,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/model/delete",
|
||||
"/user/daily/activity",
|
||||
"/user/daily/activity/aggregated",
|
||||
"/user/daily/activity/aggregated/search",
|
||||
# Endpoint restricts results to organizations the caller is ORG_ADMIN
|
||||
# of; a caller who administers none gets an empty result set.
|
||||
"/organization/daily/activity",
|
||||
|
|
@ -984,6 +990,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/user/daily/activity",
|
||||
"/team/daily/activity",
|
||||
"/team/daily/activity/aggregated",
|
||||
"/team/daily/activity/aggregated/search",
|
||||
"/tag/daily/activity",
|
||||
"/tag/list",
|
||||
"/audit",
|
||||
|
|
|
|||
|
|
@ -15,11 +15,11 @@ import re
|
|||
import time
|
||||
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeAlias
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, Unpack
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -110,7 +110,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
|
|||
team_membership_auth_cache_key,
|
||||
team_membership_reservation_cache_key,
|
||||
)
|
||||
from litellm.proxy.db.db_lookup_gate import db_lookup_gate
|
||||
from litellm.proxy.db.db_lookup_gate import bounded_db_lookup, db_lookup_gate
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.guardrails.tool_name_extraction import (
|
||||
TOOL_CAPABLE_CALL_TYPES,
|
||||
|
|
@ -223,30 +223,79 @@ class _PrismaTableHolder(Protocol[RowT_co]):
|
|||
def table(self) -> _PrismaAuthTable[RowT_co]: ...
|
||||
|
||||
|
||||
def _dictable_table(repo: _PrismaTableHolder[_PrismaDictableRow]) -> _PrismaAuthTable[_PrismaDictableRow]:
|
||||
return repo.table
|
||||
class _FindOneKwargs(TypedDict):
|
||||
where: ReadOnly[Required[Mapping[str, object]]]
|
||||
include: ReadOnly[NotRequired[Mapping[str, object] | None]]
|
||||
|
||||
|
||||
class _FindManyKwargs(TypedDict):
|
||||
where: ReadOnly[NotRequired[Mapping[str, object] | None]]
|
||||
include: ReadOnly[NotRequired[Mapping[str, object] | None]]
|
||||
take: ReadOnly[NotRequired[int | None]]
|
||||
|
||||
|
||||
class _DeadlineBoundedTable(Generic[RowT_co]):
|
||||
"""Every read on the wrapped table fails with ``DBLookupDeadlineExceeded`` once
|
||||
``PROXY_DB_LOOKUP_DEADLINE_SECONDS`` passes, so a stalled database fails the
|
||||
request fast instead of parking it in the pod until it fills its memory."""
|
||||
|
||||
__slots__ = ("_lookup", "_table")
|
||||
|
||||
def __init__(self, table: _PrismaAuthTable[RowT_co], lookup: str) -> None:
|
||||
self._table: Final = table
|
||||
self._lookup: Final = lookup
|
||||
|
||||
async def find_unique(
|
||||
self,
|
||||
**kwargs: Unpack[_FindOneKwargs], # kwargs-ok: typed pass-through that forwards exactly what the caller passed
|
||||
) -> RowT_co | None:
|
||||
return await bounded_db_lookup(self._table.find_unique(**kwargs), name=self._lookup)
|
||||
|
||||
async def find_first(
|
||||
self,
|
||||
**kwargs: Unpack[_FindOneKwargs], # kwargs-ok: typed pass-through that forwards exactly what the caller passed
|
||||
) -> RowT_co | None:
|
||||
return await bounded_db_lookup(self._table.find_first(**kwargs), name=self._lookup)
|
||||
|
||||
async def find_many(
|
||||
self,
|
||||
**kwargs: Unpack[_FindManyKwargs], # kwargs-ok: typed pass-through that forwards exactly what the caller passed
|
||||
) -> Sequence[RowT_co]:
|
||||
return await bounded_db_lookup(self._table.find_many(**kwargs), name=self._lookup)
|
||||
|
||||
async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> RowT_co | None:
|
||||
return await self._table.update(where=where, data=data)
|
||||
|
||||
async def create(self, *, data: Mapping[str, object], include: Mapping[str, object] | None = None) -> RowT_co:
|
||||
return await self._table.create(data=data, include=include)
|
||||
|
||||
|
||||
def _dictable_table(repo: _PrismaTableHolder[_PrismaDictableRow], lookup: str) -> _PrismaAuthTable[_PrismaDictableRow]:
|
||||
return _DeadlineBoundedTable(repo.table, lookup)
|
||||
|
||||
|
||||
def _jwt_key_mapping_table(
|
||||
repo: _PrismaTableHolder[_PrismaJWTKeyMappingRow],
|
||||
) -> _PrismaAuthTable[_PrismaJWTKeyMappingRow]:
|
||||
return repo.table
|
||||
return _DeadlineBoundedTable(repo.table, "jwt_key_mapping")
|
||||
|
||||
|
||||
def _model_dump_table(repo: _PrismaTableHolder[_PrismaModelDumpRow]) -> _PrismaAuthTable[_PrismaModelDumpRow]:
|
||||
return repo.table
|
||||
def _model_dump_table(
|
||||
repo: _PrismaTableHolder[_PrismaModelDumpRow], lookup: str
|
||||
) -> _PrismaAuthTable[_PrismaModelDumpRow]:
|
||||
return _DeadlineBoundedTable(repo.table, lookup)
|
||||
|
||||
|
||||
def _team_table(repo: _PrismaTableHolder[_PrismaTeamRow]) -> _PrismaAuthTable[_PrismaTeamRow]:
|
||||
return repo.table
|
||||
return _DeadlineBoundedTable(repo.table, "team")
|
||||
|
||||
|
||||
def _vector_store_table(repo: _PrismaTableHolder[_PrismaVectorStoreRow]) -> _PrismaAuthTable[_PrismaVectorStoreRow]:
|
||||
return repo.table
|
||||
return _DeadlineBoundedTable(repo.table, "vector_store")
|
||||
|
||||
|
||||
def _user_table(repo: _PrismaTableHolder[_PrismaUserRow]) -> _PrismaAuthTable[_PrismaUserRow]:
|
||||
return repo.table
|
||||
return _DeadlineBoundedTable(repo.table, "user")
|
||||
|
||||
|
||||
class _VectorStorePermissionsRow(Protocol):
|
||||
|
|
@ -257,7 +306,7 @@ class _VectorStorePermissionsRow(Protocol):
|
|||
def _object_permission_table(
|
||||
repo: _PrismaTableHolder[_VectorStorePermissionsRow],
|
||||
) -> _PrismaAuthTable[_VectorStorePermissionsRow]:
|
||||
return repo.table
|
||||
return _DeadlineBoundedTable(repo.table, "object_permission")
|
||||
|
||||
|
||||
class _PrismaTagRow(Protocol):
|
||||
|
|
@ -1422,7 +1471,7 @@ async def get_default_end_user_budget(
|
|||
|
||||
# Fetch from database
|
||||
try:
|
||||
budget_record: Final = await _dictable_table(BudgetRepository(prisma_client)).find_unique(
|
||||
budget_record: Final = await _dictable_table(BudgetRepository(prisma_client), "budget").find_unique(
|
||||
where={"budget_id": default_budget_id} # mutable-ok: prisma where clause
|
||||
)
|
||||
|
||||
|
|
@ -1483,7 +1532,7 @@ async def get_team_member_default_budget(
|
|||
return cached_budget
|
||||
|
||||
try:
|
||||
budget_record: Final = await _dictable_table(BudgetRepository(prisma_client)).find_unique(
|
||||
budget_record: Final = await _dictable_table(BudgetRepository(prisma_client), "budget").find_unique(
|
||||
where={"budget_id": budget_id}
|
||||
)
|
||||
except Exception:
|
||||
|
|
@ -1877,7 +1926,7 @@ async def get_end_user_object(
|
|||
|
||||
# Fetch from database
|
||||
try:
|
||||
response: Final = await _dictable_table(EndUserRepository(prisma_client)).find_unique(
|
||||
response: Final = await _dictable_table(EndUserRepository(prisma_client), "end_user").find_unique(
|
||||
where={"user_id": end_user_id},
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
)
|
||||
|
|
@ -2286,7 +2335,7 @@ async def _fetch_team_membership_from_db(
|
|||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
) -> LiteLLM_TeamMembership | None:
|
||||
_ = parent_otel_span, proxy_logging_obj
|
||||
response: Final = await _dictable_table(TeamMembershipRepository(prisma_client)).find_unique(
|
||||
response: Final = await _dictable_table(TeamMembershipRepository(prisma_client), "team_membership").find_unique(
|
||||
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
|
|
@ -3290,7 +3339,7 @@ async def get_access_object(
|
|||
|
||||
# Not in cache - fetch from DB
|
||||
try:
|
||||
response: Final = await _dictable_table(AccessGroupRepository(prisma_client)).find_unique(
|
||||
response: Final = await _dictable_table(AccessGroupRepository(prisma_client), "access_group").find_unique(
|
||||
where={"access_group_id": access_group_id}
|
||||
)
|
||||
|
||||
|
|
@ -3472,7 +3521,7 @@ async def get_org_object_by_alias(
|
|||
|
||||
# Query database by organization_alias
|
||||
try:
|
||||
orgs = await _model_dump_table(OrganizationRepository(prisma_client)).find_many(
|
||||
orgs = await _model_dump_table(OrganizationRepository(prisma_client), "organization").find_many(
|
||||
where={"organization_alias": org_alias}
|
||||
)
|
||||
|
||||
|
|
@ -3650,10 +3699,32 @@ async def _fetch_key_object_from_db_with_reconnect(
|
|||
prisma_client: PrismaClient,
|
||||
parent_otel_span: Span | None,
|
||||
proxy_logging_obj: ProxyLogging | None,
|
||||
deadline_seconds: float | None = None,
|
||||
) -> BaseModel | None:
|
||||
"""
|
||||
Fetch key object from DB and retry once if a DB connection error can be healed.
|
||||
The gate wait, the query, the reconnect, and the retry share one deadline, so a
|
||||
stalled database fails the request with ``DBLookupDeadlineExceeded`` instead of
|
||||
parking it.
|
||||
"""
|
||||
return await bounded_db_lookup(
|
||||
_fetch_key_object_from_db_unbounded(
|
||||
hashed_token=hashed_token,
|
||||
prisma_client=prisma_client,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
),
|
||||
name="key",
|
||||
deadline_seconds=deadline_seconds,
|
||||
)
|
||||
|
||||
|
||||
async def _fetch_key_object_from_db_unbounded(
|
||||
hashed_token: str,
|
||||
prisma_client: PrismaClient,
|
||||
parent_otel_span: Span | None,
|
||||
proxy_logging_obj: ProxyLogging | None,
|
||||
) -> BaseModel | None:
|
||||
async with db_lookup_gate.current():
|
||||
try:
|
||||
return await prisma_client.get_data(
|
||||
|
|
@ -3874,9 +3945,9 @@ async def get_object_permission(
|
|||
|
||||
# else, check db
|
||||
try:
|
||||
response: Final = await _dictable_table(ObjectPermissionRepository(prisma_client)).find_unique(
|
||||
where={"object_permission_id": object_permission_id}
|
||||
)
|
||||
response: Final = await _dictable_table(
|
||||
ObjectPermissionRepository(prisma_client), "object_permission"
|
||||
).find_unique(where={"object_permission_id": object_permission_id})
|
||||
|
||||
if response is None:
|
||||
return None
|
||||
|
|
@ -4008,7 +4079,9 @@ async def get_org_object(
|
|||
if include_budget_table:
|
||||
query_kwargs["include"] = {"litellm_budget_table": True}
|
||||
|
||||
response: Final = await _model_dump_table(OrganizationRepository(prisma_client)).find_unique(**query_kwargs)
|
||||
response: Final = await _model_dump_table(OrganizationRepository(prisma_client), "organization").find_unique(
|
||||
**query_kwargs
|
||||
)
|
||||
except Exception:
|
||||
# An operational failure (DB down, timeout, cache fault) is NOT the same fact as a confirmed
|
||||
# missing row, and relabelling it as "doesn't exist" made every caller unable to tell them
|
||||
|
|
@ -4073,7 +4146,7 @@ async def get_org_object_for_request(
|
|||
)
|
||||
except OrganizationNotFoundError:
|
||||
return None
|
||||
except Exception as e: # noqa: BLE001 # only a DB outage may fail auth here, anything else degrades to no org limits
|
||||
except Exception as e:
|
||||
if not PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(e):
|
||||
verbose_proxy_logger.debug("org lookup failed, continuing without org limits", exc_info=True)
|
||||
return None
|
||||
|
|
@ -5948,7 +6021,7 @@ async def get_project_object(
|
|||
return deserialized_project
|
||||
|
||||
# Fetch from DB
|
||||
project_row: Final = await _model_dump_table(ProjectRepository(prisma_client)).find_unique(
|
||||
project_row: Final = await _model_dump_table(ProjectRepository(prisma_client), "project").find_unique(
|
||||
where={"project_id": project_id},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -120,6 +120,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
|
|||
UserApiKeyCache,
|
||||
team_membership_auth_cache_key,
|
||||
)
|
||||
from litellm.proxy.db.db_lookup_gate import bounded_db_lookup
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.proxy.spend_tracking.carried_budget_state import carry_team_and_user_budget_state
|
||||
|
|
@ -735,8 +736,9 @@ async def _fetch_global_spend_with_event_coordination(
|
|||
"""
|
||||
|
||||
async def _load_global_spend() -> float | None:
|
||||
proxy_budget_row: Final = await prisma_client.db.litellm_usertable.find_unique(
|
||||
where={"user_id": LITELLM_PROXY_BUDGET_NAME}
|
||||
proxy_budget_row: Final = await bounded_db_lookup(
|
||||
prisma_client.db.litellm_usertable.find_unique(where={"user_id": LITELLM_PROXY_BUDGET_NAME}),
|
||||
name="proxy_budget",
|
||||
)
|
||||
return float(proxy_budget_row.spend) if proxy_budget_row is not None else None
|
||||
|
||||
|
|
|
|||
|
|
@ -218,7 +218,7 @@ ProxyRouteType: TypeAlias = Literal[
|
|||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
# Type alias for streaming chunk serializer (chunk after hooks + cost injection -> wire format)
|
||||
StreamChunkSerializer = Callable[[Any], str]
|
||||
StreamChunkSerializer = Callable[[object], str]
|
||||
# Type alias for streaming error serializer (ProxyException -> wire format)
|
||||
StreamErrorSerializer = Callable[[ProxyException], str]
|
||||
|
||||
|
|
@ -459,7 +459,7 @@ async def _bill_partial_streamed_spend_on_disconnect(request_data: dict, respons
|
|||
return True
|
||||
|
||||
|
||||
async def _cancel_pending_gather_tasks(tasks: list["asyncio.Task[Any]"]) -> None:
|
||||
async def _cancel_pending_gather_tasks(tasks: Sequence["asyncio.Task[object]"]) -> None:
|
||||
pending_tasks: Final = [task for task in tasks if not task.done()]
|
||||
for task in pending_tasks:
|
||||
task.cancel()
|
||||
|
|
@ -2323,7 +2323,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
return fallbacks if isinstance(fallbacks, list) and fallbacks else None
|
||||
|
||||
@staticmethod
|
||||
def _resolve_fallback_models(model: str, fallbacks: list) -> list | None:
|
||||
def _resolve_fallback_models(model: str, fallbacks: list) -> list[str] | None:
|
||||
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
|
||||
|
||||
fallback_model_group, generic_fallback_idx = get_fallback_model_group(
|
||||
|
|
@ -3145,7 +3145,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
logging_obj._on_detached_stream_failure = _on_detached_stream_failure
|
||||
|
||||
def _is_streaming_response(self, response: Any) -> bool:
|
||||
def _is_streaming_response(self, response: object) -> bool:
|
||||
"""
|
||||
Check if the response object is actually a streaming response by inspecting its type.
|
||||
|
||||
|
|
@ -3259,7 +3259,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
async def _handle_non_streaming_allm_passthrough_route(
|
||||
self,
|
||||
response: Any,
|
||||
response: _UpstreamHttpResponse,
|
||||
proxy_logging_obj: "ProxyLogging",
|
||||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
custom_headers: Mapping[str, str],
|
||||
|
|
@ -3852,7 +3852,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
@staticmethod
|
||||
async def async_streaming_data_generator(
|
||||
response: Any,
|
||||
response: object,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_data: dict,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
|
|
@ -3993,7 +3993,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
@staticmethod
|
||||
def async_sse_data_generator(
|
||||
response: Any,
|
||||
response: object,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_data: dict,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,8 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final, TypeAlias, Union
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeAlias, Union, cast
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
|
|
@ -28,7 +31,7 @@ class CustomOpenAPISpec:
|
|||
"/openai/deployments/{model}/embeddings",
|
||||
]
|
||||
|
||||
RESPONSES_API_PATHS = ["/v1/responses", "/responses"]
|
||||
RESPONSES_API_PATHS = ["/v1/responses", "/responses", "/openai/v1/responses"]
|
||||
|
||||
@staticmethod
|
||||
def _as_object(node: JsonValue) -> JsonObject:
|
||||
|
|
@ -44,26 +47,18 @@ class CustomOpenAPISpec:
|
|||
return CustomOpenAPISpec._as_object(components.setdefault("schemas", {}))
|
||||
|
||||
@staticmethod
|
||||
def get_pydantic_schema(model_class) -> JsonObject | None:
|
||||
def get_pydantic_schema(model_class: type) -> JsonObject | None:
|
||||
"""
|
||||
Get JSON schema from a Pydantic model, handling both v1 and v2 APIs.
|
||||
Get JSON schema for a request or response model class, including TypedDicts.
|
||||
|
||||
Args:
|
||||
model_class: Pydantic model class
|
||||
model_class: Pydantic model class or TypedDict
|
||||
|
||||
Returns:
|
||||
JSON schema dict or None if failed
|
||||
"""
|
||||
try:
|
||||
# Try Pydantic v2 method first
|
||||
return model_class.model_json_schema()
|
||||
except AttributeError:
|
||||
try:
|
||||
# Fallback to Pydantic v1 method
|
||||
return model_class.schema()
|
||||
except AttributeError:
|
||||
# If both methods fail, return None
|
||||
return None
|
||||
return cast(JsonObject, TypeAdapter(model_class).json_schema()) # cast-ok: pydantic returns dict[str, Any]
|
||||
except Exception as e:
|
||||
# FastAPI 0.120+ may fail schema generation for certain types (e.g., openai.Timeout)
|
||||
# Log the error and return None to skip schema generation for this model
|
||||
|
|
@ -83,13 +78,18 @@ class CustomOpenAPISpec:
|
|||
# Ensure components/schemas structure exists
|
||||
_ = CustomOpenAPISpec._components_schemas(openapi_schema)
|
||||
|
||||
# Add the schema
|
||||
CustomOpenAPISpec._move_defs_to_components(openapi_schema, {schema_name: schema_def})
|
||||
defs: Final[Mapping[str, JsonValue]] = (
|
||||
CustomOpenAPISpec._as_object(schema_def["$defs"]) if "$defs" in schema_def else MappingProxyType({})
|
||||
)
|
||||
renames: Final = CustomOpenAPISpec._move_defs_to_components(openapi_schema, defs, schema_name)
|
||||
schemas: Final = CustomOpenAPISpec._components_schemas(openapi_schema)
|
||||
schemas[schema_name] = CustomOpenAPISpec._rewrite_defs_refs(schema_def, renames)
|
||||
|
||||
@staticmethod
|
||||
def _expanded_request_field(field_name: str, field_def: JsonValue) -> JsonValue:
|
||||
expanded: Final = CustomOpenAPISpec._rewrite_defs_refs(
|
||||
CustomOpenAPISpec._expand_field_definition(CustomOpenAPISpec._as_object(field_def))
|
||||
CustomOpenAPISpec._expand_field_definition(CustomOpenAPISpec._as_object(field_def)),
|
||||
MappingProxyType({}),
|
||||
)
|
||||
if field_name != "messages":
|
||||
return expanded
|
||||
|
|
@ -127,13 +127,6 @@ class CustomOpenAPISpec:
|
|||
schema_properties = CustomOpenAPISpec._as_object(actual_schema.get("properties"))
|
||||
required_fields = actual_schema.get("required", [])
|
||||
|
||||
# Extract $defs and add them to components/schemas
|
||||
# This fixes Pydantic v2 $defs not being resolvable in Swagger/OpenAPI
|
||||
if "$defs" in actual_schema:
|
||||
CustomOpenAPISpec._move_defs_to_components(
|
||||
openapi_schema, CustomOpenAPISpec._as_object(actual_schema["$defs"])
|
||||
)
|
||||
|
||||
# Create an expanded inline schema instead of just a $ref
|
||||
# This makes Swagger UI show all individual fields in the request body editor
|
||||
expanded_schema: JsonObject = {
|
||||
|
|
@ -161,7 +154,9 @@ class CustomOpenAPISpec:
|
|||
]
|
||||
|
||||
@staticmethod
|
||||
def _move_defs_to_components(openapi_schema: JsonObject, defs: Mapping[str, JsonValue]) -> None:
|
||||
def _move_defs_to_components(
|
||||
openapi_schema: JsonObject, defs: Mapping[str, JsonValue], namespace: str
|
||||
) -> Mapping[str, str]:
|
||||
"""
|
||||
Move $defs from Pydantic v2 schema to OpenAPI components/schemas.
|
||||
This makes the definitions resolvable in Swagger/OpenAPI viewers.
|
||||
|
|
@ -169,36 +164,68 @@ class CustomOpenAPISpec:
|
|||
Args:
|
||||
openapi_schema: The OpenAPI schema dict to modify
|
||||
defs: The $defs dictionary from Pydantic schema
|
||||
namespace: Prefix used to rename defs that would overwrite an existing component
|
||||
|
||||
Returns:
|
||||
Map of original def names to renamed component names for collision cases
|
||||
"""
|
||||
if not defs:
|
||||
return
|
||||
|
||||
# Ensure components/schemas exists
|
||||
schemas: Final = CustomOpenAPISpec._components_schemas(openapi_schema)
|
||||
|
||||
# Add each definition to components/schemas
|
||||
renames: Final = CustomOpenAPISpec._fixed_renames(schemas, defs, namespace, MappingProxyType({}))
|
||||
for def_name, def_schema in defs.items():
|
||||
# Recursively rewrite any nested $defs references within this definition
|
||||
schemas[def_name] = CustomOpenAPISpec._rewrite_defs_refs(def_schema)
|
||||
|
||||
# If this definition also has $defs, process them recursively
|
||||
def_object = CustomOpenAPISpec._as_object(def_schema)
|
||||
if "$defs" in def_object:
|
||||
CustomOpenAPISpec._move_defs_to_components(
|
||||
openapi_schema, CustomOpenAPISpec._as_object(def_object["$defs"])
|
||||
)
|
||||
if def_name in schemas and def_name not in renames:
|
||||
continue
|
||||
schemas[renames.get(def_name, def_name)] = CustomOpenAPISpec._rewrite_defs_refs(def_schema, renames)
|
||||
return renames
|
||||
|
||||
@staticmethod
|
||||
def _rewritten_defs_entry(key: str, value: JsonValue) -> JsonValue:
|
||||
def _def_collisions(
|
||||
schemas: JsonObject, defs: Mapping[str, JsonValue], namespace: str, renames: Mapping[str, str]
|
||||
) -> Mapping[str, str]:
|
||||
return MappingProxyType(
|
||||
{
|
||||
name: f"{namespace}_{name}"
|
||||
for name, d in defs.items()
|
||||
if name in schemas
|
||||
and not CustomOpenAPISpec._same_shape(schemas[name], CustomOpenAPISpec._rewrite_defs_refs(d, renames))
|
||||
}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _same_shape(existing: JsonValue, incoming: JsonValue) -> bool:
|
||||
if existing == incoming:
|
||||
return True
|
||||
existing_obj: Final = CustomOpenAPISpec._as_object(existing)
|
||||
incoming_obj: Final = CustomOpenAPISpec._as_object(incoming)
|
||||
existing_props: Final = CustomOpenAPISpec._as_object(existing_obj.get("properties"))
|
||||
incoming_props: Final = CustomOpenAPISpec._as_object(incoming_obj.get("properties"))
|
||||
if not existing_props or not incoming_props:
|
||||
return False
|
||||
return existing_props.keys() == incoming_props.keys() and frozenset(
|
||||
x for x in CustomOpenAPISpec._as_array(existing_obj.get("required")) if isinstance(x, str)
|
||||
) == frozenset(x for x in CustomOpenAPISpec._as_array(incoming_obj.get("required")) if isinstance(x, str))
|
||||
|
||||
@staticmethod
|
||||
def _fixed_renames(
|
||||
schemas: JsonObject, defs: Mapping[str, JsonValue], namespace: str, renames: Mapping[str, str]
|
||||
) -> Mapping[str, str]:
|
||||
next_renames: Final = MappingProxyType(
|
||||
{**renames, **CustomOpenAPISpec._def_collisions(schemas, defs, namespace, renames)}
|
||||
)
|
||||
if next_renames == renames:
|
||||
return renames
|
||||
return CustomOpenAPISpec._fixed_renames(schemas, defs, namespace, next_renames)
|
||||
|
||||
@staticmethod
|
||||
def _rewritten_defs_entry(key: str, value: JsonValue, renames: Mapping[str, str]) -> JsonValue:
|
||||
if key == "$ref" and isinstance(value, str) and value.startswith("#/$defs/"):
|
||||
# Rewrite the reference to use components/schemas
|
||||
def_name: Final = value.replace("#/$defs/", "")
|
||||
return f"#/components/schemas/{def_name}"
|
||||
return f"#/components/schemas/{renames.get(def_name, def_name)}"
|
||||
# Recursively process nested structures
|
||||
return CustomOpenAPISpec._rewrite_defs_refs(value)
|
||||
return CustomOpenAPISpec._rewrite_defs_refs(value, renames)
|
||||
|
||||
@staticmethod
|
||||
def _rewrite_defs_refs(schema: JsonValue) -> JsonValue:
|
||||
def _rewrite_defs_refs(schema: JsonValue, renames: Mapping[str, str]) -> JsonValue:
|
||||
"""
|
||||
Recursively rewrite $ref values from #/$defs/... to #/components/schemas/...
|
||||
This converts Pydantic v2 references to OpenAPI-compatible references.
|
||||
|
|
@ -211,12 +238,12 @@ class CustomOpenAPISpec:
|
|||
"""
|
||||
if isinstance(schema, dict):
|
||||
return {
|
||||
key: CustomOpenAPISpec._rewritten_defs_entry(key, value)
|
||||
key: CustomOpenAPISpec._rewritten_defs_entry(key, value, renames)
|
||||
for key, value in schema.items()
|
||||
if key != "$defs"
|
||||
}
|
||||
if isinstance(schema, list):
|
||||
return [CustomOpenAPISpec._rewrite_defs_refs(item) for item in schema]
|
||||
return [CustomOpenAPISpec._rewrite_defs_refs(item, renames) for item in schema]
|
||||
return schema
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -13,9 +13,12 @@ from __future__ import annotations
|
|||
import re
|
||||
from collections.abc import Container, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from functools import reduce
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -28,6 +31,8 @@ CLAUDE_CODE_CLIENT: Final = "claude-code"
|
|||
_CLAUDE_CODE_ALIAS_PREFIX: Final = "claude-router-"
|
||||
_ONE_MILLION_SUFFIX: Final = "[1m]"
|
||||
_ONE_MILLION_TOKENS: Final = 1_000_000
|
||||
_ALIAS_ENTRIES: Final = TypeAdapter(Mapping[object, object])
|
||||
_NO_ALIASES: Final[Mapping[str, str]] = MappingProxyType({})
|
||||
|
||||
|
||||
def configured_display_names(
|
||||
|
|
@ -152,6 +157,77 @@ class ClaudeCodeRoutingNames:
|
|||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CallerAliases:
|
||||
"""`own` are the caller's key and team alias maps, the names `/v1/models` lists for it.
|
||||
`rewrite` are the maps `/chat/completions` rewrites its model through, in the order it
|
||||
applies them: the team's, the key's in `add_litellm_data_to_request`, then the global
|
||||
`model_alias_map` and the key's again in `common_processing_pre_call_logic`."""
|
||||
|
||||
own: tuple[object, ...]
|
||||
rewrite: tuple[object, ...]
|
||||
|
||||
|
||||
def caller_alias_maps(
|
||||
key_aliases: object,
|
||||
team_aliases: object,
|
||||
key_team_id: str | None,
|
||||
listed_team_id: str | None,
|
||||
) -> CallerAliases:
|
||||
"""Team aliases count only when listing the team the key authenticated as."""
|
||||
if listed_team_id is not None and listed_team_id != key_team_id:
|
||||
return CallerAliases((key_aliases,), (key_aliases, litellm.model_alias_map, key_aliases))
|
||||
return CallerAliases((team_aliases, key_aliases), (team_aliases, key_aliases, litellm.model_alias_map, key_aliases))
|
||||
|
||||
|
||||
def _alias_map(aliases: object) -> Mapping[str, str]:
|
||||
try:
|
||||
entries: Final = _ALIAS_ENTRIES.validate_python(aliases, strict=True)
|
||||
except ValidationError:
|
||||
return _NO_ALIASES
|
||||
return MappingProxyType(
|
||||
{alias: target for alias, target in entries.items() if isinstance(alias, str) and isinstance(target, str)}
|
||||
)
|
||||
|
||||
|
||||
def _alias_names(alias_maps: Sequence[Mapping[str, str]]) -> tuple[str, ...]:
|
||||
return tuple(dict.fromkeys(alias for aliases in alias_maps for alias in aliases))
|
||||
|
||||
|
||||
def _rewrite(model_id: str, alias_maps: Sequence[Mapping[str, str]]) -> str | None:
|
||||
target: Final = reduce(lambda name, aliases: aliases.get(name, name), alias_maps, model_id)
|
||||
return None if target == model_id else target
|
||||
|
||||
|
||||
def alias_target(model_id: str, aliases: CallerAliases, listed: Container[str] = frozenset()) -> str | None:
|
||||
"""The model group `/chat/completions` rewrites `model_id` to, else None. A `model_id`
|
||||
already `listed` keeps its own row, so it is never rewritten."""
|
||||
if model_id in listed:
|
||||
return None
|
||||
return _rewrite(model_id, tuple(_alias_map(alias_map) for alias_map in aliases.rewrite))
|
||||
|
||||
|
||||
def alias_listing_entries(
|
||||
entries: Sequence[tuple[str, str]],
|
||||
aliases: CallerAliases,
|
||||
) -> tuple[tuple[str, str], ...]:
|
||||
"""`entries` plus one `(alias, lookup_id)` row per key or team alias whose target is
|
||||
listed. An alias colliding with a listed id keeps the listed entry."""
|
||||
maps: Final = tuple(_alias_map(alias_map) for alias_map in aliases.rewrite)
|
||||
own: Final = tuple(_alias_map(alias_map) for alias_map in aliases.own)
|
||||
lookup_by_response: Final = MappingProxyType(dict(entries))
|
||||
lookup_ids: Final = frozenset(lookup_by_response.values())
|
||||
targets: Final = MappingProxyType(
|
||||
{alias: _rewrite(alias, maps) for alias in _alias_names(own) if alias not in lookup_by_response}
|
||||
)
|
||||
added: Final = tuple(
|
||||
(alias, lookup_by_response.get(target, target))
|
||||
for alias, target in targets.items()
|
||||
if target is not None and (target in lookup_by_response or target in lookup_ids)
|
||||
)
|
||||
return (*entries, *added)
|
||||
|
||||
|
||||
def claude_code_requested_group(
|
||||
requested: str,
|
||||
llm_router: Router,
|
||||
|
|
@ -218,7 +294,7 @@ class TeamModelNameTranslator:
|
|||
|
||||
@staticmethod
|
||||
def _response_to_lookup_map(
|
||||
model_names: list[str],
|
||||
model_names: Sequence[str],
|
||||
internal_to_public: dict[str, str],
|
||||
) -> dict[str, str]:
|
||||
"""Map each public response id to the first internal lookup id seen in
|
||||
|
|
@ -235,7 +311,7 @@ class TeamModelNameTranslator:
|
|||
|
||||
@staticmethod
|
||||
def listing_entries(
|
||||
model_names: list[str],
|
||||
model_names: Sequence[str],
|
||||
llm_router: Router | None,
|
||||
general_settings: Mapping[str, object],
|
||||
) -> list[tuple[str, str]]:
|
||||
|
|
|
|||
|
|
@ -486,7 +486,7 @@ class BaselineAccountingStore:
|
|||
async def _pages(
|
||||
self, db: SupportsRawQueries, scope: str, after_revision: int, withdraw_from: float | None = None
|
||||
) -> AsyncIterator[tuple[_StoredRecord, ...]]:
|
||||
cursor: float | None = None
|
||||
cursor: float | None = None # rebind-ok: keyset pagination advances after each complete timestamp group
|
||||
while page := _RECORDS.validate_python(
|
||||
tuple(await db.query_raw(_READ_PAGE, scope, after_revision, cursor, _PAGE_TIMESTAMPS, withdraw_from))
|
||||
):
|
||||
|
|
@ -627,7 +627,7 @@ async def flush_baseline_accounting(client: PrismaClient) -> None:
|
|||
more_queued: Final = bool(client.baseline_accounting_transactions)
|
||||
try:
|
||||
remaining: Final = await asyncio.wait_for(_flush_records(store, batch), timeout=5)
|
||||
except (Exception, asyncio.CancelledError) as error: # noqa: BLE001 # unknown acknowledgements can be replayed safely
|
||||
except (Exception, asyncio.CancelledError) as error:
|
||||
async with client.baseline_accounting_lock:
|
||||
client.baseline_accounting_transactions.extend(batch)
|
||||
if isinstance(error, asyncio.CancelledError):
|
||||
|
|
|
|||
|
|
@ -1,7 +1,14 @@
|
|||
import asyncio
|
||||
from typing import Final
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Final, TypeVar
|
||||
|
||||
from litellm.constants import PROXY_DB_LOOKUP_MAX_CONCURRENCY
|
||||
from litellm.constants import (
|
||||
PROXY_DB_LOOKUP_DEADLINE_SECONDS,
|
||||
PROXY_DB_LOOKUP_MAX_CONCURRENCY,
|
||||
)
|
||||
|
||||
LookupT = TypeVar("LookupT")
|
||||
|
||||
|
||||
class LoopBoundSemaphore:
|
||||
|
|
@ -20,4 +27,64 @@ class LoopBoundSemaphore:
|
|||
return self._semaphore
|
||||
|
||||
|
||||
class DBLookupDeadlineExceeded(asyncio.TimeoutError):
|
||||
def __init__(self, lookup: str, deadline_seconds: float) -> None:
|
||||
super().__init__(f"{lookup} lookup did not answer within {deadline_seconds:g}s")
|
||||
self.lookup: Final = lookup
|
||||
self.deadline_seconds: Final = deadline_seconds
|
||||
|
||||
|
||||
class DBLookupStallTracker:
|
||||
__slots__ = ("_clock", "_last_hit")
|
||||
|
||||
def __init__(self, clock: Callable[[], float] = time.monotonic) -> None:
|
||||
self._clock: Final = clock
|
||||
self._last_hit: float | None = None
|
||||
|
||||
def record_hit(self) -> None:
|
||||
self._last_hit = self._clock()
|
||||
|
||||
def clear(self) -> None:
|
||||
self._last_hit = None
|
||||
|
||||
def stalled_within(self, window_seconds: float) -> bool:
|
||||
if self._last_hit is None:
|
||||
return False
|
||||
return self._clock() - self._last_hit < window_seconds
|
||||
|
||||
|
||||
db_lookup_gate: Final = LoopBoundSemaphore(PROXY_DB_LOOKUP_MAX_CONCURRENCY)
|
||||
db_lookup_stall_tracker: Final = DBLookupStallTracker()
|
||||
|
||||
|
||||
def _consume_abandoned_lookup(task: asyncio.Future[LookupT]) -> None:
|
||||
if not task.cancelled():
|
||||
task.exception()
|
||||
|
||||
|
||||
async def bounded_db_lookup(
|
||||
lookup: Awaitable[LookupT],
|
||||
*,
|
||||
name: str,
|
||||
deadline_seconds: float | None = None,
|
||||
tracker: DBLookupStallTracker = db_lookup_stall_tracker,
|
||||
) -> LookupT:
|
||||
timeout: Final = PROXY_DB_LOOKUP_DEADLINE_SECONDS if deadline_seconds is None else deadline_seconds
|
||||
task: Final = asyncio.ensure_future(lookup)
|
||||
try:
|
||||
done, _ = await asyncio.wait({task}, timeout=timeout)
|
||||
except asyncio.CancelledError:
|
||||
task.cancel()
|
||||
raise
|
||||
if task not in done:
|
||||
task.cancel()
|
||||
task.add_done_callback(_consume_abandoned_lookup)
|
||||
tracker.record_hit()
|
||||
raise DBLookupDeadlineExceeded(name, timeout)
|
||||
try:
|
||||
return task.result()
|
||||
except DBLookupDeadlineExceeded:
|
||||
raise
|
||||
except asyncio.TimeoutError as e:
|
||||
tracker.record_hit()
|
||||
raise DBLookupDeadlineExceeded(name, timeout) from e
|
||||
|
|
|
|||
|
|
@ -12,7 +12,8 @@ import os
|
|||
import random
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from collections.abc import Callable, Coroutine, Mapping, Sequence
|
||||
from contextvars import ContextVar
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeAlias, TypeVar, cast, overload
|
||||
|
|
@ -269,6 +270,80 @@ def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager:
|
|||
return tx
|
||||
|
||||
|
||||
_daily_spend_commit_started: Final[ContextVar[asyncio.Event | None]] = ContextVar(
|
||||
"_daily_spend_commit_started", default=None
|
||||
)
|
||||
|
||||
|
||||
def _mark_daily_spend_commit_started() -> None:
|
||||
started: Final = _daily_spend_commit_started.get()
|
||||
if started is not None:
|
||||
started.set()
|
||||
|
||||
|
||||
def _mark_daily_spend_commit_finished() -> None:
|
||||
started: Final = _daily_spend_commit_started.get()
|
||||
if started is not None:
|
||||
started.clear()
|
||||
|
||||
|
||||
def _start_daily_spend_commit(
|
||||
commit_started: asyncio.Event, commit: Callable[[], Coroutine[object, object, None]]
|
||||
) -> "asyncio.Task[None]":
|
||||
token: Final = _daily_spend_commit_started.set(commit_started)
|
||||
try:
|
||||
return asyncio.ensure_future(commit())
|
||||
finally:
|
||||
_daily_spend_commit_started.reset(token)
|
||||
|
||||
|
||||
def _track_interrupted_commit(commits: set[asyncio.Task[None]], settle: Coroutine[object, object, None]) -> None:
|
||||
task: Final = asyncio.ensure_future(settle)
|
||||
commits.add(task)
|
||||
task.add_done_callback(commits.discard)
|
||||
|
||||
|
||||
async def _settle_interrupted_commits(commits: set[asyncio.Task[None]]) -> None:
|
||||
while commits:
|
||||
await asyncio.wait(tuple(commits))
|
||||
|
||||
|
||||
async def _restore_tag_spend_the_commit_left_behind(
|
||||
commit_task: "asyncio.Task[None]",
|
||||
redis_update_buffer: RedisUpdateBuffer,
|
||||
transactions: dict[str, DailyTagSpendTransaction],
|
||||
) -> None:
|
||||
await asyncio.wait({commit_task})
|
||||
if commit_task.cancelled() or commit_task.exception() is None:
|
||||
return
|
||||
await redis_update_buffer.restore_transactions_to_redis(
|
||||
daily_tag_spend_update_transactions=transactions,
|
||||
)
|
||||
|
||||
|
||||
async def _requeue_daily_spend_the_commit_left_behind(
|
||||
commit_task: "asyncio.Task[None]",
|
||||
queue: DailySpendUpdateQueue,
|
||||
entity_type: str,
|
||||
transactions: dict[str, BaseDailySpendTransaction],
|
||||
) -> None:
|
||||
await asyncio.wait({commit_task})
|
||||
if commit_task.cancelled() or not transactions:
|
||||
return
|
||||
failure: Final = commit_task.exception()
|
||||
if failure is None:
|
||||
return
|
||||
spend_log_error(
|
||||
"Spend tracking - daily %s spend commit interrupted by shutdown failed. Re-queued %d rows for the "
|
||||
"shutdown flush. Error: %s",
|
||||
entity_type,
|
||||
len(transactions),
|
||||
str(failure),
|
||||
exc=failure,
|
||||
)
|
||||
await queue.add_update(transactions)
|
||||
|
||||
|
||||
# The per-team advisory lock the team endpoints hold while changing a roster (TEAM_ADVISORY_LOCK_SQL),
|
||||
# so the roster check below cannot interleave with their writes. A row lock would deadlock with the
|
||||
# access-group endpoints, which lock a team row after an access-group lock.
|
||||
|
|
@ -391,6 +466,9 @@ class DBSpendUpdateWriter:
|
|||
self.daily_org_spend_update_queue = DailySpendUpdateQueue()
|
||||
self.daily_tag_spend_update_queue = DailySpendUpdateQueue()
|
||||
self.window_spend_update_queue = WindowSpendUpdateQueue()
|
||||
self.interrupted_tag_commits: set[asyncio.Task[None]] = (
|
||||
set()
|
||||
) # mutable-ok: same registry as DailySpendUpdateQueue.interrupted_commits
|
||||
|
||||
async def update_database(
|
||||
# LiteLLM management object fields
|
||||
|
|
@ -1606,17 +1684,24 @@ class DBSpendUpdateWriter:
|
|||
proxy_logging_obj: ProxyLogging,
|
||||
) -> None:
|
||||
transactions: Final = await queue.flush_and_get_aggregated_daily_spend_update_transactions()
|
||||
commit_task: Final = asyncio.ensure_future(
|
||||
commit(
|
||||
commit_started: Final = asyncio.Event()
|
||||
commit_task: Final = _start_daily_spend_commit(
|
||||
commit_started,
|
||||
lambda: commit(
|
||||
n_retry_times=n_retry_times,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
daily_spend_transactions=cast(dict[str, _DailySpendTransactionT], transactions),
|
||||
)
|
||||
),
|
||||
)
|
||||
try:
|
||||
await asyncio.shield(commit_task)
|
||||
except asyncio.CancelledError:
|
||||
if commit_started.is_set():
|
||||
queue.track_interrupted_commit(
|
||||
_requeue_daily_spend_the_commit_left_behind(commit_task, queue, entity_type, transactions)
|
||||
)
|
||||
raise
|
||||
commit_task.cancel()
|
||||
if transactions:
|
||||
await queue.add_update(transactions)
|
||||
|
|
@ -1841,23 +1926,36 @@ class DBSpendUpdateWriter:
|
|||
The drain is destructive, so a failed commit must push the transactions back for the next tick
|
||||
or their spend is lost permanently.
|
||||
"""
|
||||
await _settle_interrupted_commits(self.interrupted_tag_commits)
|
||||
daily_tag_spend_update_transactions: Final = (
|
||||
await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer()
|
||||
)
|
||||
if not daily_tag_spend_update_transactions:
|
||||
return
|
||||
|
||||
commit_task: Final = asyncio.ensure_future(
|
||||
DBSpendUpdateWriter.update_daily_tag_spend(
|
||||
commit_started: Final = asyncio.Event()
|
||||
commit_task: Final = _start_daily_spend_commit(
|
||||
commit_started,
|
||||
lambda: DBSpendUpdateWriter.update_daily_tag_spend(
|
||||
n_retry_times=n_retry_times,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
daily_spend_transactions=daily_tag_spend_update_transactions,
|
||||
)
|
||||
),
|
||||
)
|
||||
try:
|
||||
await asyncio.shield(commit_task)
|
||||
except BaseException: # noqa: BLE001 # a cancel must restore the drained rows before its rollback returns
|
||||
if commit_started.is_set():
|
||||
_track_interrupted_commit(
|
||||
self.interrupted_tag_commits,
|
||||
_restore_tag_spend_the_commit_left_behind(
|
||||
commit_task,
|
||||
self.redis_update_buffer,
|
||||
daily_tag_spend_update_transactions,
|
||||
),
|
||||
)
|
||||
raise
|
||||
commit_task.cancel()
|
||||
await self.redis_update_buffer.restore_transactions_to_redis(
|
||||
daily_tag_spend_update_transactions=daily_tag_spend_update_transactions,
|
||||
|
|
@ -2382,6 +2480,8 @@ class DBSpendUpdateWriter:
|
|||
sql, params = build_bulk_upsert(table=table, batch=merged_batch)
|
||||
async with _spend_update_tx(prisma_client) as transaction:
|
||||
await transaction.execute_raw(sql, *params)
|
||||
_mark_daily_spend_commit_started()
|
||||
_mark_daily_spend_commit_finished()
|
||||
except Exception as batch_error:
|
||||
if _spend_commit_failure_is_requeue_safe(batch_error):
|
||||
spend_log_error(
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import asyncio
|
||||
from collections.abc import Coroutine
|
||||
from copy import deepcopy
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -57,6 +58,18 @@ class DailySpendUpdateQueue(BaseUpdateQueue):
|
|||
self.update_queue: asyncio.Queue[dict[str, BaseDailySpendTransaction]] = asyncio.Queue(
|
||||
maxsize=LITELLM_ASYNCIO_QUEUE_MAXSIZE
|
||||
)
|
||||
self.interrupted_commits: set[asyncio.Task[None]] = (
|
||||
set()
|
||||
) # mutable-ok: registry of in-flight commit outcomes, entries leave via their done callback
|
||||
|
||||
def track_interrupted_commit(self, settle: Coroutine[object, object, None]) -> None:
|
||||
task: Final = asyncio.ensure_future(settle)
|
||||
self.interrupted_commits.add(task)
|
||||
task.add_done_callback(self.interrupted_commits.discard)
|
||||
|
||||
async def settle_interrupted_commits(self) -> None:
|
||||
while self.interrupted_commits:
|
||||
await asyncio.wait(tuple(self.interrupted_commits))
|
||||
|
||||
async def add_update(self, update: dict[str, BaseDailySpendTransaction]):
|
||||
"""Enqueue an update."""
|
||||
|
|
@ -81,6 +94,7 @@ class DailySpendUpdateQueue(BaseUpdateQueue):
|
|||
self,
|
||||
) -> dict[str, BaseDailySpendTransaction]:
|
||||
"""Get all updates from the queue and return all updates aggregated by daily_transaction_key. Works for both user and team spend updates."""
|
||||
await self.settle_interrupted_commits()
|
||||
updates: Final = await self.flush_all_updates_from_in_memory_queue()
|
||||
if len(updates) > 0:
|
||||
verbose_proxy_logger.info(
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from litellm.proxy._types import (
|
|||
ProxyErrorTypes,
|
||||
ProxyException,
|
||||
)
|
||||
from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
||||
# Bounds the __cause__/__context__ walk in find_database_service_unavailable_error_in_chain.
|
||||
|
|
@ -104,7 +105,7 @@ class PrismaDBExceptionHandler:
|
|||
"""
|
||||
import prisma.engine.errors
|
||||
|
||||
if isinstance(e, DB_CONNECTION_ERROR_TYPES):
|
||||
if isinstance(e, (*DB_CONNECTION_ERROR_TYPES, DBLookupDeadlineExceeded)):
|
||||
return True
|
||||
if isinstance(e, _exception_types(prisma.engine.errors.EngineConnectionError)):
|
||||
return True
|
||||
|
|
|
|||
|
|
@ -23,7 +23,7 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.proxy._types import Litellm_EntityType
|
||||
from litellm.proxy.db.db_lookup_gate import db_lookup_gate
|
||||
from litellm.proxy.db.db_lookup_gate import bounded_db_lookup, db_lookup_gate
|
||||
from litellm.proxy.spend_tracking.spend_counter_batch import read_batched_spend_counter, record_spend_counter_value
|
||||
from litellm.repositories.organization_repository import OrganizationRepository
|
||||
from litellm.repositories.project_repository import ProjectRepository
|
||||
|
|
@ -134,36 +134,9 @@ class SpendCounterReseed:
|
|||
if SpendCounterReseed._is_key_or_team_window_counter(counter_key):
|
||||
return None
|
||||
try:
|
||||
async with db_lookup_gate.current():
|
||||
if counter_key.startswith("spend:key:"):
|
||||
token: Final = counter_key[len("spend:key:") :]
|
||||
row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": token})
|
||||
elif counter_key.startswith("spend:team_member:"):
|
||||
suffix: Final = counter_key[len("spend:team_member:") :]
|
||||
if ":" not in suffix:
|
||||
return None
|
||||
user_id, team_id = suffix.rsplit(":", 1)
|
||||
row = await TeamMembershipRepository(prisma_client).table.find_unique(
|
||||
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}
|
||||
)
|
||||
elif counter_key.startswith("spend:team:"):
|
||||
team_id = counter_key[len("spend:team:") :]
|
||||
row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
|
||||
elif counter_key.startswith("spend:user:"):
|
||||
user_id = counter_key[len("spend:user:") :]
|
||||
row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
|
||||
elif counter_key.startswith(END_USER_COUNTER_PREFIX) or counter_key.startswith("spend:tag:"):
|
||||
return None
|
||||
elif counter_key.startswith("spend:org:"):
|
||||
org_id: Final = counter_key[len("spend:org:") :]
|
||||
row = await OrganizationRepository(prisma_client).table.find_unique(
|
||||
where={"organization_id": org_id}
|
||||
)
|
||||
elif counter_key.startswith("spend:project:"):
|
||||
project_id: Final = counter_key[len("spend:project:") :]
|
||||
row = await ProjectRepository(prisma_client).table.find_unique(where={"project_id": project_id})
|
||||
else:
|
||||
return None
|
||||
row: Final = await bounded_db_lookup(
|
||||
SpendCounterReseed._counter_row(prisma_client, counter_key), name="spend_counter"
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception("SpendCounterReseed.from_db: failed for %s", counter_key)
|
||||
return None
|
||||
|
|
@ -171,13 +144,47 @@ class SpendCounterReseed:
|
|||
return None
|
||||
return float(getattr(row, "spend", 0.0) or 0.0)
|
||||
|
||||
@staticmethod
|
||||
async def _counter_row(prisma_client: "PrismaClient", counter_key: str) -> object | None:
|
||||
async with db_lookup_gate.current():
|
||||
if counter_key.startswith("spend:key:"):
|
||||
token: Final = counter_key[len("spend:key:") :]
|
||||
return await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": token})
|
||||
if counter_key.startswith("spend:team_member:"):
|
||||
suffix: Final = counter_key[len("spend:team_member:") :]
|
||||
if ":" not in suffix:
|
||||
return None
|
||||
user_id, team_id = suffix.rsplit(":", 1)
|
||||
return await TeamMembershipRepository(prisma_client).table.find_unique(
|
||||
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}
|
||||
)
|
||||
if counter_key.startswith("spend:team:"):
|
||||
return await TeamRepository(prisma_client).table.find_unique(
|
||||
where={"team_id": counter_key[len("spend:team:") :]}
|
||||
)
|
||||
if counter_key.startswith("spend:user:"):
|
||||
return await UserRepository(prisma_client).table.find_unique(
|
||||
where={"user_id": counter_key[len("spend:user:") :]}
|
||||
)
|
||||
if counter_key.startswith("spend:org:"):
|
||||
return await OrganizationRepository(prisma_client).table.find_unique(
|
||||
where={"organization_id": counter_key[len("spend:org:") :]}
|
||||
)
|
||||
if counter_key.startswith("spend:project:"):
|
||||
return await ProjectRepository(prisma_client).table.find_unique(
|
||||
where={"project_id": counter_key[len("spend:project:") :]}
|
||||
)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def end_user_from_db(prisma_client: Optional["PrismaClient"], counter_key: str) -> float | None:
|
||||
if prisma_client is None or not counter_key.startswith(END_USER_COUNTER_PREFIX):
|
||||
return None
|
||||
where: Final[LiteLLM_EndUserTableWhereUniqueInput] = {"user_id": counter_key[len(END_USER_COUNTER_PREFIX) :]}
|
||||
try:
|
||||
row: Final = await EndUserRepository(prisma_client).table.find_unique(where=where)
|
||||
row: Final = await bounded_db_lookup(
|
||||
EndUserRepository(prisma_client).table.find_unique(where=where), name="end_user_spend"
|
||||
)
|
||||
except Exception: # noqa: BLE001 # a failed floor read falls back to the cached spend, like from_db
|
||||
verbose_proxy_logger.exception("SpendCounterReseed.end_user_from_db: failed for %s", counter_key)
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -42,6 +42,7 @@ from litellm.proxy._types import UserAPIKeyAuth
|
|||
from litellm.proxy.guardrails.anthropic_sse import (
|
||||
anthropic_sse_chunks_from_response,
|
||||
assemble_anthropic_sse_stream,
|
||||
is_anthropic_sse_stream,
|
||||
model_response_text,
|
||||
)
|
||||
from litellm.types.guardrails import (
|
||||
|
|
@ -93,6 +94,42 @@ def _json_escaped_len(text: str) -> int:
|
|||
return len(json.dumps(text).encode("utf-8")) - 2 # strip the surrounding quotes
|
||||
|
||||
|
||||
_MAX_FIRST_SSE_FRAME_BYTES: Final = 64 * 1024
|
||||
|
||||
|
||||
def _holds_complete_sse_frame(raw: bytes) -> bool:
|
||||
"""Whether ``raw`` holds one blank-line terminated SSE event, or is too large to keep joining."""
|
||||
return b"\n\n" in raw or b"\r\n\r\n" in raw or len(raw) >= _MAX_FIRST_SSE_FRAME_BYTES
|
||||
|
||||
|
||||
async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGenerator[object, None]:
|
||||
"""
|
||||
Join leading raw ``bytes`` chunks until they hold one complete SSE event, so
|
||||
the stream shape is decided on a whole frame rather than a transport fragment.
|
||||
Everything after that first frame is forwarded untouched.
|
||||
"""
|
||||
pending = b""
|
||||
try:
|
||||
async for chunk in stream:
|
||||
if not isinstance(chunk, bytes):
|
||||
yield chunk
|
||||
continue
|
||||
pending += chunk
|
||||
if _holds_complete_sse_frame(pending):
|
||||
break
|
||||
else:
|
||||
if pending:
|
||||
yield pending
|
||||
return
|
||||
except Exception:
|
||||
if pending:
|
||||
yield pending
|
||||
raise
|
||||
yield pending
|
||||
async for chunk in stream:
|
||||
yield chunk
|
||||
|
||||
|
||||
class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
||||
user_api_key_cache = None
|
||||
ad_hoc_recognizers: list[str] | None = None
|
||||
|
|
@ -1356,7 +1393,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
all_chunks: list[ModelResponseStream] = []
|
||||
passthrough_due_to_unknown_stream_shape = False
|
||||
try:
|
||||
stream: Final = response.__aiter__()
|
||||
stream: Final = _coalesce_first_sse_frame(response.__aiter__())
|
||||
async for chunk in stream:
|
||||
if isinstance(chunk, ModelResponseStream):
|
||||
if passthrough_due_to_unknown_stream_shape:
|
||||
|
|
@ -1364,7 +1401,15 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
else:
|
||||
all_chunks.append(chunk)
|
||||
elif isinstance(chunk, bytes):
|
||||
if passthrough_due_to_unknown_stream_shape or all_chunks:
|
||||
first_frame_is_anthropic = (
|
||||
not passthrough_due_to_unknown_stream_shape
|
||||
and not all_chunks
|
||||
and is_anthropic_sse_stream((chunk,))
|
||||
)
|
||||
if not first_frame_is_anthropic:
|
||||
passthrough_due_to_unknown_stream_shape = (
|
||||
passthrough_due_to_unknown_stream_shape or not all_chunks
|
||||
)
|
||||
yield chunk
|
||||
continue
|
||||
for masked_chunk in await self._mask_anthropic_sse_stream(chunk, stream, request_data):
|
||||
|
|
@ -1387,8 +1432,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
yield chunk
|
||||
if passthrough_due_to_unknown_stream_shape:
|
||||
verbose_proxy_logger.warning(
|
||||
"Presidio apply_to_output: streaming response contained unknown event objects "
|
||||
"(e.g. /v1/responses events). Output PII masking was skipped for this response."
|
||||
"Presidio apply_to_output: streaming response was not a parsed chat completion stream "
|
||||
"(raw non-Anthropic SSE passthrough or /v1/responses events). "
|
||||
"Output PII masking was skipped for this response."
|
||||
)
|
||||
return
|
||||
if not all_chunks:
|
||||
|
|
|
|||
|
|
@ -120,7 +120,7 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) ->
|
|||
_OPTIONAL_PresidioPIIMasking,
|
||||
)
|
||||
|
||||
explicit_filter_scope: Final = getattr(litellm_params, "presidio_filter_scope", None)
|
||||
explicit_filter_scope: Final = litellm_params.presidio_filter_scope
|
||||
filter_scope: Final = explicit_filter_scope or ("input" if _is_mcp_only_mode(litellm_params.mode) else "both")
|
||||
run_input: Final = filter_scope in ("input", "both")
|
||||
run_output: Final = filter_scope in ("output", "both")
|
||||
|
|
|
|||
|
|
@ -16,7 +16,7 @@ from typing_extensions import ReadOnly
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm.constants import HEALTH_CHECK_TIMEOUT_SECONDS
|
||||
from litellm.constants import HEALTH_CHECK_TIMEOUT_SECONDS, PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS
|
||||
from litellm.integrations.SlackAlerting.ms_teams import (
|
||||
MS_TEAMS_ALERT_HEADERS,
|
||||
build_ms_teams_payload,
|
||||
|
|
@ -44,6 +44,7 @@ from litellm.proxy.auth.auth_utils import (
|
|||
)
|
||||
from litellm.proxy.auth.model_checks import get_key_models
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.db.health_check_latest import (
|
||||
LatestHealthCheckRow,
|
||||
|
|
@ -1723,7 +1724,7 @@ async def _get_health_readiness_details(
|
|||
|
||||
# check DB
|
||||
if prisma_client is not None: # if db passed in, check if it's connected
|
||||
db_health_status: Final = await _db_health_readiness_check()
|
||||
db_status: Final = _readiness_db_status(await _db_health_readiness_check())
|
||||
# A configured DB that is not reachable means the worker cannot
|
||||
# serve requests that depend on persisted state (keys, budgets,
|
||||
# spend logs). Return 503 so orchestrators take this pod out of
|
||||
|
|
@ -1733,13 +1734,13 @@ async def _get_health_readiness_details(
|
|||
# report the DB state through the body instead.
|
||||
if (
|
||||
response is not None
|
||||
and db_health_status["status"] != "connected"
|
||||
and db_status != "connected"
|
||||
and not PrismaDBExceptionHandler.should_allow_request_on_db_unavailable()
|
||||
):
|
||||
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
|
||||
return {
|
||||
"status": "healthy",
|
||||
"db": db_health_status["status"],
|
||||
"db": db_status,
|
||||
"cache": cache_type,
|
||||
"litellm_version": version,
|
||||
"success_callbacks": success_callback_names,
|
||||
|
|
@ -1816,24 +1817,32 @@ def _authorize_drain_request(request: Request) -> None:
|
|||
)
|
||||
|
||||
|
||||
def _readiness_db_status(db_health_status: DBHealthCache) -> str:
|
||||
"""A pod whose pre-request lookups hit their deadline inside the stall window
|
||||
reports "stalled" even though the ping succeeds: the ping is a fresh
|
||||
connection, the stalled lookups are the ones requests actually wait on."""
|
||||
if db_health_status["status"] != "connected":
|
||||
return db_health_status["status"]
|
||||
if db_lookup_stall_tracker.stalled_within(PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS):
|
||||
return "stalled"
|
||||
return "connected"
|
||||
|
||||
|
||||
async def _resolve_public_readiness_db(response: Response) -> str:
|
||||
"""
|
||||
Return the db status string for the public probe and flip the response to
|
||||
503 when a configured DB is unreachable. Mirrors the legacy values:
|
||||
"Not connected" (no DB configured), "connected", "disconnected".
|
||||
503 when a configured DB is unreachable or stalled. Mirrors the legacy values:
|
||||
"Not connected" (no DB configured), "connected", "disconnected", plus "stalled".
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
return "Not connected"
|
||||
|
||||
db_health_status: Final = await _db_health_readiness_check()
|
||||
if (
|
||||
db_health_status["status"] != "connected"
|
||||
and not PrismaDBExceptionHandler.should_allow_request_on_db_unavailable()
|
||||
):
|
||||
db_status: Final = _readiness_db_status(await _db_health_readiness_check())
|
||||
if db_status != "connected" and not PrismaDBExceptionHandler.should_allow_request_on_db_unavailable():
|
||||
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
|
||||
return db_health_status["status"]
|
||||
return db_status
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
log_db_metrics,
|
||||
)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded
|
||||
from litellm.proxy.db.db_spend_update_writer import (
|
||||
DBSpendUpdateWriter,
|
||||
debitable_model_access_groups,
|
||||
|
|
@ -189,8 +190,8 @@ class _ProxyDBLogger(CustomLogger):
|
|||
)
|
||||
_metadata["error_information"] = _error_information
|
||||
|
||||
_metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(
|
||||
metadata=_metadata,
|
||||
_metadata = await _ProxyDBLogger._enrich_failure_metadata_unless_db_stalled(
|
||||
metadata=_metadata, original_exception=original_exception
|
||||
)
|
||||
|
||||
existing_metadata: Final[dict] = request_data.get("metadata", None) or {}
|
||||
|
|
@ -475,6 +476,12 @@ class _ProxyDBLogger(CustomLogger):
|
|||
|
||||
spend_log_error("Error in tracking cost callback - %s", str(e), exc=e)
|
||||
|
||||
@staticmethod
|
||||
async def _enrich_failure_metadata_unless_db_stalled(metadata: dict, original_exception: Exception) -> dict:
|
||||
if isinstance(original_exception, DBLookupDeadlineExceeded):
|
||||
return metadata
|
||||
return await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata)
|
||||
|
||||
@staticmethod
|
||||
async def _enrich_failure_metadata_with_key_info(metadata: dict, resolve_missing_key_identity: bool = True) -> dict:
|
||||
"""
|
||||
|
|
@ -773,7 +780,7 @@ async def _reconcile_budget_reservation_before_db_update(
|
|||
"Failed to invalidate budget reservation counters after pre-persist reconcile failed"
|
||||
)
|
||||
finally:
|
||||
budget_reservation["finalized"] = True # rebind-ok: the counter update reads the stamp off the shared dict
|
||||
budget_reservation["finalized"] = True # rebind-ok: stamps the caller's shared dict for the counter update
|
||||
|
||||
|
||||
async def _release_budget_reservation(budget_reservation: dict | None) -> None:
|
||||
|
|
|
|||
|
|
@ -173,6 +173,34 @@ def _check_passthrough_routes_caller_permission(
|
|||
)
|
||||
|
||||
|
||||
def _check_disable_global_guardrails_caller_permission(
|
||||
disable_global_guardrails: bool | None,
|
||||
metadata: Mapping[str, object] | None,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
*,
|
||||
entity: str = "key",
|
||||
existing_metadata: Mapping[str, object] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Only proxy admins may opt a key or team out of default-on guardrails, whether the
|
||||
flag is top-level or under `metadata`. Re-sending a flag that is already stored is
|
||||
not an opt-out, so non-admin edits of an already exempted object still go through.
|
||||
"""
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
|
||||
return
|
||||
requested: Final = bool(disable_global_guardrails) or (
|
||||
metadata is not None and bool(metadata.get("disable_global_guardrails"))
|
||||
)
|
||||
if not requested:
|
||||
return
|
||||
if existing_metadata is not None and existing_metadata.get("disable_global_guardrails") is True:
|
||||
return
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": f"Only proxy admins can set `disable_global_guardrails` on a {entity}."},
|
||||
)
|
||||
|
||||
|
||||
def _is_user_team_admin(user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable) -> bool:
|
||||
for member in team_obj.members_with_roles:
|
||||
if (member.user_id is not None and member.user_id == user_api_key_dict.user_id) and member.role == "admin":
|
||||
|
|
@ -500,7 +528,7 @@ def _prisma_value(value: object) -> object:
|
|||
return list(value) if isinstance(value, tuple) else value
|
||||
|
||||
|
||||
def member_budget_patch(source: BaseModel) -> dict[str, Any]:
|
||||
def member_budget_patch(source: BaseModel) -> Mapping[str, object]:
|
||||
"""Map the per-member limit fields a request actually set to their budget-table
|
||||
columns (merge-patch: a sent value updates, an explicit null clears, an absent
|
||||
field is left untouched)."""
|
||||
|
|
@ -533,7 +561,7 @@ async def _upsert_budget_and_membership(
|
|||
user_id: str,
|
||||
existing_budget_id: str | None,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
budget_patch: dict[str, Any],
|
||||
budget_patch: Mapping[str, object],
|
||||
team_default_budget_id: str | None = None,
|
||||
shared_budget_ids: frozenset[str] | None = None,
|
||||
):
|
||||
|
|
@ -596,9 +624,9 @@ async def _upsert_budget_and_membership(
|
|||
if is_shared_default and not temp_only
|
||||
else None
|
||||
)
|
||||
source: Final[Mapping[str, Any]] = source_row.model_dump() if source_row is not None else MappingProxyType({})
|
||||
source: Final[Mapping[str, object]] = source_row.model_dump() if source_row is not None else MappingProxyType({})
|
||||
|
||||
create_data: Final[dict[str, Any]] = { # mutable-ok: Prisma create payloads are dict-shaped
|
||||
create_data: Final[dict[str, object]] = { # mutable-ok: Prisma create payloads are dict-shaped
|
||||
"created_by": user_api_key_dict.user_id or "",
|
||||
"updated_by": user_api_key_dict.user_id or "",
|
||||
**MappingProxyType(
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from typing_extensions import ReadOnly, TypedDict
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
|
|
@ -87,11 +88,13 @@ from litellm.repositories.verification_token_repository import (
|
|||
VerificationTokenRepository,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
||||
DailySpendMetadata,
|
||||
SpendAnalyticsPaginatedResponse,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
|
||||
BulkUpdateUserRequest,
|
||||
BulkUpdateUserResponse,
|
||||
KeyActivitySearchWhere,
|
||||
UserListResponse,
|
||||
UserSearchWhere,
|
||||
UserUpdateResult,
|
||||
|
|
@ -2991,6 +2994,27 @@ async def get_user_daily_activity(
|
|||
)
|
||||
|
||||
|
||||
def _resolve_user_daily_activity_entity_id(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
user_id: str | None,
|
||||
) -> str | None:
|
||||
is_admin: Final = _user_has_admin_view(user_api_key_dict)
|
||||
|
||||
if is_admin:
|
||||
return user_id
|
||||
|
||||
caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict)
|
||||
effective_user_id: Final = user_id if user_id is not None else caller_user_id
|
||||
if effective_user_id != caller_user_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={ # mutable-ok: FastAPI detail payload shape
|
||||
"error": "Non-admin users can only view their own spend data."
|
||||
},
|
||||
)
|
||||
return effective_user_id
|
||||
|
||||
|
||||
@router.get(
|
||||
"/user/daily/activity/aggregated",
|
||||
tags=["Budget & Spend Tracking", "Internal User management"],
|
||||
|
|
@ -3057,20 +3081,7 @@ async def get_user_daily_activity_aggregated(
|
|||
)
|
||||
|
||||
try:
|
||||
is_admin: Final = _user_has_admin_view(user_api_key_dict)
|
||||
|
||||
if is_admin:
|
||||
entity_id = user_id # None means global view, otherwise filter by user
|
||||
else:
|
||||
caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict)
|
||||
if user_id is None:
|
||||
user_id = caller_user_id
|
||||
if user_id != caller_user_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={"error": "Non-admin users can only view their own spend data."},
|
||||
)
|
||||
entity_id = user_id
|
||||
entity_id: Final = _resolve_user_daily_activity_entity_id(user_api_key_dict, user_id)
|
||||
|
||||
return await get_daily_activity_aggregated(
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -3094,3 +3105,117 @@ async def get_user_daily_activity_aggregated(
|
|||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail={"error": f"Failed to fetch analytics: {e}"},
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/user/daily/activity/aggregated/search",
|
||||
tags=["Budget & Spend Tracking", "Internal User management"], # mutable-ok: FastAPI route tags shape
|
||||
dependencies=[Depends(user_api_key_auth)], # mutable-ok: FastAPI route dependencies shape
|
||||
response_model=SpendAnalyticsPaginatedResponse,
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def search_user_daily_activity_keys(
|
||||
search: str = fastapi.Query(
|
||||
...,
|
||||
min_length=1,
|
||||
description="Matches keys whose hash equals the value, or whose key alias or user ID contains it (case-insensitive)",
|
||||
),
|
||||
start_date: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Start date in YYYY-MM-DD format",
|
||||
),
|
||||
end_date: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="End date in YYYY-MM-DD format",
|
||||
),
|
||||
user_id: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Filter by specific user ID. Admins can filter by any user or omit for global view. Non-admins must provide their own user_id.",
|
||||
),
|
||||
timezone: int | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Timezone offset in minutes from UTC (e.g., 480 for PST). "
|
||||
"Matches JavaScript's Date.getTimezoneOffset() convention.",
|
||||
),
|
||||
include_current_utc_day: bool = fastapi.Query(
|
||||
default=False,
|
||||
description="When the range ends on the caller's current local day, extend it to "
|
||||
"today's UTC bucket so spend written after the caller's local midnight (in UTC "
|
||||
"terms) is included. Requires the timezone parameter. Historical ranges are "
|
||||
"never extended.",
|
||||
),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI dependency injection
|
||||
) -> SpendAnalyticsPaginatedResponse:
|
||||
"""
|
||||
Search verification tokens by exact token hash or by a case-insensitive substring of
|
||||
the key alias or owning user ID, then return the aggregated daily activity for the
|
||||
matches. Lets the Usage page surface keys that fell outside the top-spend subset
|
||||
the aggregated endpoint loads.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={ # mutable-ok: FastAPI detail payload shape
|
||||
"error": CommonProxyErrors.db_not_connected_error.value
|
||||
},
|
||||
)
|
||||
|
||||
if start_date is None or end_date is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail={"error": "Please provide start_date and end_date"}, # mutable-ok: FastAPI detail payload shape
|
||||
)
|
||||
|
||||
try:
|
||||
entity_id: Final = _resolve_user_daily_activity_entity_id(user_api_key_dict, user_id)
|
||||
|
||||
search_or: Final = (
|
||||
{"token": search}, # mutable-ok: prisma serializes where clauses, keep plain dicts
|
||||
{"key_alias": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
|
||||
{"user_id": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
|
||||
)
|
||||
where: Final[KeyActivitySearchWhere] = (
|
||||
{"OR": search_or} # mutable-ok: prisma where clause root
|
||||
if entity_id is None
|
||||
else {"user_id": entity_id, "OR": search_or} # mutable-ok: prisma where clause root
|
||||
)
|
||||
matched_keys: Final = await VerificationTokenRepository(prisma_client).table.find_many(
|
||||
where=where,
|
||||
take=USAGE_TOP_API_KEYS_LIMIT,
|
||||
order={"spend": "desc"}, # mutable-ok: prisma serializes order, keep it a plain dict
|
||||
)
|
||||
tokens: Final = [key.token for key in matched_keys] # mutable-ok: api_key filter union expects a list
|
||||
|
||||
if not tokens:
|
||||
return SpendAnalyticsPaginatedResponse(
|
||||
results=[], # mutable-ok: response model field shape
|
||||
metadata=DailySpendMetadata(
|
||||
api_key_limit=USAGE_TOP_API_KEYS_LIMIT,
|
||||
total_api_keys=0,
|
||||
),
|
||||
)
|
||||
|
||||
return await get_daily_activity_aggregated(
|
||||
prisma_client=prisma_client,
|
||||
table_name="litellm_dailyuserspend",
|
||||
entity_id_field="user_id",
|
||||
entity_id=entity_id,
|
||||
entity_metadata_field=None,
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
model=None,
|
||||
api_key=tokens,
|
||||
timezone_offset_minutes=timezone,
|
||||
include_current_utc_day=include_current_utc_day,
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("/user/daily/activity/aggregated/search: Exception occured - %s", e)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail={"error": f"Failed to fetch analytics: {e}"}, # mutable-ok: FastAPI detail payload shape
|
||||
)
|
||||
|
|
|
|||
|
|
@ -85,6 +85,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|||
from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks
|
||||
from litellm.proxy.hooks.model_max_budget_limiter import build_model_max_budget_usage
|
||||
from litellm.proxy.management_endpoints.common_utils import (
|
||||
_check_disable_global_guardrails_caller_permission,
|
||||
_check_passthrough_routes_caller_permission,
|
||||
_is_user_org_admin_for_team,
|
||||
_is_user_team_admin,
|
||||
|
|
@ -452,7 +453,7 @@ def _regenerate_request_as_update_request(key: str, data: RegenerateKeyRequest)
|
|||
)
|
||||
if not changed_fields:
|
||||
return None
|
||||
return UpdateKeyRequest(key=key, **changed_fields)
|
||||
return UpdateKeyRequest.model_validate(MappingProxyType({"key": key, **changed_fields}))
|
||||
|
||||
|
||||
class _LegacyDumpable(Protocol):
|
||||
|
|
@ -1224,6 +1225,7 @@ async def _common_key_generation_helper(
|
|||
# default_key_generate_params injected.
|
||||
_requested_max_budget: Final = data.max_budget
|
||||
_requested_team_id: Final = data.team_id
|
||||
_requested_metadata: Final = data.metadata # pyright: ignore[reportUnknownMemberType] # request models declare `metadata` as bare dict
|
||||
|
||||
# check if user set default key/generate params on config.yaml
|
||||
if litellm.default_key_generate_params is not None:
|
||||
|
|
@ -1311,6 +1313,11 @@ async def _common_key_generation_helper(
|
|||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
_check_disable_global_guardrails_caller_permission(
|
||||
data.disable_global_guardrails,
|
||||
_requested_metadata,
|
||||
user_api_key_dict,
|
||||
)
|
||||
|
||||
# APPLY ENTERPRISE KEY MANAGEMENT PARAMS
|
||||
try:
|
||||
|
|
@ -1966,7 +1973,7 @@ async def generate_key_fn(
|
|||
- metadata: Optional[dict] - Metadata for key, store information for key. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" }
|
||||
- guardrails: Optional[List[str]] - List of active guardrails for the key
|
||||
- policies: Optional[List[str]] - List of policy names to apply to the key. Policies define guardrails, conditions, and inheritance rules.
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key.
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. Proxy admin only.
|
||||
- throttle_on_budget_exceeded: Optional[bool] - When the key exceeds its max_budget, throttle its tpm/rpm to the global budget_exceeded_throttle_percentage instead of blocking the key entirely.
|
||||
- enable_prompt_caching: Optional[bool] - Auto-inject prompt caching breakpoints (Anthropic cache_control markers) on requests made with this key. Supported Claude models on Anthropic, Bedrock, Vertex AI, and Azure AI only.
|
||||
- permissions: Optional[dict] - key-specific permissions. Currently just used for turning off pii masking (if connected). Example - {"pii": false}
|
||||
|
|
@ -2729,6 +2736,13 @@ async def _process_single_key_update(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
_check_disable_global_guardrails_caller_permission(
|
||||
update_key_request.disable_global_guardrails,
|
||||
update_key_request.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict
|
||||
user_api_key_dict,
|
||||
existing_metadata=existing_key_row.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_VerificationToken.metadata is a bare dict
|
||||
)
|
||||
|
||||
enforce_batch_enqueued_token_limit_is_admin_only(
|
||||
data=update_key_request,
|
||||
existing_metadata=existing_key_row.metadata,
|
||||
|
|
@ -3025,6 +3039,12 @@ async def _validate_update_key_data(
|
|||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
_check_disable_global_guardrails_caller_permission(
|
||||
data.disable_global_guardrails,
|
||||
data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict
|
||||
user_api_key_dict,
|
||||
existing_metadata=existing_key_row.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_VerificationToken.metadata is a bare dict
|
||||
)
|
||||
|
||||
_validate_caller_can_change_key_ownership(
|
||||
data=data,
|
||||
|
|
@ -3328,7 +3348,7 @@ async def update_key_fn(
|
|||
- send_invite_email: Optional[bool] - Send invite email to user_id
|
||||
- guardrails: Optional[List[str]] - List of active guardrails for the key
|
||||
- policies: Optional[List[str]] - List of policy names to apply to the key. Policies define guardrails, conditions, and inheritance rules.
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key.
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. Proxy admin only.
|
||||
- throttle_on_budget_exceeded: Optional[bool] - When the key exceeds its max_budget, throttle its tpm/rpm to the global budget_exceeded_throttle_percentage instead of blocking the key entirely.
|
||||
- enable_prompt_caching: Optional[bool] - Auto-inject prompt caching breakpoints (Anthropic cache_control markers) on requests made with this key. Supported Claude models on Anthropic, Bedrock, Vertex AI, and Azure AI only.
|
||||
- prompts: Optional[List[str]] - List of prompts that the key is allowed to use.
|
||||
|
|
@ -5622,6 +5642,21 @@ async def _execute_virtual_key_regeneration(
|
|||
return response
|
||||
|
||||
|
||||
def _check_regenerate_guardrail_opt_out(
|
||||
data: RegenerateKeyRequest | None,
|
||||
existing_metadata: Mapping[str, object] | None,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> None:
|
||||
if data is None:
|
||||
return
|
||||
_check_disable_global_guardrails_caller_permission(
|
||||
data.disable_global_guardrails,
|
||||
data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict
|
||||
user_api_key_dict,
|
||||
existing_metadata=existing_metadata,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/key/{key:path}/regenerate",
|
||||
tags=["key management"],
|
||||
|
|
@ -5805,6 +5840,12 @@ async def regenerate_key_fn(
|
|||
detail={"error": f"Key {key} not found."},
|
||||
)
|
||||
|
||||
_check_regenerate_guardrail_opt_out(
|
||||
data,
|
||||
_key_in_db.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_VerificationToken.metadata is a bare dict
|
||||
user_api_key_dict,
|
||||
)
|
||||
|
||||
# check if user has permission to regenerate key
|
||||
await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ from typing_extensions import ReadOnly, TypedDict, assert_never
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -127,6 +128,7 @@ from litellm.proxy.management_endpoints.common_daily_activity import (
|
|||
get_daily_activity_aggregated,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.common_utils import (
|
||||
_check_disable_global_guardrails_caller_permission,
|
||||
_check_passthrough_routes_caller_permission,
|
||||
_is_user_org_admin_for_team,
|
||||
_is_user_team_admin,
|
||||
|
|
@ -195,6 +197,7 @@ from litellm.repositories.verification_token_repository import (
|
|||
from litellm.router import Router
|
||||
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
||||
DailySpendMetadata,
|
||||
SpendAnalyticsPaginatedResponse,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.team_endpoints import (
|
||||
|
|
@ -203,7 +206,9 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
|
|||
BulkUpdateTeamMemberPermissionsRequest,
|
||||
BulkUpdateTeamMemberPermissionsResponse,
|
||||
GetTeamMemberPermissionsResponse,
|
||||
TeamIdSearchFilter,
|
||||
TeamIdSearchMatch,
|
||||
TeamKeyActivitySearchWhere,
|
||||
TeamListItem,
|
||||
TeamListResponse,
|
||||
TeamMemberAddResult,
|
||||
|
|
@ -1416,7 +1421,7 @@ async def new_team(
|
|||
- model_max_budget: Optional[dict] - Per-model max budget every key on the team inherits unless the key sets its own for that model. Example: {"gpt-4o": {"max_budget": 10, "budget_duration": "1d"}}
|
||||
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
|
||||
- policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies)
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key.
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the team. Proxy admin only.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission.
|
||||
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
|
||||
- team_member_budget_duration: Optional[str] - The duration of the budget for the team member. Doc [here](https://docs.litellm.ai/docs/proxy/team_budgets)
|
||||
|
|
@ -1638,6 +1643,12 @@ async def new_team(
|
|||
data.members_with_roles.append(Member(role="admin", user_id=user_api_key_dict.user_id))
|
||||
|
||||
_check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team")
|
||||
_check_disable_global_guardrails_caller_permission(
|
||||
data.disable_global_guardrails,
|
||||
data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict
|
||||
user_api_key_dict,
|
||||
entity="team",
|
||||
)
|
||||
|
||||
if isinstance(data.metadata, dict):
|
||||
TeamMemberBudgetHandler.strip_system_managed_metadata_keys(data.metadata)
|
||||
|
|
@ -2172,7 +2183,7 @@ async def update_team(
|
|||
- model_max_budget: Optional[dict] - Per-model max budget every key on the team inherits unless the key sets its own for that model. Example: {"gpt-4o": {"max_budget": 10, "budget_duration": "1d"}}
|
||||
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
|
||||
- policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies)
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key.
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the team. Proxy admin only.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission.
|
||||
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
|
||||
- team_member_budget_duration: Optional[str] - The duration of the budget for the team member. Doc [here](https://docs.litellm.ai/docs/proxy/team_budgets)
|
||||
|
|
@ -2313,6 +2324,13 @@ async def update_team(
|
|||
)
|
||||
|
||||
_check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team")
|
||||
_check_disable_global_guardrails_caller_permission(
|
||||
data.disable_global_guardrails,
|
||||
data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict
|
||||
user_api_key_dict,
|
||||
entity="team",
|
||||
existing_metadata=_existing_team_metadata if isinstance(_existing_team_metadata, dict) else None, # pyright: ignore[reportUnknownArgumentType] # existing_team_row.metadata is a bare dict
|
||||
)
|
||||
|
||||
if data.soft_budget is not None:
|
||||
max_budget_to_check = data.max_budget if data.max_budget is not None else existing_team_row.max_budget
|
||||
|
|
@ -6791,6 +6809,111 @@ async def get_team_daily_activity_aggregated(
|
|||
)
|
||||
|
||||
|
||||
def _team_key_search_where(*, search: str, scope: _TeamDailyActivityScope) -> TeamKeyActivitySearchWhere:
|
||||
"""Caller scoping lives inside the same Prisma where as the search term so `take`
|
||||
never trims visible matches in favour of keys the caller is not allowed to see."""
|
||||
search_or: Final = (
|
||||
{"token": search}, # mutable-ok: prisma where clause leaf
|
||||
{"key_alias": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
|
||||
{"user_id": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
|
||||
)
|
||||
own_keys: Final = tuple(scope.api_key_filter) if isinstance(scope.api_key_filter, list) else None
|
||||
team_filter: Final[TeamIdSearchFilter | None] = (
|
||||
{ # mutable-ok: prisma where clause leaf
|
||||
"in": tuple(scope.team_ids),
|
||||
"notIn": tuple(scope.exclude_team_ids),
|
||||
}
|
||||
if scope.team_ids is not None and scope.exclude_team_ids is not None
|
||||
else {"in": tuple(scope.team_ids)} # mutable-ok: prisma where clause leaf
|
||||
if scope.team_ids is not None
|
||||
else {"notIn": tuple(scope.exclude_team_ids)} # mutable-ok: prisma where clause leaf
|
||||
if scope.exclude_team_ids is not None
|
||||
else None
|
||||
)
|
||||
if team_filter is None and own_keys is None:
|
||||
return {"OR": search_or} # mutable-ok: prisma where clause root
|
||||
if team_filter is None and own_keys is not None:
|
||||
return {"token": {"in": own_keys}, "OR": search_or} # mutable-ok: prisma where clause root
|
||||
if team_filter is not None and own_keys is None:
|
||||
return {"team_id": team_filter, "OR": search_or} # mutable-ok: prisma where clause root
|
||||
assert team_filter is not None and own_keys is not None
|
||||
return { # mutable-ok: prisma where clause root
|
||||
"team_id": team_filter,
|
||||
"token": {"in": own_keys}, # mutable-ok: prisma where clause leaf
|
||||
"OR": search_or,
|
||||
}
|
||||
|
||||
|
||||
@router.get(
|
||||
"/team/daily/activity/aggregated/search",
|
||||
response_model=SpendAnalyticsPaginatedResponse,
|
||||
tags=["team management"], # mutable-ok: FastAPI route tags shape
|
||||
)
|
||||
async def search_team_daily_activity_keys(
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
search: str = fastapi.Query(
|
||||
...,
|
||||
min_length=1,
|
||||
description="Exact token hash, or a case-insensitive substring of the key alias or owning user id",
|
||||
),
|
||||
team_ids: str | None = None,
|
||||
start_date: str | None = None,
|
||||
end_date: str | None = None,
|
||||
exclude_team_ids: str | None = None,
|
||||
timezone: int | None = None,
|
||||
) -> SpendAnalyticsPaginatedResponse:
|
||||
"""Aggregated daily team activity for the keys matching `search`, across every key the caller may
|
||||
see rather than only the top USAGE_TOP_API_KEYS_LIMIT keys by spend."""
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise _daily_activity_error(status_code=500, message=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
range_error: Final = _aggregated_date_range_error(start_date, end_date)
|
||||
if range_error is not None:
|
||||
raise _daily_activity_error(status_code=400, message=range_error)
|
||||
|
||||
scope: Final = await _resolve_team_daily_activity_scope(
|
||||
team_ids=team_ids,
|
||||
exclude_team_ids=exclude_team_ids,
|
||||
api_key=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
matched_keys: Final = await _tokens_db(prisma_client).find_many(
|
||||
where=_team_key_search_where(search=search, scope=scope),
|
||||
take=USAGE_TOP_API_KEYS_LIMIT,
|
||||
order={"spend": "desc"}, # mutable-ok: prisma serializes order, keep it a plain dict
|
||||
)
|
||||
tokens: Final = [key.token for key in matched_keys] # mutable-ok: get_daily_activity_aggregated takes list[str]
|
||||
if not tokens:
|
||||
return SpendAnalyticsPaginatedResponse(
|
||||
results=[], # mutable-ok: response model field shape
|
||||
metadata=DailySpendMetadata(api_key_limit=USAGE_TOP_API_KEYS_LIMIT, total_api_keys=0),
|
||||
)
|
||||
|
||||
return await get_daily_activity_aggregated(
|
||||
prisma_client=prisma_client,
|
||||
table_name="litellm_dailyteamspend",
|
||||
entity_id_field="team_id",
|
||||
entity_id=scope.team_ids,
|
||||
entity_metadata_field=scope.team_alias_metadata,
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
model=None,
|
||||
api_key=tokens,
|
||||
exclude_entity_ids=scope.exclude_team_ids,
|
||||
timezone_offset_minutes=timezone,
|
||||
include_entity_breakdown=True,
|
||||
)
|
||||
|
||||
|
||||
def _team_user_spend_sql(*, team_count: int, restrict_to_user: bool) -> str:
|
||||
team_placeholders: Final = ", ".join(f"${i}" for i in range(3, 3 + team_count))
|
||||
user_clause: Final = f' AND sl."user" = ${3 + team_count}' if restrict_to_user else ""
|
||||
|
|
|
|||
|
|
@ -348,7 +348,7 @@ async def _prepare_user(user: _PendingUser, prisma_client: PrismaClient) -> _Pre
|
|||
data: Final = {**dumped, "user_id": user.user_id} # mutable-ok: /user/new defaults helper mutates in place
|
||||
data_json: Final = _JSON_OBJECT.validate_python(_update_internal_new_user_params(data, user.request))
|
||||
with_permission: Final = _JSON_OBJECT.validate_python(
|
||||
await _set_object_permission(data_json=data_json, prisma_client=prisma_client) # pyright: ignore[reportUnknownArgumentType] # validated by the adapter
|
||||
await _set_object_permission(data_json=data_json, prisma_client=prisma_client)
|
||||
)
|
||||
return _PreparedUser(user, _USER_ROW.validate_python(with_permission))
|
||||
except Exception as exc: # noqa: BLE001 # any preparation failure is reported on this row only
|
||||
|
|
@ -509,7 +509,7 @@ class _TeamsData(TypedDict):
|
|||
def _default_member_budget_id(team: LiteLLM_TeamTable) -> str | None:
|
||||
metadata: Final = (
|
||||
_JSON_OBJECT.validate_python(
|
||||
team.metadata # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_TeamTable.metadata is a bare dict; validated by the adapter
|
||||
team.metadata # pyright: ignore[reportUnknownMemberType] # LiteLLM_TeamTable.metadata is a bare dict; validated by the adapter
|
||||
)
|
||||
if team.metadata # pyright: ignore[reportUnknownMemberType] # same bare dict
|
||||
else None
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue