mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_lit_7019_hoist_developer_messages
This commit is contained in:
commit
d5d70da338
1353 changed files with 52428 additions and 11445 deletions
|
|
@ -2648,6 +2648,19 @@ jobs:
|
|||
name: Start mock LLM server
|
||||
command: uv run --no-sync python tests/e2e/ui/fixtures/mock_llm_server/server.py
|
||||
background: true
|
||||
- run:
|
||||
name: Start mock Presidio server
|
||||
command: uv run --no-sync python tests/e2e/ui/fixtures/mock_presidio_server/server.py
|
||||
background: true
|
||||
- run:
|
||||
name: Wait for mock Presidio server
|
||||
command: |
|
||||
for i in $(seq 1 30); do
|
||||
if curl -sf http://127.0.0.1:8091/health >/dev/null 2>&1; then exit 0; fi
|
||||
sleep 1
|
||||
done
|
||||
echo "Mock Presidio server never answered /health on port 8091" >&2
|
||||
exit 1
|
||||
- run:
|
||||
name: Start LiteLLM proxy
|
||||
environment:
|
||||
|
|
@ -2778,6 +2791,19 @@ jobs:
|
|||
name: Start mock LLM server
|
||||
command: uv run --no-sync python tests/e2e/ui/fixtures/mock_llm_server/server.py
|
||||
background: true
|
||||
- run:
|
||||
name: Start mock Presidio server
|
||||
command: uv run --no-sync python tests/e2e/ui/fixtures/mock_presidio_server/server.py
|
||||
background: true
|
||||
- run:
|
||||
name: Wait for mock Presidio server
|
||||
command: |
|
||||
for i in $(seq 1 30); do
|
||||
if curl -sf http://127.0.0.1:8091/health >/dev/null 2>&1; then exit 0; fi
|
||||
sleep 1
|
||||
done
|
||||
echo "Mock Presidio server never answered /health on port 8091" >&2
|
||||
exit 1
|
||||
- run:
|
||||
name: Start LiteLLM proxy under a server root path
|
||||
environment:
|
||||
|
|
|
|||
38
.github/e2e-stack/assert_tests_ran.py
vendored
Normal file
38
.github/e2e-stack/assert_tests_ran.py
vendored
Normal file
|
|
@ -0,0 +1,38 @@
|
|||
import sys
|
||||
import xml.etree.ElementTree as ET
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
|
||||
def main() -> int:
|
||||
selected: Final = tuple(sys.argv[2:])
|
||||
try:
|
||||
report: Final = ET.parse(Path(sys.argv[1])).getroot()
|
||||
except (ET.ParseError, OSError):
|
||||
_ = sys.stdout.write("::error::could not read the test execution report\n")
|
||||
return 1
|
||||
cases: Final = tuple(report.iter("testcase"))
|
||||
passed: Final = frozenset(
|
||||
case.get("file") for case in cases if all(case.find(tag) is None for tag in ("skipped", "failure", "error"))
|
||||
)
|
||||
missing: Final = tuple(path for path in selected if path not in passed)
|
||||
for path in selected:
|
||||
collected: Final = sum(case.get("file") == path for case in cases)
|
||||
skipped: Final = sum(case.get("file") == path and case.find("skipped") is not None for case in cases)
|
||||
_ = sys.stdout.write(f"{path}: {collected} collected, {skipped} skipped\n")
|
||||
for case in cases:
|
||||
if case.get("file") != path or all(case.find(tag) is None for tag in ("failure", "error")):
|
||||
continue
|
||||
_ = sys.stdout.write(f" failed: {case.get('classname', '')}::{case.get('name', '')}\n")
|
||||
if (
|
||||
selected
|
||||
and not missing
|
||||
and not any(case.find(tag) is not None for case in cases for tag in ("failure", "error"))
|
||||
):
|
||||
return 0
|
||||
_ = sys.stdout.write("::error::every selected file must execute a passing test, with no failures or errors\n")
|
||||
return 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
17
.github/e2e-stack/down.sh
vendored
Executable file
17
.github/e2e-stack/down.sh
vendored
Executable file
|
|
@ -0,0 +1,17 @@
|
|||
#!/usr/bin/env bash
|
||||
set -uo pipefail
|
||||
|
||||
STACK_DIR="${E2E_STACK_DIR:-${RUNNER_TEMP:-/tmp}/litellm-e2e-stack}"
|
||||
|
||||
for pid_file in "${STACK_DIR}"/pids/*.pid; do
|
||||
[[ -f "${pid_file}" ]] || continue
|
||||
pkill -TERM -P "$(cat "${pid_file}")" 2>/dev/null
|
||||
kill -TERM "$(cat "${pid_file}")" 2>/dev/null
|
||||
rm -f "${pid_file}"
|
||||
done
|
||||
|
||||
for container in e2e-nginx e2e-valkey e2e-jaeger e2e-postgres; do
|
||||
docker rm -f "${container}" >/dev/null 2>&1
|
||||
done
|
||||
|
||||
exit 0
|
||||
49
.github/e2e-stack/secrets_to_env.py
vendored
Normal file
49
.github/e2e-stack/secrets_to_env.py
vendored
Normal file
|
|
@ -0,0 +1,49 @@
|
|||
import os
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
secrets_adapter: Final[TypeAdapter[dict[str, str]]] = TypeAdapter(dict[str, str])
|
||||
ENV_NAME: Final = re.compile(r"[A-Za-z_][A-Za-z0-9_]*")
|
||||
MIN_MASKED_LENGTH: Final = 8
|
||||
|
||||
|
||||
def main() -> int:
|
||||
env_path: Final = Path(sys.argv[1])
|
||||
try:
|
||||
secrets: Final = {
|
||||
key: value.rstrip("\r\n") for key, value in secrets_adapter.validate_json(sys.stdin.read()).items()
|
||||
}
|
||||
except (ValidationError, UnicodeError):
|
||||
_ = sys.stderr.write("expected a JSON object containing string environment values\n")
|
||||
return 1
|
||||
unusable: Final = tuple(
|
||||
key
|
||||
for key, value in secrets.items()
|
||||
if ENV_NAME.fullmatch(key) is None or any(char in value for char in "'\n\r\0")
|
||||
)
|
||||
if unusable:
|
||||
_ = sys.stderr.write(
|
||||
f"these names or values cannot be represented in both bash and dotenv: {' '.join(sorted(unusable))}\n"
|
||||
)
|
||||
return 1
|
||||
for value in secrets.values():
|
||||
if len(value) >= MIN_MASKED_LENGTH:
|
||||
_ = sys.stdout.write(f"::add-mask::{value.replace('%', '%25')}\n")
|
||||
sys.stdout.flush()
|
||||
lines: Final = tuple(f"{key}='{value}'" for key, value in secrets.items() if value)
|
||||
try:
|
||||
with os.fdopen(os.open(env_path, os.O_WRONLY | os.O_APPEND | os.O_CREAT | os.O_NOFOLLOW, 0o600), "w") as handle:
|
||||
os.fchmod(handle.fileno(), 0o600)
|
||||
_ = handle.write("\n".join(lines) + "\n")
|
||||
except OSError:
|
||||
_ = sys.stderr.write("could not write the environment file\n")
|
||||
return 1
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
44
.github/e2e-stack/select_tests.py
vendored
Normal file
44
.github/e2e-stack/select_tests.py
vendored
Normal file
|
|
@ -0,0 +1,44 @@
|
|||
import re
|
||||
import sys
|
||||
from typing import Final
|
||||
|
||||
SELECTABLE: Final = re.compile(r"^tests/e2e/([A-Za-z0-9_.-]+/)*test_[A-Za-z0-9_.-]+\.py$")
|
||||
UNSUPPORTED: Final = re.compile(
|
||||
r"^tests/e2e/(ui|claude_code|load)/"
|
||||
r"|^tests/e2e/llm_translation/realtime/test_realtime_pipecat_audio_e2e\.py$"
|
||||
r"|^tests/e2e/batches/test_managed_files_enforcement_e2e\.py$"
|
||||
r"|^tests/e2e/guardrails/test_presidio_masking_e2e\.py$"
|
||||
)
|
||||
HARNESS: Final = re.compile(
|
||||
r"^tests/e2e/[A-Za-z0-9_.-]+\.(py|ini)$"
|
||||
r"|^tests/e2e/gateway/"
|
||||
r"|^\.github/e2e-stack/"
|
||||
r"|^\.github/workflows/test-e2e-changed\.yml$"
|
||||
)
|
||||
UNEXPANDED: Final = re.compile(r"[*?\[]")
|
||||
|
||||
|
||||
def is_selectable(path: str) -> bool:
|
||||
return SELECTABLE.match(path) is not None and UNSUPPORTED.match(path) is None
|
||||
|
||||
|
||||
def select(changed: tuple[str, ...], canary: tuple[str, ...]) -> tuple[str, ...]:
|
||||
direct: Final = frozenset(path for path in changed if is_selectable(path))
|
||||
harness_changed: Final = any(HARNESS.match(path) for path in changed)
|
||||
canary_tests: Final = frozenset(path for path in canary if harness_changed and is_selectable(path))
|
||||
return tuple(sorted(direct | canary_tests))
|
||||
|
||||
|
||||
def main() -> int:
|
||||
canary: Final = tuple(sys.argv[1:])
|
||||
unexpanded: Final = tuple(path for path in canary if UNEXPANDED.search(path))
|
||||
if unexpanded:
|
||||
_ = sys.stderr.write(f"the canary paths reached the selector unexpanded: {' '.join(unexpanded)}\n")
|
||||
return 1
|
||||
changed: Final = tuple(line.strip() for line in sys.stdin if line.strip())
|
||||
_ = sys.stdout.write(" ".join(select(changed, canary)) + "\n")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
207
.github/e2e-stack/up.sh
vendored
Executable file
207
.github/e2e-stack/up.sh
vendored
Executable file
|
|
@ -0,0 +1,207 @@
|
|||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
umask 077
|
||||
|
||||
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)"
|
||||
STACK_DIR="${E2E_STACK_DIR:-${RUNNER_TEMP:-/tmp}/litellm-e2e-stack}"
|
||||
CERTS_DIR="${STACK_DIR}/certs"
|
||||
LOGS_DIR="${STACK_DIR}/logs"
|
||||
PIDS_DIR="${STACK_DIR}/pids"
|
||||
|
||||
POSTGRES_IMAGE="${E2E_POSTGRES_IMAGE:-postgres:16.6}"
|
||||
VALKEY_IMAGE="${E2E_VALKEY_IMAGE:-valkey/valkey:8.1.4@sha256:81db6d39e1bba3b3ff32bd3a1b19a6d69690f94a3954ec131277b9a26b95b3aa}"
|
||||
JAEGER_IMAGE="${E2E_JAEGER_IMAGE:-jaegertracing/jaeger:2.10.0}"
|
||||
NGINX_IMAGE="${E2E_NGINX_IMAGE:-nginx:1.29.1-alpine@sha256:42a516af16b852e33b7682d5ef8acbd5d13fe08fecadc7ed98605ba5e3b26ab8}"
|
||||
|
||||
LB_PORT="${E2E_LB_PORT:-4000}"
|
||||
GATEWAY_PORT_1="${E2E_GATEWAY_PORT_1:-4010}"
|
||||
GATEWAY_PORT_2="${E2E_GATEWAY_PORT_2:-4011}"
|
||||
BACKEND_PORT="${E2E_BACKEND_PORT:-4001}"
|
||||
REDIS_PORT="${E2E_REDIS_PORT:-6379}"
|
||||
DATABASE_HOST="${E2E_DATABASE_HOST:-127.0.0.1}"
|
||||
DATABASE_PORT="${E2E_DATABASE_PORT:-5432}"
|
||||
DATABASE_USER="${E2E_DATABASE_USER:-litellm}"
|
||||
DATABASE_PASSWORD="${E2E_DATABASE_PASSWORD:-dbpassword9090}"
|
||||
DATABASE_NAME="${E2E_DATABASE_NAME:-litellm}"
|
||||
JAEGER_OTLP_PORT="${E2E_JAEGER_OTLP_PORT:-4318}"
|
||||
JAEGER_QUERY_PORT="${E2E_JAEGER_QUERY_PORT:-16686}"
|
||||
|
||||
MASTER_KEY="${LITELLM_MASTER_KEY:-sk-e2e-$(openssl rand -hex 16)}"
|
||||
|
||||
mkdir -p "${CERTS_DIR}" "${LOGS_DIR}" "${PIDS_DIR}"
|
||||
chmod 700 "${STACK_DIR}" "${LOGS_DIR}" "${PIDS_DIR}"
|
||||
chmod 755 "${CERTS_DIR}"
|
||||
|
||||
log() { printf 'e2e-stack: %s\n' "$*"; }
|
||||
|
||||
port_open() { (exec 3<>"/dev/tcp/127.0.0.1/$1") 2>/dev/null; }
|
||||
|
||||
wait_for() {
|
||||
local label="$1" check="$2" deadline=$((SECONDS + ${3:-120}))
|
||||
until eval "${check}"; do
|
||||
if ((SECONDS >= deadline)); then
|
||||
log "timed out waiting for ${label}"
|
||||
exit 1
|
||||
fi
|
||||
sleep 2
|
||||
done
|
||||
log "${label} is up"
|
||||
}
|
||||
|
||||
if [[ -f "${REPO_ROOT}/tests/e2e/.env" ]]; then
|
||||
set -a
|
||||
source "${REPO_ROOT}/tests/e2e/.env"
|
||||
set +a
|
||||
fi
|
||||
|
||||
if [[ -z "${DD_API_KEY:-}" ]]; then
|
||||
log "DD_API_KEY is empty; the gateway config enables the datadog callback, so put a Datadog API key in tests/e2e/.env"
|
||||
exit 1
|
||||
fi
|
||||
export DD_SITE="${DD_SITE:-datadoghq.com}"
|
||||
|
||||
if ! port_open "${DATABASE_PORT}"; then
|
||||
docker run -d --name e2e-postgres -p "${DATABASE_PORT}:5432" \
|
||||
-e "POSTGRES_USER=${DATABASE_USER}" -e "POSTGRES_PASSWORD=${DATABASE_PASSWORD}" -e "POSTGRES_DB=${DATABASE_NAME}" \
|
||||
"${POSTGRES_IMAGE}" >/dev/null
|
||||
fi
|
||||
wait_for "postgres" "port_open ${DATABASE_PORT}"
|
||||
|
||||
if ! port_open "${JAEGER_QUERY_PORT}"; then
|
||||
docker run -d --name e2e-jaeger -p "${JAEGER_OTLP_PORT}:4318" -p "${JAEGER_QUERY_PORT}:16686" \
|
||||
"${JAEGER_IMAGE}" >/dev/null
|
||||
fi
|
||||
wait_for "jaeger" "curl -fs http://127.0.0.1:${JAEGER_QUERY_PORT}/api/services >/dev/null"
|
||||
|
||||
openssl genrsa -out "${CERTS_DIR}/ca.key" 2048 2>/dev/null
|
||||
openssl req -x509 -new -nodes -key "${CERTS_DIR}/ca.key" -sha256 -days 7 \
|
||||
-subj "/CN=litellm-e2e-ca" \
|
||||
-addext "basicConstraints=critical,CA:TRUE" -addext "keyUsage=critical,keyCertSign,cRLSign" \
|
||||
-out "${CERTS_DIR}/ca.crt" 2>/dev/null
|
||||
openssl genrsa -out "${CERTS_DIR}/server.key" 2048 2>/dev/null
|
||||
openssl req -new -key "${CERTS_DIR}/server.key" -subj "/CN=localhost" -out "${CERTS_DIR}/server.csr" 2>/dev/null
|
||||
openssl x509 -req -in "${CERTS_DIR}/server.csr" -CA "${CERTS_DIR}/ca.crt" -CAkey "${CERTS_DIR}/ca.key" \
|
||||
-CAcreateserial -days 7 -sha256 \
|
||||
-extfile <(printf 'basicConstraints=CA:FALSE\nkeyUsage=critical,digitalSignature,keyEncipherment\nextendedKeyUsage=serverAuth\nsubjectAltName=DNS:localhost,IP:127.0.0.1\n') \
|
||||
-out "${CERTS_DIR}/server.crt" 2>/dev/null
|
||||
chmod 644 "${CERTS_DIR}"/*.key "${CERTS_DIR}"/*.crt
|
||||
|
||||
CERTIFI_BUNDLE="$(cd "${REPO_ROOT}" && uv run --no-sync python -c 'import certifi; print(certifi.where())')"
|
||||
cat "${CERTIFI_BUNDLE}" "${CERTS_DIR}/ca.crt" > "${CERTS_DIR}/ca-bundle.pem"
|
||||
|
||||
docker rm -f e2e-valkey >/dev/null 2>&1 || true
|
||||
docker run -d --name e2e-valkey -p "${REDIS_PORT}:${REDIS_PORT}" -v "${CERTS_DIR}:/certs:ro" \
|
||||
"${VALKEY_IMAGE}" valkey-server \
|
||||
--cluster-enabled yes --port 0 --tls-port "${REDIS_PORT}" \
|
||||
--tls-cert-file /certs/server.crt --tls-key-file /certs/server.key --tls-ca-cert-file /certs/ca.crt \
|
||||
--tls-auth-clients no --cluster-announce-ip 127.0.0.1 >/dev/null
|
||||
VALKEY_CLI="docker exec e2e-valkey valkey-cli --tls --cacert /certs/ca.crt -h 127.0.0.1 -p ${REDIS_PORT}"
|
||||
wait_for "valkey" "${VALKEY_CLI} ping 2>/dev/null | grep -q PONG"
|
||||
${VALKEY_CLI} cluster addslotsrange 0 16383 >/dev/null
|
||||
wait_for "valkey cluster" "${VALKEY_CLI} cluster info 2>/dev/null | grep -q cluster_state:ok"
|
||||
|
||||
CONFIG_SOURCE="${REPO_ROOT}/tests/e2e/gateway/stage_mirror_ci_config.yml"
|
||||
CONFIG_PATH="${CONFIG_SOURCE}"
|
||||
if [[ "${REDIS_PORT}" != "6379" ]]; then
|
||||
CONFIG_PATH="${STACK_DIR}/litellm-config.yml"
|
||||
sed "s/port: 6379/port: ${REDIS_PORT}/" "${CONFIG_SOURCE}" > "${CONFIG_PATH}"
|
||||
fi
|
||||
|
||||
SERVER_ENV=(
|
||||
"LITELLM_MASTER_KEY=${MASTER_KEY}"
|
||||
"DATABASE_HOST=${DATABASE_HOST}"
|
||||
"DATABASE_PORT=${DATABASE_PORT}"
|
||||
"DATABASE_USER=${DATABASE_USER}"
|
||||
"DATABASE_PASSWORD=${DATABASE_PASSWORD}"
|
||||
"DATABASE_NAME=${DATABASE_NAME}"
|
||||
"DISABLE_SCHEMA_UPDATE=true"
|
||||
"REDIS_HOST=127.0.0.1"
|
||||
"REDIS_PORT=${REDIS_PORT}"
|
||||
"REDIS_CLUSTER_NODES=[{\"host\":\"127.0.0.1\",\"port\":${REDIS_PORT}}]"
|
||||
"CONFIG_FILE_PATH=${CONFIG_PATH}"
|
||||
"STORE_MODEL_IN_DB=True"
|
||||
"OTEL_EXPORTER_OTLP_PROTOCOL=http/protobuf"
|
||||
"OTEL_EXPORTER_OTLP_ENDPOINT=http://127.0.0.1:${JAEGER_OTLP_PORT}"
|
||||
"SSL_CERT_FILE=${CERTS_DIR}/ca-bundle.pem"
|
||||
"PYTHONPATH=${REPO_ROOT}"
|
||||
)
|
||||
if [[ -n "${VERTEXAI_CREDENTIALS:-}" ]]; then
|
||||
printf '%s' "${VERTEXAI_CREDENTIALS}" > "${STACK_DIR}/vertex-adc.json"
|
||||
SERVER_ENV+=("GOOGLE_APPLICATION_CREDENTIALS=${STACK_DIR}/vertex-adc.json")
|
||||
fi
|
||||
|
||||
cd "${REPO_ROOT}"
|
||||
|
||||
log "running migrations"
|
||||
env "${SERVER_ENV[@]}" uv run --no-sync python migrations/run.py >"${LOGS_DIR}/migrations.log" 2>&1
|
||||
|
||||
start_server() {
|
||||
local name="$1"; shift
|
||||
env "${SERVER_ENV[@]}" "$@" >"${LOGS_DIR}/${name}.log" 2>&1 &
|
||||
echo $! > "${PIDS_DIR}/${name}.pid"
|
||||
}
|
||||
|
||||
start_server backend uv run --no-sync uvicorn backend.main:app --host 0.0.0.0 --port "${BACKEND_PORT}"
|
||||
start_server gateway-1 uv run --no-sync uvicorn gateway.main:app --workers 1 --host 0.0.0.0 --port "${GATEWAY_PORT_1}"
|
||||
start_server gateway-2 uv run --no-sync uvicorn gateway.main:app --workers 1 --host 0.0.0.0 --port "${GATEWAY_PORT_2}"
|
||||
|
||||
if [[ "$(uname)" == "Linux" ]]; then
|
||||
NGINX_UPSTREAM_HOST=127.0.0.1
|
||||
NGINX_DOCKER_ARGS=(--network host)
|
||||
else
|
||||
NGINX_UPSTREAM_HOST=host.docker.internal
|
||||
NGINX_DOCKER_ARGS=(-p "${LB_PORT}:${LB_PORT}")
|
||||
fi
|
||||
|
||||
cat > "${STACK_DIR}/nginx.conf" <<EOF
|
||||
events {}
|
||||
http {
|
||||
map \$http_upgrade \$connection_upgrade {
|
||||
default upgrade;
|
||||
'' close;
|
||||
}
|
||||
upstream litellm_gateways {
|
||||
server ${NGINX_UPSTREAM_HOST}:${GATEWAY_PORT_1};
|
||||
server ${NGINX_UPSTREAM_HOST}:${GATEWAY_PORT_2};
|
||||
}
|
||||
server {
|
||||
listen ${LB_PORT};
|
||||
client_max_body_size 100m;
|
||||
location / {
|
||||
proxy_pass http://litellm_gateways;
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Host \$host;
|
||||
proxy_set_header X-Forwarded-For \$proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto \$scheme;
|
||||
proxy_set_header Upgrade \$http_upgrade;
|
||||
proxy_set_header Connection \$connection_upgrade;
|
||||
proxy_buffering off;
|
||||
proxy_read_timeout 600s;
|
||||
proxy_send_timeout 600s;
|
||||
}
|
||||
}
|
||||
}
|
||||
EOF
|
||||
|
||||
docker rm -f e2e-nginx >/dev/null 2>&1 || true
|
||||
docker run -d --name e2e-nginx "${NGINX_DOCKER_ARGS[@]}" \
|
||||
-v "${STACK_DIR}/nginx.conf:/etc/nginx/nginx.conf:ro" "${NGINX_IMAGE}" >/dev/null
|
||||
|
||||
wait_for "backend" "curl -fs http://127.0.0.1:${BACKEND_PORT}/health/liveliness >/dev/null" 300
|
||||
wait_for "gateway-1" "curl -fs http://127.0.0.1:${GATEWAY_PORT_1}/health/liveliness >/dev/null" 300
|
||||
wait_for "gateway-2" "curl -fs http://127.0.0.1:${GATEWAY_PORT_2}/health/liveliness >/dev/null" 300
|
||||
wait_for "load balancer" "curl -fs http://127.0.0.1:${LB_PORT}/health/liveliness >/dev/null" 60
|
||||
|
||||
cat > "${STACK_DIR}/stack.env" <<EOF
|
||||
LITELLM_PROXY_URL=http://127.0.0.1:${LB_PORT}
|
||||
LITELLM_CONTROL_PLANE_URL=http://127.0.0.1:${BACKEND_PORT}
|
||||
LITELLM_PROXY_REPLICA_URLS=http://127.0.0.1:${GATEWAY_PORT_1},http://127.0.0.1:${GATEWAY_PORT_2}
|
||||
LITELLM_MASTER_KEY=${MASTER_KEY}
|
||||
REDIS_HOST=127.0.0.1
|
||||
REDIS_PORT=${REDIS_PORT}
|
||||
E2E_OTEL_QUERY_URL=http://127.0.0.1:${JAEGER_QUERY_PORT}
|
||||
SSL_CERT_FILE=${CERTS_DIR}/ca-bundle.pem
|
||||
DATABASE_URL=postgresql://${DATABASE_USER}:${DATABASE_PASSWORD}@${DATABASE_HOST}:${DATABASE_PORT}/${DATABASE_NAME}
|
||||
EOF
|
||||
|
||||
log "stack is up; pytest env written to ${STACK_DIR}/stack.env"
|
||||
21
.github/workflows/_test-unit-base.yml
vendored
21
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -55,17 +55,14 @@ on:
|
|||
permissions:
|
||||
contents: read
|
||||
|
||||
env:
|
||||
UV_PYTHON: "3.12"
|
||||
|
||||
jobs:
|
||||
run:
|
||||
name: ${{ matrix.python-version == '3.12' && 'Run tests' || format('Run tests (Python {0})', matrix.python-version) }}
|
||||
name: Run tests
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: ${{ inputs.job-timeout-minutes }}
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
python-version: ["3.10", "3.11", "3.12", "3.13", "3.14"]
|
||||
env:
|
||||
UV_PYTHON: ${{ matrix.python-version }}
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: read
|
||||
|
|
@ -88,7 +85,7 @@ jobs:
|
|||
timeout-minutes: 3
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
python-version: ${{ env.UV_PYTHON }}
|
||||
|
||||
- name: Set up uv
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
|
|
@ -103,9 +100,9 @@ jobs:
|
|||
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
|
||||
with:
|
||||
path: ${{ env.UV_CACHE_DIR }}
|
||||
key: ${{ runner.os }}-uv-downloads-py${{ matrix.python-version }}-${{ hashFiles('uv.lock') }}
|
||||
key: ${{ runner.os }}-uv-downloads-py${{ env.UV_PYTHON }}-${{ hashFiles('uv.lock') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-uv-downloads-py${{ matrix.python-version }}-
|
||||
${{ runner.os }}-uv-downloads-py${{ env.UV_PYTHON }}-
|
||||
|
||||
- name: Cache the Rust build
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
|
|
@ -139,7 +136,7 @@ jobs:
|
|||
WORKERS: ${{ inputs.workers }}
|
||||
RERUNS: ${{ inputs.reruns }}
|
||||
DIST: ${{ inputs.dist }}
|
||||
COVERAGE_CORE: ${{ contains(fromJSON('["3.10", "3.11"]'), matrix.python-version) && 'ctrace' || 'sysmon' }}
|
||||
COVERAGE_CORE: sysmon
|
||||
run: |
|
||||
if [ "${WORKERS}" = "0" ]; then
|
||||
uv run --no-sync pytest ${TEST_PATH:?} \
|
||||
|
|
@ -166,7 +163,7 @@ jobs:
|
|||
fi
|
||||
|
||||
- name: Save coverage report
|
||||
if: always() && matrix.python-version == '3.12' && steps.changes.outputs.decision != 'skip'
|
||||
if: always() && steps.changes.outputs.decision != 'skip'
|
||||
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
|
||||
with:
|
||||
name: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
|
||||
|
|
|
|||
2
.github/workflows/auto-close-duplicates.yml
vendored
2
.github/workflows/auto-close-duplicates.yml
vendored
|
|
@ -22,7 +22,7 @@ on:
|
|||
permissions: {}
|
||||
|
||||
jobs:
|
||||
test:
|
||||
sweep-tests:
|
||||
if: github.event_name == 'pull_request'
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
|
|
|
|||
45
.github/workflows/cost-map-guard.yml
vendored
Normal file
45
.github/workflows/cost-map-guard.yml
vendored
Normal file
|
|
@ -0,0 +1,45 @@
|
|||
name: Cost map guard
|
||||
|
||||
on: # zizmor: ignore[dangerous-triggers] runs the base branch's code only; the PR's cost map files are read as data and never executed
|
||||
pull_request_target:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
cost-map-guard:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
- name: Fetch the pull request head and its merge base
|
||||
id: revisions
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
run: |
|
||||
merge_base="$(gh api "repos/${GITHUB_REPOSITORY}/compare/${BASE_SHA}...${HEAD_SHA}" --jq '.merge_base_commit.sha')"
|
||||
git fetch --no-tags --depth=1 origin "$merge_base" "$HEAD_SHA"
|
||||
echo "merge_base=$merge_base" >> "$GITHUB_OUTPUT"
|
||||
- name: Set up uv
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
- name: Run the guard
|
||||
env:
|
||||
MERGE_BASE: ${{ steps.revisions.outputs.merge_base }}
|
||||
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
HEAD_REF: ${{ github.event.pull_request.head.ref }}
|
||||
run: |
|
||||
uv run --frozen python ci_cd/cost_map_guard.py --base "$MERGE_BASE" --head "$HEAD_SHA" --head-ref "$HEAD_REF"
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
name: Publish basedpyright base counts
|
||||
|
||||
# Every commit on litellm_internal_staging is some branch's future merge-base.
|
||||
# Every commit on main or litellm_internal_staging can become a future merge-base.
|
||||
# Publishing its per-rule basedpyright counts as an artifact lets
|
||||
# scripts/type_check_gate.py download them in seconds instead of paying a
|
||||
# 60-110s second basedpyright pass on every fresh worktree or moved merge-base.
|
||||
|
|
@ -10,13 +10,13 @@ name: Publish basedpyright base counts
|
|||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
ref:
|
||||
description: "Ref to compute and publish base counts for"
|
||||
description: "Ref to compute and publish base counts for (defaults to the workflow run's commit)"
|
||||
required: false
|
||||
default: litellm_internal_staging
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
130
.github/workflows/report-rust-release-wheel.yml
vendored
130
.github/workflows/report-rust-release-wheel.yml
vendored
|
|
@ -1,130 +0,0 @@
|
|||
name: Report LiteLLM Rust release wheel
|
||||
|
||||
on: # zizmor: ignore[dangerous-triggers] reporter executes no PR code and consumes no PR artifacts or outputs
|
||||
workflow_run:
|
||||
workflows:
|
||||
- LiteLLM Rust
|
||||
types:
|
||||
- completed
|
||||
|
||||
permissions: {}
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.workflow_run.pull_requests[0].number || github.event.workflow_run.id }}
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
report-release-wheel:
|
||||
name: report release wheel
|
||||
if: >-
|
||||
github.event.workflow_run.event == 'pull_request' &&
|
||||
github.event.workflow_run.path == '.github/workflows/test-rust.yml' &&
|
||||
github.event.workflow_run.head_repository.full_name == github.repository &&
|
||||
github.event.workflow_run.pull_requests[0].number != null
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
permissions:
|
||||
issues: write # PR comments use the issues API
|
||||
pull-requests: read # Current-head validation rejects stale workflow runs
|
||||
|
||||
steps:
|
||||
- name: Link release wheel report on PR
|
||||
uses: actions/github-script@f28e40c7f34bde8b3046d885e986cb6290c5673b # v7.1.0
|
||||
env:
|
||||
COMMENT_MARKER: "<!-- litellm-release-wheel-size -->"
|
||||
with:
|
||||
script: |
|
||||
const marker = process.env.COMMENT_MARKER;
|
||||
const workflowRun = context.payload.workflow_run;
|
||||
const allowedConclusions = new Set([
|
||||
"action_required",
|
||||
"cancelled",
|
||||
"failure",
|
||||
"neutral",
|
||||
"skipped",
|
||||
"stale",
|
||||
"startup_failure",
|
||||
"success",
|
||||
"timed_out",
|
||||
]);
|
||||
if (
|
||||
!allowedConclusions.has(workflowRun.conclusion) ||
|
||||
workflowRun.event !== "pull_request" ||
|
||||
workflowRun.path !== ".github/workflows/test-rust.yml" ||
|
||||
workflowRun.head_repository?.full_name !==
|
||||
`${context.repo.owner}/${context.repo.repo}` ||
|
||||
workflowRun.pull_requests?.length !== 1
|
||||
) {
|
||||
throw new Error("unexpected source workflow");
|
||||
}
|
||||
const pullRequest = workflowRun.pull_requests[0];
|
||||
const pullRequestNumber = pullRequest.number;
|
||||
const headSha = workflowRun.head_sha;
|
||||
const runId = workflowRun.id;
|
||||
if (
|
||||
!Number.isSafeInteger(pullRequestNumber) ||
|
||||
pullRequestNumber <= 0 ||
|
||||
!Number.isSafeInteger(runId) ||
|
||||
runId <= 0 ||
|
||||
!/^[0-9a-f]{40}$/.test(headSha) ||
|
||||
pullRequest.head?.sha !== headSha
|
||||
) {
|
||||
throw new Error("invalid source workflow metadata");
|
||||
}
|
||||
const runUrl =
|
||||
`${context.serverUrl}/${context.repo.owner}/${context.repo.repo}` +
|
||||
`/actions/runs/${runId}`;
|
||||
const result =
|
||||
workflowRun.conclusion === "success"
|
||||
? "successfully"
|
||||
: `with \`${workflowRun.conclusion}\``;
|
||||
const body = [
|
||||
marker,
|
||||
"## LiteLLM Rust workflow",
|
||||
"",
|
||||
`Workflow completed ${result} for \`${headSha}\``,
|
||||
"",
|
||||
`[View workflow run](${runUrl})`,
|
||||
].join("\n");
|
||||
const comments = await github.paginate(github.rest.issues.listComments, {
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
issue_number: pullRequestNumber,
|
||||
per_page: 100,
|
||||
});
|
||||
const existing = comments.find(
|
||||
(comment) =>
|
||||
comment.user?.login === "github-actions[bot]" &&
|
||||
comment.body?.startsWith(marker),
|
||||
);
|
||||
const currentPullRequest = (
|
||||
await github.rest.pulls.get({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
pull_number: pullRequestNumber,
|
||||
})
|
||||
).data;
|
||||
if (
|
||||
currentPullRequest.state !== "open" ||
|
||||
currentPullRequest.head.repo?.full_name !==
|
||||
`${context.repo.owner}/${context.repo.repo}` ||
|
||||
currentPullRequest.head.sha !== headSha
|
||||
) {
|
||||
core.info("source workflow no longer matches the current pull request head");
|
||||
return;
|
||||
}
|
||||
if (existing) {
|
||||
await github.rest.issues.updateComment({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
comment_id: existing.id,
|
||||
body,
|
||||
});
|
||||
} else {
|
||||
await github.rest.issues.createComment({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
issue_number: pullRequestNumber,
|
||||
body,
|
||||
});
|
||||
}
|
||||
|
|
@ -13,10 +13,12 @@ jobs:
|
|||
sync_together_ai_models:
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
BASE_BRANCH: ${{ github.event.repository.default_branch }}
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
ref: litellm_internal_staging
|
||||
ref: ${{ env.BASE_BRANCH }}
|
||||
persist-credentials: false
|
||||
- name: Set up uv
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
|
|
@ -63,6 +65,6 @@ jobs:
|
|||
gh pr create --title "feat(models): sync together_ai model registry" \
|
||||
--body-file "$RUNNER_TEMP/pr_body.md" \
|
||||
--head "$branch" \
|
||||
--base litellm_internal_staging
|
||||
--base "$BASE_BRANCH"
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GH_TOKEN || github.token }}
|
||||
|
|
|
|||
9
.github/workflows/test-code-quality.yml
vendored
9
.github/workflows/test-code-quality.yml
vendored
|
|
@ -74,6 +74,15 @@ jobs:
|
|||
- name: check_workflow_startup_safety
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/check_workflow_startup_safety.py
|
||||
|
||||
- name: check_workflow_job_name_collisions
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/check_workflow_job_name_collisions.py
|
||||
|
||||
- name: test_workflow_job_name_collisions
|
||||
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_workflow_job_name_collisions.py
|
||||
|
||||
- name: test_e2e_changed_gate
|
||||
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_e2e_changed_gate.py
|
||||
|
||||
- name: router_code_coverage
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/router_code_coverage.py
|
||||
|
||||
|
|
|
|||
240
.github/workflows/test-e2e-changed.yml
vendored
Normal file
240
.github/workflows/test-e2e-changed.yml
vendored
Normal file
|
|
@ -0,0 +1,240 @@
|
|||
name: e2e-changed-tests
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
|
||||
concurrency:
|
||||
group: e2e-changed-${{ github.event.pull_request.number }}
|
||||
cancel-in-progress: true
|
||||
|
||||
permissions: {}
|
||||
|
||||
jobs:
|
||||
detect:
|
||||
name: Detect changed e2e tests
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: read
|
||||
outputs:
|
||||
tests: ${{ steps.changed.outputs.tests }}
|
||||
any: ${{ steps.changed.outputs.any }}
|
||||
steps:
|
||||
- name: Checkout the selector and the canary suite
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
sparse-checkout: |
|
||||
.github/e2e-stack
|
||||
tests/e2e/access_control
|
||||
persist-credentials: false
|
||||
ref: ${{ github.sha }}
|
||||
|
||||
- name: List the e2e test files this PR added or modified
|
||||
id: changed
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
REPO: ${{ github.repository }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number }}
|
||||
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
run: |
|
||||
gh api "repos/${REPO}/pulls/${PR_NUMBER}" \
|
||||
--jq 'select(.head.sha == env.HEAD_SHA and .changed_files < 3000) | .head.sha' \
|
||||
| grep -Fxq "${HEAD_SHA}"
|
||||
files="$(gh api "repos/${REPO}/pulls/${PR_NUMBER}/files" --paginate \
|
||||
--jq '.[] | select(.status != "removed") | .filename')"
|
||||
gh api "repos/${REPO}/pulls/${PR_NUMBER}" --jq '.head.sha' | grep -Fxq "${HEAD_SHA}"
|
||||
tests="$(printf '%s\n' "${files}" \
|
||||
| python3 .github/e2e-stack/select_tests.py tests/e2e/access_control/test_*.py)"
|
||||
echo "tests=${tests}" >> "${GITHUB_OUTPUT}"
|
||||
if [ -n "${tests}" ]; then
|
||||
echo "any=true" >> "${GITHUB_OUTPUT}"
|
||||
echo "selected e2e tests: ${tests}"
|
||||
else
|
||||
echo "any=false" >> "${GITHUB_OUTPUT}"
|
||||
echo "no changed e2e test files supported by this stack; nothing to run"
|
||||
fi
|
||||
|
||||
run:
|
||||
name: Run changed e2e tests against the stage-mirror stack
|
||||
needs: detect
|
||||
if: needs.detect.outputs.any == 'true' && github.event.pull_request.head.repo.full_name == github.repository
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 90
|
||||
environment: e2e-changed
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
services:
|
||||
postgres:
|
||||
image: postgres:16.6
|
||||
env:
|
||||
POSTGRES_USER: litellm
|
||||
POSTGRES_PASSWORD: dbpassword9090
|
||||
POSTGRES_DB: litellm
|
||||
ports:
|
||||
- 5432:5432
|
||||
options: >-
|
||||
--health-cmd "pg_isready -U litellm"
|
||||
--health-interval 5s
|
||||
--health-timeout 5s
|
||||
--health-retries 10
|
||||
jaeger:
|
||||
image: jaegertracing/jaeger:2.10.0
|
||||
ports:
|
||||
- 4318:4318
|
||||
- 16686:16686
|
||||
steps:
|
||||
- name: Validate configuration
|
||||
env:
|
||||
ROLE: ${{ vars.E2E_AWS_ROLE_TO_ASSUME }}
|
||||
run: test -n "${ROLE}" || { echo "::error::Set repo variable E2E_AWS_ROLE_TO_ASSUME to an OIDC role with read access to the e2e secrets"; exit 1; }
|
||||
|
||||
- name: Checkout
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
ref: ${{ github.sha }}
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.13"
|
||||
|
||||
- name: Set up uv
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Cache the Rust build
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen \
|
||||
--extra proxy --extra proxy-runtime --extra extra_proxy \
|
||||
--extra semantic-router --extra bedrock-realtime \
|
||||
--group ci --group proxy-dev --group e2e-dev
|
||||
uv pip install "pipecat-ai[openai]==1.4.0"
|
||||
|
||||
- name: Cache Prisma binaries
|
||||
uses: ./.github/actions/cache-prisma-binaries
|
||||
|
||||
- name: Generate Prisma client
|
||||
run: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Install Playwright chromium
|
||||
run: uv run --no-sync playwright install --with-deps chromium
|
||||
|
||||
- name: Configure AWS credentials
|
||||
id: aws
|
||||
uses: aws-actions/configure-aws-credentials@e7f100cf4c008499ea8adda475de1042d6975c7b # v6.2.0
|
||||
with:
|
||||
role-to-assume: ${{ vars.E2E_AWS_ROLE_TO_ASSUME }}
|
||||
aws-region: us-east-1
|
||||
role-session-name: litellm-e2e-changed-${{ github.run_id }}
|
||||
role-duration-seconds: 900
|
||||
output-env-credentials: false
|
||||
output-credentials: true
|
||||
|
||||
- name: Fetch provider credentials from AWS Secrets Manager
|
||||
env:
|
||||
AWS_ACCESS_KEY_ID: ${{ steps.aws.outputs.aws-access-key-id }}
|
||||
AWS_SECRET_ACCESS_KEY: ${{ steps.aws.outputs.aws-secret-access-key }}
|
||||
AWS_SESSION_TOKEN: ${{ steps.aws.outputs.aws-session-token }}
|
||||
AWS_DEFAULT_REGION: us-east-1
|
||||
run: |
|
||||
umask 077
|
||||
aws secretsmanager get-secret-value --secret-id litellm-e2e-changed-provider-keys \
|
||||
--query SecretString --output text \
|
||||
| uv run --no-sync python .github/e2e-stack/secrets_to_env.py tests/e2e/.env
|
||||
aws secretsmanager get-secret-value --secret-id litellm-e2e-changed-license \
|
||||
--query SecretString --output text \
|
||||
| jq -R -s '{"LITELLM_LICENSE": .}' \
|
||||
| uv run --no-sync python .github/e2e-stack/secrets_to_env.py tests/e2e/.env
|
||||
|
||||
- name: Boot the stage-mirror stack
|
||||
id: boot
|
||||
run: |
|
||||
umask 077
|
||||
if ! bash .github/e2e-stack/up.sh > "${RUNNER_TEMP}/e2e-boot.log" 2>&1; then
|
||||
echo "::error::stage-mirror stack failed to boot; raw logs are not published"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Export stack environment
|
||||
run: |
|
||||
master_key="$(grep '^LITELLM_MASTER_KEY=' "${RUNNER_TEMP}/litellm-e2e-stack/stack.env" | cut -d= -f2-)"
|
||||
echo "::add-mask::${master_key}"
|
||||
cat "${RUNNER_TEMP}/litellm-e2e-stack/stack.env" >> "${GITHUB_ENV}"
|
||||
|
||||
- name: Run the selected tests three times
|
||||
env:
|
||||
TESTS: ${{ needs.detect.outputs.tests }}
|
||||
E2E_FIXTURE_MODE: live
|
||||
run: |
|
||||
umask 077
|
||||
read -r -a test_files <<< "${TESTS}"
|
||||
for pass in 1 2 3; do
|
||||
report="${RUNNER_TEMP}/e2e-pass-${pass}.xml"
|
||||
log="${RUNNER_TEMP}/e2e-pass-${pass}.log"
|
||||
echo "::group::pass ${pass} of 3"
|
||||
set +e
|
||||
uv run --no-sync pytest "${test_files[@]}" --rootdir=. -v -p no:cacheprovider \
|
||||
-o junit_family=xunit1 --junitxml="${report}" > "${log}" 2>&1
|
||||
status=$?
|
||||
uv run --no-sync python .github/e2e-stack/assert_tests_ran.py "${report}" "${test_files[@]}"
|
||||
verified=$?
|
||||
set -e
|
||||
grep -E '^=+ .* in [0-9.]+s( \([0-9:]+\))? =+$' "${log}" | tail -n 1
|
||||
echo "::endgroup::"
|
||||
if [ "${status}" = "5" ]; then
|
||||
echo "::error::the selected files collected no runnable tests, so nothing was verified"
|
||||
exit 1
|
||||
fi
|
||||
if [ "${status}" != "0" ]; then
|
||||
echo "::error::pass ${pass} of 3 failed with exit code ${status}"
|
||||
exit "${status}"
|
||||
fi
|
||||
if [ "${verified}" != "0" ]; then
|
||||
echo "::error::pass ${pass} of 3 did not verify every selected file"
|
||||
exit 1
|
||||
fi
|
||||
echo "pass ${pass} of 3 passed"
|
||||
done
|
||||
|
||||
- name: Stop the stack
|
||||
if: always() && steps.boot.outcome != 'skipped'
|
||||
run: bash .github/e2e-stack/down.sh
|
||||
|
||||
- name: Remove credentials and raw output
|
||||
if: always()
|
||||
run: |
|
||||
rm -f tests/e2e/.env "${RUNNER_TEMP}/e2e-boot.log" "${RUNNER_TEMP}"/e2e-pass-*.log "${RUNNER_TEMP}"/e2e-pass-*.xml
|
||||
rm -rf "${RUNNER_TEMP}/litellm-e2e-stack"
|
||||
|
||||
gate:
|
||||
name: e2e-changed-tests
|
||||
needs: [detect, run]
|
||||
if: always()
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
steps:
|
||||
- name: Require three successful passes when tests changed
|
||||
env:
|
||||
DETECT_RESULT: ${{ needs.detect.result }}
|
||||
ANY_TESTS: ${{ needs.detect.outputs.any }}
|
||||
RUN_RESULT: ${{ needs.run.result }}
|
||||
run: |
|
||||
if [ "${DETECT_RESULT}" != "success" ]; then
|
||||
echo "::error::changed-test detection did not succeed"
|
||||
exit 1
|
||||
fi
|
||||
if [ "${ANY_TESTS}" = "false" ]; then
|
||||
echo "no changed e2e test files supported by this stack; nothing to run"
|
||||
exit 0
|
||||
fi
|
||||
if [ "${ANY_TESTS}" != "true" ] || [ "${RUN_RESULT}" != "success" ]; then
|
||||
echo "::error::selected e2e tests require an approved, successful run; fork PRs must run from a reviewed same-repository branch"
|
||||
exit 1
|
||||
fi
|
||||
1
.github/workflows/test-litellm-ui-unit.yml
vendored
1
.github/workflows/test-litellm-ui-unit.yml
vendored
|
|
@ -12,6 +12,7 @@ on:
|
|||
- "litellm_**"
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
|
||||
concurrency:
|
||||
|
|
|
|||
37
.github/workflows/test-model-map.yml
vendored
37
.github/workflows/test-model-map.yml
vendored
|
|
@ -1,37 +0,0 @@
|
|||
name: Validate model_prices_and_context_window.json
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
|
||||
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
|
||||
|
||||
jobs:
|
||||
validate-model-prices-json:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Validate model_prices_and_context_window.json
|
||||
run: |
|
||||
jq empty model_prices_and_context_window.json
|
||||
|
||||
- name: Set up uv
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Check model_prices_and_context_window.schema.json is in sync
|
||||
run: |
|
||||
uv run --frozen python ci_cd/generate_model_prices_schema.py --check
|
||||
111
.github/workflows/test-rust.yml
vendored
111
.github/workflows/test-rust.yml
vendored
|
|
@ -7,6 +7,7 @@ on:
|
|||
- ".cargo/**"
|
||||
- "pyproject.toml"
|
||||
- "rust-toolchain.toml"
|
||||
- ".github/actions/setup-uv-with-retries/**"
|
||||
- ".github/scripts/smoke_test_native_wheel.py"
|
||||
- ".github/scripts/verify_linux_native_wheel.py"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
|
|
@ -22,6 +23,7 @@ on:
|
|||
- ".cargo/**"
|
||||
- "pyproject.toml"
|
||||
- "rust-toolchain.toml"
|
||||
- ".github/actions/setup-uv-with-retries/**"
|
||||
- ".github/scripts/smoke_test_native_wheel.py"
|
||||
- ".github/scripts/verify_linux_native_wheel.py"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
|
|
@ -34,102 +36,89 @@ concurrency:
|
|||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
env:
|
||||
CARGO_TERM_COLOR: always
|
||||
|
||||
jobs:
|
||||
rust-checks:
|
||||
name: rustfmt, clippy, test
|
||||
rust-lint:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
defaults:
|
||||
run:
|
||||
working-directory: litellm-rust
|
||||
env:
|
||||
CARGO_TERM_COLOR: always
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Rust
|
||||
run: rustup toolchain install
|
||||
- run: rustup toolchain install --no-self-update
|
||||
|
||||
- name: Cache Cargo registry and target
|
||||
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
|
||||
- run: cargo fmt --check
|
||||
|
||||
- uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
|
||||
with:
|
||||
path: |
|
||||
~/.cargo/registry
|
||||
~/.cargo/git
|
||||
litellm-rust/target
|
||||
key: ${{ runner.os }}-cargo-${{ hashFiles('rust-toolchain.toml', 'litellm-rust/Cargo.lock') }}
|
||||
key: ${{ runner.os }}-cargo-${{ github.job }}-${{ hashFiles('rust-toolchain.toml', '.cargo/**', 'litellm-rust/Cargo.lock') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-cargo-
|
||||
${{ runner.os }}-cargo-${{ github.job }}-
|
||||
|
||||
- name: Check Rust formatting
|
||||
run: cargo fmt --check
|
||||
- run: cargo clippy --workspace --all-targets --locked -- -D warnings
|
||||
|
||||
- name: Run Clippy
|
||||
run: cargo clippy --workspace --all-targets --locked -- -D warnings
|
||||
- run: cargo clippy -p litellm-core --all-targets --features bedrock-auth --locked -- -D warnings
|
||||
|
||||
- name: Run Clippy with Bedrock auth
|
||||
run: cargo clippy -p litellm-core --all-targets --features bedrock-auth --locked -- -D warnings
|
||||
- run: cargo clippy -p litellm-ai-gateway --all-targets --all-features --locked -- -D warnings
|
||||
|
||||
- name: Run Clippy with all gateway features
|
||||
run: cargo clippy -p litellm-ai-gateway --all-targets --all-features --locked -- -D warnings
|
||||
|
||||
- name: Run Rust tests
|
||||
run: cargo test --workspace --locked
|
||||
|
||||
- name: Run core tests with Bedrock auth
|
||||
run: cargo test -p litellm-core --features bedrock-auth --locked
|
||||
|
||||
# Not --all-features: python-config links libpython, which this job does not install.
|
||||
- name: Run gateway tests with the server feature
|
||||
run: cargo test -p litellm-ai-gateway --features server --locked
|
||||
|
||||
release-wheel:
|
||||
name: release wheel
|
||||
rust-test:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
permissions:
|
||||
contents: read
|
||||
env:
|
||||
CARGO_TERM_COLOR: always
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
- uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
- uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Set up Rust
|
||||
run: rustup toolchain install
|
||||
- run: rustup toolchain install --no-self-update
|
||||
|
||||
- name: Build release wheel
|
||||
run: uv build --wheel --out-dir dist
|
||||
- uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
|
||||
with:
|
||||
path: |
|
||||
~/.cargo/registry
|
||||
~/.cargo/git
|
||||
litellm-rust/target
|
||||
key: ${{ runner.os }}-cargo-${{ github.job }}-${{ hashFiles('rust-toolchain.toml', '.cargo/**', 'litellm-rust/Cargo.lock') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-cargo-${{ github.job }}-
|
||||
|
||||
- name: Build panic contract wheel
|
||||
run: >-
|
||||
- run: cargo test --workspace --locked
|
||||
working-directory: litellm-rust
|
||||
|
||||
- run: cargo test -p litellm-core --features bedrock-auth --locked
|
||||
working-directory: litellm-rust
|
||||
|
||||
- run: cargo test -p litellm-ai-gateway --features server --locked
|
||||
working-directory: litellm-rust
|
||||
|
||||
- run: uv build --wheel --out-dir dist
|
||||
|
||||
- run: python .github/scripts/verify_linux_native_wheel.py dist/*.whl
|
||||
env:
|
||||
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
|
||||
|
||||
- run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl
|
||||
|
||||
- run: >-
|
||||
uv build --wheel --out-dir panic-dist
|
||||
--config-setting "maturin.build-args=--features panic-test,extension-module"
|
||||
|
||||
- name: Smoke-test native panic unwinding
|
||||
run: python .github/scripts/smoke_test_native_wheel.py panic-dist/*.whl
|
||||
|
||||
- name: Verify stripped native extension
|
||||
env:
|
||||
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
|
||||
run: python .github/scripts/verify_linux_native_wheel.py dist/*.whl
|
||||
|
||||
- name: Test native route wheel
|
||||
run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl
|
||||
- run: python .github/scripts/smoke_test_native_wheel.py panic-dist/*.whl
|
||||
|
|
|
|||
31
.github/workflows/test-terraform-modules.yml
vendored
31
.github/workflows/test-terraform-modules.yml
vendored
|
|
@ -4,6 +4,7 @@ on:
|
|||
push:
|
||||
paths:
|
||||
- "terraform/litellm/aws/**"
|
||||
- "terraform/litellm/gcp/**"
|
||||
- ".github/workflows/test-terraform-modules.yml"
|
||||
pull_request:
|
||||
branches:
|
||||
|
|
@ -13,6 +14,7 @@ on:
|
|||
- "litellm_**"
|
||||
paths:
|
||||
- "terraform/litellm/aws/**"
|
||||
- "terraform/litellm/gcp/**"
|
||||
- ".github/workflows/test-terraform-modules.yml"
|
||||
|
||||
permissions:
|
||||
|
|
@ -52,3 +54,32 @@ jobs:
|
|||
# Plan-only, mock_provider-backed: no AWS credentials, no API calls.
|
||||
- name: test
|
||||
run: terraform test
|
||||
|
||||
gcp-module:
|
||||
name: fmt, validate, test (gcp)
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
defaults:
|
||||
run:
|
||||
working-directory: terraform/litellm/gcp
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- uses: hashicorp/setup-terraform@b9cd54a3c349d3f38e8881555d616ced269862dd # v3.1.2
|
||||
with:
|
||||
terraform_version: 1.13.3
|
||||
terraform_wrapper: false
|
||||
|
||||
- name: fmt
|
||||
run: terraform fmt -recursive -check -diff
|
||||
|
||||
- name: init
|
||||
run: terraform init -backend=false -input=false
|
||||
|
||||
- name: validate
|
||||
run: terraform validate
|
||||
|
||||
- name: test
|
||||
run: terraform test
|
||||
|
|
|
|||
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -116,6 +116,7 @@ jobs:
|
|||
tests/test_litellm/rerank_api
|
||||
tests/test_litellm/rust_bridge
|
||||
tests/test_litellm/sandbox
|
||||
tests/test_litellm/skills
|
||||
tests/test_litellm/test_router
|
||||
tests/test_litellm/vector_stores
|
||||
tests/test_litellm/videos
|
||||
|
|
|
|||
|
|
@ -149,7 +149,7 @@ graph TD
|
|||
| `parallel_request_limiter` | `proxy/hooks/parallel_request_limiter_v3.py` | Rate limiting per key/user |
|
||||
| `cache_control_check` | `proxy/hooks/cache_control_check.py` | Cache validation |
|
||||
| `responses_id_security` | `proxy/hooks/responses_id_security.py` | Response ID validation |
|
||||
| `litellm_skills` | `proxy/hooks/skills_injection.py` | Skills injection |
|
||||
| `litellm_skills` | `proxy/hooks/litellm_skills/main.py` | Skills injection |
|
||||
|
||||
To add a new proxy hook, implement `CustomLogger` and register in `PROXY_HOOKS`.
|
||||
|
||||
|
|
@ -220,20 +220,20 @@ graph LR
|
|||
| Job | Interval | Purpose | Key Files |
|
||||
|-----|----------|---------|-----------|
|
||||
| `update_spend` | 60s | Batch write spend logs to PostgreSQL | `proxy/db/db_spend_update_writer.py` |
|
||||
| `reset_budget` | 10-12min | Reset budgets for keys/users/teams | `proxy/management_helpers/budget_reset_job.py` |
|
||||
| `reset_budget` | 10-12min | Reset budgets for keys/users/teams | `proxy/common_utils/reset_budget_job.py` |
|
||||
| `add_deployment` | 10s | Sync new model deployments from DB | `proxy/proxy_server.py` (`ProxyConfig`) |
|
||||
| `cleanup_old_spend_logs` | cron/interval | Delete old spend logs | `proxy/management_helpers/spend_log_cleanup.py` |
|
||||
| `check_batch_cost` | 30min | Calculate costs for batch jobs | `proxy/management_helpers/check_batch_cost_job.py` |
|
||||
| `check_responses_cost` | 30min | Calculate costs for responses API | `proxy/management_helpers/check_responses_cost_job.py` |
|
||||
| `process_rotations` | 1hr | Auto-rotate API keys | `proxy/management_helpers/key_rotation_manager.py` |
|
||||
| `cleanup_old_spend_logs` | cron/interval | Delete old spend logs | `proxy/db/db_transaction_queue/spend_log_cleanup.py` |
|
||||
| `check_batch_cost` | 30min | Calculate costs for batch jobs | `enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py` |
|
||||
| `check_responses_cost` | 30min | Calculate costs for responses API | `enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py` |
|
||||
| `process_rotations` | 1hr | Auto-rotate API keys | `proxy/common_utils/key_rotation_manager.py` |
|
||||
| `_run_background_health_check` | continuous | Health check model deployments | `proxy/proxy_server.py` |
|
||||
| `send_weekly_spend_report` | weekly | Slack spend alerts | `proxy/utils.py` (`SlackAlerting`) |
|
||||
| `send_monthly_spend_report` | monthly | Slack spend alerts | `proxy/utils.py` (`SlackAlerting`) |
|
||||
|
||||
**Cost Attribution Flow:**
|
||||
1. LLM response returns to `utils.py` wrapper after `litellm.acompletion()` completes
|
||||
2. `update_response_metadata()` (`llm_response_utils/response_metadata.py`) is called
|
||||
3. `logging_obj._response_cost_calculator()` (`litellm_logging.py`) calculates cost via `litellm.completion_cost()` (`cost_calculator.py`)
|
||||
2. `update_response_metadata()` (`litellm_core_utils/llm_response_utils/response_metadata.py`) is called
|
||||
3. `logging_obj._response_cost_calculator()` (`litellm_core_utils/litellm_logging.py`) calculates cost via `litellm.completion_cost()` (`cost_calculator.py`)
|
||||
4. Cost is stored in `response._hidden_params["response_cost"]`
|
||||
5. `proxy/common_request_processing.py` extracts cost from `hidden_params` and adds to response headers (`x-litellm-response-cost`)
|
||||
6. `logging_obj.async_success_handler()` triggers callbacks including `_ProxyDBLogger.async_log_success_event()`
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ Never test structure of code only function of it
|
|||
|
||||
End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `CLAUDE.md`
|
||||
|
||||
When creating PRs, don't set base to `main`. `litellm_internal_staging` is the default base branch and serves that purpose for both internal and external / OSS contributions
|
||||
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
|
||||
|
||||
|
|
@ -37,7 +37,7 @@ Same applies for filing bug reports and feature requests, with .github/ISSUE_TEM
|
|||
|
||||
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
|
||||
|
||||
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, just tell me which URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, what fields to fill out, etc. along with the other commands to run in an ordered list, and I'll do it myself and post the screenshots after you make the PR
|
||||
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
|
||||
|
||||
If you ever write any human-facing text (pull requests, issues, commit messages, discussion posts, github comments, release notes, docs, etc.), always follow these guidelines to sound less AI-y:
|
||||
- don't use emojis
|
||||
|
|
@ -52,7 +52,7 @@ Don't hesitate to use values in .env to get needed API keys and other secrets, a
|
|||
|
||||
Python max line length is 120, not 88
|
||||
|
||||
When you fix violations gated by `ruff-strict-budget.json`, `type-discipline-budget.json`, or `basedpyright-code-budget.json`, run `make lint-budget-update` and commit the lowered limits so the ceilings ratchet down instead of leaving stale headroom. It measures the working tree, so it must contain exactly the fixes you're committing
|
||||
Never edit or commit `ruff-strict-budget.json`, `type-discipline-budget.json`, `basedpyright-code-budget.json`, or `test-quality-budget.json` on a PR branch, and don't run `make lint-budget-update` there. A scheduled Devin automation lowers the limits on the default branch in its own PR by exactly what landed since the last ratchet, so concurrent PRs don't fight over the same `"limit"` lines. Keep the hosted automation's target in sync when the repository default changes. If your branch already carries a budget edit, drop it before opening the PR
|
||||
|
||||
`make check` (f.k.a. `make pre-commit`, which still works identically as an alias) saves its complete output to a log file in .git (overwriting previous logs) and prints that path as its first and last output lines. To inspect a run, read or grep that log instead of re-running the multi-minute checks just to see a different slice
|
||||
|
||||
|
|
@ -70,7 +70,7 @@ When referencing or running models (coding, QA'ing, writing docs, writing tests,
|
|||
|
||||
Always pull before starting any work. The checkout or worktree may be sitting on a stale branch
|
||||
|
||||
If you're an internal contributor, when creating a new PR, the typical flow is to branch off litellm_internal_staging and create a branch prefixed with litellm_. Do not create a branch prefixed with claude/ and generally do not have / in your branch names
|
||||
If you're an internal contributor, when creating a new PR, the typical flow is to branch off the repository's current default branch and create a branch prefixed with litellm_. Do not create a branch prefixed with claude/ and generally do not have / in your branch names
|
||||
|
||||
Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages. Never use a `claude/` prefix or put a `/` in a branch name. Do not add "Generated with Claude Code" (or any similar attribution) to PR descriptions or comments. Do not create a new PR/branch off the existing PR to fix/add something that is related and could've just been committed directly to the existing PR's branch
|
||||
|
||||
|
|
|
|||
|
|
@ -315,10 +315,12 @@ Ensure the UI builds successfully before submitting your PR:
|
|||
npm run build
|
||||
```
|
||||
|
||||
Local lint and budget checks follow origin's current default branch. They refresh it from the remote instead of trusting cached `origin/HEAD`. For an intentional comparison against another branch or commit, use `make check BASE_REF=<ref>` or the standalone gate's `--base <ref>` option. An explicit ref can also be used offline once it has been fetched locally. Without an override, unavailable remote metadata stops the check
|
||||
|
||||
## Submitting Your PR
|
||||
|
||||
1. **Push your branch**: `git push origin your-feature-branch`
|
||||
2. **Create a PR**: Go to GitHub and open a pull request against [`litellm_internal_staging`](https://github.com/BerriAI/litellm/tree/litellm_internal_staging), which is the default base branch. Do not target `main`.
|
||||
2. **Create a PR**: Go to GitHub and open a pull request against the repository's current default branch. Run `python3 scripts/default_branch.py --branch` to check its name
|
||||
3. **Fill out the PR template**: Provide clear description of changes
|
||||
4. **Wait for review**: Maintainers will review and provide feedback
|
||||
5. **Address feedback**: Make requested changes and push updates
|
||||
|
|
|
|||
|
|
@ -67,6 +67,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
|
|||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13
|
||||
|
||||
# Copy full source tree
|
||||
|
|
@ -89,6 +90,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
|
|||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13
|
||||
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
|
|
|
|||
60
Makefile
60
Makefile
|
|
@ -34,7 +34,7 @@ help:
|
|||
@echo " make lint-basedpyright-budget-update - Ratchet basedpyright limits down by what this branch fixed"
|
||||
@echo " make lint-format - Check ruff format formatting (matches CI)"
|
||||
@echo " make lint-ruff-budget - Gate the codebase total of each strict ruff rule against its limit"
|
||||
@echo " make lint-gate - Strict ruff gate in CI-parity mode (fetches staging, simulates the merge)"
|
||||
@echo " make lint-gate - Strict ruff gate in CI-parity mode (fetches the default branch, simulates the merge)"
|
||||
@echo " make lint-ruff-budget-update - Ratchet ruff-strict-budget.json limits down by what this branch fixed"
|
||||
@echo " make lint-test-quality - Gate the test suite against test-quality-budget.json"
|
||||
@echo " make lint-budget-update - Ratchet all budgets down (ruff + type-discipline + test quality + basedpyright)"
|
||||
|
|
@ -60,6 +60,9 @@ help:
|
|||
|
||||
UV := uv
|
||||
UV_RUN := $(UV) run --no-sync
|
||||
BASE_REF ?=
|
||||
export BASE_REF
|
||||
RESOLVE_BASE = python3 scripts/default_branch.py --base "$(BASE_REF)"
|
||||
|
||||
# Machine-wide slot queue for the heavy targets below; python3 + stdlib only, so
|
||||
# it runs before any venv exists. See scripts/gate_slot_lock.py.
|
||||
|
|
@ -67,7 +70,7 @@ GATE_SLOT_LOCK := python3 scripts/gate_slot_lock.py
|
|||
|
||||
LINT_DEP_INSTALL ?= install-dev
|
||||
LINT_E2E_DEP_INSTALL ?= lint-install
|
||||
LINT_DEP_BASE ?= lint-fetch-base
|
||||
LINT_DEP_BASE ?=
|
||||
LINT_JOBS := $(shell sysctl -n hw.ncpu 2>/dev/null || nproc 2>/dev/null || echo 4)
|
||||
LINT_OUTPUT_SYNC := $(if $(filter output-sync,$(.FEATURES)),--output-sync=target,)
|
||||
|
||||
|
|
@ -130,10 +133,8 @@ format: install-dev
|
|||
format-check: install-dev
|
||||
cd litellm && $(UV_RUN) ruff format --check --exclude '/enterprise/' . && cd ..
|
||||
|
||||
# Single fetch of the PR base so the delta-based gates below share one network round
|
||||
# trip instead of each re-fetching when chained from `lint`.
|
||||
lint-fetch-base:
|
||||
git fetch origin litellm_internal_staging
|
||||
@$(RESOLVE_BASE)
|
||||
|
||||
# Mirror test-linting.yml's lint job environment: the proxy-dev group plus a generated
|
||||
# Prisma client, so `basedpyright tests/e2e` resolves the same modules CI does. The
|
||||
|
|
@ -150,7 +151,9 @@ lint-install:
|
|||
# recursively, so 'litellm/*.py' covers nested modules and the top-level files that
|
||||
# CI's 'litellm/**/*.py' skips, which makes this target a superset of the CI step.
|
||||
lint-format-check-changed: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
@files=$$(git diff --name-only --diff-filter=ACMR origin/litellm_internal_staging...HEAD -- 'litellm/*.py' | grep -v '^litellm/enterprise/' || true); \
|
||||
@base_ref=$$($(RESOLVE_BASE)) && \
|
||||
changed=$$(git diff --name-only --diff-filter=ACMR "$$base_ref...HEAD" -- 'litellm/*.py') && \
|
||||
files=$$(printf '%s\n' "$$changed" | grep -v '^litellm/enterprise/' || true) || exit $$?; \
|
||||
if [ -z "$$files" ]; then \
|
||||
echo "No changed litellm Python files to format-check."; \
|
||||
else \
|
||||
|
|
@ -167,7 +170,9 @@ lint-ruff: $(LINT_DEP_INSTALL)
|
|||
# https://github.com/astral-sh/ruff/discussions/10977
|
||||
# https://github.com/astral-sh/ruff/discussions/4049
|
||||
lint-format-changed: install-dev
|
||||
@git diff origin/main --unified=0 --no-color -- '*.py' | \
|
||||
@base_ref=$$($(RESOLVE_BASE)) && \
|
||||
diff=$$(git diff "$$base_ref" --unified=0 --no-color -- '*.py') && \
|
||||
printf '%s\n' "$$diff" | \
|
||||
perl -ne '\
|
||||
if (/^diff --git a\/(.*) b\//) { $$file = $$1; } \
|
||||
if (/^@@ .* \+(\d+)(?:,(\d+))? @@/) { \
|
||||
|
|
@ -182,20 +187,22 @@ lint-format-changed: install-dev
|
|||
done
|
||||
|
||||
lint-ruff-dev: install-dev
|
||||
@tmpfile=$$(mktemp /tmp/ruff-dev.XXXXXX) && \
|
||||
@base_ref=$$($(RESOLVE_BASE)) || exit $$?; \
|
||||
tmpfile=$$(mktemp /tmp/ruff-dev.XXXXXX) && \
|
||||
cd litellm && \
|
||||
($(UV_RUN) ruff check . --output-format=pylint || true) > "$$tmpfile" && \
|
||||
$(UV_RUN) diff-quality --violations=pylint "$$tmpfile" --compare-branch=origin/main && \
|
||||
$(UV_RUN) diff-quality --violations=pylint "$$tmpfile" --compare-branch="$$base_ref" && \
|
||||
cd .. ; \
|
||||
rm -f "$$tmpfile"
|
||||
|
||||
lint-ruff-FULL-dev: install-dev
|
||||
@files=$$(git diff --name-only origin/main -- '*.py'); \
|
||||
@base_ref=$$($(RESOLVE_BASE)) && \
|
||||
files=$$(git diff --name-only "$$base_ref" -- '*.py') || exit $$?; \
|
||||
if [ -n "$$files" ]; then echo "$$files" | xargs $(UV_RUN) ruff check; \
|
||||
else echo "No changed .py files to check."; fi
|
||||
|
||||
lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
$(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging
|
||||
$(UV_RUN) python scripts/type_check_gate.py --base "$(BASE_REF)"
|
||||
|
||||
lint-e2e-basedpyright: $(LINT_E2E_DEP_INSTALL)
|
||||
$(UV_RUN) basedpyright tests/e2e
|
||||
|
|
@ -203,37 +210,37 @@ lint-e2e-basedpyright: $(LINT_E2E_DEP_INSTALL)
|
|||
# Type-discipline budget (mutable collections / casts / type guards / kwargs /
|
||||
# unexplained suppressions), the test-linting.yml step `make lint` used to omit.
|
||||
lint-type-discipline: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
$(UV_RUN) python scripts/type_discipline_gate.py --base origin/litellm_internal_staging
|
||||
$(UV_RUN) python scripts/type_discipline_gate.py --base "$(BASE_REF)"
|
||||
|
||||
# Test-quality budget (zero-assert / mock-echo tests, sys.path.insert, raw env writes,
|
||||
# litellm module-global mutation, credential-gated skips, conftest snapshot
|
||||
# inventory), counted across tests/ the same delta-vs-base way.
|
||||
lint-test-quality: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
$(UV_RUN) python scripts/test_quality_gate.py --base origin/litellm_internal_staging
|
||||
$(UV_RUN) python scripts/test_quality_gate.py --base "$(BASE_REF)"
|
||||
|
||||
# --update lowers each limit by what this branch fixed since its branch point, so
|
||||
# it needs the base ref fetched to resolve the merge-base.
|
||||
lint-basedpyright-budget-update: install-dev lint-fetch-base
|
||||
$(UV_RUN) python scripts/type_check_gate.py --update
|
||||
lint-basedpyright-budget-update: install-dev
|
||||
$(UV_RUN) python scripts/type_check_gate.py --update --base "$(BASE_REF)"
|
||||
|
||||
lint-format: format-check
|
||||
|
||||
lint-ruff-budget: install-dev
|
||||
$(UV_RUN) python scripts/ruff_strict_gate.py
|
||||
$(UV_RUN) python scripts/ruff_strict_gate.py --base "$(BASE_REF)"
|
||||
|
||||
# Strict gate, invoked the same way CI does in test-linting.yml so a local pass
|
||||
# means the CI check will pass too.
|
||||
lint-gate: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
$(UV_RUN) python scripts/ruff_strict_gate.py --base origin/litellm_internal_staging
|
||||
$(UV_RUN) python scripts/ruff_strict_gate.py --base "$(BASE_REF)"
|
||||
|
||||
lint-ruff-budget-update: install-dev lint-fetch-base
|
||||
$(UV_RUN) python scripts/ruff_strict_gate.py --update
|
||||
lint-ruff-budget-update: install-dev
|
||||
$(UV_RUN) python scripts/ruff_strict_gate.py --update --base "$(BASE_REF)"
|
||||
|
||||
lint-type-discipline-budget-update: install-dev lint-fetch-base
|
||||
$(UV_RUN) python scripts/type_discipline_gate.py --update
|
||||
lint-type-discipline-budget-update: install-dev
|
||||
$(UV_RUN) python scripts/type_discipline_gate.py --update --base "$(BASE_REF)"
|
||||
|
||||
lint-test-quality-budget-update: install-dev lint-fetch-base
|
||||
$(UV_RUN) python scripts/test_quality_gate.py --update
|
||||
lint-test-quality-budget-update: install-dev
|
||||
$(UV_RUN) python scripts/test_quality_gate.py --update --base "$(BASE_REF)"
|
||||
|
||||
# Ratchet all budgets in one shot (ruff strict + type-discipline + test quality + basedpyright)
|
||||
lint-budget-update: lint-ruff-budget-update lint-type-discipline-budget-update lint-test-quality-budget-update lint-basedpyright-budget-update
|
||||
|
|
@ -249,14 +256,15 @@ check-import-safety: $(LINT_DEP_INSTALL)
|
|||
# runs the diff-scoped ruff format check, whole-tree ruff check, the strict-rule /
|
||||
# type-discipline / basedpyright budgets as a delta vs the base, then the circular-import
|
||||
# and import-safety checks. Steps that compare against the base resolve it the same way CI
|
||||
# does (merge-base with origin/litellm_internal_staging). Setup (env sync, Prisma client,
|
||||
# does (merge-base with origin's current default branch). Setup (env sync, Prisma client,
|
||||
# base fetch) runs once up front; the checks themselves are independent, so a sub-make
|
||||
# fans them out with -j and the fast ones finish under basedpyright's shadow.
|
||||
lint:
|
||||
@$(GATE_SLOT_LOCK) $(MAKE) lint-inner
|
||||
|
||||
lint-inner: lint-install lint-fetch-base
|
||||
$(MAKE) -j $(LINT_JOBS) $(LINT_OUTPUT_SYNC) LINT_DEP_INSTALL= LINT_E2E_DEP_INSTALL= LINT_DEP_BASE= lint-checks
|
||||
lint-inner: lint-install
|
||||
@base_ref=$$($(RESOLVE_BASE)) && \
|
||||
$(MAKE) BASE_REF="$$base_ref" -j $(LINT_JOBS) $(LINT_OUTPUT_SYNC) LINT_DEP_INSTALL= LINT_E2E_DEP_INSTALL= LINT_DEP_BASE= lint-checks
|
||||
|
||||
lint-checks: lint-format-check-changed lint-ruff lint-gate lint-type-discipline lint-test-quality lint-basedpyright lint-e2e-basedpyright check-circular-imports check-import-safety
|
||||
|
||||
|
|
|
|||
|
|
@ -84,7 +84,7 @@
|
|||
"limit": 56
|
||||
},
|
||||
"reportPrivateUsage": {
|
||||
"limit": 1808
|
||||
"limit": 1804
|
||||
},
|
||||
"reportRedeclaration": {
|
||||
"limit": 8
|
||||
|
|
@ -105,13 +105,13 @@
|
|||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38283
|
||||
"limit": 38271
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19584
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 29829
|
||||
"limit": 29814
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 110
|
||||
|
|
@ -135,7 +135,7 @@
|
|||
"limit": 21
|
||||
},
|
||||
"reportUnusedFunction": {
|
||||
"limit": 138
|
||||
"limit": 136
|
||||
},
|
||||
"reportUnusedImport": {
|
||||
"limit": 542
|
||||
|
|
|
|||
146
ci_cd/cost_map_guard.py
Normal file
146
ci_cd/cost_map_guard.py
Normal file
|
|
@ -0,0 +1,146 @@
|
|||
"""Guard the cost map on pull requests.
|
||||
|
||||
Every pull request gets the file checks: the three cost map files parse, the backup copy matches the root file,
|
||||
and the JSON schema is in sync and validates the map. Pull requests from the cost map sync bot (branches named
|
||||
litellm_cost_map_sync_*) additionally may only touch those three files and may only add or update models.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import subprocess
|
||||
import sys
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
from generate_model_prices_schema import SPECIAL_ROOT_KEYS, build_schema, render, validation_errors
|
||||
|
||||
COST_MAP_PATH: Final = "model_prices_and_context_window.json"
|
||||
BACKUP_PATH: Final = "litellm/model_prices_and_context_window_backup.json"
|
||||
SCHEMA_PATH: Final = "model_prices_and_context_window.schema.json"
|
||||
GUARDED_PATHS: Final = (COST_MAP_PATH, BACKUP_PATH, SCHEMA_PATH)
|
||||
BOT_BRANCH_PREFIX: Final = "litellm_cost_map_sync_"
|
||||
|
||||
CostMap = dict[str, object]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Snapshot:
|
||||
cost_map: str
|
||||
backup: str
|
||||
schema: str
|
||||
|
||||
|
||||
def _parse_object(text: str, path: str) -> CostMap | str:
|
||||
try:
|
||||
parsed: Final = json.loads(text)
|
||||
except json.JSONDecodeError as error:
|
||||
return f"{path} is not valid JSON: {error}"
|
||||
return parsed if isinstance(parsed, dict) else f"{path} must be a JSON object at the root"
|
||||
|
||||
|
||||
def _rendered_schema(cost_map: CostMap) -> str:
|
||||
try:
|
||||
return render(build_schema(cost_map))
|
||||
except SystemExit as error:
|
||||
return str(error)
|
||||
|
||||
|
||||
def _file_failures(head: Snapshot, head_map: CostMap) -> tuple[str, ...]:
|
||||
schema_text: Final = _rendered_schema(head_map)
|
||||
if not schema_text.startswith("{"):
|
||||
return (schema_text,)
|
||||
backup_failure: Final = (
|
||||
()
|
||||
if head.backup == head.cost_map
|
||||
else (f"{BACKUP_PATH} differs from {COST_MAP_PATH}; copy the root file over it",)
|
||||
)
|
||||
schema_failure: Final = (
|
||||
()
|
||||
if head.schema == schema_text
|
||||
else (
|
||||
f"{SCHEMA_PATH} is out of sync with {COST_MAP_PATH}; "
|
||||
"run `python ci_cd/generate_model_prices_schema.py` and commit the result",
|
||||
)
|
||||
)
|
||||
return (
|
||||
*backup_failure,
|
||||
*schema_failure,
|
||||
*(
|
||||
f"{COST_MAP_PATH} does not validate against its schema: {error}"
|
||||
for error in validation_errors(head_map, json.loads(schema_text))[:20]
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _entries(cost_map: CostMap) -> dict[str, dict[str, object]]:
|
||||
return {key: entry for key, entry in cost_map.items() if isinstance(entry, dict)}
|
||||
|
||||
|
||||
def _bot_failures(base: Snapshot, head_map: CostMap, changed_files: Sequence[str]) -> tuple[str, ...]:
|
||||
base_map: Final = _parse_object(base.cost_map, COST_MAP_PATH)
|
||||
if isinstance(base_map, str):
|
||||
return (f"merge base: {base_map}",)
|
||||
base_entries: Final = _entries(base_map)
|
||||
head_entries: Final = _entries(head_map)
|
||||
removed_fields: Final = tuple(
|
||||
f"{key}.{field}"
|
||||
for key, entry in base_entries.items()
|
||||
if key in head_entries
|
||||
for field in entry
|
||||
if field not in head_entries[key]
|
||||
)
|
||||
return (
|
||||
*(
|
||||
f"bot PRs may only change the cost map files, not {path}"
|
||||
for path in changed_files
|
||||
if path not in GUARDED_PATHS
|
||||
),
|
||||
*(f"bot PRs may not remove models: {key}" for key in base_map if key not in head_map),
|
||||
*(f"bot PRs may not remove fields: {ref}" for ref in removed_fields),
|
||||
*(
|
||||
f"bot PRs may not change {key}"
|
||||
for key in sorted(SPECIAL_ROOT_KEYS)
|
||||
if base_map.get(key) != head_map.get(key)
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def guard_failures(base: Snapshot, head: Snapshot, changed_files: Sequence[str], bot: bool) -> tuple[str, ...]:
|
||||
head_map: Final = _parse_object(head.cost_map, COST_MAP_PATH)
|
||||
if isinstance(head_map, str):
|
||||
return (head_map,)
|
||||
return (*_file_failures(head, head_map), *(_bot_failures(base, head_map, changed_files) if bot else ()))
|
||||
|
||||
|
||||
def _git(*args: str) -> str:
|
||||
result: Final = subprocess.run(("git", *args), check=False, capture_output=True, text=True)
|
||||
return result.stdout if result.returncode == 0 else ""
|
||||
|
||||
|
||||
def snapshot(revision: str) -> Snapshot:
|
||||
return Snapshot(*(_git("show", f"{revision}:{path}") for path in GUARDED_PATHS))
|
||||
|
||||
|
||||
def main(argv: Sequence[str]) -> int:
|
||||
parser: Final = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--base", required=True, help="merge base of the pull request")
|
||||
parser.add_argument("--head", required=True, help="head commit of the pull request")
|
||||
parser.add_argument("--head-ref", required=True, help="head branch name of the pull request")
|
||||
args: Final = parser.parse_args(argv)
|
||||
bot: Final = args.head_ref.startswith(BOT_BRANCH_PREFIX)
|
||||
changed_files: Final = tuple(_git("diff", "--name-only", args.base, args.head).splitlines())
|
||||
failures: Final = guard_failures(snapshot(args.base), snapshot(args.head), changed_files, bot)
|
||||
contract: Final = "bot contract enforced" if bot else "human PR, file checks only"
|
||||
if failures:
|
||||
print(f"cost map guard failed ({contract}):")
|
||||
print("\n".join(f"- {failure}" for failure in failures))
|
||||
return 1
|
||||
print(f"cost map guard passed ({contract})")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main(sys.argv[1:]))
|
||||
|
|
@ -58,6 +58,11 @@ OBJECT_KEYS: dict[str, JsonSchema] = {
|
|||
}
|
||||
|
||||
ARRAY_KEYS: dict[str, JsonSchema] = {
|
||||
"supported_audio_formats": {
|
||||
"type": "array",
|
||||
"description": "Audio container formats the model can return.",
|
||||
"items": {"type": "string", "enum": ["mp3", "wav"]},
|
||||
},
|
||||
"supported_endpoints": {
|
||||
"type": "array",
|
||||
"description": "OpenAI-style API routes this model can be called through, e.g. /v1/chat/completions.",
|
||||
|
|
@ -231,6 +236,10 @@ def string_key_schemas(modes: tuple) -> dict[str, JsonSchema]:
|
|||
},
|
||||
"comment": STRING,
|
||||
"audio_transcription_config": STRING,
|
||||
"vertex_ai_audio_api": {
|
||||
"type": "string",
|
||||
"enum": ["lyria_predict", "lyria_interactions"],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -6,12 +6,9 @@ import subprocess
|
|||
import sys
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
import testing.postgresql
|
||||
|
||||
from typing import Final
|
||||
|
||||
DESTRUCTIVE_PATTERN = re.compile(r"\bDROP\s+(COLUMN|TABLE|INDEX)\b", re.IGNORECASE)
|
||||
DEFAULT_BASE_BRANCH = "litellm_internal_staging"
|
||||
|
||||
|
||||
def _find_destructive_statements(sql: str) -> list:
|
||||
|
|
@ -94,31 +91,57 @@ def _print_stale_branch_refusal(base_branch: str, behind: int) -> None:
|
|||
print(banner, file=out)
|
||||
|
||||
|
||||
def _check_branch_freshness(root_dir: Path, base_branch: str) -> None:
|
||||
def _default_base_branch(root_dir: Path) -> str:
|
||||
try:
|
||||
result: Final = subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
str(Path(__file__).resolve().parents[1] / "scripts" / "default_branch.py"),
|
||||
"--repo-root",
|
||||
str(root_dir),
|
||||
"--branch",
|
||||
],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=90,
|
||||
)
|
||||
except (OSError, subprocess.SubprocessError) as exc:
|
||||
_print_freshness_failure(
|
||||
"default branch",
|
||||
"Could not discover origin's default branch. Pass --base-branch <name> to choose one.",
|
||||
exc.stderr if isinstance(exc, subprocess.CalledProcessError) else str(exc),
|
||||
)
|
||||
sys.exit(3)
|
||||
return result.stdout.strip()
|
||||
|
||||
|
||||
def _check_branch_freshness(root_dir: Path, base_branch: str | None = None) -> None:
|
||||
"""Fetch origin/<base_branch> and exit 3 if HEAD is behind it."""
|
||||
resolved_branch: Final = base_branch or _default_base_branch(root_dir)
|
||||
cwd = str(root_dir)
|
||||
try:
|
||||
subprocess.run(
|
||||
["git", "fetch", "origin", base_branch],
|
||||
["git", "fetch", "origin", f"+refs/heads/{resolved_branch}:refs/remotes/origin/{resolved_branch}"],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=cwd,
|
||||
)
|
||||
except FileNotFoundError:
|
||||
_print_freshness_failure(base_branch, "git executable not found on PATH")
|
||||
_print_freshness_failure(resolved_branch, "git executable not found on PATH")
|
||||
sys.exit(3)
|
||||
except subprocess.CalledProcessError as e:
|
||||
_print_freshness_failure(
|
||||
base_branch,
|
||||
f"`git fetch origin {base_branch}` failed",
|
||||
resolved_branch,
|
||||
f"`git fetch origin {resolved_branch}` failed",
|
||||
e.stderr or "",
|
||||
)
|
||||
sys.exit(3)
|
||||
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["git", "rev-list", "--count", f"HEAD..origin/{base_branch}"],
|
||||
["git", "rev-list", "--count", f"HEAD..origin/{resolved_branch}"],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
|
|
@ -127,23 +150,23 @@ def _check_branch_freshness(root_dir: Path, base_branch: str) -> None:
|
|||
behind = int(result.stdout.strip())
|
||||
except subprocess.CalledProcessError as e:
|
||||
_print_freshness_failure(
|
||||
base_branch,
|
||||
f"`git rev-list HEAD..origin/{base_branch}` failed",
|
||||
resolved_branch,
|
||||
f"`git rev-list HEAD..origin/{resolved_branch}` failed",
|
||||
e.stderr or "",
|
||||
)
|
||||
sys.exit(3)
|
||||
except ValueError:
|
||||
_print_freshness_failure(
|
||||
base_branch,
|
||||
resolved_branch,
|
||||
"could not parse commit count from `git rev-list`",
|
||||
)
|
||||
sys.exit(3)
|
||||
|
||||
if behind > 0:
|
||||
_print_stale_branch_refusal(base_branch, behind)
|
||||
_print_stale_branch_refusal(resolved_branch, behind)
|
||||
sys.exit(3)
|
||||
|
||||
print(f"Branch freshness OK: up to date with origin/{base_branch}.")
|
||||
print(f"Branch freshness OK: up to date with origin/{resolved_branch}.")
|
||||
|
||||
|
||||
def _print_destructive_refusal(destructive_lines: list) -> None:
|
||||
|
|
@ -198,7 +221,7 @@ def _print_destructive_refusal(destructive_lines: list) -> None:
|
|||
def create_migration(
|
||||
migration_name: str = None,
|
||||
allow_destructive: bool = False,
|
||||
base_branch: str = DEFAULT_BASE_BRANCH,
|
||||
base_branch: str | None = None,
|
||||
skip_freshness_check: bool = False,
|
||||
):
|
||||
"""
|
||||
|
|
@ -211,7 +234,7 @@ def create_migration(
|
|||
DROP COLUMN, DROP TABLE, or DROP INDEX statements. Without this
|
||||
flag, the script exits non-zero and prints guidance.
|
||||
base_branch (str): Branch to check freshness against
|
||||
(default: "litellm_internal_staging").
|
||||
(default: origin's current default branch).
|
||||
skip_freshness_check (bool): Skip the "branch is up to date" check.
|
||||
Only for intentional migrations against an older base.
|
||||
"""
|
||||
|
|
@ -225,6 +248,8 @@ def create_migration(
|
|||
else:
|
||||
_check_branch_freshness(root_dir, base_branch)
|
||||
|
||||
import testing.postgresql
|
||||
|
||||
try:
|
||||
migrations_dir = (
|
||||
root_dir / "litellm-proxy-extras" / "litellm_proxy_extras" / "migrations"
|
||||
|
|
@ -342,9 +367,8 @@ if __name__ == "__main__":
|
|||
)
|
||||
parser.add_argument(
|
||||
"--base-branch",
|
||||
default=DEFAULT_BASE_BRANCH,
|
||||
help=(
|
||||
f"Branch to check freshness against (default: {DEFAULT_BASE_BRANCH}). "
|
||||
"Branch to check freshness against (default: origin's current default branch). "
|
||||
"The script fetches origin/<base-branch> and refuses to run if HEAD "
|
||||
"is behind it."
|
||||
),
|
||||
|
|
|
|||
|
|
@ -65,6 +65,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
|
|||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13
|
||||
|
||||
# Copy full source tree
|
||||
|
|
@ -87,6 +88,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
|
|||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13
|
||||
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
|
|
|
|||
|
|
@ -71,6 +71,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
|
|||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13
|
||||
|
||||
# Copy full source tree
|
||||
|
|
@ -99,6 +100,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
|
|||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13 \
|
||||
--no-sources-package litellm-proxy-extras; \
|
||||
else \
|
||||
|
|
@ -109,6 +111,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
|
|||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13; \
|
||||
fi
|
||||
|
||||
|
|
|
|||
|
|
@ -433,9 +433,9 @@ _default_detect_secrets_config = {
|
|||
"name": "ZendeskSecretKeyDetector",
|
||||
"path": _custom_plugins_path + "/zendesk_secret_key.py",
|
||||
},
|
||||
{"name": "Base64HighEntropyString", "limit": 3.0},
|
||||
{"name": "Base64HighEntropyString", "limit": 4.5},
|
||||
{"name": "HexHighEntropyString", "limit": 3.0},
|
||||
]
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -466,16 +466,19 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
|
|||
|
||||
os.remove(temp_file.name)
|
||||
|
||||
detected_secrets = []
|
||||
for file in secrets.files:
|
||||
for found_secret in secrets[file]:
|
||||
if found_secret.secret_value is None:
|
||||
continue
|
||||
detected_secrets.append(
|
||||
{"type": found_secret.type, "value": found_secret.secret_value}
|
||||
)
|
||||
|
||||
return detected_secrets
|
||||
return [
|
||||
{"type": found_secret.type, "value": found_secret.secret_value}
|
||||
for file in sorted(secrets.files)
|
||||
for found_secret in sorted(
|
||||
secrets[file],
|
||||
key=lambda secret: (
|
||||
-len(secret.secret_value or ""),
|
||||
secret.type,
|
||||
secret.secret_value or "",
|
||||
),
|
||||
)
|
||||
if found_secret.secret_value is not None
|
||||
]
|
||||
|
||||
def redact_text(self, text: str, source: str = "message") -> str:
|
||||
"""Replace every detected secret in ``text`` with ``[REDACTED]`` and
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ This plugin searches for OpenAI API Keys.
|
|||
"""
|
||||
|
||||
import re
|
||||
from collections.abc import Generator
|
||||
|
||||
from detect_secrets.plugins.base import RegexBasedDetector
|
||||
|
||||
|
|
@ -16,4 +17,16 @@ class OpenAIApiKeyDetector(RegexBasedDetector):
|
|||
|
||||
@property
|
||||
def denylist(self) -> list[re.Pattern]:
|
||||
return [re.compile(r"""(sk-[a-zA-Z0-9]{5,})""")]
|
||||
return [
|
||||
re.compile(
|
||||
r"((?:(?<![a-zA-Z0-9])|(?<=%[0-9A-Fa-f]{2}))"
|
||||
r"sk[-_]"
|
||||
r"[a-zA-Z0-9_-]{5,}"
|
||||
r"(?![a-zA-Z0-9_-]))"
|
||||
)
|
||||
]
|
||||
|
||||
def analyze_string(self, string: str) -> Generator[str, None, None]:
|
||||
# the digit check lives outside the regex: a lookahead re-scans the token
|
||||
# from every `sk` inside it, which is quadratic on `-sk-sk-sk-...` input
|
||||
yield from (match for match in super().analyze_string(string) if re.search(r"[0-9]", match))
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.64"
|
||||
version = "0.1.65"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.1.64"
|
||||
version = "0.1.65"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -47,6 +47,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13
|
||||
|
||||
# Stage 2 — copy source and install the project + workspace members.
|
||||
|
|
@ -59,6 +60,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13
|
||||
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "per_server_oauth_discovery" BOOLEAN NOT NULL DEFAULT false;
|
||||
|
|
@ -0,0 +1,3 @@
|
|||
ALTER TABLE "LiteLLM_AutoRouterSession"
|
||||
ADD COLUMN IF NOT EXISTS "classifier_cost" DOUBLE PRECISION NOT NULL DEFAULT 0,
|
||||
ADD COLUMN IF NOT EXISTS "classifier_cost_recorded_turns" INTEGER NOT NULL DEFAULT 0;
|
||||
|
|
@ -343,6 +343,7 @@ model LiteLLM_MCPServerTable {
|
|||
delegate_auth_to_upstream Boolean @default(false)
|
||||
oauth_passthrough Boolean @default(false)
|
||||
dcr_bridge Boolean?
|
||||
per_server_oauth_discovery Boolean @default(false)
|
||||
is_byok Boolean @default(false)
|
||||
byok_description String[] @default([])
|
||||
byok_api_key_help_url String?
|
||||
|
|
@ -1508,6 +1509,8 @@ model LiteLLM_AutoRouterSession {
|
|||
total_tokens BigInt @default(0)
|
||||
spend Float @default(0)
|
||||
saved_spend Float @default(0)
|
||||
classifier_cost Float @default(0)
|
||||
classifier_cost_recorded_turns Int @default(0)
|
||||
tier_turns Json @default("{}")
|
||||
|
||||
@@id([api_key, session_id, router_name])
|
||||
|
|
|
|||
|
|
@ -48,7 +48,7 @@ uv run --with testing.postgresql python ci_cd/run_migration.py "your_migration_n
|
|||
|
||||
## What It Does
|
||||
|
||||
1. **Verifies the current branch is up to date with `origin/litellm_internal_staging`** (see [Branch freshness](#branch-freshness-check))
|
||||
1. **Verifies the current branch is up to date with origin's current default branch** (see [Branch freshness](#branch-freshness-check))
|
||||
2. Creates temp PostgreSQL DB
|
||||
3. Applies existing migrations
|
||||
4. Compares with `schema.prisma`
|
||||
|
|
@ -57,11 +57,11 @@ uv run --with testing.postgresql python ci_cd/run_migration.py "your_migration_n
|
|||
|
||||
## Branch Freshness Check
|
||||
|
||||
Before generating anything, `run_migration.py` runs `git fetch origin <base>` and refuses to proceed if `HEAD` is behind `origin/<base>`. Default base is `litellm_internal_staging` (the branch PRs target). A previous incident saw a stale branch silently drop production columns; freshness is the first-line defense.
|
||||
Before generating anything, `run_migration.py` runs `git fetch origin <base>` and refuses to proceed if `HEAD` is behind `origin/<base>`. The default base is discovered from origin's advertised HEAD on each run, so an existing clone follows a default-branch change without trusting cached `origin/HEAD`. If discovery or fetching fails, migration generation stops. A previous incident saw a stale branch silently drop production columns; freshness is the first-line defense.
|
||||
|
||||
Flags:
|
||||
|
||||
- `--base-branch <name>` — check against a different base (e.g. `main`). Default is `litellm_internal_staging`.
|
||||
- `--base-branch <name>` — check against a different base (e.g. a release branch). Defaults to origin's current default branch
|
||||
- `--skip-freshness-check` — bypass entirely. Only for intentional migrations against an older base.
|
||||
|
||||
When the guard fires:
|
||||
|
|
@ -69,8 +69,9 @@ When the guard fires:
|
|||
1. Update your branch:
|
||||
|
||||
```bash
|
||||
git fetch origin && git rebase origin/litellm_internal_staging
|
||||
# or git merge origin/litellm_internal_staging — whichever matches your workflow
|
||||
base_branch=$(python3 scripts/default_branch.py --branch) &&
|
||||
git fetch origin "+refs/heads/$base_branch:refs/remotes/origin/$base_branch" &&
|
||||
git rebase "origin/$base_branch"
|
||||
```
|
||||
2. Re-run `run_migration.py`.
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.93"
|
||||
version = "0.4.94"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.93"
|
||||
version = "0.4.94"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
2
litellm-rust/Cargo.lock
generated
2
litellm-rust/Cargo.lock
generated
|
|
@ -1415,6 +1415,8 @@ dependencies = [
|
|||
"litellm-config",
|
||||
"litellm-core",
|
||||
"reqwest",
|
||||
"rustls 0.23.42",
|
||||
"rustls-native-certs",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2 0.10.9",
|
||||
|
|
|
|||
|
|
@ -28,6 +28,8 @@ pythonize = "0.29.0"
|
|||
rand = "0.8"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "multipart", "rustls-tls", "http2", "stream"] }
|
||||
rstest = "0.26.1"
|
||||
rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"] }
|
||||
rustls-native-certs = "0.8"
|
||||
serde = { version = "1.0", features = ["derive"] }
|
||||
serde_json = { version = "1.0", features = ["float_roundtrip"] }
|
||||
sha2 = "0.10"
|
||||
|
|
|
|||
|
|
@ -20,6 +20,10 @@ litellm-config.workspace = true
|
|||
# reqwest (rustls + json) is used by io/ocr and ships realtime logs to the
|
||||
# Python proxy callbacks API.
|
||||
reqwest.workspace = true
|
||||
# rustls and its root store are direct dependencies so `io::tls` can build the
|
||||
# one TLS config the outbound dials use; see that module for why it has to.
|
||||
rustls.workspace = true
|
||||
rustls-native-certs.workspace = true
|
||||
# `sync` powers the bounded mpsc channel the realtime logger drains.
|
||||
tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "time", "sync"] }
|
||||
tokio-tungstenite.workspace = true
|
||||
|
|
|
|||
|
|
@ -3,3 +3,4 @@ pub mod ocr;
|
|||
pub mod realtime;
|
||||
pub mod realtime_pool;
|
||||
pub mod responses_ws;
|
||||
pub(crate) mod tls;
|
||||
|
|
|
|||
|
|
@ -23,10 +23,12 @@ use tokio_tungstenite::tungstenite::Message;
|
|||
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
|
||||
use tokio_tungstenite::tungstenite::http::HeaderValue;
|
||||
use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION;
|
||||
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async};
|
||||
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
|
||||
|
||||
use litellm_core::providers::openai::realtime::transformation::OPENAI_REALTIME_CONFIG;
|
||||
|
||||
use crate::io::tls::connect_upstream;
|
||||
|
||||
/// Environment variable holding the OpenAI API key (last-resort fallback).
|
||||
const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY";
|
||||
|
||||
|
|
@ -84,7 +86,7 @@ pub(crate) async fn dial_upstream(
|
|||
.map_err(|err| Error::Auth(err.to_string()))?,
|
||||
);
|
||||
|
||||
let (upstream, _response) = connect_async(request)
|
||||
let (upstream, _response) = connect_upstream(request)
|
||||
.await
|
||||
.map_err(|err| Error::Network(err.to_string()))?;
|
||||
Ok(upstream)
|
||||
|
|
@ -284,6 +286,33 @@ mod tests {
|
|||
serde_json::from_str(raw).expect("valid event json")
|
||||
}
|
||||
|
||||
/// The realtime dial has to reach a `wss://` upstream without a process-wide
|
||||
/// crypto provider installed, which is what dialing through `io::tls` buys.
|
||||
#[tokio::test]
|
||||
async fn dial_upstream_over_wss_reports_an_error_instead_of_panicking() {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("bind a loopback port");
|
||||
let port = listener
|
||||
.local_addr()
|
||||
.expect("read the bound address")
|
||||
.port();
|
||||
tokio::spawn(async move {
|
||||
while let Ok((stream, _peer)) = listener.accept().await {
|
||||
drop(stream);
|
||||
}
|
||||
});
|
||||
|
||||
let result = dial_upstream(
|
||||
"gpt-realtime",
|
||||
"sk-test",
|
||||
Some(&format!("wss://127.0.0.1:{port}")),
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::Network(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_api_key_prefers_param_then_blank_falls_through() {
|
||||
assert_eq!(resolve_api_key(Some("sk-test")).unwrap(), "sk-test");
|
||||
|
|
|
|||
|
|
@ -14,7 +14,9 @@ use tokio_tungstenite::tungstenite::Message;
|
|||
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
|
||||
use tokio_tungstenite::tungstenite::http::HeaderValue;
|
||||
use tokio_tungstenite::tungstenite::http::header::{AUTHORIZATION, HeaderName};
|
||||
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async};
|
||||
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
|
||||
|
||||
use crate::io::tls::connect_upstream;
|
||||
|
||||
use crate::constants::{
|
||||
DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS, DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS,
|
||||
|
|
@ -49,14 +51,14 @@ impl ResponsesWebSocketConnection {
|
|||
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
|
||||
request.headers_mut().insert(header_name, header_value);
|
||||
}
|
||||
let connect = connect_async(request);
|
||||
let connect = connect_upstream(request);
|
||||
let result = match timeout {
|
||||
Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| {
|
||||
Error::Network("Responses WebSocket connection timed out".to_string())
|
||||
})?,
|
||||
None => connect.await,
|
||||
};
|
||||
let (socket, _) = result.map_err(|error| match error {
|
||||
let (socket, _) = result.map_err(|error| match *error {
|
||||
tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http {
|
||||
status: response.status().as_u16(),
|
||||
body: String::new(),
|
||||
|
|
@ -138,13 +140,13 @@ async fn dial_upstream(
|
|||
);
|
||||
let result = tokio::time::timeout(
|
||||
Duration::from_secs(DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS),
|
||||
connect_async(request),
|
||||
connect_upstream(request),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| Error::Network("Responses WebSocket connection timed out".to_string()))?;
|
||||
result
|
||||
.map(|(socket, _)| socket)
|
||||
.map_err(|error| match error {
|
||||
.map_err(|error| match *error {
|
||||
tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http {
|
||||
status: response.status().as_u16(),
|
||||
body: String::new(),
|
||||
|
|
@ -324,6 +326,29 @@ mod tests {
|
|||
use tokio::net::TcpListener;
|
||||
use tokio_tungstenite::accept_async;
|
||||
|
||||
/// The Responses dial has to reach a `wss://` upstream without a process-wide
|
||||
/// crypto provider installed, which is what dialing through `io::tls` buys.
|
||||
#[tokio::test]
|
||||
async fn dial_upstream_over_wss_reports_an_error_instead_of_panicking() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("bind a loopback port");
|
||||
let port = listener
|
||||
.local_addr()
|
||||
.expect("read the bound address")
|
||||
.port();
|
||||
tokio::spawn(async move {
|
||||
while let Ok((stream, _peer)) = listener.accept().await {
|
||||
drop(stream);
|
||||
}
|
||||
});
|
||||
|
||||
let result =
|
||||
dial_upstream("gpt-5", "sk-test", Some(&format!("wss://127.0.0.1:{port}"))).await;
|
||||
|
||||
assert!(matches!(result, Err(Error::Network(_))));
|
||||
}
|
||||
|
||||
async fn websocket_base() -> (String, tokio::task::JoinHandle<()>) {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
|
||||
let address = listener.local_addr().expect("local address");
|
||||
|
|
|
|||
80
litellm-rust/crates/ai-gateway/src/io/tls.rs
Normal file
80
litellm-rust/crates/ai-gateway/src/io/tls.rs
Normal file
|
|
@ -0,0 +1,80 @@
|
|||
//! Outbound WebSocket dials over a TLS config this crate builds once and owns.
|
||||
//!
|
||||
//! `reqwest/rustls-tls` enables `rustls/ring` and `litellm-core`'s `bedrock-auth`
|
||||
//! enables `rustls/aws-lc-rs`, so the bare `ClientConfig::builder()` that
|
||||
//! `tokio-tungstenite` uses when handed no connector panics rather than guess
|
||||
//! between them. Naming ring on a connector of our own settles that for these
|
||||
//! dials without touching the process-wide default, and building the config
|
||||
//! once keeps the platform trust store, which `tokio-tungstenite` would
|
||||
//! otherwise re-read on every dial, off the dial path.
|
||||
|
||||
use std::io;
|
||||
use std::sync::{Arc, OnceLock};
|
||||
|
||||
use rustls::{ClientConfig, RootCertStore};
|
||||
use tokio::net::TcpStream;
|
||||
use tokio_tungstenite::tungstenite::Error;
|
||||
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
|
||||
use tokio_tungstenite::tungstenite::error::TlsError;
|
||||
use tokio_tungstenite::tungstenite::handshake::client::Response;
|
||||
use tokio_tungstenite::{
|
||||
Connector, MaybeTlsStream, WebSocketStream, connect_async_tls_with_config,
|
||||
};
|
||||
|
||||
static TLS_CONFIG: OnceLock<Arc<ClientConfig>> = OnceLock::new();
|
||||
|
||||
fn build_config() -> Result<ClientConfig, Box<Error>> {
|
||||
let native = rustls_native_certs::load_native_certs();
|
||||
let roots = {
|
||||
let mut store = RootCertStore::empty();
|
||||
let (added, _ignored) = store.add_parsable_certificates(native.certs);
|
||||
if added == 0 {
|
||||
return Err(Box::new(Error::Io(io::Error::other(format!(
|
||||
"no usable native root certificates: {:?}",
|
||||
native.errors
|
||||
)))));
|
||||
}
|
||||
store
|
||||
};
|
||||
|
||||
ClientConfig::builder_with_provider(Arc::new(rustls::crypto::ring::default_provider()))
|
||||
.with_safe_default_protocol_versions()
|
||||
.map(|builder| builder.with_root_certificates(roots).with_no_client_auth())
|
||||
.map_err(|error| Box::new(Error::Tls(TlsError::Rustls(error))))
|
||||
}
|
||||
|
||||
fn tls_config() -> Result<Arc<ClientConfig>, Box<Error>> {
|
||||
if let Some(config) = TLS_CONFIG.get() {
|
||||
return Ok(Arc::clone(config));
|
||||
}
|
||||
let built = Arc::new(build_config()?);
|
||||
Ok(Arc::clone(TLS_CONFIG.get_or_init(|| built)))
|
||||
}
|
||||
|
||||
pub(crate) async fn connect_upstream<R>(
|
||||
request: R,
|
||||
) -> Result<(WebSocketStream<MaybeTlsStream<TcpStream>>, Response), Box<Error>>
|
||||
where
|
||||
R: IntoClientRequest + Unpin,
|
||||
{
|
||||
let request = request.into_client_request().map_err(Box::new)?;
|
||||
let connector = match request.uri().scheme_str() {
|
||||
Some("wss") => Some(Connector::Rustls(tls_config()?)),
|
||||
_ => None,
|
||||
};
|
||||
connect_async_tls_with_config(request, None, false, connector)
|
||||
.await
|
||||
.map_err(Box::new)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::build_config;
|
||||
|
||||
#[test]
|
||||
fn builds_a_usable_config_with_both_provider_features_enabled() {
|
||||
let config = build_config().expect("a client config");
|
||||
|
||||
assert!(!config.crypto_provider().cipher_suites.is_empty());
|
||||
}
|
||||
}
|
||||
|
|
@ -265,6 +265,12 @@ impl CallLifecycleHooks<PreparedOcrRequest, PreparedOcrRequest, Value> for OcrLi
|
|||
Box::pin(async move { Ok(request) })
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "success_callback",
|
||||
target = "litellm::function_trace",
|
||||
level = "trace",
|
||||
skip_all
|
||||
)]
|
||||
fn async_log_success_event<'a>(
|
||||
&'a self,
|
||||
context: &'a CallLifecycleContext,
|
||||
|
|
@ -288,6 +294,12 @@ impl CallLifecycleHooks<PreparedOcrRequest, PreparedOcrRequest, Value> for OcrLi
|
|||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "failure_callback",
|
||||
target = "litellm::function_trace",
|
||||
level = "trace",
|
||||
skip_all
|
||||
)]
|
||||
fn async_log_failure_event<'a>(
|
||||
&'a self,
|
||||
context: &'a CallLifecycleContext,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,48 @@
|
|||
//! Guards the wiring, not just the helper: a `wss://` dial through the public
|
||||
//! API has to resolve its own crypto provider, in a test binary where nothing
|
||||
//! has installed a process-wide one, and has to leave it uninstalled.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection;
|
||||
use tokio::net::TcpListener;
|
||||
|
||||
async fn dead_tls_server() -> u16 {
|
||||
let listener = TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("bind a loopback port");
|
||||
let port = listener
|
||||
.local_addr()
|
||||
.expect("read the bound address")
|
||||
.port();
|
||||
|
||||
tokio::spawn(async move {
|
||||
while let Ok((stream, _peer)) = listener.accept().await {
|
||||
drop(stream);
|
||||
}
|
||||
});
|
||||
|
||||
port
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn dialing_wss_returns_an_error_instead_of_panicking() {
|
||||
let port = dead_tls_server().await;
|
||||
|
||||
let result = ResponsesWebSocketConnection::connect_url(
|
||||
&format!("wss://127.0.0.1:{port}/"),
|
||||
&HashMap::new(),
|
||||
Some(Duration::from_secs(10)),
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"a plain TCP server cannot finish a TLS handshake"
|
||||
);
|
||||
assert!(
|
||||
rustls::crypto::CryptoProvider::get_default().is_none(),
|
||||
"the dial settles its provider on its own connector, not process-wide"
|
||||
);
|
||||
}
|
||||
|
|
@ -11,9 +11,13 @@ use litellm_ai_gateway::integrations::custom_logger::{
|
|||
use litellm_ai_gateway::integrations::types::RequestMetadata;
|
||||
use litellm_ai_gateway::ocr::{OcrRequest, ocr};
|
||||
use litellm_core::error::Error;
|
||||
#[cfg(feature = "trace-parity")]
|
||||
use litellm_core::observability::FunctionTrace;
|
||||
use serde_json::{Map, Value, json};
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::{TcpListener, TcpStream};
|
||||
#[cfg(feature = "trace-parity")]
|
||||
use tracing::instrument::WithSubscriber;
|
||||
|
||||
async fn read_http_headers(socket: &mut TcpStream) -> String {
|
||||
let mut request = Vec::new();
|
||||
|
|
@ -320,14 +324,17 @@ async fn ocr_lifecycle_runs_pre_during_and_success_hooks() {
|
|||
GuardrailEventHook::PreCall,
|
||||
GuardrailEventHook::DuringCall,
|
||||
]));
|
||||
let response = ocr(OcrRequest {
|
||||
#[cfg(feature = "trace-parity")]
|
||||
let trace = FunctionTrace::default();
|
||||
let api_base = format!("http://{addr}");
|
||||
let call = ocr(OcrRequest {
|
||||
model: "mistral-ocr-latest",
|
||||
document: json!({
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/doc.pdf"
|
||||
}),
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some(&format!("http://{addr}")),
|
||||
api_base: Some(&api_base),
|
||||
custom_llm_provider: Some("mistral"),
|
||||
extra_headers: None,
|
||||
optional_params: Map::new(),
|
||||
|
|
@ -339,9 +346,10 @@ async fn ocr_lifecycle_runs_pre_during_and_success_hooks() {
|
|||
..Default::default()
|
||||
},
|
||||
litellm_call_id: Some("ocr-call-1"),
|
||||
})
|
||||
.await
|
||||
.expect("ocr request succeeds");
|
||||
});
|
||||
#[cfg(feature = "trace-parity")]
|
||||
let call = call.with_subscriber(trace.dispatcher());
|
||||
let response = call.await.expect("ocr request succeeds");
|
||||
|
||||
assert_eq!(response["pages"][0]["markdown"], "ok");
|
||||
assert_eq!(
|
||||
|
|
@ -359,6 +367,16 @@ async fn ocr_lifecycle_runs_pre_during_and_success_hooks() {
|
|||
error_kind: None,
|
||||
}]
|
||||
);
|
||||
#[cfg(feature = "trace-parity")]
|
||||
assert_eq!(
|
||||
trace
|
||||
.events()
|
||||
.iter()
|
||||
.filter(|event| event.function.ends_with("_callback"))
|
||||
.map(|event| event.function)
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["success_callback"]
|
||||
);
|
||||
|
||||
let request = server.await.expect("server task completes");
|
||||
assert!(request.contains(r#""guarded_pre":true"#), "{request}");
|
||||
|
|
@ -388,14 +406,17 @@ async fn ocr_lifecycle_runs_failure_hook_on_provider_error() {
|
|||
});
|
||||
|
||||
let logger = Arc::new(RecordingOcrLogger::default());
|
||||
let err = ocr(OcrRequest {
|
||||
#[cfg(feature = "trace-parity")]
|
||||
let trace = FunctionTrace::default();
|
||||
let api_base = format!("http://{addr}");
|
||||
let call = ocr(OcrRequest {
|
||||
model: "mistral-ocr-latest",
|
||||
document: json!({
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/doc.pdf"
|
||||
}),
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some(&format!("http://{addr}")),
|
||||
api_base: Some(&api_base),
|
||||
custom_llm_provider: Some("mistral"),
|
||||
extra_headers: None,
|
||||
optional_params: Map::new(),
|
||||
|
|
@ -404,9 +425,10 @@ async fn ocr_lifecycle_runs_failure_hook_on_provider_error() {
|
|||
guardrails: Vec::new(),
|
||||
request_metadata: RequestMetadata::default(),
|
||||
litellm_call_id: Some("ocr-call-2"),
|
||||
})
|
||||
.await
|
||||
.expect_err("provider error propagates");
|
||||
});
|
||||
#[cfg(feature = "trace-parity")]
|
||||
let call = call.with_subscriber(trace.dispatcher());
|
||||
let err = call.await.expect_err("provider error propagates");
|
||||
|
||||
assert!(matches!(err, Error::Http { status: 500, .. }));
|
||||
server.await.expect("server task completes");
|
||||
|
|
@ -421,6 +443,16 @@ async fn ocr_lifecycle_runs_failure_hook_on_provider_error() {
|
|||
error_kind: Some("HttpError".to_string()),
|
||||
}]
|
||||
);
|
||||
#[cfg(feature = "trace-parity")]
|
||||
assert_eq!(
|
||||
trace
|
||||
.events()
|
||||
.iter()
|
||||
.filter(|event| event.function.ends_with("_callback"))
|
||||
.map(|event| event.function)
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["failure_callback"]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use std::fmt::Display;
|
||||
use std::future::Future;
|
||||
|
||||
use litellm_core::observability::{FunctionTrace, FunctionTraceEvent};
|
||||
|
|
@ -6,17 +7,32 @@ use tracing::instrument::WithSubscriber;
|
|||
|
||||
#[derive(Serialize)]
|
||||
pub(crate) struct TracedResponse<T> {
|
||||
response: T,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
response: Option<T>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
error: Option<String>,
|
||||
trace: Vec<FunctionTraceEvent>,
|
||||
}
|
||||
|
||||
pub(crate) async fn capture<T, E>(
|
||||
future: impl Future<Output = Result<T, E>>,
|
||||
) -> Result<TracedResponse<T>, E> {
|
||||
) -> Result<TracedResponse<T>, E>
|
||||
where
|
||||
E: Display,
|
||||
{
|
||||
let trace = FunctionTrace::default();
|
||||
let response = future.with_subscriber(trace.dispatcher()).await?;
|
||||
Ok(TracedResponse {
|
||||
response,
|
||||
trace: trace.events(),
|
||||
let result = future.with_subscriber(trace.dispatcher()).await;
|
||||
let events = trace.events();
|
||||
Ok(match result {
|
||||
Ok(response) => TracedResponse {
|
||||
response: Some(response),
|
||||
error: None,
|
||||
trace: events,
|
||||
},
|
||||
Err(error) => TracedResponse {
|
||||
response: None,
|
||||
error: Some(error.to_string()),
|
||||
trace: events,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -486,10 +486,11 @@ asyncio.run(exercise())
|
|||
let code = CString::new(
|
||||
r#"
|
||||
result = routes.echo("traced")
|
||||
assert result == {
|
||||
"response": "traced",
|
||||
"trace": [{"function": "execute_echo", "depth": 0}],
|
||||
}
|
||||
assert result["response"] == "traced", result
|
||||
assert [event["function"] for event in result["trace"]] == ["execute_echo"], result
|
||||
failure = routes.echo("error")
|
||||
assert failure["error"] == "invalid request: synthetic error", failure
|
||||
assert [event["function"] for event in failure["trace"]] == ["execute_echo"], failure
|
||||
"#,
|
||||
)
|
||||
.expect("Python source should not contain null bytes");
|
||||
|
|
|
|||
|
|
@ -495,6 +495,7 @@ public_model_groups: Optional[List[str]] = None
|
|||
public_agent_groups: Optional[List[str]] = None
|
||||
agent_search_embedding_model: Optional[str] = None
|
||||
mcp_tool_search: Optional[Mapping[str, object]] = None
|
||||
skill_search_embedding_model: Optional[str] = None
|
||||
# Supports both old format (Dict[str, str]) and new format (Dict[str, Dict[str, Any]])
|
||||
# New format: { "displayName": { "url": "...", "index": 0 } }
|
||||
# Old format: { "displayName": "url" } (for backward compatibility)
|
||||
|
|
@ -2001,6 +2002,9 @@ if TYPE_CHECKING:
|
|||
from .llms.hosted_vllm.responses.transformation import (
|
||||
HostedVLLMResponsesAPIConfig as HostedVLLMResponsesAPIConfig,
|
||||
)
|
||||
from .llms.fireworks_ai.responses.transformation import (
|
||||
FireworksAIResponsesAPIConfig as FireworksAIResponsesAPIConfig,
|
||||
)
|
||||
from .llms.github_copilot.chat.transformation import (
|
||||
GithubCopilotConfig as GithubCopilotConfig,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -229,7 +229,7 @@ def _module_attribute(module: ModuleType, attr_name: str) -> object:
|
|||
return attribute["value"]
|
||||
|
||||
|
||||
def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], category: str) -> object:
|
||||
def _generic_lazy_import(name: str, import_map: Mapping[str, tuple[str, str]], category: str) -> object:
|
||||
"""
|
||||
Generic function that handles lazy importing for most attributes.
|
||||
|
||||
|
|
|
|||
|
|
@ -237,6 +237,7 @@ LLM_CONFIG_NAMES: Final = (
|
|||
"XAIResponsesAPIConfig",
|
||||
"LiteLLMProxyResponsesAPIConfig",
|
||||
"HostedVLLMResponsesAPIConfig",
|
||||
"FireworksAIResponsesAPIConfig",
|
||||
"VolcEngineResponsesAPIConfig",
|
||||
"PerplexityResponsesConfig",
|
||||
"DatabricksResponsesAPIConfig",
|
||||
|
|
@ -957,6 +958,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
|
|||
".llms.hosted_vllm.responses.transformation",
|
||||
"HostedVLLMResponsesAPIConfig",
|
||||
),
|
||||
"FireworksAIResponsesAPIConfig": (
|
||||
".llms.fireworks_ai.responses.transformation",
|
||||
"FireworksAIResponsesAPIConfig",
|
||||
),
|
||||
"VolcEngineResponsesAPIConfig": (
|
||||
".llms.volcengine.responses.transformation",
|
||||
"VolcEngineResponsesAPIConfig",
|
||||
|
|
|
|||
|
|
@ -25,6 +25,27 @@ class BatchCostUsageResult:
|
|||
failed_requests: int
|
||||
|
||||
|
||||
_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"})
|
||||
_TERMINAL_BATCH_STATUSES: Final = _COMPLETED_BATCH_STATUSES | frozenset({"failed", "cancelled", "expired"})
|
||||
|
||||
|
||||
def batch_cost_is_final(batch: Batch) -> bool:
|
||||
"""Whether this retrieve of the batch is the one to account its cost from.
|
||||
|
||||
A batch still in flight has nothing to price, and a "completed" batch can report
|
||||
no output_file_id for a moment before the output populates; pricing either records
|
||||
$0 under the batch's single spend row and pins it there. Final means a completed
|
||||
batch whose output file has arrived or whose counts prove no line succeeded, or
|
||||
any other terminal status (failed, cancelled, expired).
|
||||
"""
|
||||
if batch.status not in _TERMINAL_BATCH_STATUSES:
|
||||
return False
|
||||
if batch.status not in _COMPLETED_BATCH_STATUSES or batch.output_file_id is not None:
|
||||
return True
|
||||
request_counts: Final = batch.request_counts
|
||||
return request_counts is not None and request_counts.total > 0 and request_counts.completed == 0
|
||||
|
||||
|
||||
async def calculate_batch_cost_and_usage(
|
||||
file_content_dictionary: list[dict],
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],
|
||||
|
|
|
|||
|
|
@ -5,8 +5,11 @@ Handler for transforming /chat/completions api requests to litellm.responses req
|
|||
import json
|
||||
import os
|
||||
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, Union, cast, get_args
|
||||
|
||||
from openai.types.chat import ChatCompletion
|
||||
from openai.types.responses import Response
|
||||
from openai.types.responses.custom_tool_param import CustomToolParam
|
||||
from openai.types.responses.response_input_param import (
|
||||
FunctionCallOutput,
|
||||
|
|
@ -33,7 +36,7 @@ from litellm.responses.sse_output_recovery import (
|
|||
record_output_item_chunk,
|
||||
record_output_text_chunk,
|
||||
)
|
||||
from litellm.responses.utils import normalize_responses_api_stream_options
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils, normalize_responses_api_stream_options
|
||||
from litellm.types.llms.openai import (
|
||||
REASONING_EFFORT,
|
||||
ChatCompletionAnnotation,
|
||||
|
|
@ -43,6 +46,7 @@ from litellm.types.llms.openai import (
|
|||
ChatCompletionToolParamFunctionChunk,
|
||||
Reasoning,
|
||||
ResponsesAPIOptionalRequestParams,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
|
||||
|
|
@ -54,7 +58,7 @@ if TYPE_CHECKING:
|
|||
)
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm import LiteLLMLoggingObj, ModelResponse
|
||||
from litellm import LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.types.llms.openai import (
|
||||
ALL_RESPONSES_API_TOOL_PARAMS,
|
||||
|
|
@ -69,6 +73,28 @@ if TYPE_CHECKING:
|
|||
from litellm.types.utils import Choices
|
||||
|
||||
|
||||
_CHAT_COMPLETION_FIELDS: Final = frozenset((*ModelResponse.model_fields, "usage"))
|
||||
_RESPONSES_API_ONLY_FIELDS: Final = frozenset((*Response.model_fields, *ResponsesAPIResponse.model_fields)) - frozenset(
|
||||
ChatCompletion.model_fields
|
||||
)
|
||||
|
||||
|
||||
def _provider_metadata(response_fields: Mapping[str, object] | None) -> Mapping[str, object]:
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key, value in (response_fields.items() if response_fields else ())
|
||||
if value is not None and key not in _CHAT_COMPLETION_FIELDS and key not in _RESPONSES_API_ONLY_FIELDS
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _upstream_response_id(response_id: str | None) -> str | None:
|
||||
if response_id is None:
|
||||
return None
|
||||
return ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(response_id)
|
||||
|
||||
|
||||
class _ReasoningSummaryText(TypedDict):
|
||||
type: str
|
||||
text: str
|
||||
|
|
@ -904,6 +930,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_response.usage),
|
||||
)
|
||||
|
||||
model_response.id = _upstream_response_id(raw_response.id) or raw_response.id
|
||||
for key, value in _provider_metadata(raw_response.model_extra).items():
|
||||
setattr(model_response, key, value)
|
||||
|
||||
# Preserve hidden params from the ResponsesAPIResponse, especially the headers
|
||||
# which contain important provider information like x-request-id
|
||||
raw_response_hidden_params: Final = getattr(raw_response, "_hidden_params", {})
|
||||
|
|
@ -1359,14 +1389,16 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
if event_type == "response.created":
|
||||
# Initial response creation event
|
||||
verbose_logger.debug("Chat provider: response.created -> %s", parsed_chunk)
|
||||
created_response: Final = parsed_chunk.get("response")
|
||||
return ModelResponseStream(
|
||||
id=_upstream_response_id(created_response.get("id")) if created_response else None,
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(content=""),
|
||||
finish_reason=None,
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
elif event_type == "response.output_item.added":
|
||||
# New output item added
|
||||
|
|
@ -1534,6 +1566,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
from litellm.responses.utils import ResponseAPILoggingUtils
|
||||
|
||||
usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(response_data.get("usage"))
|
||||
provider_metadata: Final = _provider_metadata(response_data)
|
||||
return ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
|
|
@ -1546,6 +1579,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
)
|
||||
],
|
||||
usage=usage,
|
||||
provider_specific_fields=dict(provider_metadata) or None, # mutable-ok: field is typed dict
|
||||
)
|
||||
else:
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -40,7 +40,7 @@ ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
|
|||
"router_general_settings",
|
||||
"ignore_invalid_deployments",
|
||||
"fallback_access_check",
|
||||
"heuristic_v2_router_limit",
|
||||
"auto_router_capability_limit",
|
||||
}
|
||||
)
|
||||
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
|
||||
|
|
@ -89,6 +89,7 @@ LITELLM_MAX_STREAMING_DURATION_SECONDS: Final = (
|
|||
# Data URIs exceeding this are replaced with a size placeholder.
|
||||
# Set to 0 to disable truncation.
|
||||
MAX_BASE64_LENGTH_FOR_LOGGING: Final = int(os.getenv("MAX_BASE64_LENGTH_FOR_LOGGING", 64))
|
||||
BASE64_TRUNCATION_OFFLOAD_THRESHOLD_CHARS: Final = 256 * 1024
|
||||
REDACTED_BY_LITELLM: Final = "redacted-by-litellm"
|
||||
# in-memory stand-in handed to provider converters for redacted arguments; never stored
|
||||
REDACTED_TOOL_CALL_ARGUMENTS_PLACEHOLDER: Final = "{}"
|
||||
|
|
@ -215,6 +216,9 @@ MAX_CALLBACKS: Final = get_env_int("LITELLM_MAX_CALLBACKS", 100)
|
|||
# so the deployment-level hook does not re-run them for the same request
|
||||
PRE_CALL_EXECUTED_GUARDRAILS_KEY: Final = "_pre_call_executed_guardrails"
|
||||
|
||||
# Attribute stamped on log_guardrail_information wrappers so __init_subclass__ does not wrap them again
|
||||
LOGS_GUARDRAIL_INFORMATION_MARKER: Final = "_litellm_logs_guardrail_information"
|
||||
|
||||
# Generic fallback for unknown models
|
||||
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET: Final = int(
|
||||
os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET", 128)
|
||||
|
|
@ -1898,6 +1902,16 @@ HTTP_FRAMING_HEADERS: Final[frozenset[str]] = frozenset(
|
|||
}
|
||||
)
|
||||
|
||||
PROVIDER_REQUEST_ID_HEADERS: Final[tuple[str, ...]] = (
|
||||
"x-amzn-requestid",
|
||||
"x-request-id",
|
||||
"request-id",
|
||||
"x-ms-request-id",
|
||||
"apim-request-id",
|
||||
"x-goog-request-id",
|
||||
"cf-ray",
|
||||
)
|
||||
|
||||
# Browser-facing security headers that a malicious or misconfigured upstream
|
||||
# provider must not be able to set on the proxy's own response.
|
||||
BROWSER_SECURITY_HEADERS: Final[frozenset[str]] = frozenset(
|
||||
|
|
|
|||
|
|
@ -81,6 +81,7 @@ from litellm.llms.together_ai.cost_calculator import (
|
|||
get_model_params_and_category,
|
||||
has_together_registry_pricing,
|
||||
)
|
||||
from litellm.llms.vertex_ai.common_utils import get_vertex_ai_lyria_generation_cost
|
||||
from litellm.llms.vertex_ai.cost_calculator import (
|
||||
cost_per_character as google_cost_per_character,
|
||||
)
|
||||
|
|
@ -496,6 +497,13 @@ def cost_per_token(
|
|||
|
||||
# see this https://learn.microsoft.com/en-us/azure/ai-services/openai/concepts/models
|
||||
if call_type == "speech" or call_type == "aspeech":
|
||||
lyria_generation_cost: Final = (
|
||||
get_vertex_ai_lyria_generation_cost(model=model_without_prefix)
|
||||
if custom_llm_provider in ("vertex_ai", "vertex_ai_beta")
|
||||
else None
|
||||
)
|
||||
if lyria_generation_cost is not None:
|
||||
return 0.0, lyria_generation_cost
|
||||
speech_model_info = litellm.get_model_info(model=model_without_prefix, custom_llm_provider=custom_llm_provider)
|
||||
cost_metric: Final = select_cost_metric_for_model(speech_model_info)
|
||||
prompt_cost: float = 0.0
|
||||
|
|
|
|||
|
|
@ -338,6 +338,7 @@ class Timeout(openai.APITimeoutError):
|
|||
num_retries: int | None = None,
|
||||
headers: dict | None = None,
|
||||
exception_status_code: int | None = None,
|
||||
response: httpx.Response | None = None,
|
||||
):
|
||||
request: Final = httpx.Request(
|
||||
method="POST",
|
||||
|
|
@ -352,6 +353,8 @@ class Timeout(openai.APITimeoutError):
|
|||
self.max_retries = max_retries
|
||||
self.num_retries = num_retries
|
||||
self.headers = headers
|
||||
if response is not None:
|
||||
self.response = response
|
||||
|
||||
# custom function to convert to str
|
||||
def __str__(self):
|
||||
|
|
|
|||
|
|
@ -16,23 +16,32 @@ import asyncio
|
|||
import os
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from typing import Final, TypeVar
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.batch_utils import (
|
||||
BatchSendCancelled,
|
||||
send_batch_with_413_split,
|
||||
undelivered_after_http_error,
|
||||
)
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
MaskedHTTPStatusError,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.integrations.azure_sentinel import AZURE_SENTINEL_MAX_PAYLOAD_SIZE_BYTES
|
||||
from litellm.types.utils import StandardAuditLogPayload, StandardLoggingPayload
|
||||
|
||||
DEFAULT_AZURE_AUTHORITY_HOST: Final = "https://login.microsoftonline.com"
|
||||
DEFAULT_AZURE_MONITOR_SCOPE: Final = "https://monitor.azure.com/.default"
|
||||
|
||||
_QueuedPayload = TypeVar("_QueuedPayload", StandardLoggingPayload, StandardAuditLogPayload)
|
||||
|
||||
MONITOR_SCOPE_BY_AUTHORITY_HOST: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
"login.microsoftonline.com": DEFAULT_AZURE_MONITOR_SCOPE,
|
||||
|
|
@ -153,6 +162,8 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
asyncio.create_task(self.periodic_flush())
|
||||
self.log_queue: list[StandardLoggingPayload] = []
|
||||
self.audit_log_queue: list[StandardAuditLogPayload] = []
|
||||
self.logs_awaiting_retry = False
|
||||
self.audit_logs_awaiting_retry = False
|
||||
|
||||
@staticmethod
|
||||
def _normalize_authority_host(authority_host: str) -> str:
|
||||
|
|
@ -245,8 +256,8 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
|
||||
self.log_queue.append(standard_logging_payload)
|
||||
|
||||
if len(self.log_queue) >= self.batch_size:
|
||||
await self.async_send_batch()
|
||||
if len(self.log_queue) >= self.batch_size and not self.logs_awaiting_retry:
|
||||
await self._threshold_send_logs()
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Azure Sentinel Layer Error - %s\n%s", e, traceback.format_exc())
|
||||
|
|
@ -275,8 +286,8 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
|
||||
self.log_queue.append(standard_logging_payload)
|
||||
|
||||
if len(self.log_queue) >= self.batch_size:
|
||||
await self.async_send_batch()
|
||||
if len(self.log_queue) >= self.batch_size and not self.logs_awaiting_retry:
|
||||
await self._threshold_send_logs()
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Azure Sentinel Layer Error - %s\n%s", e, traceback.format_exc())
|
||||
|
|
@ -298,12 +309,24 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
|
||||
self.audit_log_queue.append(audit_log)
|
||||
|
||||
if len(self.audit_log_queue) >= self.batch_size:
|
||||
await self.async_send_audit_batch()
|
||||
if len(self.audit_log_queue) >= self.batch_size and not self.audit_logs_awaiting_retry:
|
||||
await self._threshold_send_audit_logs()
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Azure Sentinel Audit Log Layer Error - %s\n%s", e, traceback.format_exc())
|
||||
|
||||
async def _threshold_send_logs(self) -> None:
|
||||
async with self.flush_lock:
|
||||
if self.logs_awaiting_retry or len(self.log_queue) < self.batch_size:
|
||||
return
|
||||
await self.async_send_batch()
|
||||
|
||||
async def _threshold_send_audit_logs(self) -> None:
|
||||
async with self.flush_lock:
|
||||
if self.audit_logs_awaiting_retry or len(self.audit_log_queue) < self.batch_size:
|
||||
return
|
||||
await self.async_send_audit_batch()
|
||||
|
||||
async def async_send_batch(self):
|
||||
"""
|
||||
Sends the batch of logs to Azure Monitor Logs Ingestion API
|
||||
|
|
@ -311,67 +334,110 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
Raises:
|
||||
Raises a NON Blocking verbose_logger.exception if an error occurs
|
||||
"""
|
||||
await self._async_send_batch_to_api(
|
||||
log_queue=self.log_queue,
|
||||
api_endpoint=self.api_endpoint,
|
||||
log_type="logs",
|
||||
)
|
||||
batch_to_send: Final = tuple(self.log_queue)
|
||||
self.log_queue = [] # mutable-ok: queue ownership is detached before the async send
|
||||
try:
|
||||
undelivered: Final = await self._async_send_batch_to_api(
|
||||
log_queue=batch_to_send,
|
||||
api_endpoint=self.api_endpoint,
|
||||
log_type="logs",
|
||||
)
|
||||
except BatchSendCancelled as cancelled:
|
||||
self.log_queue = self._requeue(cancelled.undelivered, self.log_queue, "logs")
|
||||
self.logs_awaiting_retry = bool(self.log_queue)
|
||||
raise asyncio.CancelledError() from cancelled
|
||||
except asyncio.CancelledError:
|
||||
self.log_queue = self._requeue(batch_to_send, self.log_queue, "logs")
|
||||
self.logs_awaiting_retry = bool(self.log_queue)
|
||||
raise
|
||||
self.log_queue = self._requeue(undelivered, self.log_queue, "logs")
|
||||
self.logs_awaiting_retry = bool(undelivered) and bool(self.log_queue)
|
||||
|
||||
async def async_send_audit_batch(self):
|
||||
"""
|
||||
Sends the batch of audit logs to Azure Monitor Logs Ingestion API
|
||||
"""
|
||||
await self._async_send_batch_to_api(
|
||||
log_queue=self.audit_log_queue,
|
||||
api_endpoint=self.audit_api_endpoint,
|
||||
log_type="audit logs",
|
||||
batch_to_send: Final = tuple(self.audit_log_queue)
|
||||
self.audit_log_queue = [] # mutable-ok: queue ownership is detached before the async send
|
||||
try:
|
||||
undelivered: Final = await self._async_send_batch_to_api(
|
||||
log_queue=batch_to_send,
|
||||
api_endpoint=self.audit_api_endpoint,
|
||||
log_type="audit logs",
|
||||
)
|
||||
except BatchSendCancelled as cancelled:
|
||||
self.audit_log_queue = self._requeue(cancelled.undelivered, self.audit_log_queue, "audit logs")
|
||||
self.audit_logs_awaiting_retry = bool(self.audit_log_queue)
|
||||
raise asyncio.CancelledError() from cancelled
|
||||
except asyncio.CancelledError:
|
||||
self.audit_log_queue = self._requeue(batch_to_send, self.audit_log_queue, "audit logs")
|
||||
self.audit_logs_awaiting_retry = bool(self.audit_log_queue)
|
||||
raise
|
||||
self.audit_log_queue = self._requeue(undelivered, self.audit_log_queue, "audit logs")
|
||||
self.audit_logs_awaiting_retry = bool(undelivered) and bool(self.audit_log_queue)
|
||||
|
||||
def _requeue(
|
||||
self,
|
||||
undelivered: tuple[_QueuedPayload, ...],
|
||||
queue: list[_QueuedPayload],
|
||||
log_type: str,
|
||||
) -> list[_QueuedPayload]:
|
||||
merged: Final = [*undelivered, *queue] # mutable-ok: queue trimming returns a mutable logger queue
|
||||
overflow: Final = len(merged) - self.max_queue_size
|
||||
if overflow <= 0:
|
||||
return merged
|
||||
|
||||
verbose_logger.warning(
|
||||
"Azure Sentinel: %s queue exceeded max_queue_size=%s, dropped %s oldest records",
|
||||
log_type,
|
||||
self.max_queue_size,
|
||||
overflow,
|
||||
)
|
||||
return merged[overflow:]
|
||||
|
||||
async def _async_send_batch_to_api(
|
||||
self,
|
||||
log_queue: list[StandardLoggingPayload | StandardAuditLogPayload],
|
||||
log_queue: tuple[_QueuedPayload, ...],
|
||||
api_endpoint: str,
|
||||
log_type: str,
|
||||
) -> None:
|
||||
) -> tuple[_QueuedPayload, ...]:
|
||||
if not log_queue:
|
||||
return ()
|
||||
|
||||
verbose_logger.debug("Azure Sentinel - about to flush %s %s", len(log_queue), log_type)
|
||||
try:
|
||||
if not log_queue:
|
||||
return
|
||||
|
||||
verbose_logger.debug("Azure Sentinel - about to flush %s %s", len(log_queue), log_type)
|
||||
|
||||
# Get OAuth2 token
|
||||
bearer_token: Final = await self._get_oauth_token()
|
||||
except MaskedHTTPStatusError as e:
|
||||
return undelivered_after_http_error(log_queue, e.status_code, "Azure Sentinel OAuth token", str(e))
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Azure Sentinel Error getting OAuth token - %s", e)
|
||||
return tuple(log_queue)
|
||||
|
||||
# Convert log queue to JSON array format expected by Logs Ingestion API
|
||||
# Each log entry should be a JSON object in the array
|
||||
body: Final = safe_dumps(log_queue)
|
||||
headers: Final = {
|
||||
"Authorization": f"Bearer {bearer_token}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
# Set headers for Logs Ingestion API
|
||||
headers: Final = {
|
||||
"Authorization": f"Bearer {bearer_token}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
# Send the request
|
||||
response = await self.async_httpx_client.post(url=api_endpoint, data=body.encode("utf-8"), headers=headers)
|
||||
|
||||
if response.status_code not in [200, 204]:
|
||||
verbose_logger.error(
|
||||
"Azure Sentinel API error: status_code=%s, response=%s",
|
||||
response.status_code,
|
||||
response.text,
|
||||
)
|
||||
raise Exception(f"Failed to send logs to Azure Sentinel: {response.status_code} - {response.text}")
|
||||
|
||||
verbose_logger.debug(
|
||||
"Azure Sentinel: Response from API status_code: %s",
|
||||
response.status_code,
|
||||
async def _send_batch(batch: Sequence[_QueuedPayload]):
|
||||
body: Final = safe_dumps(batch)
|
||||
return await self.async_httpx_client.post(
|
||||
url=api_endpoint,
|
||||
data=body.encode("utf-8"),
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Azure Sentinel Error sending batch API - %s\n%s", e, traceback.format_exc())
|
||||
finally:
|
||||
log_queue.clear()
|
||||
return await send_batch_with_413_split(
|
||||
batch=log_queue,
|
||||
send_batch=_send_batch,
|
||||
exceeds_limits=lambda batch: (
|
||||
len(batch) > self.batch_size
|
||||
or len(safe_dumps(batch).encode("utf-8")) > AZURE_SENTINEL_MAX_PAYLOAD_SIZE_BYTES
|
||||
),
|
||||
success_status_codes=frozenset({200, 204}),
|
||||
integration_name="Azure Sentinel",
|
||||
drop_error_message="Azure Sentinel API Error - Payload too large for a single record",
|
||||
non_success_handler=undelivered_after_http_error,
|
||||
)
|
||||
|
||||
async def flush_queue(self):
|
||||
if self.flush_lock is None:
|
||||
|
|
|
|||
160
litellm/integrations/batch_utils.py
Normal file
160
litellm/integrations/batch_utils.py
Normal file
|
|
@ -0,0 +1,160 @@
|
|||
import asyncio
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from typing import Final, Generic, TypeVar
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError
|
||||
|
||||
_BatchItem = TypeVar("_BatchItem")
|
||||
|
||||
_RETRYABLE_CLIENT_STATUS_CODES: Final = frozenset({408, 429})
|
||||
|
||||
|
||||
def is_retryable_status(status_code: int) -> bool:
|
||||
return not 400 <= status_code < 500 or status_code in _RETRYABLE_CLIENT_STATUS_CODES
|
||||
|
||||
|
||||
def undelivered_after_http_error(
|
||||
batch: Sequence[_BatchItem],
|
||||
status_code: int,
|
||||
integration_name: str,
|
||||
detail: str,
|
||||
) -> tuple[_BatchItem, ...]:
|
||||
"""The records to requeue after a non-2xx: all of them on a status a retry can clear, none on
|
||||
a 4xx that would only repeat, since retaining those retries a misconfiguration forever."""
|
||||
if is_retryable_status(status_code):
|
||||
verbose_logger.error(
|
||||
"%s API error: status_code=%s, will retry %s records - %s",
|
||||
integration_name,
|
||||
status_code,
|
||||
len(batch),
|
||||
detail,
|
||||
)
|
||||
return tuple(batch)
|
||||
verbose_logger.error(
|
||||
"%s API error: status_code=%s is not retryable, dropped %s records - %s",
|
||||
integration_name,
|
||||
status_code,
|
||||
len(batch),
|
||||
detail,
|
||||
)
|
||||
return ()
|
||||
|
||||
|
||||
def requeue_after_http_error(
|
||||
batch: Sequence[_BatchItem],
|
||||
status_code: int,
|
||||
integration_name: str,
|
||||
detail: str,
|
||||
) -> tuple[_BatchItem, ...]:
|
||||
verbose_logger.error(
|
||||
"%s API error: status_code=%s, will retry %s records - %s",
|
||||
integration_name,
|
||||
status_code,
|
||||
len(batch),
|
||||
detail,
|
||||
)
|
||||
return tuple(batch)
|
||||
|
||||
|
||||
class BatchSendCancelled(asyncio.CancelledError, Generic[_BatchItem]):
|
||||
"""Cancellation of a batch send, carrying only the records the destination never accepted.
|
||||
|
||||
A batch split under the size cap is delivered in pieces, so requeueing all of it after a
|
||||
cancellation partway through would send the accepted pieces a second time.
|
||||
"""
|
||||
|
||||
def __init__(self, undelivered: tuple[_BatchItem, ...]) -> None:
|
||||
super().__init__()
|
||||
self.undelivered: Final = undelivered
|
||||
|
||||
|
||||
async def _keep_the_remainder_on_cancel(
|
||||
send: Awaitable[tuple[_BatchItem, ...]],
|
||||
remainder: Sequence[_BatchItem],
|
||||
) -> tuple[_BatchItem, ...]:
|
||||
try:
|
||||
return await send
|
||||
except BatchSendCancelled as cancelled:
|
||||
raise BatchSendCancelled((*cancelled.undelivered, *remainder)) from cancelled
|
||||
|
||||
|
||||
async def send_batch_with_413_split(
|
||||
batch: Sequence[_BatchItem],
|
||||
send_batch: Callable[[Sequence[_BatchItem]], Awaitable[httpx.Response]],
|
||||
exceeds_limits: Callable[[Sequence[_BatchItem]], bool],
|
||||
success_status_codes: frozenset[int],
|
||||
integration_name: str,
|
||||
drop_error_message: str,
|
||||
non_success_handler: Callable[
|
||||
[Sequence[_BatchItem], int, str, str], tuple[_BatchItem, ...]
|
||||
] = requeue_after_http_error,
|
||||
) -> tuple[_BatchItem, ...]:
|
||||
async def _halve() -> tuple[_BatchItem, ...]:
|
||||
midpoint: Final = len(batch) // 2
|
||||
left_batch: Final = batch[:midpoint]
|
||||
right_batch: Final = batch[midpoint:]
|
||||
left_undelivered: Final = await _keep_the_remainder_on_cancel(
|
||||
send_batch_with_413_split(
|
||||
batch=left_batch,
|
||||
send_batch=send_batch,
|
||||
exceeds_limits=exceeds_limits,
|
||||
success_status_codes=success_status_codes,
|
||||
integration_name=integration_name,
|
||||
drop_error_message=drop_error_message,
|
||||
non_success_handler=non_success_handler,
|
||||
),
|
||||
right_batch,
|
||||
)
|
||||
if left_undelivered:
|
||||
return (*left_undelivered, *right_batch)
|
||||
return await send_batch_with_413_split(
|
||||
batch=right_batch,
|
||||
send_batch=send_batch,
|
||||
exceeds_limits=exceeds_limits,
|
||||
success_status_codes=success_status_codes,
|
||||
integration_name=integration_name,
|
||||
drop_error_message=drop_error_message,
|
||||
non_success_handler=non_success_handler,
|
||||
)
|
||||
|
||||
async def _handle_413() -> tuple[_BatchItem, ...]:
|
||||
if len(batch) == 1:
|
||||
verbose_logger.error(drop_error_message)
|
||||
return ()
|
||||
return await _halve()
|
||||
|
||||
if not batch:
|
||||
return ()
|
||||
|
||||
try:
|
||||
oversized: Final = exceeds_limits(batch)
|
||||
except Exception as e: # noqa: BLE001 # any record that cannot be serialized is isolated and dropped alone
|
||||
if len(batch) > 1:
|
||||
return await _halve()
|
||||
verbose_logger.exception("%s dropped a record that cannot be serialized - %s", integration_name, e)
|
||||
return ()
|
||||
if oversized and len(batch) > 1:
|
||||
return await _halve()
|
||||
|
||||
try:
|
||||
response: Final = await send_batch(batch)
|
||||
except MaskedHTTPStatusError as e:
|
||||
if e.status_code == 413:
|
||||
return await _handle_413()
|
||||
return non_success_handler(batch, e.status_code, integration_name, str(e))
|
||||
except asyncio.CancelledError as cancelled:
|
||||
raise BatchSendCancelled(tuple(batch)) from cancelled
|
||||
except Exception as e:
|
||||
verbose_logger.exception("%s Error sending batch API - %s", integration_name, e)
|
||||
return tuple(batch)
|
||||
|
||||
if response.status_code == 413:
|
||||
return await _handle_413()
|
||||
if response.status_code not in success_status_codes:
|
||||
return non_success_handler(batch, response.status_code, integration_name, response.text)
|
||||
|
||||
verbose_logger.debug("%s delivered %s records, status_code=%s", integration_name, len(batch), response.status_code)
|
||||
return ()
|
||||
|
|
@ -97,7 +97,11 @@ class CloudZeroStreamer:
|
|||
continue
|
||||
|
||||
# Convert lists back to DataFrames
|
||||
return {date_key: pl.DataFrame(records) for date_key, records in daily_batches.items() if records}
|
||||
return {
|
||||
date_key: pl.DataFrame(records, infer_schema_length=None)
|
||||
for date_key, records in daily_batches.items()
|
||||
if records
|
||||
}
|
||||
|
||||
def _parse_and_convert_timestamp(self, timestamp_str: str) -> datetime:
|
||||
"""Parse timestamp string and convert to UTC."""
|
||||
|
|
|
|||
|
|
@ -95,7 +95,7 @@ class CBFTransformer:
|
|||
if len(cbf_data) > 0:
|
||||
console.print(f"[green]✓ Successfully transformed {len(cbf_data):,} records[/green]")
|
||||
|
||||
return pl.DataFrame(cbf_data)
|
||||
return pl.DataFrame(cbf_data, infer_schema_length=None)
|
||||
|
||||
def _create_cbf_record(self, row: dict[str, object]) -> CBFRecord:
|
||||
"""Create a single CBF record from LiteLLM daily spend row."""
|
||||
|
|
|
|||
|
|
@ -46,6 +46,7 @@ dc: Final = DualCache()
|
|||
|
||||
from litellm.constants import (
|
||||
GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS,
|
||||
LOGS_GUARDRAIL_INFORMATION_MARKER,
|
||||
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
|
||||
)
|
||||
from litellm.exceptions import (
|
||||
|
|
@ -151,6 +152,13 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
records_own_guardrail_information: ClassVar[bool] = False
|
||||
|
||||
def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks
|
||||
super().__init_subclass__(**kwargs)
|
||||
own_apply_guardrail: Final = cls.__dict__.get("apply_guardrail")
|
||||
if own_apply_guardrail is None or LOGS_GUARDRAIL_INFORMATION_MARKER in vars(own_apply_guardrail):
|
||||
return
|
||||
cls.apply_guardrail = log_guardrail_information(own_apply_guardrail)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
guardrail_name: str | None = None,
|
||||
|
|
@ -940,6 +948,23 @@ class CustomGuardrail(CustomLogger):
|
|||
"""
|
||||
return False
|
||||
|
||||
def _suppressed_by_auto_router_compression(self) -> bool:
|
||||
"""True when an auto router's own compression policy suppresses this guardrail.
|
||||
|
||||
Reads request-scoped state set by `arm_pre_call`, never request metadata. The
|
||||
caller controls metadata, and metadata reaches spend logs the caller can read,
|
||||
so a suppression list carried there would be one a request could replay to
|
||||
switch off a PII or content-filter guardrail for itself.
|
||||
"""
|
||||
name: Final = self.guardrail_name
|
||||
if not name:
|
||||
return False
|
||||
from litellm.proxy.guardrails.auto_router_compression import (
|
||||
suppressed_compression_guardrails,
|
||||
)
|
||||
|
||||
return name in suppressed_compression_guardrails()
|
||||
|
||||
def should_run_guardrail(
|
||||
self,
|
||||
data,
|
||||
|
|
@ -948,6 +973,9 @@ class CustomGuardrail(CustomLogger):
|
|||
"""
|
||||
Returns True if the guardrail should be run on the event_type
|
||||
"""
|
||||
if self._suppressed_by_auto_router_compression():
|
||||
return False
|
||||
|
||||
requested_guardrails: Final = self.get_guardrail_from_metadata(data)
|
||||
disable_global_guardrail: Final = self.get_disable_global_guardrail(data)
|
||||
opted_out_global_guardrails: Final = self.get_opted_out_global_guardrails_from_metadata(data)
|
||||
|
|
@ -1559,4 +1587,5 @@ def log_guardrail_information(func):
|
|||
return async_wrapper(*args, **kwargs)
|
||||
return sync_wrapper(*args, **kwargs)
|
||||
|
||||
vars(wrapper)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True # rebind-ok: stamps the wrapper this call just built
|
||||
return wrapper
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ from typing_extensions import ReadOnly, TypedDict
|
|||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.integrations.batch_utils import BatchSendCancelled, requeue_after_http_error, send_batch_with_413_split
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
from litellm.integrations.datadog.datadog_handler import (
|
||||
get_datadog_base_url_from_env,
|
||||
|
|
@ -43,7 +44,6 @@ from litellm.integrations.datadog.datadog_mock_client import (
|
|||
)
|
||||
from litellm.litellm_core_utils.dd_tracing import tracer
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
MaskedHTTPStatusError,
|
||||
_get_httpx_client,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
|
|
@ -396,6 +396,9 @@ class DataDogLogger(
|
|||
if self.is_mock_mode:
|
||||
verbose_logger.debug("[DATADOG MOCK] Batch of %s events successfully mocked", len(batch_to_send))
|
||||
|
||||
except BatchSendCancelled as cancelled:
|
||||
self.log_queue = list(cancelled.undelivered) + self.log_queue # mutable-ok: logger queue remains appendable
|
||||
raise asyncio.CancelledError() from cancelled
|
||||
except Exception as e:
|
||||
self.log_queue = batch_to_send + self.log_queue
|
||||
verbose_logger.exception("Datadog Error sending batch API - %s\n%s", e, traceback.format_exc())
|
||||
|
|
@ -413,53 +416,16 @@ class DataDogLogger(
|
|||
that could not be delivered because of a non-413 (transient) error, so the caller
|
||||
re-queues only those and never the events already accepted by Datadog.
|
||||
"""
|
||||
pending: Final[list[list]] = [batch]
|
||||
while pending:
|
||||
chunk = pending.pop()
|
||||
if not chunk:
|
||||
continue
|
||||
if len(chunk) > 1 and self._exceeds_intake_limits(chunk):
|
||||
mid = len(chunk) // 2
|
||||
pending.append(chunk[mid:])
|
||||
pending.append(chunk[:mid])
|
||||
continue
|
||||
try:
|
||||
response = await self.async_send_compressed_data(chunk)
|
||||
except Exception as e:
|
||||
if isinstance(e, MaskedHTTPStatusError) and e.status_code == 413:
|
||||
response = e.response
|
||||
else:
|
||||
verbose_logger.exception("Datadog Error sending batch API - %s", e)
|
||||
return self._undelivered(chunk, pending)
|
||||
|
||||
if response.status_code == 413:
|
||||
if len(chunk) == 1:
|
||||
verbose_logger.error(DD_ERRORS.DATADOG_413_ERROR.value)
|
||||
continue
|
||||
mid = len(chunk) // 2
|
||||
pending.append(chunk[mid:])
|
||||
pending.append(chunk[:mid])
|
||||
continue
|
||||
|
||||
if response.status_code != 202:
|
||||
verbose_logger.error(
|
||||
"Datadog: unexpected response status_code=%s, text=%s",
|
||||
response.status_code,
|
||||
response.text,
|
||||
)
|
||||
return self._undelivered(chunk, pending)
|
||||
|
||||
verbose_logger.debug(
|
||||
"Datadog: delivered %s events, status_code=%s, text=%s",
|
||||
len(chunk),
|
||||
response.status_code,
|
||||
response.text,
|
||||
)
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _undelivered(chunk: list, pending: list[list]) -> list:
|
||||
return chunk + [event for remaining in reversed(pending) for event in remaining]
|
||||
undelivered: Final = await send_batch_with_413_split(
|
||||
batch=batch,
|
||||
send_batch=self.async_send_compressed_data,
|
||||
exceeds_limits=self._exceeds_intake_limits,
|
||||
success_status_codes=frozenset({202}),
|
||||
integration_name="Datadog",
|
||||
drop_error_message=DD_ERRORS.DATADOG_413_ERROR.value,
|
||||
non_success_handler=requeue_after_http_error,
|
||||
)
|
||||
return list(undelivered) # mutable-ok: caller prepends records to the logger queue
|
||||
|
||||
@staticmethod
|
||||
def _exceeds_intake_limits(chunk: Sequence[DatadogPayload]) -> bool:
|
||||
|
|
@ -606,7 +572,7 @@ class DataDogLogger(
|
|||
)
|
||||
return dd_payload
|
||||
|
||||
async def async_send_compressed_data(self, data: list) -> Response:
|
||||
async def async_send_compressed_data(self, data: Sequence[DatadogPayload]) -> Response:
|
||||
"""
|
||||
Async helper to send compressed data to datadog self.intake_url
|
||||
|
||||
|
|
|
|||
|
|
@ -61,9 +61,9 @@ _MAX_CONCURRENT_SHADOW_TASKS: Final = 16
|
|||
_MAX_JUDGE_RESPONSE_CHARS: Final = 8_000
|
||||
_MAX_JUDGE_PROMPT_CHARS: Final = 24_000
|
||||
|
||||
# The judge answers with a small JSON object; a tighter budget truncates the JSON
|
||||
# mid-object and the attempt is lost to an error row.
|
||||
JUDGE_MAX_OUTPUT_TOKENS: Final = 1500
|
||||
# Covers the judge's reasoning tokens as well as its small JSON answer: a judge deployment
|
||||
# carrying an elevated reasoning_effort spends a tight cap before it ever answers.
|
||||
JUDGE_MAX_OUTPUT_TOKENS: Final = 4096
|
||||
|
||||
_MAX_ERROR_CHARS: Final = 500
|
||||
|
||||
|
|
@ -419,6 +419,20 @@ def _failure_detail(e: BaseException) -> str:
|
|||
return f"{type(e).__name__}{location}: {e}"
|
||||
|
||||
|
||||
def _judge_reply_shape(response: object) -> str:
|
||||
"""How an unparseable judge reply was shaped. The parser's own message cannot separate a
|
||||
judge that answered with nothing from one truncated mid-object, and those want opposite
|
||||
fixes. Shape only, never the reply text: the judge quotes the sampled turns it compares,
|
||||
and no attempt row carries sampled content today."""
|
||||
read: Final = _chat_message_reader(response)
|
||||
if read is None:
|
||||
return "unreadable judge reply"
|
||||
content: Final = read("content")
|
||||
served: Final = str(_field_reader(response)("model") or "unknown")
|
||||
body: Final = f"{len(str(content))} chars" if content else "no content"
|
||||
return f"finish_reason={_chat_finish_reason(response)}, content={body}, model={served}"
|
||||
|
||||
|
||||
def _call_cost(response: object) -> float:
|
||||
"""Price one eval-arm call with the figure the spend pipeline bills: the router client
|
||||
stamps _hidden_params.response_cost from the deployment's own pricing, which the public
|
||||
|
|
@ -1266,7 +1280,9 @@ class ShadowEvalLogger(CustomLogger):
|
|||
verdict: Final = PairwiseVerdict.model_validate(parse_json_verdict(raw))
|
||||
except Exception as e: # noqa: BLE001 # malformed verdicts become error rows
|
||||
verbose_logger.debug("shadow_eval: unparseable judge verdict: %s", e)
|
||||
return _CallFailure(f"unparseable judge verdict: {e}", cost=_call_cost(response))
|
||||
return _CallFailure(
|
||||
f"unparseable judge verdict: {e}; {_judge_reply_shape(response)}", cost=_call_cost(response)
|
||||
)
|
||||
return _JudgeVerdict(
|
||||
preference=_unmask_preference(verdict.preference, real_is_a),
|
||||
confidence=max(0.0, min(1.0, verdict.confidence)),
|
||||
|
|
|
|||
|
|
@ -14,21 +14,48 @@ from typing import Final
|
|||
from litellm.types.utils import API_ROUTE_TO_CALL_TYPES, CallTypes
|
||||
|
||||
|
||||
def _segment_matches(route_segment: str, pattern_segment: str) -> bool:
|
||||
"""
|
||||
Match one concrete path segment against one pattern segment.
|
||||
A bare placeholder ({param}) matches any segment; a placeholder with a
|
||||
literal suffix ({model}:generateContent) requires the segment to end with
|
||||
that suffix and have a non-empty value before it.
|
||||
"""
|
||||
if not pattern_segment.startswith("{"):
|
||||
return route_segment == pattern_segment
|
||||
placeholder_end: Final = pattern_segment.find("}")
|
||||
if placeholder_end == -1:
|
||||
return route_segment == pattern_segment
|
||||
literal_suffix: Final = pattern_segment[placeholder_end + 1 :]
|
||||
if not literal_suffix:
|
||||
return True
|
||||
return route_segment.endswith(literal_suffix) and len(route_segment) > len(literal_suffix)
|
||||
|
||||
|
||||
def _pattern_tail_spans_segments(pattern_tail: str) -> bool:
|
||||
"""
|
||||
Whether the pattern's last segment is a suffixed placeholder
|
||||
({model}:generateContent) that may absorb extra route segments, mirroring
|
||||
FastAPI's {model_name:path} converter for slash-containing model names.
|
||||
"""
|
||||
return pattern_tail.startswith("{") and "}" in pattern_tail and not pattern_tail.endswith("}")
|
||||
|
||||
|
||||
def _route_matches_pattern(route: str, pattern: str) -> bool:
|
||||
"""
|
||||
Return True if the concrete route matches the pattern.
|
||||
Pattern segments like {param} match any single path segment.
|
||||
Pattern segments like {param} match any single path segment, and a
|
||||
suffixed placeholder in the last segment may span multiple segments.
|
||||
"""
|
||||
route_parts: Final = route.strip("/").split("/")
|
||||
pattern_parts: Final = pattern.strip("/").split("/")
|
||||
if len(route_parts) != len(pattern_parts):
|
||||
if len(route_parts) < len(pattern_parts):
|
||||
return False
|
||||
for r, p in zip(route_parts, pattern_parts):
|
||||
if p.startswith("{") and p.endswith("}"):
|
||||
continue
|
||||
if r != p:
|
||||
return False
|
||||
return True
|
||||
if len(route_parts) > len(pattern_parts) and not _pattern_tail_spans_segments(pattern_parts[-1]):
|
||||
return False
|
||||
head_count: Final = len(pattern_parts) - 1
|
||||
merged_parts: Final = (*route_parts[:head_count], "/".join(route_parts[head_count:]))
|
||||
return all(_segment_matches(r, p) for r, p in zip(merged_parts, pattern_parts))
|
||||
|
||||
|
||||
def get_call_types_for_route(route: str) -> Sequence[CallTypes] | None:
|
||||
|
|
|
|||
|
|
@ -860,6 +860,7 @@ def _map_bedrock_exception(
|
|||
message=mantle_context_window_message,
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
response=getattr(original_exception, "response", None),
|
||||
)
|
||||
if (
|
||||
"too many tokens" in error_str
|
||||
|
|
@ -873,6 +874,7 @@ def _map_bedrock_exception(
|
|||
message=f"BedrockException: Context Window Error - {error_str}",
|
||||
model=model,
|
||||
llm_provider="bedrock",
|
||||
response=getattr(original_exception, "response", None),
|
||||
)
|
||||
elif "Conversation blocks and tool result blocks cannot be provided in the same turn." in error_str:
|
||||
raise BadRequestError(
|
||||
|
|
@ -924,12 +926,14 @@ def _map_bedrock_exception(
|
|||
message=f"BedrockException: Timeout Error - {error_str}",
|
||||
model=model,
|
||||
llm_provider="bedrock",
|
||||
response=getattr(original_exception, "response", None),
|
||||
)
|
||||
elif "Could not process image" in error_str:
|
||||
raise litellm.InternalServerError(
|
||||
message=f"BedrockException - {error_str}",
|
||||
model=model,
|
||||
llm_provider="bedrock",
|
||||
response=getattr(original_exception, "response", None),
|
||||
)
|
||||
elif hasattr(original_exception, "status_code"):
|
||||
if original_exception.status_code == 500:
|
||||
|
|
@ -937,10 +941,7 @@ def _map_bedrock_exception(
|
|||
message=f"BedrockException - {original_exception.message}",
|
||||
llm_provider="bedrock",
|
||||
model=model,
|
||||
response=httpx.Response(
|
||||
status_code=500,
|
||||
request=httpx.Request(method="POST", url="https://api.openai.com/v1/"),
|
||||
),
|
||||
response=getattr(original_exception, "response", None),
|
||||
)
|
||||
elif original_exception.status_code == 401:
|
||||
raise AuthenticationError(
|
||||
|
|
@ -969,6 +970,7 @@ def _map_bedrock_exception(
|
|||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
litellm_debug_info=extra_information,
|
||||
response=getattr(original_exception, "response", None),
|
||||
)
|
||||
elif original_exception.status_code == 422:
|
||||
raise BadRequestError(
|
||||
|
|
@ -1001,6 +1003,7 @@ def _map_bedrock_exception(
|
|||
llm_provider=custom_llm_provider,
|
||||
litellm_debug_info=extra_information,
|
||||
exception_status_code=original_exception.status_code,
|
||||
response=getattr(original_exception, "response", None),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -21,13 +21,10 @@ AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset(
|
|||
}
|
||||
)
|
||||
|
||||
# The per-deployment Rust opt-in.
|
||||
RUST_KWARG_KEY: Final = "rust"
|
||||
|
||||
# Keys `completion()` forwards from its own kwargs into `get_litellm_params`,
|
||||
# which are otherwise invisible to it because that call site passes explicit
|
||||
# named arguments rather than `**kwargs`.
|
||||
FORWARDED_KWARGS_KEYS: Final = AWS_CREDENTIAL_KWARGS_KEYS | frozenset({RUST_KWARG_KEY})
|
||||
FORWARDED_KWARGS_KEYS: Final = AWS_CREDENTIAL_KWARGS_KEYS
|
||||
|
||||
# Pre-define optional kwargs keys as frozenset for O(1) lookups
|
||||
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
|
||||
|
|
@ -58,10 +55,6 @@ OPTIONAL_KWARGS_KEYS: Final = (
|
|||
"itpm",
|
||||
"otpm",
|
||||
"use_xai_oauth",
|
||||
# The per-deployment Rust opt-in. `all_litellm_params` keeps it out
|
||||
# of the provider body; this keeps it *in* litellm_params, which is
|
||||
# where the chat completions handlers read it from.
|
||||
RUST_KWARG_KEY,
|
||||
}
|
||||
)
|
||||
| AWS_CREDENTIAL_KWARGS_KEYS
|
||||
|
|
|
|||
|
|
@ -53,12 +53,14 @@ class GetModelCostMap:
|
|||
|
||||
_backup_model_count: int = -1 # -1 = not yet loaded
|
||||
|
||||
@staticmethod
|
||||
def read_local_model_cost_map_text() -> str:
|
||||
return files("litellm").joinpath("model_prices_and_context_window_backup.json").read_text(encoding="utf-8")
|
||||
|
||||
@staticmethod
|
||||
def load_local_model_cost_map() -> dict:
|
||||
"""Load the local backup model cost map bundled with the package."""
|
||||
content: Final = json.loads(
|
||||
files("litellm").joinpath("model_prices_and_context_window_backup.json").read_text(encoding="utf-8")
|
||||
)
|
||||
content: Final = json.loads(GetModelCostMap.read_local_model_cost_map_text())
|
||||
return content
|
||||
|
||||
@classmethod
|
||||
|
|
|
|||
|
|
@ -6,14 +6,13 @@ import base64
|
|||
from collections.abc import Awaitable, Callable
|
||||
from typing import TYPE_CHECKING, Final, Literal
|
||||
|
||||
from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS
|
||||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, DocumentType
|
||||
from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS, LlmProviders
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
# Minimal PDF for health checks - base64 encoded 1-page PDF with just "test"
|
||||
TEST_PDF_URL = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y="
|
||||
|
||||
# Minimal image for health checks - base64 encoded 512x512 blue circle on a white background PNG
|
||||
TEST_IMAGE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAAgAAAAIACAIAAAB7GkOtAAAJk0lEQVR42u3VQREAIRADwVWCOmTjBVzwSLorCri6nbkAVBpPACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAAAgAAAIAZdY+HgEBgIRr/meeGgGA8EMvDAgAuPh6gACAi68HCAA4+mKAAICjLwYIALj7SoAAgLuvBAgA7r4pAQKAu29KgADg7psSIAA4/SYDCADuvikBAoDTbzKAAOD0mwwgADj9JgMIAE6/yQACgNNvMoAA4PSbDCAAOP0mAwgATr/JAAKA668BIAA4/TKAAOD0mwwgALj+pgEIAE6/yQACgOtvGoAA4PSbDCAAuP6mAQgATr/JAAKA628agADg9JsMIAC4/qYBCACuv2kAAoDrbxqAAOD0mwwgALj+pgEIAK6/aQACgOtvGoAA4PqbBiAArr+ZBiAATr+ZDCAArr+ZBiAArr+ZBiAArr+ZBiAArr+ZBggArr+ZBggArr+ZBggArr+ZBggArr+ZBggArr+ZBggArr+ZBggAAmAmAAKA62+mAQKA62+mAQKA62/m1xYAXH/TAAQA1980AAHA9TcNQAAEwEwAEADX30wDEADX30wDEADX30wDEADX30wDEAABMBMABMD1N9MABMD1N9MABEAAzAQAAXD9zTQAAXD9zTRAABAAMwEQAFx/Mw0QAFx/Mw0QAATATAAEANffTAMEANffTAMEAAEwEwABwPU30wABQADMBEAAXH8z0wABcP3NTAMEQADMTAAEwPU3Mw0QAAEwMwEQANffTAMQAAEwEwAEwPU30wAEQADMBAABcP3NNAABEAAzAUAAXH8zDUAABMBMABAA199MAwQAATATAAHA9TfTAAFAAMwEQABw/c00QAAQADMBEAAEwEwABMD1NzMNEAABMDMBEADX38w0QAAEwMwEQAAEwMwEQABcfzPTAAEQADMTAAEQADMTAAFw/c1MAwRAAMxMAARAAMxMAATA9TczDRAAATATAARAAMwEAAFw/c00AAEQADMBQAAEwEwAEADX30wDBAABMBMAAUAAzARAABAAMwEQAFx/Mw0QAAEwMwEQAAEwMwEQAAEwMwEQANffzDRAAATAzARAAATAzARAAATAzARAAATAzARAAFx/M9MAARAAMxMAARAAMxMAARAAMxMAARAAMxMAAXD9zUwDBEAAzEwABEAAzAQAARAAMwFAAATATAAQAAEwEwABQADMBEAAEAAzARAAXH8zDRAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAA19/MNEAANMDM9UcABMBMABAAATATAAHwBAJgJgACgACYCYAAIABmAiAA+IvMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAANMDMXH8BEAAzEwABEAAzEwABEAAzEwABEAAzEwABEAAzEwABEAAzAUAABMBMABAADTBz/REAATATAAFAAMwEQAAQADMBEAAEwEwABAANMHP9BUAAzEwABEAAzEwABEAAzEwABEAAzEwABEADzMz1FwABMDMBEAABMDMBEAABMDMBEAANMDPXXwAEwMwEQAAEwMwEQAAEwMwEQAA0wMxcfwEQADMTAAEQADMBQAA0wMz1RwAEwEwAEAABMBMABEADzFx/AUAAzARAABAAMwEQADTAzPUXAATATAAEAAEwEwABQAPMXH8BEAAzEwABEAAzEwAB0AAzc/0FQADMTAAEQAPMzPUXAAEwMwEQAAEwMwEQAA0wM9dfAATAzARAADTAzFx/ARAAMwFAADTAzPVHAATATAAQAA0wc/0RAAEwEwAEQAPMXH8EQADMBAAB0AAz118AEAAzARAANMDM9RcABMBMAAQADTBz/QUAATATAAFAA8xcfwFAA8xcfwFAAMwEQADQADPXXwAEwMwEQAA0wMxcfwHQADNz/QVAAMxMAARAA8xcfwRAA8xcfwRAAMwEAAHQADPXHwHQADPXHwEQADMBQAA0wMz1RwA0wMz1RwAEwEwAEAANMHP9BQANMHP9BQANMHP9BQANMHP9BQABMBMAAUADzFx/AUADzFx/AUADzFx/AUADzFx/AUADzFx/AUADzPVHABAAEwAEAA0w1x8BQAPM9UcA0ABz/REANMBcfwQADTDXHwHQADPXHwHQADPXHwHQADPXHwHQADPXHwHQADPXHwGQATOnHwHQADPXHwHQADPXXwDQADPXXwDQADPXXwDQADPXXwCQAXP6EQA0wFx/BAANMNcfAUADzPVHAJABc/oRADTAXH8EABkwpx8BQAPM9UcAkAFz+hEANMBcfwQAGTCnHwFAA8z1RwCQAXP6EQBkwJx+BAANMNcfAUAGzOlHAJABc/oRAGTA6QcBQAacfhAAZMDpRwBABpx+BABkwOlHAEAGnH4EAJTA3UcAQAacfgQAlMDdRwBACdx9BACUwN1HAEAJ3H0EAJTA3UcAQAwcfQQAmmLgsyIA0NIDHw4BgJYe+DQIAISHwVMjAJDQDI+AAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACACAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAIAABNHpialFcmLajuAAAAAElFTkSuQmCC"
|
||||
|
|
@ -29,6 +28,14 @@ def get_image_file_for_health_check() -> bytes:
|
|||
return base64.b64decode(TEST_IMAGE_BASE64)
|
||||
|
||||
|
||||
def _ocr_health_check_document(model: str, custom_llm_provider: str) -> DocumentType:
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
provider: Final = next((known for known in LlmProviders if known.value == custom_llm_provider), None)
|
||||
config: Final = ProviderConfigManager.get_provider_ocr_config(model=model, provider=provider) if provider else None
|
||||
return (config or BaseOCRConfig()).get_health_check_document()
|
||||
|
||||
|
||||
class HealthCheckHelpers:
|
||||
@staticmethod
|
||||
async def ahealth_check_wildcard_models(
|
||||
|
|
@ -247,9 +254,6 @@ class HealthCheckHelpers:
|
|||
),
|
||||
"ocr": lambda: litellm.aocr(
|
||||
**_filter_model_params(model_params=model_params),
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": TEST_PDF_URL,
|
||||
},
|
||||
document=_ocr_health_check_document(model=model, custom_llm_provider=custom_llm_provider),
|
||||
),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -36,12 +36,13 @@ from litellm._logging import (
|
|||
verbose_logger,
|
||||
)
|
||||
from litellm._uuid import uuid
|
||||
from litellm.batches.batch_utils import _handle_completed_batch
|
||||
from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final
|
||||
from litellm.caching.caching import DualCache, InMemoryCache
|
||||
from litellm.caching.caching_handler import LLMCachingHandler
|
||||
from litellm.constants import (
|
||||
DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT,
|
||||
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT,
|
||||
PROVIDER_REQUEST_ID_HEADERS,
|
||||
SENTRY_DENYLIST,
|
||||
SENTRY_PII_DENYLIST,
|
||||
)
|
||||
|
|
@ -78,7 +79,10 @@ from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
|
|||
from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import (
|
||||
InteractionsUsageObjectTransformation,
|
||||
)
|
||||
from litellm.litellm_core_utils.logging_utils import truncate_base64_in_messages
|
||||
from litellm.litellm_core_utils.logging_utils import (
|
||||
truncate_base64_in_messages,
|
||||
truncate_base64_in_messages_async,
|
||||
)
|
||||
from litellm.litellm_core_utils.model_param_helper import ModelParamHelper
|
||||
from litellm.litellm_core_utils.redact_messages import (
|
||||
redact_message_input_output_from_custom_logger,
|
||||
|
|
@ -252,6 +256,30 @@ _in_memory_loggers: Final[list[CustomLogger]] = []
|
|||
|
||||
_STANDARD_LOGGING_METADATA_KEYS: Final[frozenset[str]] = frozenset(StandardLoggingMetadata.__annotations__.keys())
|
||||
|
||||
|
||||
def _get_provider_request_id(original_exception: Exception) -> str | None:
|
||||
try:
|
||||
error_response: Final = getattr(original_exception, "response", None)
|
||||
header_sources: Final = (
|
||||
_get_response_headers(original_exception),
|
||||
getattr(error_response, "headers", None),
|
||||
getattr(original_exception, "litellm_response_headers", None),
|
||||
)
|
||||
return next(
|
||||
(
|
||||
str(value)
|
||||
for expected_header_name in PROVIDER_REQUEST_ID_HEADERS
|
||||
for headers in header_sources
|
||||
if isinstance(headers, Mapping)
|
||||
for header_name, value in headers.items()
|
||||
if isinstance(header_name, str) and header_name.lower() == expected_header_name and value
|
||||
),
|
||||
None,
|
||||
)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
### GLOBAL VARIABLES ###
|
||||
|
||||
# Cache custom pricing keys as frozenset for O(1) lookups instead of looping through 49 keys
|
||||
|
|
@ -538,6 +566,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.standard_built_in_tools_params: StandardBuiltInToolsParams = (
|
||||
self.initialize_standard_built_in_tools_params(kwargs)
|
||||
)
|
||||
self.truncated_messages_for_logging: str | list | dict | None = None # mutable-ok: logged messages shape
|
||||
## TIME TO FIRST TOKEN LOGGING ##
|
||||
self.completion_start_time: datetime.datetime | None = None
|
||||
self._llm_caching_handler: LLMCachingHandler | None = None
|
||||
|
|
@ -1820,6 +1849,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
and litellm_params.get(CallTypes.aanthropic_messages.value, False) is not True
|
||||
and litellm_params.get(CallTypes.agenerate_content.value, False) is not True
|
||||
and litellm_params.get(CallTypes.agenerate_content_stream.value, False) is not True
|
||||
and litellm_params.get(CallTypes.arealtime.value, False) is not True
|
||||
)
|
||||
|
||||
def _is_assembled_stream_success(self, result=None) -> bool:
|
||||
|
|
@ -1913,7 +1943,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
two paths cannot mutate it at the same time. ``prefer_async_handlers`` only
|
||||
bypasses the sync-SDK-only shortcut (e.g. ``async for`` on a stream from
|
||||
``completion()``); legacy string callbacks still run via
|
||||
``executor.submit(failure_handler)`` when configured.
|
||||
``executor.submit(failure_handler)`` when configured, and still get submitted
|
||||
when the awaiting task is cancelled (e.g. the event loop shuts down right after
|
||||
the request failed).
|
||||
"""
|
||||
litellm_params: Final = self.model_call_details.get("litellm_params", {}) or {}
|
||||
sync_sdk: Final = self._is_sync_litellm_request(litellm_params)
|
||||
|
|
@ -1922,12 +1954,11 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.failure_handler(exception, traceback_exception)
|
||||
return
|
||||
|
||||
await self.async_failure_handler(exception, traceback_exception)
|
||||
|
||||
if not self._should_run_sync_failure_callbacks_for_async_calls():
|
||||
return
|
||||
|
||||
executor.submit(self.failure_handler, exception, traceback_exception)
|
||||
try:
|
||||
await self.async_failure_handler(exception, traceback_exception)
|
||||
finally:
|
||||
if self._should_run_sync_failure_callbacks_for_async_calls():
|
||||
executor.submit(self.failure_handler, exception, traceback_exception)
|
||||
|
||||
def should_run_logging(
|
||||
self,
|
||||
|
|
@ -2893,13 +2924,6 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
): # polling job will query these frequently, don't spam db logs
|
||||
return
|
||||
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
_is_base64_encoded_unified_file_id,
|
||||
)
|
||||
|
||||
# check if file id is a unified file id
|
||||
is_base64_unified_file_id: Final = _is_base64_encoded_unified_file_id(result.id)
|
||||
|
||||
batch_cost: Final = kwargs.get("batch_cost", None)
|
||||
batch_usage = kwargs.get("batch_usage", None)
|
||||
batch_models = kwargs.get("batch_models", None)
|
||||
|
|
@ -2907,9 +2931,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
batch_failed_requests: Final = kwargs.get("batch_failed_requests", None)
|
||||
has_explicit_batch_data: Final = all(x is not None for x in (batch_cost, batch_usage, batch_models))
|
||||
|
||||
should_compute_batch_data: Final = (
|
||||
not is_base64_unified_file_id or not has_explicit_batch_data and result.status == "completed"
|
||||
)
|
||||
should_compute_batch_data: Final = not has_explicit_batch_data and batch_cost_is_final(result)
|
||||
if has_explicit_batch_data:
|
||||
result._hidden_params["response_cost"] = batch_cost
|
||||
result._hidden_params["batch_models"] = batch_models
|
||||
|
|
@ -2932,6 +2954,11 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
result._hidden_params["batch_failed_requests"] = batch_result.failed_requests # pyright: ignore[reportPrivateUsage] # rebind-ok: same pattern as above
|
||||
result.usage = batch_result.usage
|
||||
|
||||
self.truncated_messages_for_logging = await truncate_base64_in_messages_async(
|
||||
StandardLoggingPayloadSetup.append_system_prompt_messages(
|
||||
kwargs=self.model_call_details, messages=self.model_call_details.get("messages")
|
||||
)
|
||||
)
|
||||
start_time, end_time, result = self._success_handler_helper_fn(
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
|
|
@ -3224,8 +3251,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.model_call_details = {}
|
||||
|
||||
if (
|
||||
self.model_call_details.get("log_event_type") == "failed_api_call"
|
||||
and self.model_call_details.get("exception") is exception
|
||||
self.model_call_details.get("exception") is exception
|
||||
and self.model_call_details.get("standard_logging_object") is not None
|
||||
):
|
||||
return start_time, self.model_call_details["end_time"]
|
||||
|
|
@ -3908,11 +3934,12 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
served_id: Final = _provider_response_id(result)
|
||||
try:
|
||||
return LiteLLMResponsesTransformationHandler().transform_response(
|
||||
translated: Final = LiteLLMResponsesTransformationHandler().transform_response(
|
||||
model=self.model,
|
||||
raw_response=result,
|
||||
model_response=litellm.ModelResponse(id=_provider_response_id(result)),
|
||||
model_response=litellm.ModelResponse(id=served_id),
|
||||
logging_obj=self,
|
||||
request_data={},
|
||||
messages=[],
|
||||
|
|
@ -3920,6 +3947,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
litellm_params={},
|
||||
encoding=litellm.encoding,
|
||||
)
|
||||
translated.id = served_id or translated.id
|
||||
return translated
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"Responses API -> ModelResponse translation failed for "
|
||||
|
|
@ -3927,7 +3956,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
"usage-only ModelResponse to keep the spend_logs row.",
|
||||
str(e),
|
||||
)
|
||||
model_response: Final = litellm.ModelResponse(id=_provider_response_id(result))
|
||||
model_response: Final = litellm.ModelResponse(id=served_id)
|
||||
model_response.model = self.model
|
||||
usage: Final = getattr(result, "usage", None)
|
||||
if usage is not None and ResponseAPILoggingUtils._is_response_api_usage(usage):
|
||||
|
|
@ -5660,13 +5689,15 @@ class StandardLoggingPayloadSetup:
|
|||
rate_limit_category: Final = validate_rate_limit_category(getattr(original_exception, "category", None))
|
||||
rate_limit_type: Final = validate_rate_limit_type(getattr(original_exception, "rate_limit_type", None))
|
||||
budget_error: Final = original_exception if isinstance(original_exception, BudgetExceededError) else None
|
||||
provider_request_id: Final = _get_provider_request_id(original_exception) if original_exception else None
|
||||
|
||||
return StandardLoggingPayloadErrorInformation(
|
||||
error_code=error_status,
|
||||
error_class=error_class,
|
||||
llm_provider=_llm_provider_in_exception,
|
||||
traceback=traceback_info,
|
||||
error_message=error_message,
|
||||
traceback=_redact_string(traceback_info),
|
||||
error_message=_redact_string(error_message),
|
||||
error_provider_request_id=provider_request_id,
|
||||
error_rate_limit_category=rate_limit_category,
|
||||
error_rate_limit_type=rate_limit_type,
|
||||
error_budget_entity_type=budget_error.entity_type if budget_error else None,
|
||||
|
|
@ -5871,6 +5902,7 @@ def _get_status_fields(
|
|||
# Mapping for legacy guardrail status values to new GuardrailStatus values
|
||||
GUARDRAIL_STATUS_MAP: Final[dict[str, GuardrailStatus]] = {
|
||||
"success": "success",
|
||||
"guardrail_flagged": "guardrail_flagged",
|
||||
"blocked": "guardrail_intervened", # legacy
|
||||
"guardrail_intervened": "guardrail_intervened", # direct
|
||||
"failure": "guardrail_failed_to_respond", # legacy
|
||||
|
|
@ -5892,6 +5924,7 @@ def _get_status_fields(
|
|||
GUARDRAIL_STATUS_SEVERITY: Final[tuple[GuardrailStatus, ...]] = (
|
||||
"not_run",
|
||||
"success",
|
||||
"guardrail_flagged",
|
||||
"guardrail_failed_to_respond",
|
||||
"guardrail_intervened",
|
||||
)
|
||||
|
|
@ -6201,9 +6234,13 @@ def get_standard_logging_object_payload(
|
|||
model_id=_model_id,
|
||||
requester_ip_address=clean_metadata.get("requester_ip_address", None),
|
||||
user_agent=clean_metadata.get("user_agent", None),
|
||||
messages=truncate_base64_in_messages(
|
||||
StandardLoggingPayloadSetup.append_system_prompt_messages(
|
||||
kwargs=kwargs, messages=kwargs.get("messages")
|
||||
messages=(
|
||||
logging_obj.truncated_messages_for_logging
|
||||
if logging_obj.truncated_messages_for_logging is not None
|
||||
else truncate_base64_in_messages(
|
||||
StandardLoggingPayloadSetup.append_system_prompt_messages(
|
||||
kwargs=kwargs, messages=kwargs.get("messages")
|
||||
)
|
||||
)
|
||||
),
|
||||
response=final_response_obj,
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from litellm.types.utils import (
|
|||
PromptTokensDetailsWrapper,
|
||||
ServiceTier,
|
||||
Usage,
|
||||
text_tokens_without_nested_reasoning,
|
||||
)
|
||||
from litellm.utils import get_model_info
|
||||
|
||||
|
|
@ -513,6 +514,7 @@ def _get_token_base_cost(
|
|||
current_time: datetime | None = None,
|
||||
*,
|
||||
threshold_is_inclusive: bool = False,
|
||||
missing_cache_read_uses_input: bool = False,
|
||||
) -> tuple[float, float, float, float, float]:
|
||||
"""
|
||||
Return prompt cost, completion cost, and cache costs for a given model and usage.
|
||||
|
|
@ -523,6 +525,9 @@ def _get_token_base_cost(
|
|||
`threshold_is_inclusive` switches that comparison to >=, for providers such as xAI
|
||||
that bill the higher tier once the prompt reaches the threshold.
|
||||
|
||||
`missing_cache_read_uses_input` resolves an absent cache-read rate to the resolved
|
||||
input rate instead of 0.0; an explicit 0.0 rate stays a real price either way.
|
||||
|
||||
Returns:
|
||||
Tuple[float, float, float, float] - (prompt_cost, completion_cost, cache_creation_cost, cache_read_cost)
|
||||
"""
|
||||
|
|
@ -550,29 +555,16 @@ def _get_token_base_cost(
|
|||
float,
|
||||
_get_cost_per_unit(model_info, "cache_creation_input_token_cost_above_1hr"),
|
||||
)
|
||||
cache_read_cost = cast(float, _get_cost_per_unit(model_info, cache_read_cost_key))
|
||||
cache_read_cost = _get_cost_per_unit(model_info, cache_read_cost_key, default_value=None)
|
||||
|
||||
## CHECK IF ABOVE THRESHOLD
|
||||
# Optimization: collect threshold keys first to avoid sorting all model_info keys.
|
||||
# Most models don't have threshold pricing, so we can return early.
|
||||
# Exclude service_tier-specific variants (e.g. input_cost_per_token_above_200k_tokens_priority)
|
||||
# so that the threshold detection loop only processes standard keys. The
|
||||
# service_tier-specific above-threshold key is resolved later via _get_service_tier_cost_key.
|
||||
threshold_keys: Final = [
|
||||
k for k in model_info if k.startswith("input_cost_per_token_above_") and not k.endswith(_SERVICE_TIER_SUFFIXES)
|
||||
]
|
||||
if not threshold_keys:
|
||||
return _apply_off_peak_to_base_costs(
|
||||
model_info,
|
||||
current_time,
|
||||
(
|
||||
prompt_base_cost,
|
||||
completion_base_cost,
|
||||
cache_creation_cost,
|
||||
cache_creation_cost_above_1hr,
|
||||
cache_read_cost,
|
||||
),
|
||||
)
|
||||
|
||||
# Only sort the threshold keys (typically 1-2 keys instead of 66+)
|
||||
threshold: float | None = None
|
||||
|
|
@ -661,10 +653,7 @@ def _get_token_base_cost(
|
|||
),
|
||||
)
|
||||
|
||||
cache_read_cost = cast(
|
||||
float,
|
||||
_get_cost_per_unit(model_info, cache_read_tiered_key, cache_read_cost),
|
||||
)
|
||||
cache_read_cost = _get_cost_per_unit(model_info, cache_read_tiered_key, cache_read_cost)
|
||||
|
||||
break
|
||||
except (IndexError, ValueError):
|
||||
|
|
@ -672,6 +661,17 @@ def _get_token_base_cost(
|
|||
except Exception:
|
||||
continue
|
||||
|
||||
if cache_read_cost is None:
|
||||
cache_read_cost = (
|
||||
_off_peak_rate(
|
||||
_open_off_peak_block(model_info, current_time) or MappingProxyType({}),
|
||||
"input_cost_per_token",
|
||||
prompt_base_cost,
|
||||
)
|
||||
if missing_cache_read_uses_input
|
||||
else 0.0
|
||||
)
|
||||
|
||||
return _apply_off_peak_to_base_costs(
|
||||
model_info,
|
||||
current_time,
|
||||
|
|
@ -860,7 +860,7 @@ def parse_completion_tokens_details(usage: Usage) -> CompletionTokensDetailsResu
|
|||
)
|
||||
or 0
|
||||
)
|
||||
text_tokens: Final = (
|
||||
reported_text_tokens: Final = (
|
||||
cast(
|
||||
int | None,
|
||||
getattr(usage.completion_tokens_details, "text_tokens", None),
|
||||
|
|
@ -882,6 +882,12 @@ def parse_completion_tokens_details(usage: Usage) -> CompletionTokensDetailsResu
|
|||
or 0
|
||||
)
|
||||
video_tokens: Final = _coerce_token_count(getattr(usage.completion_tokens_details, "video_tokens", 0))
|
||||
text_tokens: Final = text_tokens_without_nested_reasoning(
|
||||
completion_tokens=usage.completion_tokens,
|
||||
text_tokens=reported_text_tokens,
|
||||
reasoning_tokens=reasoning_tokens,
|
||||
other_modality_tokens=audio_tokens + image_tokens + video_tokens,
|
||||
)
|
||||
|
||||
return CompletionTokensDetailsResult(
|
||||
audio_tokens=audio_tokens,
|
||||
|
|
@ -1409,6 +1415,57 @@ def get_token_type_cost_breakdown(
|
|||
)
|
||||
|
||||
|
||||
def calculate_prompt_caching_savings(
|
||||
model_info: ModelInfo,
|
||||
usage: Usage,
|
||||
custom_llm_provider: str | None,
|
||||
service_tier: str | None = None,
|
||||
data_residency: str | None = None,
|
||||
vertex_location: str | None = None,
|
||||
billed_at: datetime | None = None,
|
||||
) -> float:
|
||||
"""Read discount minus write premium, using the biller's rate and TTL resolution.
|
||||
|
||||
Missing reads and unpublished (missing/zero) writes claim no saving or premium;
|
||||
explicit zero reads remain free. An unpublished 1h price uses the ordinary write rate.
|
||||
``billed_at`` is the request's completion time, so off-peak windows resolve as the
|
||||
biller saw them rather than at the later spend write.
|
||||
"""
|
||||
prompt_base_cost, _, cache_creation_cost, cache_creation_cost_above_1hr, cache_read_cost = _get_token_base_cost(
|
||||
model_info=model_info,
|
||||
usage=usage,
|
||||
service_tier=service_tier,
|
||||
current_time=billed_at,
|
||||
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
|
||||
missing_cache_read_uses_input=True,
|
||||
)
|
||||
write_rate: Final = cache_creation_cost or prompt_base_cost
|
||||
write_rate_1h: Final = cache_creation_cost_above_1hr or write_rate
|
||||
prompt_tokens_details: Final = parse_prompt_tokens_details(usage)
|
||||
cache_read_tokens: Final = max(prompt_tokens_details["cache_hit_tokens"], 0)
|
||||
cache_creation_tokens: Final = max(prompt_tokens_details["cache_creation_tokens"], 0)
|
||||
details: Final = prompt_tokens_details["cache_creation_token_details"]
|
||||
cache_creation_details: Final = (
|
||||
CacheCreationTokenDetails(
|
||||
ephemeral_5m_input_tokens=max(details.ephemeral_5m_input_tokens or 0, 0),
|
||||
ephemeral_1h_input_tokens=max(details.ephemeral_1h_input_tokens or 0, 0),
|
||||
)
|
||||
if details is not None
|
||||
else None
|
||||
)
|
||||
read_discount: Final = cache_read_tokens * max(prompt_base_cost - cache_read_cost, 0.0)
|
||||
write_premium: Final = calculate_cache_writing_cost(
|
||||
cache_creation_tokens=cache_creation_tokens,
|
||||
cache_creation_token_details=cache_creation_details,
|
||||
cache_creation_cost_above_1hr=write_rate_1h - prompt_base_cost,
|
||||
cache_creation_cost=write_rate - prompt_base_cost,
|
||||
)
|
||||
uplift: Final = _get_regional_uplift_multiplier(model_info, data_residency) * get_vertex_regional_endpoint_uplift(
|
||||
model_info, vertex_location
|
||||
)
|
||||
return (read_discount - write_premium) * uplift
|
||||
|
||||
|
||||
def calculate_image_response_cost_from_usage(
|
||||
model: str,
|
||||
image_response: ImageResponse,
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
|
||||
def get_response_headers(_response_headers: dict | None = None) -> dict:
|
||||
def get_response_headers(_response_headers: Mapping[str, str] | None = None) -> dict:
|
||||
"""
|
||||
|
||||
Sets the Appropriate OpenAI headers for the response and forward all headers as llm_provider-{header}
|
||||
|
|
@ -31,7 +32,7 @@ def get_response_headers(_response_headers: dict | None = None) -> dict:
|
|||
return {**llm_provider_headers, **openai_headers}
|
||||
|
||||
|
||||
def _get_llm_provider_headers(response_headers: dict) -> dict:
|
||||
def _get_llm_provider_headers(response_headers: Mapping[str, str]) -> dict:
|
||||
"""
|
||||
Adds a llm_provider-{header} to all headers that are not already prefixed with llm_provider
|
||||
|
||||
|
|
|
|||
|
|
@ -3,12 +3,15 @@ import functools
|
|||
import inspect
|
||||
import re
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MAX_BASE64_LENGTH_FOR_LOGGING
|
||||
from litellm.constants import (
|
||||
BASE64_TRUNCATION_OFFLOAD_THRESHOLD_CHARS,
|
||||
MAX_BASE64_LENGTH_FOR_LOGGING,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
|
|
@ -141,6 +144,39 @@ def truncate_base64_in_messages(
|
|||
return messages
|
||||
|
||||
|
||||
_StringTree = str | Sequence["_StringTree"] | Mapping[str, "_StringTree"] | None
|
||||
|
||||
|
||||
def _iter_string_leaves(value: _StringTree) -> Iterator[str]:
|
||||
stack: Final[list[_StringTree]] = [value] # mutable-ok: explicit stack, recursive functions are banned in litellm/
|
||||
while stack:
|
||||
match stack.pop():
|
||||
case str() as text:
|
||||
yield text
|
||||
case Mapping() as mapping:
|
||||
stack.extend(mapping.values())
|
||||
case Sequence() as items:
|
||||
stack.extend(items)
|
||||
case None:
|
||||
pass
|
||||
|
||||
|
||||
async def truncate_base64_in_messages_async(
|
||||
messages: str | list | dict | None, # mutable-ok: same contract as truncate_base64_in_messages
|
||||
) -> str | list | dict | None: # mutable-ok: same contract as truncate_base64_in_messages
|
||||
"""
|
||||
Same result as truncate_base64_in_messages, but payloads whose string content
|
||||
reaches BASE64_TRUNCATION_OFFLOAD_THRESHOLD_CHARS are scanned in a worker
|
||||
thread so the regex pass over multi-MB base64 images does not block the event loop.
|
||||
"""
|
||||
if messages is None or MAX_BASE64_LENGTH_FOR_LOGGING <= 0:
|
||||
return messages
|
||||
total_chars: Final = sum(len(leaf) for leaf in _iter_string_leaves(messages))
|
||||
if total_chars < BASE64_TRUNCATION_OFFLOAD_THRESHOLD_CHARS:
|
||||
return truncate_base64_in_messages(messages)
|
||||
return await asyncio.to_thread(truncate_base64_in_messages, messages)
|
||||
|
||||
|
||||
# Global service logger instance to avoid recreating it
|
||||
_service_logger = None
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,11 @@
|
|||
Helper functions to handle images passed in messages
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from httpx import Response
|
||||
|
|
@ -11,9 +15,11 @@ import litellm
|
|||
from litellm import verbose_logger
|
||||
from litellm.caching.caching import InMemoryCache
|
||||
from litellm.constants import MAX_IMAGE_URL_DOWNLOAD_SIZE_MB
|
||||
from litellm.litellm_core_utils.url_utils import async_safe_get, safe_get
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get, safe_get
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
MAX_IMGS_IN_MEMORY: Final = 10
|
||||
MAX_CONCURRENT_REMOTE_MEDIA_FETCHES: Final = 20
|
||||
|
||||
in_memory_cache: Final = InMemoryCache(max_size_in_memory=MAX_IMGS_IN_MEMORY)
|
||||
|
||||
|
|
@ -72,6 +78,14 @@ def _process_image_response(response: Response, url: str) -> str:
|
|||
return result
|
||||
|
||||
|
||||
def _rejected_image_fetch(url: str, verdict: SSRFError) -> "litellm.ImageFetchError":
|
||||
verbose_logger.warning("Image fetch of %s rejected before any request went out: %s", url, verdict)
|
||||
return litellm.ImageFetchError(
|
||||
"Error: Unable to fetch image from URL. The proxy could not resolve this host or its URL policy rejected it; "
|
||||
f"an admin can check the proxy log and `user_url_allowed_hosts` in general_settings. url={url}"
|
||||
)
|
||||
|
||||
|
||||
async def async_convert_url_to_base64(url: str) -> str:
|
||||
if url.startswith("data:") and ";base64," in url:
|
||||
return url
|
||||
|
|
@ -93,6 +107,8 @@ async def async_convert_url_to_base64(url: str) -> str:
|
|||
return _process_image_response(response, url)
|
||||
except litellm.ImageFetchError:
|
||||
raise
|
||||
except SSRFError as e:
|
||||
raise _rejected_image_fetch(url, e) from e
|
||||
except Exception:
|
||||
pass
|
||||
raise litellm.ImageFetchError(f"Error: Unable to fetch image from URL after 3 attempts. url={url}")
|
||||
|
|
@ -119,8 +135,192 @@ def convert_url_to_base64(url: str) -> str:
|
|||
return _process_image_response(response, url)
|
||||
except litellm.ImageFetchError:
|
||||
raise
|
||||
except SSRFError as e:
|
||||
raise _rejected_image_fetch(url, e) from e
|
||||
except Exception as e:
|
||||
verbose_logger.exception(e)
|
||||
raise litellm.ImageFetchError(
|
||||
f"Error: Unable to fetch image from URL after 3 attempts. url={url}",
|
||||
)
|
||||
|
||||
|
||||
_REMOTE_URL_PREFIXES: Final = ("http://", "https://")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RemoteImage:
|
||||
part: Mapping[str, object]
|
||||
image_url: Mapping[str, object] | None
|
||||
url: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RemoteFile:
|
||||
part: Mapping[str, object]
|
||||
file: Mapping[str, object]
|
||||
url: str
|
||||
|
||||
|
||||
def _as_mapping(value: object) -> Mapping[str, object] | None:
|
||||
return value if isinstance(value, Mapping) else None # pyright: ignore[reportUnknownVariableType] # fields are parsed one by one
|
||||
|
||||
|
||||
def _remote_url(candidate: object) -> str | None:
|
||||
return candidate if isinstance(candidate, str) and candidate.startswith(_REMOTE_URL_PREFIXES) else None
|
||||
|
||||
|
||||
_ANTHROPIC_MEDIA_BLOCK_TYPES: Final = frozenset({"document", "image"})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RemoteSource:
|
||||
part: Mapping[str, object]
|
||||
source: Mapping[str, object]
|
||||
url: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RemoteMedia:
|
||||
url: str
|
||||
fields: Mapping[str, object]
|
||||
|
||||
|
||||
_NO_FIELDS: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
|
||||
def inline_every_remote_url(_media: RemoteMedia) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
def _parse_remote_image(fields: Mapping[str, object]) -> _RemoteImage | None:
|
||||
if fields.get("type") != "image_url":
|
||||
return None
|
||||
image_url: Final = fields.get("image_url")
|
||||
image_url_fields: Final = _as_mapping(image_url)
|
||||
url: Final = _remote_url(image_url_fields.get("url") if image_url_fields is not None else image_url)
|
||||
return _RemoteImage(fields, image_url_fields, url) if url is not None else None
|
||||
|
||||
|
||||
def _parse_remote_file(fields: Mapping[str, object]) -> _RemoteFile | None:
|
||||
file: Final = _as_mapping(fields.get("file")) if fields.get("type") == "file" else None
|
||||
url: Final = _remote_url(file.get("file_id")) if file is not None else None
|
||||
return _RemoteFile(fields, file, url) if file is not None and url is not None else None
|
||||
|
||||
|
||||
def _parse_remote_source(fields: Mapping[str, object]) -> _RemoteSource | None:
|
||||
source: Final = _as_mapping(fields.get("source")) if fields.get("type") in _ANTHROPIC_MEDIA_BLOCK_TYPES else None
|
||||
url: Final = _remote_url(source.get("url")) if source is not None and source.get("type") == "url" else None
|
||||
return _RemoteSource(fields, source, url) if source is not None and url is not None else None
|
||||
|
||||
|
||||
def _parse_remote_part(part: object) -> _RemoteImage | _RemoteFile | _RemoteSource | None:
|
||||
fields: Final = _as_mapping(part)
|
||||
if fields is None:
|
||||
return None
|
||||
return _parse_remote_image(fields) or _parse_remote_file(fields) or _parse_remote_source(fields)
|
||||
|
||||
|
||||
def _remote_media(remote: _RemoteImage | _RemoteFile | _RemoteSource) -> RemoteMedia:
|
||||
match remote:
|
||||
case _RemoteImage(_, image_url, url):
|
||||
return RemoteMedia(url, image_url if image_url is not None else _NO_FIELDS)
|
||||
case _RemoteFile(_, file, url):
|
||||
return RemoteMedia(url, file)
|
||||
case _RemoteSource(_, source, url):
|
||||
return RemoteMedia(url, source)
|
||||
|
||||
|
||||
_PDF_FORMAT: Final = MappingProxyType({"format": "application/pdf"})
|
||||
|
||||
|
||||
def _inferred_format(file: Mapping[str, object], url: str) -> Mapping[str, str]:
|
||||
return _PDF_FORMAT if "format" not in file and url.lower().endswith(".pdf") else MappingProxyType({})
|
||||
|
||||
|
||||
def _inlined_image_url(image_url: Mapping[str, object] | None, data_url: str) -> Mapping[str, object] | str:
|
||||
return {**image_url, "url": data_url} if image_url is not None else data_url # mutable-ok: json-serialized part
|
||||
|
||||
|
||||
def _inlined_file(file: Mapping[str, object], url: str, data_url: str) -> Mapping[str, object]:
|
||||
kept: Final = {k: v for k, v in file.items() if k != "file_id"} # mutable-ok: json-serialized message part
|
||||
return {**kept, **_inferred_format(file, url), "file_data": data_url} # mutable-ok: json-serialized part
|
||||
|
||||
|
||||
def _base64_source(url: str, data_url: str) -> Mapping[str, str]:
|
||||
fetched_media_type, data = data_url.removeprefix("data:").split(";base64,", 1)
|
||||
media_type: Final = "application/pdf" if url.lower().endswith(".pdf") else fetched_media_type
|
||||
return {"type": "base64", "media_type": media_type, "data": data} # mutable-ok: json-serialized message part
|
||||
|
||||
|
||||
def _inline(remote: _RemoteImage | _RemoteFile | _RemoteSource, data_url: str) -> Mapping[str, object]:
|
||||
match remote:
|
||||
case _RemoteImage(part, image_url, _):
|
||||
return {**part, "image_url": _inlined_image_url(image_url, data_url)} # mutable-ok: json-serialized part
|
||||
case _RemoteFile(part, file, url):
|
||||
return {**part, "file": _inlined_file(file, url, data_url)} # mutable-ok: json-serialized message part
|
||||
case _RemoteSource(part, _, url):
|
||||
return {**part, "source": _base64_source(url, data_url)} # mutable-ok: json-serialized message part
|
||||
|
||||
|
||||
def _content_parts(message: Mapping[str, object]) -> tuple[object, ...]:
|
||||
content: Final = message.get("content")
|
||||
return tuple(content) if isinstance(content, list) else () # pyright: ignore[reportUnknownVariableType, reportUnknownArgumentType] # parts are parsed one by one
|
||||
|
||||
|
||||
def _inline_part(part: object, data_urls: Mapping[str, str], should_inline: Callable[[RemoteMedia], bool]) -> object:
|
||||
remote: Final = _parse_remote_part(part)
|
||||
if remote is None or not should_inline(_remote_media(remote)):
|
||||
return part
|
||||
data_url: Final = data_urls.get(remote.url)
|
||||
return _inline(remote, data_url) if data_url is not None else part
|
||||
|
||||
|
||||
def _inline_message(
|
||||
message: AllMessageValues, data_urls: Mapping[str, str], should_inline: Callable[[RemoteMedia], bool]
|
||||
) -> AllMessageValues:
|
||||
parts: Final = _content_parts(message)
|
||||
if not parts:
|
||||
return message
|
||||
inlined_parts: Final = [ # mutable-ok: content must stay a list for the transforms' isinstance checks
|
||||
_inline_part(part, data_urls, should_inline) for part in parts
|
||||
]
|
||||
inlined_message: Final = {**message, "content": inlined_parts} # mutable-ok: json-serialized message
|
||||
return inlined_message # pyright: ignore[reportReturnType] # the same message with its remote parts inlined
|
||||
|
||||
|
||||
async def _fetch_data_url(url: str, in_flight: asyncio.Semaphore) -> str:
|
||||
async with in_flight:
|
||||
return await async_convert_url_to_base64(url)
|
||||
|
||||
|
||||
async def _fetch_data_urls(remote_urls: tuple[str, ...]) -> tuple[str, ...]:
|
||||
in_flight: Final = asyncio.Semaphore(MAX_CONCURRENT_REMOTE_MEDIA_FETCHES)
|
||||
fetches: Final = tuple(asyncio.create_task(_fetch_data_url(url, in_flight)) for url in remote_urls)
|
||||
try:
|
||||
return tuple(await asyncio.gather(*fetches))
|
||||
except BaseException:
|
||||
for fetch in fetches:
|
||||
fetch.cancel()
|
||||
await asyncio.gather(*fetches, return_exceptions=True)
|
||||
raise
|
||||
|
||||
|
||||
async def async_inline_remote_media(
|
||||
messages: list[AllMessageValues], # mutable-ok: every transform_request takes list[AllMessageValues]
|
||||
should_inline: Callable[[RemoteMedia], bool] = inline_every_remote_url,
|
||||
) -> list[AllMessageValues]: # mutable-ok: every transform_request takes list[AllMessageValues]
|
||||
remote_urls: Final = tuple(
|
||||
dict.fromkeys(
|
||||
remote.url
|
||||
for message in messages
|
||||
for part in _content_parts(message)
|
||||
if (remote := _parse_remote_part(part)) is not None and should_inline(_remote_media(remote))
|
||||
)
|
||||
)
|
||||
if not remote_urls:
|
||||
return messages
|
||||
data_urls: Final = await _fetch_data_urls(remote_urls)
|
||||
inlined: Final = MappingProxyType(dict(zip(remote_urls, data_urls, strict=True)))
|
||||
return [ # mutable-ok: transform_request takes a list
|
||||
_inline_message(message, inlined, should_inline) for message in messages
|
||||
]
|
||||
|
|
|
|||
|
|
@ -29,3 +29,11 @@ def websocket_close_reason(message: str, fallback: str) -> str:
|
|||
if len(encoded) <= WEBSOCKET_CLOSE_REASON_MAX_BYTES:
|
||||
return message
|
||||
return encoded[:WEBSOCKET_CLOSE_REASON_MAX_BYTES].decode("utf-8", errors="ignore")
|
||||
|
||||
|
||||
def client_close_code(upstream_code: int) -> int:
|
||||
from websockets.frames import EXTERNAL_CLOSE_CODES, CloseCode
|
||||
|
||||
if upstream_code in EXTERNAL_CLOSE_CODES or 3000 <= upstream_code < 5000:
|
||||
return upstream_code
|
||||
return int(CloseCode.INTERNAL_ERROR)
|
||||
|
|
|
|||
|
|
@ -1,12 +1,15 @@
|
|||
import asyncio
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, TypedDict, cast
|
||||
import traceback
|
||||
from collections.abc import Coroutine, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum, auto
|
||||
from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, TypedDict, cast
|
||||
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._logging import redact_internal_details_from_client_message, verbose_logger
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
|
||||
from litellm.types.llms.openai import (
|
||||
|
|
@ -19,9 +22,11 @@ from litellm.types.llms.openai import (
|
|||
from litellm.types.realtime import ALL_DELTA_TYPES
|
||||
|
||||
from .litellm_logging import Logging as LiteLLMLogging
|
||||
from .realtime_errors import client_close_code, realtime_error_event, websocket_close_reason
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
|
|
@ -30,8 +35,30 @@ else:
|
|||
CLIENT_CONNECTION_CLASS = Any
|
||||
|
||||
|
||||
class _ClientWebSocketExceptions(Protocol):
|
||||
ConnectionClosed: type[Exception]
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY: Final = "realtime_session_success_logged"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BackendClose:
|
||||
code: int
|
||||
reason: str
|
||||
|
||||
@property
|
||||
def message(self) -> str:
|
||||
if not self.reason:
|
||||
return f"upstream websocket closed with code {self.code}"
|
||||
return f"upstream websocket closed with code {self.code}: {self.reason}"
|
||||
|
||||
|
||||
class ClientLoopExit(Enum):
|
||||
CLIENT_DISCONNECTED = auto()
|
||||
BACKEND_CLOSED = auto()
|
||||
|
||||
|
||||
def backend_close_from(error: "ConnectionClosed") -> BackendClose:
|
||||
if error.rcvd is None:
|
||||
return BackendClose(code=1006, reason=str(error))
|
||||
return BackendClose(code=error.rcvd.code, reason=error.rcvd.reason)
|
||||
|
||||
|
||||
class _ASGIScope(TypedDict, total=False):
|
||||
|
|
@ -69,10 +96,13 @@ class _ScopedWebSocket(Protocol):
|
|||
|
||||
|
||||
class _ClientWebSocket(_ScopedWebSocket, Protocol):
|
||||
exceptions: _ClientWebSocketExceptions
|
||||
|
||||
async def send_text(self, data: str) -> None: ...
|
||||
async def receive_text(self) -> str: ...
|
||||
async def close(self, code: int = 1000, reason: str | None = None) -> None: ...
|
||||
|
||||
|
||||
class _LoggingWorker(Protocol):
|
||||
def ensure_initialized_and_enqueue(self, async_coroutine: Coroutine[object, object, None]) -> None: ...
|
||||
|
||||
|
||||
def _decode_json_object(payload: str) -> Mapping[str, object]:
|
||||
|
|
@ -108,11 +138,14 @@ class RealTimeStreaming:
|
|||
backend_uses_beta_protocol: bool | None = None,
|
||||
force_transcription_model: str | None = None,
|
||||
event_normalizer: RealtimeEventNormalizer | None = None,
|
||||
logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER,
|
||||
):
|
||||
self.websocket: _ClientWebSocket = websocket
|
||||
self.backend_ws = backend_ws
|
||||
self.logging_obj = logging_obj
|
||||
self._logging_worker = logging_worker
|
||||
self.messages: list[OpenAIRealtimeEvents] = []
|
||||
self._backend_sent_frames: bool = False
|
||||
self.input_message: dict = {}
|
||||
self.input_messages: list[dict[str, str]] = []
|
||||
self.session_tools: list[dict] = []
|
||||
|
|
@ -388,9 +421,10 @@ class RealTimeStreaming:
|
|||
# Route through the bounded logging worker (per-coroutine timeout +
|
||||
# concurrency cap) instead of a bare create_task, so a slow callback
|
||||
# can't leave suspended tasks pinning each call's response in memory.
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
self._logging_worker.ensure_initialized_and_enqueue(
|
||||
self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True)
|
||||
)
|
||||
self.logging_obj.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
|
||||
|
||||
async def _send_to_backend(self, message: str) -> bool:
|
||||
"""Send a message to the backend WebSocket.
|
||||
|
|
@ -1035,60 +1069,84 @@ class RealTimeStreaming:
|
|||
return True
|
||||
return False
|
||||
|
||||
async def backend_to_client_send_messages(self):
|
||||
async def _relay_backend_messages(self) -> NoReturn:
|
||||
while True:
|
||||
try:
|
||||
raw_response = await self.backend_ws.recv(decode=False)
|
||||
except TypeError:
|
||||
raw_response = await self.backend_ws.recv()
|
||||
self._backend_sent_frames = True
|
||||
|
||||
if isinstance(raw_response, bytes):
|
||||
try:
|
||||
raw_response = raw_response.decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
verbose_logger.warning("Received non-UTF-8 binary frame from backend, skipping.")
|
||||
continue
|
||||
|
||||
if self.provider_config:
|
||||
try:
|
||||
await self._handle_provider_config_message(raw_response)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error processing backend message, skipping: %s", e)
|
||||
continue
|
||||
else:
|
||||
event = self._parse_backend_event(raw_response)
|
||||
if event is None:
|
||||
await self.websocket.send_text(raw_response)
|
||||
continue
|
||||
|
||||
if self._should_drop_event_from_client(event):
|
||||
continue
|
||||
|
||||
if await self._handle_raw_backend_message(event, raw_response):
|
||||
continue
|
||||
|
||||
event = self._normalize_event_for_ga_client(event)
|
||||
self.store_message(event)
|
||||
|
||||
if not self._client_wants_beta:
|
||||
await self.websocket.send_text(json.dumps(event))
|
||||
continue
|
||||
|
||||
translated = self._translate_event_to_beta(event)
|
||||
if translated is None:
|
||||
continue
|
||||
await self.websocket.send_text(json.dumps(translated))
|
||||
|
||||
async def backend_to_client_send_messages(self) -> BackendClose:
|
||||
import websockets
|
||||
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
raw_response = await self.backend_ws.recv(decode=False)
|
||||
except TypeError:
|
||||
raw_response = await self.backend_ws.recv()
|
||||
|
||||
if isinstance(raw_response, bytes):
|
||||
try:
|
||||
raw_response = raw_response.decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
verbose_logger.warning("Received non-UTF-8 binary frame from backend, skipping.")
|
||||
continue
|
||||
|
||||
if self.provider_config:
|
||||
try:
|
||||
await self._handle_provider_config_message(raw_response)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error processing backend message, skipping: %s", e)
|
||||
continue
|
||||
else:
|
||||
event = self._parse_backend_event(raw_response)
|
||||
if event is None:
|
||||
await self.websocket.send_text(raw_response)
|
||||
continue
|
||||
|
||||
if self._should_drop_event_from_client(event):
|
||||
continue
|
||||
|
||||
if await self._handle_raw_backend_message(event, raw_response):
|
||||
continue
|
||||
|
||||
event = self._normalize_event_for_ga_client(event)
|
||||
self.store_message(event)
|
||||
|
||||
if not self._client_wants_beta:
|
||||
await self.websocket.send_text(json.dumps(event))
|
||||
continue
|
||||
|
||||
translated = self._translate_event_to_beta(event)
|
||||
if translated is None:
|
||||
continue
|
||||
await self.websocket.send_text(json.dumps(translated))
|
||||
|
||||
await self._relay_backend_messages()
|
||||
except websockets.exceptions.ConnectionClosed as e:
|
||||
verbose_logger.exception("Connection closed in backend to client send messages - %s", e)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error in backend to client send messages: %s", e)
|
||||
finally:
|
||||
close: Final = backend_close_from(e)
|
||||
self._flush_unbilled_transcription_usage()
|
||||
if self._backend_refused_session(close):
|
||||
await self.log_backend_refusal(e)
|
||||
else:
|
||||
await self.log_messages()
|
||||
return close
|
||||
except asyncio.CancelledError:
|
||||
self._flush_unbilled_transcription_usage()
|
||||
await self.log_messages()
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error in backend to client send messages: %s", e)
|
||||
self._flush_unbilled_transcription_usage()
|
||||
await self.log_messages()
|
||||
return BackendClose(code=1011, reason="proxy failed while relaying the upstream websocket")
|
||||
|
||||
def _backend_refused_session(self, close: BackendClose) -> bool:
|
||||
return close.code != 1000 and not self._backend_sent_frames
|
||||
|
||||
async def log_backend_refusal(self, error: Exception) -> None:
|
||||
if not self.logging_obj:
|
||||
return
|
||||
self._logging_worker.ensure_initialized_and_enqueue(
|
||||
self.logging_obj.dispatch_failure_handlers(error, traceback.format_exc(), prefer_async_handlers=True)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _detect_beta_header(websocket: _ScopedWebSocket) -> bool:
|
||||
|
|
@ -1243,11 +1301,22 @@ class RealTimeStreaming:
|
|||
item["content"] = new_content
|
||||
return item
|
||||
|
||||
async def client_ack_messages(self):
|
||||
async def _receive_client_message(self) -> str | None:
|
||||
try:
|
||||
return await self.websocket.receive_text()
|
||||
except Exception as e: # noqa: BLE001 # whatever the client socket raises, the client is gone
|
||||
verbose_logger.debug("Client disconnected: %s", e)
|
||||
return None
|
||||
|
||||
async def client_ack_messages(self) -> ClientLoopExit:
|
||||
import websockets
|
||||
|
||||
client_event: _ClientEventFrame
|
||||
try:
|
||||
while True:
|
||||
message = await self.websocket.receive_text()
|
||||
message = await self._receive_client_message()
|
||||
if message is None:
|
||||
return ClientLoopExit.CLIENT_DISCONNECTED
|
||||
|
||||
## GUARDRAIL: intercept conversation.item.create for text-based injection.
|
||||
guardrail_turn_detection_injected = False
|
||||
|
|
@ -1481,23 +1550,38 @@ class RealTimeStreaming:
|
|||
if guardrail_turn_detection_injected and sent:
|
||||
self._guardrail_turn_detection_update_sent = True
|
||||
|
||||
except websockets.exceptions.ConnectionClosed as e:
|
||||
verbose_logger.debug("Backend closed while forwarding a client message: %s", e)
|
||||
return ClientLoopExit.BACKEND_CLOSED
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Error in client ack messages: %s", e)
|
||||
return ClientLoopExit.CLIENT_DISCONNECTED
|
||||
|
||||
async def bidirectional_forward(self):
|
||||
async def bidirectional_forward(self) -> None:
|
||||
forward_task: Final = asyncio.create_task(self.backend_to_client_send_messages())
|
||||
client_task: Final = asyncio.create_task(self.client_ack_messages())
|
||||
try:
|
||||
await self.client_ack_messages()
|
||||
except self.websocket.exceptions.ConnectionClosed:
|
||||
verbose_logger.debug("Connection closed")
|
||||
forward_task.cancel()
|
||||
await asyncio.wait((forward_task, client_task), return_when=asyncio.FIRST_COMPLETED)
|
||||
if client_task.done() and client_task.result() is ClientLoopExit.CLIENT_DISCONNECTED:
|
||||
return
|
||||
await self._close_client(await forward_task)
|
||||
finally:
|
||||
if not forward_task.done():
|
||||
forward_task.cancel()
|
||||
try:
|
||||
await forward_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
forward_task.cancel()
|
||||
client_task.cancel()
|
||||
await asyncio.gather(forward_task, client_task, return_exceptions=True)
|
||||
|
||||
async def _close_client(self, close: BackendClose) -> None:
|
||||
redacted_message: Final = redact_internal_details_from_client_message(close.message)
|
||||
redacted_reason: Final = redact_internal_details_from_client_message(close.reason)
|
||||
try:
|
||||
if close.code != 1000:
|
||||
await self.websocket.send_text(realtime_error_event(redacted_message, error_type="server_error"))
|
||||
await self.websocket.close(
|
||||
code=client_close_code(close.code),
|
||||
reason=websocket_close_reason(redacted_reason, fallback=redacted_message),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # the client may already be gone; the session is over either way
|
||||
verbose_logger.debug("Could not relay the upstream close to the client: %s", e)
|
||||
|
||||
|
||||
def client_sent_openai_beta_realtime_header(websocket: _ScopedWebSocket) -> bool:
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ Admins can opt out via two ``litellm`` globals (wired from proxy config):
|
|||
check but still resolve DNS and still rewrite HTTP to the resolved IP.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import socket
|
||||
from ipaddress import ip_address, ip_network
|
||||
from typing import Any, Final, Protocol
|
||||
|
|
@ -471,7 +472,7 @@ async def async_safe_get(client: Any, url: str, **kwargs: Any) -> httpx.Response
|
|||
kwargs.pop("follow_redirects", None)
|
||||
headers_view: Final[_CallerHeadersView] = {"headers": kwargs.pop("headers", {})}
|
||||
for _ in range(_MAX_REDIRECTS):
|
||||
validated_url, original_host = validate_url(url)
|
||||
validated_url, original_host = await asyncio.to_thread(validate_url, url)
|
||||
response = await fetcher.get(
|
||||
validated_url,
|
||||
headers={**headers_view["headers"], "Host": original_host},
|
||||
|
|
|
|||
|
|
@ -1411,6 +1411,25 @@ def flatten_unencrypted_web_search_results_in_anthropic_messages( # mutable-ok:
|
|||
return [_flatten_web_search_results_in_message(m) for m in messages] # mutable-ok: JSON wire format
|
||||
|
||||
|
||||
def _without_provider_specific_fields(block: object) -> object:
|
||||
if not isinstance(block, dict) or "provider_specific_fields" not in block:
|
||||
return block
|
||||
return {k: v for k, v in block.items() if k != "provider_specific_fields"} # mutable-ok: JSON wire format
|
||||
|
||||
|
||||
def _strip_provider_specific_fields_in_message(message: object) -> object:
|
||||
if not isinstance(message, dict) or not isinstance(message.get("content"), list):
|
||||
return message
|
||||
content: Final = [_without_provider_specific_fields(b) for b in message["content"]] # mutable-ok: JSON wire format
|
||||
return {**message, "content": content} # mutable-ok: JSON wire format
|
||||
|
||||
|
||||
def strip_provider_specific_fields_from_anthropic_messages(
|
||||
messages: Sequence[object],
|
||||
) -> Sequence[object]:
|
||||
return [_strip_provider_specific_fields_in_message(m) for m in messages] # mutable-ok: JSON wire format
|
||||
|
||||
|
||||
def _normalized_cache_control(cache_control: object) -> dict[str, str] | None: # mutable-ok: JSON wire format
|
||||
if not isinstance(cache_control, Mapping):
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from typing import (
|
|||
Final,
|
||||
Literal,
|
||||
Protocol,
|
||||
cast, # noqa: TID251 # rebuilt message_delta dict spans the ContentBlockDelta/MessageBlockDelta union
|
||||
get_args,
|
||||
)
|
||||
|
||||
|
|
@ -100,6 +101,10 @@ class _CombinedChunkSplitter:
|
|||
@staticmethod
|
||||
def _is_combined(chunk: "ModelResponseStream") -> bool:
|
||||
"""True if ``chunk`` carries response content AND a finish_reason."""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
|
||||
openai_chat_refusal_text,
|
||||
)
|
||||
|
||||
choices: Final = _optional_attr_sequence(chunk, "choices")
|
||||
if not choices:
|
||||
return False
|
||||
|
|
@ -114,6 +119,7 @@ class _CombinedChunkSplitter:
|
|||
or _optional_attr(delta, "tool_calls")
|
||||
or _optional_attr(delta, "reasoning_content")
|
||||
or _optional_attr(delta, "thinking_blocks")
|
||||
or openai_chat_refusal_text(delta)
|
||||
)
|
||||
|
||||
_PAYLOAD_FIELD_GROUPS: "tuple[tuple[str, ...], ...]" = (
|
||||
|
|
@ -305,6 +311,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
# Synthesized compaction block from compact_20260112 polyfill (streaming).
|
||||
self.compaction_block = compaction_block
|
||||
self.iterations_usage = iterations_usage
|
||||
self._refusal_text: str = ""
|
||||
self.sent_compaction_block: bool = False
|
||||
# Per-phase flags so the compaction block's start/delta/stop events
|
||||
# are emitted (and the public state machine is advanced) in
|
||||
|
|
@ -572,6 +579,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
current_content_block_index=self.current_content_block_index,
|
||||
applied_edits=(self.applied_edits if is_final_chunk and not will_merge_into_held else None),
|
||||
)
|
||||
processed_chunk = self._with_refusal_stop_details(processed_chunk)
|
||||
|
||||
# Check if this is a usage chunk and we have a held stop_reason chunk
|
||||
if will_merge_into_held:
|
||||
|
|
@ -806,6 +814,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
current_content_block_index=self.current_content_block_index,
|
||||
applied_edits=(self.applied_edits if is_final_chunk and not will_merge_into_held else None),
|
||||
)
|
||||
processed_chunk = self._with_refusal_stop_details(processed_chunk)
|
||||
|
||||
# Check if this is a usage chunk and we have a held stop_reason chunk
|
||||
if will_merge_into_held:
|
||||
|
|
@ -993,6 +1002,31 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
def _increment_content_block_index(self):
|
||||
self.current_content_block_index += 1
|
||||
|
||||
def _with_refusal_stop_details(
|
||||
self,
|
||||
processed_chunk: ContentBlockDelta | MessageBlockDelta,
|
||||
) -> ContentBlockDelta | MessageBlockDelta:
|
||||
if processed_chunk.get("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
|
||||
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),
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _delta_has_content(processed_chunk: Mapping[str, object]) -> bool:
|
||||
"""Return True if a translated chunk carries a non-empty
|
||||
|
|
@ -1035,6 +1069,9 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
@staticmethod
|
||||
def _is_blank_delta(chunk: "ModelResponseStream") -> bool:
|
||||
from litellm.llms.anthropic.common_utils import is_empty_unsigned_thinking_block
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
|
||||
openai_chat_refusal_text,
|
||||
)
|
||||
|
||||
choice: Final = chunk.choices[0]
|
||||
if choice.finish_reason is not None:
|
||||
|
|
@ -1044,6 +1081,8 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
return False
|
||||
if getattr(delta, "content", None):
|
||||
return False
|
||||
if openai_chat_refusal_text(delta):
|
||||
return False
|
||||
if getattr(delta, "reasoning_content", None):
|
||||
return False
|
||||
# thinking_blocks whose entries are all empty AND unsigned must not
|
||||
|
|
@ -1067,13 +1106,19 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
- Different content types in the response
|
||||
- Specific markers in the content
|
||||
"""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
|
||||
openai_chat_refusal_text,
|
||||
)
|
||||
|
||||
from .transformation import LiteLLMAnthropicMessagesAdapter
|
||||
|
||||
# Example logic - customize based on your needs:
|
||||
# If chunk indicates a tool call
|
||||
if chunk.choices[0].finish_reason is not None:
|
||||
return False
|
||||
|
||||
refusal_text: Final = openai_chat_refusal_text(chunk.choices[0].delta)
|
||||
if refusal_text is not None:
|
||||
self._refusal_text = self._refusal_text + refusal_text
|
||||
|
||||
(
|
||||
block_type,
|
||||
content_block_start,
|
||||
|
|
|
|||
|
|
@ -117,6 +117,10 @@ from litellm.llms.anthropic.common_utils import (
|
|||
from litellm.llms.anthropic.experimental_pass_through.context_management import (
|
||||
PolyfillResult,
|
||||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
|
||||
openai_chat_refusal_text,
|
||||
refusal_stop_details,
|
||||
)
|
||||
from litellm.types.llms.anthropic import (
|
||||
ANTHROPIC_HOSTED_TOOLS,
|
||||
AllAnthropicPassThroughMessageValues,
|
||||
|
|
@ -1314,6 +1318,8 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
new_content.append(
|
||||
AnthropicResponseContentBlockText(type="text", text=choice.message.content).model_dump()
|
||||
)
|
||||
if (refusal_text := openai_chat_refusal_text(choice.message)) is not None:
|
||||
new_content.append(AnthropicResponseContentBlockText(type="text", text=refusal_text).model_dump())
|
||||
# Handle tool calls (in parallel to text content)
|
||||
if choice.message.tool_calls is not None and len(choice.message.tool_calls) > 0:
|
||||
for tool_call in choice.message.tool_calls:
|
||||
|
|
@ -1346,7 +1352,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
# Add provider_specific_fields if signature is present
|
||||
if provider_specific_fields:
|
||||
tool_use_block.provider_specific_fields = provider_specific_fields
|
||||
new_content.append(tool_use_block.model_dump())
|
||||
new_content.append(tool_use_block.model_dump(exclude_none=True))
|
||||
|
||||
return new_content
|
||||
|
||||
|
|
@ -1472,14 +1478,23 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
choices=response.choices,
|
||||
tool_name_mapping=tool_name_mapping,
|
||||
)
|
||||
refusal_text: Final = next(
|
||||
(text for choice in response.choices if (text := openai_chat_refusal_text(choice.message)) is not None),
|
||||
None,
|
||||
)
|
||||
|
||||
if polyfill_result is not None and polyfill_result.compaction_block is not None:
|
||||
anthropic_content.insert(0, polyfill_result.compaction_block)
|
||||
|
||||
## extract finish reason
|
||||
anthropic_finish_reason: Final = self._translate_openai_finish_reason_to_anthropic(
|
||||
translated_finish_reason: Final = self._translate_openai_finish_reason_to_anthropic(
|
||||
openai_finish_reason=response.choices[0].finish_reason
|
||||
)
|
||||
anthropic_finish_reason: Final = (
|
||||
"refusal"
|
||||
if refusal_text is not None and translated_finish_reason != "max_tokens"
|
||||
else translated_finish_reason
|
||||
)
|
||||
# extract usage
|
||||
usage: Final[Usage] = getattr(response, "usage")
|
||||
anthropic_usage: Final = self._translate_openai_usage_to_anthropic_usage(usage)
|
||||
|
|
@ -1501,6 +1516,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
usage=anthropic_usage,
|
||||
content=anthropic_content,
|
||||
stop_reason=anthropic_finish_reason,
|
||||
stop_details=(refusal_stop_details(refusal_text) if anthropic_finish_reason == "refusal" else None),
|
||||
)
|
||||
|
||||
applied_edits: Final = polyfill_result.applied_edits_for_response() if polyfill_result else None
|
||||
|
|
@ -1541,7 +1557,9 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
"signature": thought_sig,
|
||||
}
|
||||
return "tool_use", cast("ContentBlockContentBlockDict", tool_block)
|
||||
elif choice.delta.content is not None and len(choice.delta.content) > 0:
|
||||
elif (choice.delta.content is not None and len(choice.delta.content) > 0) or openai_chat_refusal_text(
|
||||
choice.delta
|
||||
) is not None:
|
||||
return "text", TextBlock(type="text", text="")
|
||||
elif isinstance(choice, StreamingChoices) and hasattr(choice.delta, "thinking_blocks"):
|
||||
thinking_blocks = choice.delta.thinking_blocks or []
|
||||
|
|
@ -1613,7 +1631,10 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
elif reasoning_content:
|
||||
return "thinking_delta", ContentThinkingBlockDelta(type="thinking_delta", thinking=reasoning_content)
|
||||
else:
|
||||
return "text_delta", ContentTextBlockDelta(type="text_delta", text=text)
|
||||
refusal_text: Final = "".join(
|
||||
refusal for choice in choices if (refusal := openai_chat_refusal_text(choice.delta)) is not None
|
||||
)
|
||||
return "text_delta", ContentTextBlockDelta(type="text_delta", text=text + refusal_text)
|
||||
|
||||
def translate_streaming_openai_response_to_anthropic(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from litellm.llms.anthropic.common_utils import (
|
|||
flatten_unencrypted_web_search_results_in_anthropic_messages,
|
||||
sanitize_tool_use_ids_in_anthropic_messages,
|
||||
strip_empty_content_blocks_from_anthropic_messages,
|
||||
strip_provider_specific_fields_from_anthropic_messages,
|
||||
)
|
||||
from litellm.llms.base_llm.anthropic_messages.transformation import (
|
||||
BaseAnthropicMessagesConfig,
|
||||
|
|
@ -650,7 +651,7 @@ def anthropic_messages_handler(
|
|||
|
||||
return base_llm_http_handler.anthropic_messages_handler(
|
||||
model=model,
|
||||
messages=messages,
|
||||
messages=strip_provider_specific_fields_from_anthropic_messages(messages),
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
anthropic_messages_optional_request_params=dict(anthropic_messages_optional_request_params),
|
||||
_is_async=is_async,
|
||||
|
|
|
|||
|
|
@ -1,8 +1,11 @@
|
|||
from collections.abc import Mapping
|
||||
from collections.abc import Iterable, Mapping, Sequence
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING, Any, Final, cast, get_type_hints
|
||||
|
||||
from litellm.types.llms.anthropic import AnthropicMessagesRequestOptionalParams
|
||||
from litellm.types.llms.anthropic import (
|
||||
AnthropicMessagesRequestOptionalParams,
|
||||
AnthropicStopDetails,
|
||||
)
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
)
|
||||
|
|
@ -25,6 +28,69 @@ def get_safeguard_refusal_stop_details(response: object) -> Mapping[str, Any] |
|
|||
return stop_details if isinstance(stop_details, dict) else None
|
||||
|
||||
|
||||
def refusal_stop_details(explanation: str | None) -> AnthropicStopDetails:
|
||||
"""The ``stop_details`` object accompanying a translated ``stop_reason: "refusal"``."""
|
||||
return AnthropicStopDetails(type="refusal", category=None, explanation=explanation)
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def _mapping_str_field(container: object, key: str) -> str | None:
|
||||
value: Final = _mapping_field(container, key)
|
||||
return value if isinstance(value, str) and value else None
|
||||
|
||||
|
||||
def openai_chat_refusal_text(message_or_delta: object) -> str | None:
|
||||
"""
|
||||
Refusal text carried by an OpenAI Chat Completions message or streaming delta,
|
||||
read from ``refusal`` or from the ``provider_specific_fields`` LiteLLM parks it
|
||||
in, or None when the turn is not a refusal.
|
||||
"""
|
||||
refusal: Final = getattr(message_or_delta, "refusal", None)
|
||||
if isinstance(refusal, str) and refusal:
|
||||
return refusal
|
||||
return _mapping_str_field(getattr(message_or_delta, "provider_specific_fields", None), "refusal")
|
||||
|
||||
|
||||
def _responses_message_refusal_text(item: object) -> str | None:
|
||||
from openai.types.responses import ResponseOutputMessage, ResponseOutputRefusal
|
||||
|
||||
if isinstance(item, ResponseOutputMessage):
|
||||
return next(
|
||||
(part.refusal for part in item.content if isinstance(part, ResponseOutputRefusal) and part.refusal),
|
||||
None,
|
||||
)
|
||||
raw_parts: Final = _mapping_field(item, "content")
|
||||
if _mapping_str_field(item, "type") != "message" or not isinstance(raw_parts, Sequence):
|
||||
return None
|
||||
return next(
|
||||
(
|
||||
refusal
|
||||
for part in cast(Sequence[object], raw_parts) # cast-ok: members re-validated below
|
||||
if _mapping_str_field(part, "type") == "refusal"
|
||||
and isinstance(refusal := _mapping_str_field(part, "refusal"), str)
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def responses_output_refusal_text(output: Iterable[object]) -> str | None:
|
||||
"""
|
||||
Refusal text carried by an OpenAI Responses ``output`` list, in typed
|
||||
(``ResponseOutputRefusal``) or raw-dictionary shape, or None when none of the
|
||||
output messages refused.
|
||||
"""
|
||||
return next(
|
||||
(text for item in output if (text := _responses_message_refusal_text(item)) is not None),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def safeguard_refusal_error(model: str, stop_details: Mapping[str, object]) -> "ContentPolicyViolationError":
|
||||
"""The exception a safeguard-refused Anthropic response converts into so the
|
||||
content-policy fallback chain can re-dispatch it."""
|
||||
|
|
|
|||
|
|
@ -1,13 +1,18 @@
|
|||
# What is this?
|
||||
## Translates OpenAI call to Anthropic `/v1/messages` format
|
||||
import asyncio
|
||||
import json
|
||||
import traceback
|
||||
from collections import deque
|
||||
from collections.abc import AsyncIterator, Mapping
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
|
||||
refusal_stop_details,
|
||||
responses_output_refusal_text,
|
||||
)
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage
|
||||
|
||||
from .transformation import LiteLLMAnthropicToResponsesAPIAdapter
|
||||
|
|
@ -49,6 +54,8 @@ class AnthropicResponsesStreamWrapper:
|
|||
self._sent_message_start = False
|
||||
self._sent_message_stop = False
|
||||
self._chunk_queue: deque[dict[str, object]] = deque()
|
||||
self._refusal_text: str = ""
|
||||
self._sync_responses_iterator: Iterator[object] | None = None
|
||||
|
||||
def _make_message_start(self) -> dict[str, object]:
|
||||
return {
|
||||
|
|
@ -131,6 +138,24 @@ class AnthropicResponsesStreamWrapper:
|
|||
)
|
||||
return
|
||||
|
||||
if event_type == "response.refusal.delta":
|
||||
delta = getattr(event, "delta", "") or (event.get("delta", "") if isinstance(event, dict) else "")
|
||||
if not isinstance(delta, str) or not delta:
|
||||
return
|
||||
self._refusal_text = self._refusal_text + delta
|
||||
item_id = getattr(event, "item_id", None) or (event.get("item_id") if isinstance(event, dict) else None)
|
||||
block_idx = self._item_id_to_block_index.get(item_id, -1) if item_id else self._current_block_index
|
||||
if block_idx < 0:
|
||||
block_idx = self._open_block(item_id, {"type": "text", "text": ""})
|
||||
self._chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": block_idx,
|
||||
"delta": {"type": "text_delta", "text": delta},
|
||||
}
|
||||
)
|
||||
return
|
||||
|
||||
# ---- text delta ----
|
||||
if event_type == "response.output_text.delta":
|
||||
item_id = getattr(event, "item_id", None) or (event.get("item_id") if isinstance(event, dict) else None)
|
||||
|
|
@ -215,34 +240,47 @@ class AnthropicResponsesStreamWrapper:
|
|||
response_obj: Final = getattr(event, "response", None) or (
|
||||
event.get("response") if isinstance(event, dict) else None
|
||||
)
|
||||
stop_reason = "end_turn"
|
||||
anthropic_usage: AnthropicUsage = AnthropicUsage(input_tokens=0, output_tokens=0)
|
||||
|
||||
if response_obj is not None:
|
||||
status: Final = getattr(response_obj, "status", None)
|
||||
if status == "incomplete":
|
||||
stop_reason = "max_tokens"
|
||||
anthropic_usage = (
|
||||
LiteLLMAnthropicToResponsesAPIAdapter.translate_responses_api_usage_to_anthropic_usage(
|
||||
getattr(response_obj, "usage", None)
|
||||
)
|
||||
output: Final = (getattr(response_obj, "output", None) or ()) if response_obj is not None else ()
|
||||
refusal_text: Final = responses_output_refusal_text(output) or (self._refusal_text or None)
|
||||
status: Final = getattr(response_obj, "status", None) if response_obj is not None else None
|
||||
has_tool_call: Final = any(
|
||||
getattr(item, "type", None) == "function_call"
|
||||
or (isinstance(item, dict) and item.get("type") == "function_call")
|
||||
for item in output
|
||||
)
|
||||
stop_reason: Final = (
|
||||
"max_tokens"
|
||||
if status == "incomplete"
|
||||
else "refusal"
|
||||
if refusal_text is not None
|
||||
else "tool_use"
|
||||
if has_tool_call
|
||||
else "end_turn"
|
||||
)
|
||||
anthropic_usage: Final[AnthropicUsage] = (
|
||||
LiteLLMAnthropicToResponsesAPIAdapter.translate_responses_api_usage_to_anthropic_usage(
|
||||
getattr(response_obj, "usage", None)
|
||||
)
|
||||
if response_obj is not None
|
||||
else AnthropicUsage(input_tokens=0, output_tokens=0)
|
||||
)
|
||||
|
||||
# Check if tool_use was in the output to override stop_reason
|
||||
if response_obj is not None:
|
||||
output: Final = getattr(response_obj, "output", []) or []
|
||||
for out_item in output:
|
||||
out_type = getattr(out_item, "type", None) or (
|
||||
out_item.get("type") if isinstance(out_item, dict) else None
|
||||
)
|
||||
if out_type == "function_call":
|
||||
stop_reason = "tool_use"
|
||||
break
|
||||
message_delta_payload: Final = { # mutable-ok: fresh message_delta payload built per chunk
|
||||
"stop_reason": stop_reason,
|
||||
"stop_sequence": None,
|
||||
**(
|
||||
{ # mutable-ok: fresh message_delta stop_details entry built per chunk
|
||||
"stop_details": refusal_stop_details(refusal_text)
|
||||
}
|
||||
if stop_reason == "refusal"
|
||||
else {} # mutable-ok: empty spread placeholder for non-refusal stop
|
||||
),
|
||||
}
|
||||
|
||||
self._chunk_queue.append(
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": stop_reason, "stop_sequence": None},
|
||||
"delta": message_delta_payload,
|
||||
"usage": dict(anthropic_usage),
|
||||
}
|
||||
)
|
||||
|
|
@ -266,10 +304,20 @@ class AnthropicResponsesStreamWrapper:
|
|||
|
||||
# Consume the upstream stream
|
||||
try:
|
||||
async for event in self.responses_stream:
|
||||
self._process_event(event)
|
||||
if self._chunk_queue:
|
||||
return self._chunk_queue.popleft()
|
||||
if hasattr(self.responses_stream, "__aiter__"):
|
||||
async for event in self.responses_stream:
|
||||
self._process_event(event)
|
||||
if self._chunk_queue:
|
||||
return self._chunk_queue.popleft()
|
||||
else:
|
||||
if self._sync_responses_iterator is None:
|
||||
self._sync_responses_iterator = iter(self.responses_stream)
|
||||
sync_iterator: Final = self._sync_responses_iterator
|
||||
missing: Final = object()
|
||||
while (event := await asyncio.to_thread(next, sync_iterator, missing)) is not missing:
|
||||
self._process_event(event)
|
||||
if self._chunk_queue:
|
||||
return self._chunk_queue.popleft()
|
||||
except StopAsyncIteration:
|
||||
pass
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -19,6 +19,10 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
from litellm.litellm_core_utils.reasoning_effort_utils import (
|
||||
reasoning_effort_from_thinking_budget,
|
||||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
|
||||
refusal_stop_details,
|
||||
responses_output_refusal_text,
|
||||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.utils import (
|
||||
is_reasoning_auto_summary_enabled,
|
||||
prompt_cache_key_from_user_id,
|
||||
|
|
@ -624,6 +628,9 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
|
||||
content: Final[list[dict[str, object]]] = []
|
||||
stop_reason: AnthropicFinishReason = "end_turn"
|
||||
refusal_text: Final = responses_output_refusal_text(
|
||||
cast(Iterable[object], response.output) # cast-ok: output items re-validated per item
|
||||
)
|
||||
|
||||
for item in response.output:
|
||||
if isinstance(item, ResponseReasoningItem):
|
||||
|
|
@ -631,10 +638,17 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
|
||||
elif isinstance(item, ResponseOutputMessage):
|
||||
for part in item.content:
|
||||
if getattr(part, "type", None) == "output_text":
|
||||
part_type = getattr(part, "type", None)
|
||||
if part_type == "output_text":
|
||||
content.append(
|
||||
AnthropicResponseContentBlockText(type="text", text=getattr(part, "text", "")).model_dump()
|
||||
)
|
||||
elif part_type == "refusal":
|
||||
content.append(
|
||||
AnthropicResponseContentBlockText(
|
||||
type="text", text=getattr(part, "refusal", "") or ""
|
||||
).model_dump()
|
||||
)
|
||||
|
||||
elif isinstance(item, ResponseFunctionToolCall):
|
||||
try:
|
||||
|
|
@ -647,18 +661,28 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
id=item.call_id or item.id or "",
|
||||
name=item.name,
|
||||
input=input_data,
|
||||
).model_dump()
|
||||
).model_dump(exclude_none=True)
|
||||
)
|
||||
stop_reason = "tool_use"
|
||||
|
||||
elif isinstance(item, dict):
|
||||
item_type = item.get("type")
|
||||
if item_type == "message":
|
||||
for part in item.get("content", []):
|
||||
if isinstance(part, dict) and part.get("type") == "output_text":
|
||||
content.append(
|
||||
AnthropicResponseContentBlockText(type="text", text=part.get("text", "")).model_dump()
|
||||
)
|
||||
for part in item.get("content", ()):
|
||||
if isinstance(part, dict):
|
||||
part_type = part.get("type")
|
||||
if part_type == "output_text":
|
||||
content.append(
|
||||
AnthropicResponseContentBlockText(
|
||||
type="text", text=part.get("text", "")
|
||||
).model_dump()
|
||||
)
|
||||
elif part_type == "refusal":
|
||||
content.append(
|
||||
AnthropicResponseContentBlockText(
|
||||
type="text", text=part.get("refusal", "") or ""
|
||||
).model_dump()
|
||||
)
|
||||
elif item_type == "reasoning":
|
||||
content.extend(
|
||||
self._thinking_blocks_from_reasoning_item(
|
||||
|
|
@ -676,13 +700,13 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
id=item.get("call_id") or item.get("id", ""),
|
||||
name=item.get("name", ""),
|
||||
input=input_data,
|
||||
).model_dump()
|
||||
).model_dump(exclude_none=True)
|
||||
)
|
||||
stop_reason = "tool_use"
|
||||
|
||||
# status -> stop_reason override
|
||||
if response.status == "incomplete":
|
||||
stop_reason = "max_tokens"
|
||||
elif refusal_text is not None:
|
||||
stop_reason = "refusal"
|
||||
|
||||
anthropic_usage: Final = self.translate_responses_api_usage_to_anthropic_usage(response.usage)
|
||||
|
||||
|
|
@ -695,4 +719,5 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
usage=anthropic_usage,
|
||||
content=content,
|
||||
stop_reason=stop_reason,
|
||||
stop_details=(refusal_stop_details(refusal_text) if stop_reason == "refusal" else None),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
from litellm.llms.azure.common_utils import BaseAzureLLM
|
||||
from litellm.llms.azure_ai.common_utils import is_foundry_model_inference_base
|
||||
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
|
||||
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
|
||||
from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error
|
||||
from litellm.llms.openai.openai import OpenAIConfig
|
||||
from litellm.llms.xai.chat.transformation import XAIChatConfig
|
||||
|
|
@ -42,12 +43,37 @@ NON_OPENAI_SPEC_MESSAGE_FIELDS: Final = (
|
|||
)
|
||||
|
||||
|
||||
class AzureAIGPT5Config(OpenAIGPT5Config):
|
||||
@classmethod
|
||||
def _model_map_lookup_name(cls, model: str) -> str:
|
||||
"""Normalise a Foundry routing name to its cost-map key, when the map has one.
|
||||
|
||||
A Foundry deployment and its OpenAI-hosted namesake are different products with
|
||||
different capabilities, so ``azure_ai/<model>`` is the entry to read whenever the map
|
||||
carries it. Most gpt-5-family names have no ``azure_ai/`` row, though, and prefixing
|
||||
those anyway costs them every flag: ``get_llm_provider`` re-resolves an ``azure_ai/``
|
||||
name to the azure provider when a global AZURE_AI_API_BASE points at an
|
||||
openai.azure.com host, ``azure/<model>`` is not a key either, so the lookup lands
|
||||
nowhere and every effort answer degrades to False. A missing key defers to the base
|
||||
resolver instead.
|
||||
"""
|
||||
prefixed: Final = model if model.startswith("azure_ai/") else f"azure_ai/{model}"
|
||||
return prefixed if prefixed in litellm.model_cost else super()._model_map_lookup_name(model)
|
||||
|
||||
|
||||
azureAIGPT5Config: Final = AzureAIGPT5Config()
|
||||
|
||||
|
||||
class AzureAIStudioConfig(OpenAIConfig):
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
model_supports_tool_choice = True # azure ai supports this by default
|
||||
if not supports_tool_choice(model=f"azure_ai/{model}"):
|
||||
model_supports_tool_choice = False
|
||||
supported_params = super().get_supported_openai_params(model)
|
||||
supported_params = (
|
||||
azureAIGPT5Config.get_supported_openai_params(model)
|
||||
if azureAIGPT5Config.is_model_gpt_5_model(model)
|
||||
else super().get_supported_openai_params(model)
|
||||
)
|
||||
if not model_supports_tool_choice:
|
||||
filtered_supported_params: Final = []
|
||||
for param in supported_params:
|
||||
|
|
@ -61,6 +87,27 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
|
||||
return supported_params
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict[str, object], # mutable-ok: OpenAIConfig.map_openai_params signature
|
||||
optional_params: dict[str, object], # mutable-ok: OpenAIConfig.map_openai_params signature
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict[str, object]: # mutable-ok: OpenAIConfig.map_openai_params signature
|
||||
if not azureAIGPT5Config.is_model_gpt_5_model(model):
|
||||
return super().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=drop_params,
|
||||
)
|
||||
return azureAIGPT5Config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=drop_params,
|
||||
)
|
||||
|
||||
def _supports_stop_reason(self, model: str) -> bool:
|
||||
"""
|
||||
Check if the model supports stop tokens.
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Azure AI OCR module."""
|
||||
|
||||
from .cohere_parse_transformation import AzureAICohereParseConfig
|
||||
from .common_utils import get_azure_ai_ocr_config
|
||||
from .document_intelligence.transformation import (
|
||||
AzureDocumentIntelligenceOCRConfig,
|
||||
|
|
@ -7,6 +8,7 @@ from .document_intelligence.transformation import (
|
|||
from .transformation import AzureAIOCRConfig
|
||||
|
||||
__all__ = [
|
||||
"AzureAICohereParseConfig",
|
||||
"AzureAIOCRConfig",
|
||||
"AzureDocumentIntelligenceOCRConfig",
|
||||
"get_azure_ai_ocr_config",
|
||||
|
|
|
|||
91
litellm/llms/azure_ai/ocr/cohere_parse_transformation.py
Normal file
91
litellm/llms/azure_ai/ocr/cohere_parse_transformation.py
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
"""Cohere Parse served from Azure AI Foundry (`/providers/cohere/v2/parse`)."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.image_handling import (
|
||||
async_convert_url_to_base64,
|
||||
convert_url_to_base64,
|
||||
)
|
||||
from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers
|
||||
from litellm.llms.cohere.ocr.transformation import COHERE_PARSE_PATH, CohereParseConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
AZURE_AI_API_KEY_ENV_VAR: Final = "AZURE_AI_API_KEY"
|
||||
AZURE_AI_API_BASE_ENV_VAR: Final = "AZURE_AI_API_BASE"
|
||||
AZURE_AI_COHERE_PROVIDER_PATH: Final = "/providers/cohere"
|
||||
AZURE_AI_MODELS_PATH_SUFFIX: Final = "/models"
|
||||
|
||||
|
||||
class AzureAICohereParseConfig(CohereParseConfig):
|
||||
"""Same request and response shape as Cohere Parse, behind Azure AI auth and URL layout.
|
||||
|
||||
Foundry cannot fetch external URLs, so remote images are inlined as base64 data URIs.
|
||||
"""
|
||||
|
||||
def get_api_key_env_var(self) -> str | None:
|
||||
return AZURE_AI_API_KEY_ENV_VAR
|
||||
|
||||
def _llm_provider(self) -> str:
|
||||
return "azure_ai"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Mapping[str, str],
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.validate_environment signature
|
||||
) -> dict[str, str]: # mutable-ok: BaseOCRConfig signature
|
||||
resolved_base: Final = api_base or get_secret_str(AZURE_AI_API_BASE_ENV_VAR)
|
||||
if resolved_base is None:
|
||||
raise ValueError(
|
||||
f"Missing Azure AI API Base - Set {AZURE_AI_API_BASE_ENV_VAR} environment variable "
|
||||
"or pass api_base parameter"
|
||||
)
|
||||
resolved_key: Final = api_key or get_secret_str(AZURE_AI_API_KEY_ENV_VAR)
|
||||
return { # mutable-ok: BaseOCRConfig signature
|
||||
**get_azure_ai_auth_headers(api_key=resolved_key, litellm_params=litellm_params),
|
||||
"Content-Type": "application/json",
|
||||
**headers,
|
||||
}
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.get_complete_url signature
|
||||
) -> str:
|
||||
resolved_base: Final = api_base or get_secret_str(AZURE_AI_API_BASE_ENV_VAR)
|
||||
if resolved_base is None:
|
||||
raise ValueError(
|
||||
f"Missing Azure AI API Base - Set {AZURE_AI_API_BASE_ENV_VAR} environment variable "
|
||||
"or pass api_base parameter"
|
||||
)
|
||||
url: Final = httpx.URL(resolved_base)
|
||||
if not url.is_absolute_url:
|
||||
raise ValueError(
|
||||
"Azure AI API Base must be an absolute URL including scheme (e.g. "
|
||||
f"'https://<resource>.services.ai.azure.com'). Got api_base={resolved_base!r}."
|
||||
)
|
||||
path: Final = url.path.rstrip("/")
|
||||
if path.endswith(COHERE_PARSE_PATH):
|
||||
return str(url.copy_with(path=path))
|
||||
if path.endswith(f"{AZURE_AI_COHERE_PROVIDER_PATH}/v2"):
|
||||
return str(url.copy_with(path=f"{path}/parse"))
|
||||
return str(
|
||||
url.copy_with(
|
||||
path=f"{path.removesuffix(AZURE_AI_MODELS_PATH_SUFFIX)}{AZURE_AI_COHERE_PROVIDER_PATH}{COHERE_PARSE_PATH}"
|
||||
)
|
||||
)
|
||||
|
||||
def _resolve_image_url_sync(self, image_url: str) -> str:
|
||||
return convert_url_to_base64(image_url)
|
||||
|
||||
async def _resolve_image_url_async(self, image_url: str) -> str:
|
||||
return await async_convert_url_to_base64(image_url)
|
||||
|
|
@ -24,6 +24,11 @@ def is_azure_document_intelligence_model(model: str) -> bool:
|
|||
return "doc-intelligence" in lowered or "documentintelligence" in lowered
|
||||
|
||||
|
||||
def is_azure_cohere_parse_model(model: str) -> bool:
|
||||
lowered: Final = model.lower()
|
||||
return "cohere" in lowered and "parse" in lowered
|
||||
|
||||
|
||||
def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
|
||||
"""
|
||||
Determine which Azure AI OCR configuration to use based on the model name.
|
||||
|
|
@ -46,6 +51,7 @@ def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
|
|||
>>> get_azure_ai_ocr_config("azure_ai/pixtral-12b-2409")
|
||||
<AzureAIOCRConfig object>
|
||||
"""
|
||||
from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig
|
||||
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
|
||||
AzureDocumentIntelligenceOCRConfig,
|
||||
)
|
||||
|
|
@ -56,6 +62,10 @@ def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
|
|||
verbose_logger.debug("Routing %s to Azure Document Intelligence OCR config", model)
|
||||
return AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
if is_azure_cohere_parse_model(model):
|
||||
verbose_logger.debug("Routing %s to Azure AI Cohere Parse config", model)
|
||||
return AzureAICohereParseConfig()
|
||||
|
||||
# Default to Mistral-based OCR for other azure_ai models
|
||||
verbose_logger.debug("Routing %s to Azure AI (Mistral) OCR config", model)
|
||||
return AzureAIOCRConfig()
|
||||
|
|
|
|||
|
|
@ -416,6 +416,10 @@ class BaseConfig(ABC):
|
|||
def has_custom_stream_wrapper(self) -> bool:
|
||||
return False
|
||||
|
||||
@property
|
||||
def uses_async_transform_request(self) -> bool:
|
||||
return False
|
||||
|
||||
@property
|
||||
def supports_stream_param_in_request_body(self) -> bool:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import types
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import httpx
|
||||
|
|
@ -102,6 +103,24 @@ class BaseImageEditConfig(ABC):
|
|||
) -> tuple[dict, RequestFiles]:
|
||||
pass
|
||||
|
||||
async def async_transform_image_edit_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str | None,
|
||||
image: FileTypes | None,
|
||||
image_edit_optional_request_params: Mapping[str, object],
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: Mapping[str, str],
|
||||
) -> tuple[dict, RequestFiles]:
|
||||
return self.transform_image_edit_request(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
image=image,
|
||||
image_edit_optional_request_params=dict(image_edit_optional_request_params),
|
||||
litellm_params=litellm_params,
|
||||
headers=dict(headers),
|
||||
)
|
||||
|
||||
def finalize_image_edit_request_data(self, data: dict, resolved_request_url: str) -> dict:
|
||||
"""
|
||||
Last pass on the request dict after ``transform_image_edit_request``, using the
|
||||
|
|
|
|||
|
|
@ -33,6 +33,8 @@ OCR_REQUEST_FORMAT_HEADER: Final = "x-req-format"
|
|||
|
||||
PROVIDER_NATIVE_RESPONSE_KEY: Final = "provider_native_response"
|
||||
|
||||
HEALTH_CHECK_PDF_DATA_URI: Final = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y="
|
||||
|
||||
|
||||
def parse_ocr_request_format(value: object) -> OCRRequestFormat:
|
||||
if value == "litellm":
|
||||
|
|
@ -142,6 +144,16 @@ class BaseOCRConfig:
|
|||
"""
|
||||
return None
|
||||
|
||||
def supports_rust_bridge(self) -> bool:
|
||||
"""Whether the Rust OCR bridge may serve this config when it is enabled for the provider."""
|
||||
return True
|
||||
|
||||
def get_health_check_document(self) -> DocumentType:
|
||||
return { # mutable-ok: litellm.aocr rejects any document that is not a dict
|
||||
"type": "document_url",
|
||||
"document_url": HEALTH_CHECK_PDF_DATA_URI,
|
||||
}
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
|
|
|
|||
|
|
@ -667,7 +667,12 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
|||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise BedrockError(status_code=response.status_code, message=str(response.read()))
|
||||
raise BedrockError(
|
||||
status_code=response.status_code,
|
||||
message=str(response.read()),
|
||||
headers=response.headers,
|
||||
response=response,
|
||||
)
|
||||
|
||||
# LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -690,6 +695,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
|||
raise BedrockError(
|
||||
status_code=response.status_code,
|
||||
message=f"AgentCore: Failed to read/parse JSON response body: {e}",
|
||||
headers=response.headers,
|
||||
)
|
||||
parsed: Final = self._parse_json_response(response_json)
|
||||
|
||||
|
|
@ -880,7 +886,12 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
|||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise BedrockError(status_code=response.status_code, message=str(await response.aread()))
|
||||
raise BedrockError(
|
||||
status_code=response.status_code,
|
||||
message=str(await response.aread()),
|
||||
headers=response.headers,
|
||||
response=response,
|
||||
)
|
||||
|
||||
# LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -903,6 +914,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
|||
raise BedrockError(
|
||||
status_code=response.status_code,
|
||||
message=f"AgentCore: Failed to read/parse JSON response body: {e}",
|
||||
headers=response.headers,
|
||||
)
|
||||
parsed: Final = self._parse_json_response(response_json)
|
||||
|
||||
|
|
@ -1031,6 +1043,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
|||
raise BedrockError(
|
||||
message=f"Error processing response: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
def validate_environment(
|
||||
|
|
@ -1046,7 +1059,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
|||
return headers
|
||||
|
||||
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException:
|
||||
return BedrockError(status_code=status_code, message=error_message)
|
||||
return BedrockError(status_code=status_code, message=error_message, headers=headers)
|
||||
|
||||
def should_fake_stream(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -22,7 +22,7 @@ from litellm.types.utils import ModelResponse
|
|||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token
|
||||
from ..common_utils import BedrockError, _get_all_bedrock_regions
|
||||
from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text
|
||||
from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call
|
||||
|
||||
|
||||
|
|
@ -66,7 +66,12 @@ def make_sync_call(
|
|||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise BedrockError(status_code=response.status_code, message=str(response.read()))
|
||||
raise BedrockError(
|
||||
status_code=response.status_code,
|
||||
message=str(response.read()),
|
||||
headers=response.headers,
|
||||
response=response,
|
||||
)
|
||||
|
||||
if fake_stream:
|
||||
model_response: Final[ModelResponse] = litellm.AmazonConverseConfig()._transform_response(
|
||||
|
|
@ -247,7 +252,12 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code: Final = err.response.status_code
|
||||
raise BedrockError(status_code=error_code, message=err.response.text)
|
||||
raise BedrockError(
|
||||
status_code=error_code,
|
||||
message=error_response_text(err.response),
|
||||
headers=err.response.headers,
|
||||
response=err.response,
|
||||
)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
|
||||
|
|
@ -594,7 +604,12 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code: Final = err.response.status_code
|
||||
raise BedrockError(status_code=error_code, message=err.response.text)
|
||||
raise BedrockError(
|
||||
status_code=error_code,
|
||||
message=error_response_text(err.response),
|
||||
headers=err.response.headers,
|
||||
response=err.response,
|
||||
)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
|
||||
|
|
|
|||
|
|
@ -2255,6 +2255,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
raise BedrockError(
|
||||
message=f"Error converting to valid response block={e}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues",
|
||||
status_code=422,
|
||||
headers=response.headers,
|
||||
)
|
||||
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -470,6 +470,7 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM):
|
|||
raise BedrockError(
|
||||
message=f"Error processing response: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
def validate_environment(
|
||||
|
|
@ -485,7 +486,7 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM):
|
|||
return headers
|
||||
|
||||
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException:
|
||||
return BedrockError(status_code=status_code, message=error_message)
|
||||
return BedrockError(status_code=status_code, message=error_message, headers=headers)
|
||||
|
||||
def should_fake_stream(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -42,6 +42,7 @@ from litellm.types.utils import GenericStreamingChunk as GChunk
|
|||
from ..common_utils import (
|
||||
BedrockError,
|
||||
build_bedrock_stream_error,
|
||||
error_response_text,
|
||||
get_bedrock_response_stream_shape,
|
||||
get_bedrock_tool_name,
|
||||
)
|
||||
|
|
@ -184,7 +185,12 @@ async def make_call(
|
|||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise BedrockError(status_code=response.status_code, message=response.text)
|
||||
raise BedrockError(
|
||||
status_code=response.status_code,
|
||||
message=error_response_text(response),
|
||||
headers=response.headers,
|
||||
response=response,
|
||||
)
|
||||
|
||||
if fake_stream:
|
||||
model_response: Final[ModelResponse] = litellm.AmazonConverseConfig()._transform_response(
|
||||
|
|
@ -228,9 +234,16 @@ async def make_call(
|
|||
)
|
||||
|
||||
return completion_stream, response.headers
|
||||
except BedrockError:
|
||||
raise
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code: Final = err.response.status_code
|
||||
raise BedrockError(status_code=error_code, message=err.response.text)
|
||||
raise BedrockError(
|
||||
status_code=error_code,
|
||||
message=error_response_text(err.response),
|
||||
headers=err.response.headers,
|
||||
response=err.response,
|
||||
)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
except Exception as e:
|
||||
|
|
@ -270,7 +283,12 @@ def make_sync_call(
|
|||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise BedrockError(status_code=response.status_code, message=response.text)
|
||||
raise BedrockError(
|
||||
status_code=response.status_code,
|
||||
message=error_response_text(response),
|
||||
headers=response.headers,
|
||||
response=response,
|
||||
)
|
||||
|
||||
if fake_stream:
|
||||
model_response: Final[ModelResponse] = litellm.AmazonConverseConfig()._transform_response(
|
||||
|
|
@ -314,9 +332,16 @@ def make_sync_call(
|
|||
)
|
||||
|
||||
return completion_stream, response.headers
|
||||
except BedrockError:
|
||||
raise
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code: Final = err.response.status_code
|
||||
raise BedrockError(status_code=error_code, message=err.response.text)
|
||||
raise BedrockError(
|
||||
status_code=error_code,
|
||||
message=error_response_text(err.response),
|
||||
headers=err.response.headers,
|
||||
response=err.response,
|
||||
)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -247,4 +247,4 @@ class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig):
|
|||
|
||||
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BedrockError:
|
||||
"""Return the appropriate error class for Bedrock."""
|
||||
return BedrockError(status_code=status_code, message=error_message)
|
||||
return BedrockError(status_code=status_code, message=error_message, headers=headers)
|
||||
|
|
|
|||
|
|
@ -182,4 +182,4 @@ class AmazonBedrockOpenAIConfig(OpenAIGPTConfig, BaseAWSLLM):
|
|||
|
||||
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BedrockError:
|
||||
"""Return the appropriate error class for Bedrock."""
|
||||
return BedrockError(status_code=status_code, message=error_message)
|
||||
return BedrockError(status_code=status_code, message=error_message, headers=headers)
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue