Merge remote-tracking branch 'origin/litellm_internal_staging' into pr37112-head

This commit is contained in:
mateo-berri 2026-08-20 11:18:30 -07:00
commit 059aec8887
1557 changed files with 123528 additions and 43006 deletions

View file

@ -1025,7 +1025,7 @@ jobs:
name: Run tests
command: |
mkdir -p test-results
TEST_FILES=$(circleci tests glob "tests/agent_tests/**/test_*.py" | grep -v "^tests/agent_tests/local_only_agent_tests/")
TEST_FILES=$(circleci tests glob "tests/agent_tests/test_*.py")
echo "$TEST_FILES" | circleci tests run \
--verbose \
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \

View file

@ -1,15 +1,17 @@
#!/usr/bin/env bash
set -uo pipefail
category="${1:?usage: classify_changes.sh <backend|client>}"
category="${1:?usage: classify_changes.sh <backend|client|ui>}"
has_client=false
has_backend=false
has_ci=false
while IFS= read -r file || [ -n "$file" ]; do
[ -n "$file" ] || continue
case "$file" in
ui/* | tests/e2e/ui/*) has_client=true ;;
docs/* | *.md | *.mdx) : ;;
.github/* | .circleci/*) has_ci=true; has_backend=true ;;
*) has_backend=true ;;
esac
done
@ -21,6 +23,9 @@ case "$category" in
client)
{ [ "$has_client" = true ] || [ "$has_backend" = true ]; } && echo run || echo skip
;;
ui)
{ [ "$has_client" = true ] || [ "$has_ci" = true ]; } && echo run || echo skip
;;
*)
echo run
;;

2
.github/CODEOWNERS vendored
View file

@ -1,3 +1,5 @@
/ui/ @yuneng-jiang @ryan-crabbe-berri
/litellm/proxy/_experimental/out/ @yuneng-jiang @ryan-crabbe-berri
/ui/litellm-dashboard/src/lib/http/schema.d.ts
/model_prices_and_context_window.json @mateo-berri
/litellm/model_prices_and_context_window_backup.json @mateo-berri

View file

@ -1,48 +0,0 @@
name: "Detect backend-relevant changes"
description: >-
Classify the pull request's changed files with .circleci/scripts/classify_changes.sh
and expose decision=run|skip. decision=skip means only ui/**, **.md or **.mdx files
changed, so callers can short-circuit expensive steps while the job still completes
successfully and satisfies its required status check. The decision defaults to run for
any non pull_request event or whenever the changed set cannot be resolved, so tests are
never skipped when the classification is uncertain.
outputs:
decision:
description: "run when backend-relevant files changed, otherwise skip"
value: ${{ steps.classify.outputs.decision }}
runs:
using: composite
steps:
- id: classify
shell: bash
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
run: |
set -uo pipefail
if [ -z "${BASE_SHA:-}" ]; then
echo "detect-backend-changes: not a pull_request event; running job"
echo "decision=run" >> "${GITHUB_OUTPUT}"
exit 0
fi
if ! git fetch --no-tags --depth=1 origin "${BASE_SHA}" >/dev/null 2>&1; then
echo "detect-backend-changes: could not fetch base ${BASE_SHA}; running job"
echo "decision=run" >> "${GITHUB_OUTPUT}"
exit 0
fi
changed="$(git diff --name-only "${BASE_SHA}" HEAD 2>/dev/null)" || {
echo "detect-backend-changes: git diff failed; running job"
echo "decision=run" >> "${GITHUB_OUTPUT}"
exit 0
}
if [ -z "${changed}" ]; then
echo "detect-backend-changes: no changed files vs ${BASE_SHA}; skipping job"
echo "decision=skip" >> "${GITHUB_OUTPUT}"
exit 0
fi
echo "detect-backend-changes: changed files vs ${BASE_SHA}:"
printf '%s\n' "${changed}" | sed 's/^/ /'
decision="$(printf '%s\n' "${changed}" | bash .circleci/scripts/classify_changes.sh backend)" || decision="run"
echo "detect-backend-changes: decision=${decision}"
echo "decision=${decision}" >> "${GITHUB_OUTPUT}"

View file

@ -0,0 +1,41 @@
name: "Detect relevant changes"
description: >-
Classify the pull request's changed files with .circleci/scripts/classify_changes.sh
and expose decision=run|skip for one category. backend means anything outside ui/,
docs/ and markdown; ui means the dashboard sources alone. decision=skip lets callers
short-circuit expensive steps while the job still completes successfully and satisfies
its required status check, which a paths: filter cannot do because a workflow that
never starts never reports. The file list comes from the pull request itself rather
than from a git diff, because the checked-out merge ref is recomputed as the base
branch advances and would otherwise attribute the base branch's own commits to the
pull request. The decision defaults to run for any non pull_request event or whenever
the changed set cannot be resolved, so jobs are never skipped when the classification
is uncertain.
inputs:
category:
description: "Which classification to apply: backend, client or ui"
required: false
default: backend
github-token:
description: "Token used to list the pull request's files; needs pull-requests: read"
required: false
default: ${{ github.token }}
outputs:
decision:
description: "run when category-relevant files changed, otherwise skip"
value: ${{ steps.classify.outputs.decision }}
runs:
using: composite
steps:
- id: classify
shell: bash
env:
GH_TOKEN: ${{ inputs.github-token }}
CATEGORY: ${{ inputs.category }}
REPO: ${{ github.repository }}
PR_NUMBER: ${{ github.event.pull_request.number }}
CHANGED_FILE_COUNT: ${{ github.event.pull_request.changed_files }}
run: bash "${GITHUB_ACTION_PATH}/../../scripts/detect_changes.sh"

View file

@ -4,6 +4,25 @@ description: >-
by a job nor listed here, so every entry below is a decision on the record.
test_paths:
- reason: >-
The caching suite in tests/local_testing, which runs nowhere. Every job that globs that
directory either deselects it (local_testing_part1 and part2 carry `-k "... and not caching
and not cache"`) or keeps only another keyword (langfuse, router, assistants), and no job
names these files the way redis_caching_unit_tests names test_dual_cache.py. Measured
2026-08-20 by collecting the directory under each job's own selector: 118 tests across
these eight files are selected by none of them. Listed so the gap is a decision rather
than an accident, and so the --slices guard has a baseline to ratchet down from. Revisit
when tests/local_testing is ported off CircleCI, where the keyless part of this suite
belongs in a real job
paths:
- tests/local_testing/test_cache_preset_key.py
- tests/local_testing/test_caching.py
- tests/local_testing/test_caching_handler.py
- tests/local_testing/test_disk_cache_unit_tests.py
- tests/local_testing/test_gcs_cache_unit_tests.py
- tests/local_testing/test_prompt_caching.py
- tests/local_testing/test_responses_stream_cache_keys.py
- tests/local_testing/test_unit_test_caching.py
- reason: >-
The end-to-end suite runs against a deployed proxy from its own in-cluster rig rather than
from a pull request; it needs a live gateway and provider credentials no PR job holds
@ -21,72 +40,27 @@ test_paths:
- tests/documentation_tests/test_requests_lib_usage.py
- tests/documentation_tests/test_standard_logging_payload.py
- reason: >-
Sibling files here are executed by name from the code-quality workflow; this one is referenced
by no job
Named like a test but shaped like a benchmark: it fetches live image URLs, times aiohttp
against httpx, prints the ratio, and asserts nothing, so pytest cannot collect it (its
functions take arguments, not fixtures) and running it beside its siblings in the
code-quality workflow would add a network dependency for a number nothing reads. Exempt
as a script rather than as an unresolved gap; revisit by deleting it once the aiohttp
choice it informed is settled
paths:
- tests/code_coverage_tests/test_aio_http_image_conversion.py
- reason: >-
A second mirror of the package tree living beside tests/test_litellm, which is the mirror the
repo convention names; only test_no_hardcoded_secrets.py is invoked, from the linting
workflow, and whether this directory should exist at all is unresolved
What is left of a second mirror that sat beside tests/test_litellm and ran nowhere. Its
other 30 files moved into the real mirror on 2026-08-20 and now run; these four cannot,
because each shares a filename with a live test whose contents are disjoint from it, so
landing them means merging test bodies rather than moving a file. Measured on the same
date: test_common_utils.py holds 15 tests the live file does not, test_oci_chat_transformation
13, test_deepseek_chat_transformation 12, and test_discoverable_endpoints 5. Revisit by
merging each into its twin, which is a content review, not a move
paths:
- tests/litellm/a2a_protocol/providers/pydantic_ai_agents/test_pydantic_ai_agent_headers.py
- tests/litellm/a2a_protocol/providers/pydantic_ai_agents/test_pydantic_ai_agent_transformation.py
- tests/litellm/integrations/helicone/test_helicone_gemini.py
- tests/litellm/litellm_core_utils/test_json_schema_validation.py
- tests/litellm/llms/anthropic/test_anthropic_reasoning_effort.py
- tests/litellm/llms/anthropic/test_anthropic_schema_filter.py
- tests/litellm/llms/azure/test_azure_embedding.py
- tests/litellm/llms/bedrock/embed/test_embedding.py
- tests/litellm/llms/bedrock/test_nova_imported_models.py
- tests/litellm/llms/deepseek/chat/test_deepseek_chat_transformation.py
- tests/litellm/llms/gradient_ai/chat/test_gradient_ai_chat_transformation.py
- tests/litellm/llms/oci/chat/test_oci_chat_transformation.py
- tests/litellm/llms/openai_like/test_abliteration_provider.py
- tests/litellm/llms/openai_like/test_assemblyai_provider.py
- tests/litellm/llms/openai_like/test_empiriolabs_provider.py
- tests/litellm/llms/vertex_ai/agent_engine/test_transformation.py
- tests/litellm/llms/vertex_ai/gemini/test_transformation.py
- tests/litellm/llms/vertex_ai/text_to_speech/test_transformation.py
- tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py
- tests/litellm/proxy/agent_endpoints/test_agent_rbac.py
- tests/litellm/proxy/common_utils/test_rbac_utils.py
- tests/litellm/proxy/management_endpoints/test_common_utils.py
- tests/litellm/proxy/management_endpoints/test_cost_estimate_endpoint.py
- tests/litellm/proxy/test_claude_code_marketplace.py
- tests/litellm/proxy/test_init_litellm_callbacks.py
- tests/litellm/proxy/test_prisma_engine_watchdog.py
- tests/litellm/proxy/vector_store_endpoints/test_vector_store_rbac.py
- tests/litellm/test_bedrock_extended_beta_models.py
- tests/litellm/test_bedrock_nemotron_super.py
- tests/litellm/test_proxy_auth.py
- tests/litellm/test_router_retry_backoff_headers.py
- tests/litellm/test_sambanova_model_metadata.py
- tests/litellm/test_stream_chunk_builder_images.py
- reason: >-
Legacy proxy suite superseded by the proxy shards; no job invokes it and whether it still
describes supported behaviour is unresolved
paths:
- tests/old_proxy_tests/tests/test_anthropic_context_caching.py
- tests/old_proxy_tests/tests/test_anthropic_sdk.py
- tests/old_proxy_tests/tests/test_async.py
- tests/old_proxy_tests/tests/test_gemini_context_caching.py
- tests/old_proxy_tests/tests/test_langchain_embedding.py
- tests/old_proxy_tests/tests/test_langchain_request.py
- tests/old_proxy_tests/tests/test_llamaindex.py
- tests/old_proxy_tests/tests/test_mistral_sdk.py
- tests/old_proxy_tests/tests/test_openai_embedding.py
- tests/old_proxy_tests/tests/test_openai_exception_request.py
- tests/old_proxy_tests/tests/test_openai_request.py
- tests/old_proxy_tests/tests/test_openai_request_with_traceparent.py
- tests/old_proxy_tests/tests/test_openai_simple_embedding.py
- tests/old_proxy_tests/tests/test_openai_tts_request.py
- tests/old_proxy_tests/tests/test_pass_through_langfuse.py
- tests/old_proxy_tests/tests/test_q.py
- tests/old_proxy_tests/tests/test_simple_traceparent_openai.py
- tests/old_proxy_tests/tests/test_vertex_sdk_forward_headers.py
- tests/old_proxy_tests/tests/test_vtx_embedding.py
- tests/old_proxy_tests/tests/test_vtx_sdk_embedding.py
- reason: >-
No job invokes this suite and its files mix pure transformation tests with ones driving live
vendor vector stores, so assigning them needs a per-file decision
@ -116,6 +90,14 @@ test_paths:
- tests/load_tests/test_otel_load_test.py
- tests/load_tests/test_vertex_embeddings_load_test.py
- tests/load_tests/test_vertex_load_tests.py
- reason: >-
A local-only agent rig: test_a2a_completion_bridge.py needs a LangGraph server on
localhost:2024 and test_a2a.py drives a live A2A endpoint, so neither can run in a
pull request job. Until 2026-08-20 the CircleCI agent job hid them behind a grep -v
that this census could not see; the glob now excludes them structurally and this entry
is the decision on the record. Revisit when the A2A bridge gets a recorded-wire fixture
paths:
- tests/agent_tests/local_only_agent_tests
- reason: >-
Third-party integration tests that skip themselves without OCI configuration or sandbox
credentials, neither of which a pull request job holds
@ -124,14 +106,14 @@ test_paths:
- tests/integration/test_oci_integration.py
- tests/integration/test_oci_proxy_integration.py
- reason: >-
Two prompt-factory tests sitting at the top level of tests/ instead of under the
tests/test_litellm mirror the shards enumerate; they need moving rather than a shard entry
paths:
- tests/litellm_core_utils/test_anthropic_dedup_factory.py
- tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py
- reason: >-
A unit test for the proxy-extras package that no job invokes, while the package's other tests
live under tests/proxy_migration_tests
A unit test for the proxy-extras package that no job invokes, while the package's other
tests live under tests/proxy_migration_tests. Measured 2026-08-20: 24 of its 28 tests pass
and the 4 in TestMigrationSQLIdempotency fail, because 13 migrations from 2026-03 onward use
bare CREATE TABLE, ADD COLUMN, CREATE INDEX and ADD CONSTRAINT rather than the guarded forms
this file requires. It also matches those keywords inside SQL comments, so two further
migrations are reported that are in fact fine. Wiring it up means deciding what to do about
the 13 first, and they cannot simply be edited: Prisma checksums an applied migration, so a
changed one breaks migrate deploy for existing installs
paths:
- tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py

View file

@ -53,7 +53,8 @@ After: the same request comes back with real token counts, so the dashboard show
**Please complete all items before asking a LiteLLM maintainer to review your PR**
- [ ] I have added meaningful tests
- [ ] My PR passes all CI/CD checks (e.g., lint, format, unit tests)
- [ ] The handful of test files covering my change pass locally, e.g. `uv run pytest tests/test_litellm/<your_test_file>.py -v`. Leave the suites (`make test-unit-*`, `make test-unit`) to CI: it finishes in ~15 minutes where a laptop takes an hour or more
- [ ] My PR passes all required CI/CD checks (e.g., lint, schema.d.ts sync check, etc.)
- [ ] My PR's scope is as isolated as possible; it only solves 1 specific problem
- [ ] I have received a Greptile **Confidence Score of at least 4/5** before requesting a maintainer review (Greptile reviews automatically once the PR is opened; only comment `@greptileai` to re-request a review after pushing changes)

View file

@ -1,10 +1,13 @@
from __future__ import annotations
import ast
import pathlib
import re
import sys
import warnings
from collections.abc import Iterable, Mapping, Sequence
from dataclasses import dataclass
from typing import Final
import yaml
@ -25,6 +28,15 @@ DOCKERFILE_TOKEN_RE = re.compile(r"[A-Za-z0-9_./-]*Dockerfile[A-Za-z0-9_.-]*")
COMMENT_RE = re.compile(r"^\s*#.*$", re.MULTILINE)
GLOB_CHARS = frozenset("*?")
# Trees whose jobs are sharded with no catch-all bucket, so every child that holds
# tests has to be named by some shard or it runs nowhere. A child listed here is
# itself decomposed one level deeper and is checked through its own entry.
SHARDED_ROOTS: tuple[str, ...] = (
"tests/proxy_unit_tests",
"tests/test_litellm",
"tests/test_litellm/proxy",
)
@dataclass(frozen=True, slots=True)
class AllowEntry:
@ -107,20 +119,34 @@ def _built_dockerfile_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
)
def _glob_to_regex(token: str) -> re.Pattern[str]:
parts = re.split(r"(\*\*/|\*\*|\*|\?)", token)
def _glob_to_regex(token: str, *, subtree: bool) -> re.Pattern[str]:
parts = re.split(r"(\*\*/|\*\*|\*|\?|\[[^\]]*\])", token)
translated = "".join(
{"**/": r"(?:.*/)?", "**": r".*", "*": r"[^/]*", "?": r"[^/]"}.get(part, re.escape(part)) for part in parts
{"**/": r"(?:.*/)?", "**": r".*", "*": r"[^/]*", "?": r"[^/]"}.get(part)
or (part if part.startswith("[") and part.endswith("]") else re.escape(part))
for part in parts
)
return re.compile(rf"{translated}(?:/.*)?$")
return re.compile(rf"{translated}(?:/.*)?$" if subtree else rf"{translated}$")
def _token_covers(token: str, relative_path: str) -> bool:
if GLOB_CHARS & set(token):
return _glob_to_regex(token).match(relative_path) is not None
return _glob_to_regex(token, subtree=True).match(relative_path) is not None
return relative_path == token or relative_path.startswith(f"{token}/")
def _token_names(token: str, relative_path: str) -> bool:
"""Whether the token names this path itself, rather than merely containing it.
A sharded tree has no catch-all bucket, so the ancestor token the census is happy
with (`tests/x` standing in for everything below it) is exactly what would let a
newly added child ride along without a shard.
"""
if GLOB_CHARS & set(token):
return _glob_to_regex(token, subtree=False).match(relative_path) is not None
return token == relative_path
def _test_files() -> tuple[str, ...]:
return tuple(
sorted(
@ -166,6 +192,174 @@ def _describe(paths: tuple[str, ...]) -> str:
return f"{len(paths)} test file(s) invoked by no job: {names}{suffix}"
GLOB_CALL_RE = re.compile(r'circleci tests glob "([^"]+)"')
KEYWORD_RE = re.compile(r"-k\s+\\?[\"']([^\"'\\]+)")
@dataclass(frozen=True, slots=True)
class Slice:
"""One job's selection: the files it globs, narrowed by its `-k` expression."""
job: str
globs: tuple[str, ...]
named: frozenset[str]
required: tuple[str, ...]
excluded: tuple[str, ...]
understood: bool
def claims(self, relative_path: str, inner_names: frozenset[str]) -> bool:
"""Whether this job runs any test in the file.
The question is deliberately per-file, not per-test. An excluded term is only
honoured when it appears in the path, because that is the case where it takes
the whole module with it; a term matching one function inside drops that test
and leaves the file claimed. Losing a whole file is the failure worth a gate,
and answering per-test would mean a baseline of test ids that churns on every
rename.
"""
if relative_path in self.named:
return True
if not any(_token_covers(glob, relative_path) for glob in self.globs):
return False
if not self.understood:
return True # a `-k` this parser cannot model is assumed to claim everything
if any(term.lower() in relative_path.lower() for term in self.excluded):
return False
return not self.required or any(
term.lower() in name.lower() for term in self.required for name in inner_names
)
def _strings(node: object) -> Iterable[str]:
if isinstance(node, str):
yield node
elif isinstance(node, dict):
for value in node.values():
yield from _strings(value)
elif isinstance(node, list):
for value in node:
yield from _strings(value)
def _keyword_terms(
expressions: Sequence[str], *, attributable: bool = True
) -> tuple[tuple[str, ...], tuple[str, ...], bool]:
"""A `-k` expression as (required, excluded, understood).
Only flat `and` chains of bare terms are modelled. Anything with `or`, parentheses
or negation of a group is left unmodelled, and its job is then treated as claiming
every file it globs, so an unparsed selector can never raise a false alarm.
`attributable` is False when a job runs several pytest commands, since a selector
read out of the job's text cannot then be tied to the glob it belongs to, and
pairing one command's exclusion with another's glob would invent a gap.
"""
terms: Final = tuple(part.strip() for expression in expressions for part in expression.split(" and "))
if not attributable and terms:
return (), (), False
if any(("or " in term) or ("(" in term) or (term.startswith("not ") and " " in term[4:]) for term in terms):
return (), (), False
return (
tuple(term for term in terms if term and not term.startswith("not ")),
tuple(term[4:].strip() for term in terms if term.startswith("not ")),
True,
)
def _slices() -> tuple[Slice, ...]:
if not CIRCLECI_CONFIG.exists():
return ()
jobs: Final = yaml.safe_load(CIRCLECI_CONFIG.read_text()).get("jobs", {})
return tuple(
Slice(job=job, globs=globs, named=named, required=required, excluded=excluded, understood=understood)
for job, body in jobs.items()
for text in ("\n".join(_strings(body)),)
if "pytest" in text
for globs in (tuple(GLOB_CALL_RE.findall(text)),)
for named in (frozenset(TEST_TOKEN_RE.findall(text)) & frozenset(_test_files()),)
for required, excluded, understood in (
_keyword_terms(tuple(KEYWORD_RE.findall(text)), attributable=len(globs) < 2),
)
if globs or named
)
def _matchable_names(relative_path: str) -> frozenset[str]:
"""Every name a `-k` term can match for this file: its path, plus the names inside it.
pytest matches a keyword against an item's own name and each of its parents', so a
positive term hits a file when it appears in the path or in a class or function name.
"""
try:
with warnings.catch_warnings():
warnings.simplefilter("ignore") # test files carry stray escapes; their names still parse
tree: Final = ast.parse((REPO_ROOT / relative_path).read_text())
except (OSError, SyntaxError):
return frozenset({relative_path})
return frozenset({relative_path}) | frozenset(
node.name
for node in ast.walk(tree)
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef))
)
def _deselected_everywhere(allowlist: Allowlist) -> tuple[Finding, ...]:
slices: Final = _slices()
globbed: Final = tuple(
path
for path in _test_files()
if any(_token_covers(glob, path) for slice_ in slices for glob in slice_.globs)
)
return tuple(
Finding(
subject=path,
detail="globbed by a job, then deselected by every one of their -k expressions",
)
for path in globbed
if not allowlist.covers_test(path)
and not any(slice_.claims(path, _matchable_names(path)) for slice_ in slices)
)
def _holds_tests(directory: pathlib.Path) -> bool:
return any(directory.rglob("test_*.py"))
def _shard_children(root: str, repo_root: pathlib.Path = REPO_ROOT) -> tuple[str, ...]:
"""Children of a sharded root that carry tests, so each one needs its own shard.
A directory earns an entry by containing a test file rather than by being named
`test_*`, which is what keeps fixture directories (`test_configs`, `expected_*`)
out without a hand-maintained list of exceptions.
"""
return tuple(
sorted(
child.relative_to(repo_root).as_posix()
for child in (repo_root / root).iterdir()
if not child.name.startswith(".")
and (
_holds_tests(child)
if child.is_dir()
else child.name.startswith("test_") and child.suffix == ".py"
)
)
)
def _unassigned_shard_children(
tokens: frozenset[str],
roots: tuple[str, ...] = SHARDED_ROOTS,
repo_root: pathlib.Path = REPO_ROOT,
) -> tuple[Finding, ...]:
return tuple(
Finding(subject=child, detail=f"holds tests but no shard of {root} names it")
for root in roots
if (repo_root / root).is_dir()
for child in _shard_children(root, repo_root)
if child not in roots and not any(_token_names(token, child) for token in tokens)
)
def _uncovered_dockerfiles(allowlist: Allowlist, tokens: frozenset[str]) -> tuple[Finding, ...]:
return tuple(
Finding(subject=relative_path, detail="built by no job")
@ -229,7 +423,43 @@ def _report(title: str, findings: tuple[Finding, ...], remedy: str) -> None:
_write("")
def _check_slices() -> int:
findings: Final = _deselected_everywhere(_load_allowlist())
if findings:
_report(
"test files a -k expression removes from every job that globs them",
findings,
"Give each one a job whose -k keeps it, or list it in "
".github/ci-coverage-allowlist.yml with the reason it may stay unrun.",
)
return 1
_write(f"OK: no test file is globbed by a job and then deselected by every -k across {len(_slices())} slices.")
return 0
def _check_shards() -> int:
findings = _unassigned_shard_children(_invoked_test_tokens(_all_scalars()))
if findings:
_report(
"test directories and files that no shard claims",
findings,
"Add each to the shard it belongs to. A directory that is itself split across "
"several shards belongs in SHARDED_ROOTS instead, so its own children get checked.",
)
return 1
counted = sum(len(_shard_children(root)) for root in SHARDED_ROOTS if (REPO_ROOT / root).is_dir())
_write(f"OK: all {counted} test children across {len(SHARDED_ROOTS)} sharded trees are assigned to a shard.")
return 0
def main() -> int:
if "--shards" in sys.argv[1:]:
return _check_shards()
if "--slices" in sys.argv[1:]:
return _check_slices()
allowlist = _load_allowlist()
scalars = _all_scalars()

42
.github/scripts/detect_changes.sh vendored Executable file
View file

@ -0,0 +1,42 @@
#!/usr/bin/env bash
set -uo pipefail
readonly API_FILE_CEILING=3000
readonly CATEGORY="${CATEGORY:-backend}"
decide() {
echo "detect-changes[${CATEGORY}]: decision=$1"
[ -z "${GITHUB_OUTPUT:-}" ] || echo "decision=$1" >>"${GITHUB_OUTPUT}"
exit 0
}
run_full() {
echo "detect-changes[${CATEGORY}]: $1; running job"
decide run
}
here="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
classify="${here}/../../.circleci/scripts/classify_changes.sh"
[ -n "${PR_NUMBER:-}" ] || run_full "not a pull_request event"
[ -n "${REPO:-}" ] || run_full "no repository in the environment"
case "${CHANGED_FILE_COUNT:-}" in
'' | *[!0-9]*) run_full "the event payload carries no changed_files count" ;;
esac
[ "${CHANGED_FILE_COUNT}" -le "${API_FILE_CEILING}" ] ||
run_full "PR #${PR_NUMBER} changes ${CHANGED_FILE_COUNT} files, past the ${API_FILE_CEILING}-file listing ceiling"
changed="$(gh api "repos/${REPO}/pulls/${PR_NUMBER}/files" --paginate --jq '.[].filename')" ||
run_full "could not list the files on PR #${PR_NUMBER}"
[ -n "${changed}" ] || run_full "the API listed no files on PR #${PR_NUMBER}"
echo "detect-changes[${CATEGORY}]: files changed by PR #${PR_NUMBER}:"
printf '%s\n' "${changed}" | sed 's/^/ /'
decision="$(printf '%s\n' "${changed}" | bash "${classify}" "${CATEGORY}")" ||
run_full "classify_changes.sh failed"
case "${decision}" in
run | skip) decide "${decision}" ;;
*) run_full "classify_changes.sh printed an unexpected decision: ${decision}" ;;
esac

15
.github/scripts/select_ui_test_scope.sh vendored Executable file
View file

@ -0,0 +1,15 @@
#!/usr/bin/env bash
set -uo pipefail
has_file=false
has_file_outside_src=false
while IFS= read -r file || [ -n "$file" ]; do
[ -n "$file" ] || continue
has_file=true
case "$file" in
src/*) ;;
*) has_file_outside_src=true ;;
esac
done
{ [ "$has_file" = true ] && [ "$has_file_outside_src" = false ]; } && echo related || echo full

View file

@ -1,557 +0,0 @@
#!/usr/bin/env python3
"""One-shot 7-day heads-up sweep for the Agent Shin rollout.
Posts a friendly "the OSS triage bot kicks in next Monday" comment on every
open external PR/issue that currently *would* fail the new rubric — i.e.,
every PR/issue Agent Shin would close once the rollout completes. The point
is to give contributors a full week to fix their description before the bot
ever takes a destructive action, so nobody is surprised by an auto-close.
The script is designed to run **exactly once** at rollout, fired by a manual
``workflow_dispatch`` (``dry_run=false``) on the heads-up workflow. Re-runs
are safe: every comment is stamped with the hidden ``HEADS_UP_MARKER`` and
PRs/issues that already carry the marker are skipped.
Dry-run vs. real run
--------------------
Defaults to dry-run. Passing ``--close`` flips into real mode. Every GitHub
mutation goes through ``_agent_shin_actions``, which has a one-line
``if dry_run: log else: do_it`` per call, so the only difference between a
dry-run preview and the real run is the call site that actually hits the
GitHub API.
Local preview::
python3 .github/scripts/triage_rollout_heads_up.py --repo BerriAI/litellm
Real run (the manual rollout dispatch uses this)::
python3 .github/scripts/triage_rollout_heads_up.py --repo BerriAI/litellm --close
"""
from __future__ import annotations
import argparse
import datetime as dt
import json
import os
import sys
from pathlib import Path
from typing import Any
# Make the sibling triage_with_llm + _agent_shin_actions importable when this
# script is invoked directly (the GitHub workflow does `python3 .github/scripts/...`).
_SCRIPTS_DIR = Path(__file__).resolve().parent
if str(_SCRIPTS_DIR) not in sys.path:
sys.path.insert(0, str(_SCRIPTS_DIR))
from _agent_shin_actions import maybe_post_comment # noqa: E402
from agent_shin_shared import ( # noqa: E402
AGENT_SHIN_DEFAULT_BOT_LOGIN,
ALLOWLIST_LOGINS,
list_open_items,
)
from triage_with_llm import ( # noqa: E402
DEFAULT_MODEL,
call_llm_judge,
fetch_issue,
fetch_pr,
gh,
is_internal_contributor,
review_gate,
triage,
)
# Hidden marker so re-runs skip PRs/issues we've already notified. Distinct from
# the within-grace / ready / regressed markers so it can't be confused with the
# steady-state lifecycle comments.
HEADS_UP_MARKER = "<!-- agent-shin:rollout-heads-up -->"
# Placeholder until the litellm-docs PR ships. The rollout blog post explains
# the new rubric, the 7-day grace, and how to recover after an auto-close.
# TODO(docs): replace with the canonical URL once the litellm-docs PR merges.
ROLLOUT_BLOG_URL = "https://docs.litellm.ai/docs/agent_shin_triage_rollout"
# Default cutoff is one week from "now". Computed at runtime so the wording
# stays correct even if the rollout is merged later than planned. The user can
# override with --close-on YYYY-MM-DD when running the script manually.
DEFAULT_GRACE_DAYS = 7
# The daily auto-close sweeps (close_low_quality_prs.yml at 09:00 UTC and
# review_gate.yml at 09:30 UTC) are what actually close a still-failing item,
# so the deadline we promise contributors has to name that wall-clock moment.
ACTIVATION_TIME_UTC = "09:00 UTC"
def _format_cutoff(cutoff: dt.date) -> str:
"""Human-readable, timezone-explicit cutoff, e.g. ``Monday, June 1, 2026
(09:00 UTC)`` — the moment a still-failing PR/issue gets closed."""
return (
f"{cutoff.strftime('%A, %B')} {cutoff.day}, {cutoff.year} "
f"({ACTIVATION_TIME_UTC})"
)
def _rubric_section_pr() -> str:
return (
"**Going forward, every external PR needs ONE of:**\n"
"\n"
"- A linked GitHub issue using a closing keyword: "
"`Fixes #1234`, `Closes #1234`, or `Resolves #1234`, OR\n"
"- All three of: a clear **problem description**, **expected vs. "
"actual behavior**, and **end-to-end QA proof** (at least one of a "
"short screen recording / video, before/after screenshots, or the "
"exact commands you ran with their real output; mocked or stubbed "
"runs don't count).\n"
"\n"
"PRs also need a **Greptile confidence score of 4/5 or higher** before "
"the bot will tag them `ready for review`. You can `@greptileai` to "
"request a fresh review at any time, including after the PR is closed."
)
def _rubric_section_issue() -> str:
return (
"**Going forward, every external issue needs:**\n"
"\n"
"- For **bug reports**: end-to-end evidence of the bug (at least one "
"of a screen recording / video, a screenshot, or the exact commands "
"you ran with their real output / traceback) plus expected vs. actual "
"behavior. Written steps with no run output don't count, and mocked "
"or stubbed runs don't count.\n"
"- For **feature requests**: a clear description of the proposed "
"feature plus a use case + concrete example (config, API call, UI "
"flow, or scenario showing what's blocked today)."
)
def _description_only_note(kind: str) -> str:
noun = "PR" if kind == "pr" else "issue"
return (
f"⚠️ **The requirements must live in the {noun} *description*, not in "
"comments.** Some PRs/issues collect 100+ comments from humans and "
"bots; reading the entire thread on every triage run would balloon "
"GitHub API usage (we'd start getting 429'd) and blow out the LLM "
"judge's context. The bot only reads the description, so anything "
"you add as a comment will be invisible to it."
)
def _missing_section(verdict: dict, greptile_score: int | None) -> str:
"""Bullet list of what's currently missing on this PR/issue.
Combines the LLM judge's `missing` list (rubric items) with a Greptile
shortfall (for PRs) so the contributor sees one list of things to fix.
"""
missing = list(verdict.get("missing") or [])
if greptile_score is not None and greptile_score < 4:
missing.insert(
0,
f"Greptile's most recent review scored this PR {greptile_score}/5 "
"(below the 4/5 bar Agent Shin will require).",
)
if not missing:
return (
"_The bot couldn't articulate a specific missing piece; see the "
"rubric link above and double-check the description includes all "
"of it before the rollout._"
)
bullets = "\n".join(f"- {m}" for m in missing)
return f"**What this one is currently missing:**\n\n{bullets}"
def _recovery_section(kind: str) -> str:
if kind == "pr":
return (
"**If the bot closes this PR after the rollout:** update the "
"description with the missing pieces, then either open a fresh "
"PR or comment `@agent-shin reconsider` on the closed PR. If "
"Greptile re-scores you at 4/5 or higher I'll reopen and tag "
"the PR `ready for review`. (`@greptileai` works on closed PRs "
"too; a fresh review is one of the signals that lifts you back "
"into the queue.) This is **not** us losing interest in your "
"change; far from it. We just need open PRs to be a list of "
"things a maintainer can act on, so we can get to yours faster."
)
return (
"**If the bot closes this issue after the rollout:** edit the issue "
"description to add the missing pieces, then comment `@agent-shin "
"reconsider` on the closed issue. I'll re-evaluate and, if the rubric "
"is met, reopen it. (GitHub doesn't let external authors reopen an "
"issue a maintainer or bot closed, so the comment is the reliable "
"path.) This is **not** us saying the bug isn't real or the request "
"isn't useful; it's so the remaining open issues are a list of things "
"a maintainer can act on."
)
def format_heads_up_comment(
*, kind: str, verdict: dict, greptile_score: int | None, cutoff: dt.date
) -> str:
"""Compose the friendly 7-day heads-up comment posted on a failing PR/issue."""
noun = "PR" if kind == "pr" else "issue"
rubric = _rubric_section_pr() if kind == "pr" else _rubric_section_issue()
cutoff_str = _format_cutoff(cutoff)
explanation = (verdict.get("explanation") or "").strip()
explanation_block = (
f"> _(The judge's note for this one: {explanation})_\n\n" if explanation else ""
)
return (
"🚅 **Heads-up: we're turning on the OSS triage bot in "
f"{DEFAULT_GRACE_DAYS} days, on {cutoff_str}.**\n"
"\n"
"We're rolling out **Agent Shin**, an LLM-as-judge triage bot for "
f"external {noun}s. Once it's live, the bot reads each open "
f"{noun}'s description, scores it against a small rubric, and "
f"auto-closes any {noun} that's missing the basics, with a single "
f"comment explaining what's missing and how to recover. Full "
f"context: [Agent Shin rollout blog post]({ROLLOUT_BLOG_URL}).\n"
"\n"
f"{rubric}\n"
"\n"
f"{_description_only_note(kind)}\n"
"\n"
f"{_missing_section(verdict, greptile_score)}\n"
"\n"
f"{explanation_block}"
"**Timeline (you have a week):**\n"
"\n"
f"- We turn the bot on in {DEFAULT_GRACE_DAYS} days, on "
f"**{cutoff_str}**. You have until then to update this {noun}'s "
"description with the missing pieces above.\n"
f"- If this {noun} still fails the rubric at **{cutoff_str}**, "
"we'll close it.\n"
f"- From then on the bot runs daily, and every {noun} that fails "
"the rubric gets a **2-hour lifetime**: one warning comment, then "
"auto-close 2 hours later.\n"
"\n"
f"{_recovery_section(kind)}\n"
"\n"
f"{HEADS_UP_MARKER}"
)
def _list_open_numbers(repo: str, kind: str) -> list[int]:
"""Return every open PR or issue number in ``repo``.
Delegates to ``list_open_items`` so the full backlog is fetched (no cap)
and the `gh {pr,issue} list` invocation stays in one shared place. ``gh
issue list`` would include PRs, but ``list_open_items`` uses the dedicated
command per kind, so the two never mix.
"""
return [
item["number"] for item in list_open_items(kind, repo=repo, fields="number")
]
def _has_heads_up_marker(item: dict) -> bool:
"""Cheap fast-path: check the PR/issue body itself for the marker.
The marker is appended to the *comment* we post, not the body, so this
will only fire if the body literally contains the marker text. We still
do the comment-marker check separately below; this body check just lets
us short-circuit for PRs/issues that quote the marker for any reason.
"""
body = item.get("body") or ""
return HEADS_UP_MARKER in body
def _comments_have_marker(repo: str, number: int) -> bool:
"""True if the bot already posted a comment carrying the marker.
Used for idempotency: a re-run skips items the previous run notified.
Filters by author (matching the sibling marker-checks in
``triage_with_llm._has_marker`` and
``agent_shin_shared.seconds_since_latest_marker_comment``) so a
contributor who quotes the heads-up via GitHub's "Quote reply" — which
preserves HTML comments in the raw markdown — can't trick the
idempotency check into silently skipping a real heads-up.
Comments live on the unified issues endpoint regardless of whether the
item is a PR or an issue, so no ``kind`` argument is required here.
"""
expected_login = (
os.environ.get("AGENT_SHIN_BOT_LOGIN") or AGENT_SHIN_DEFAULT_BOT_LOGIN
).lower()
raw = gh(
"api",
"--paginate",
f"repos/{repo}/issues/{number}/comments?per_page=100",
)
for line in raw.splitlines():
line = line.strip()
if not line:
continue
try:
payload = json.loads(line)
except json.JSONDecodeError:
continue
comments = payload if isinstance(payload, list) else [payload]
for comment in comments:
author = ((comment.get("user") or {}).get("login") or "").lower()
if author != expected_login:
continue
if HEADS_UP_MARKER in (comment.get("body") or ""):
return True
return False
def _evaluate_pr(*, repo: str, number: int, model: str, judge: Any = None) -> dict:
"""Run the future PR rubric (review_gate) in dry-run and return the result."""
return review_gate(
repo=repo,
number=number,
close=False, # we only want the verdict, never act here
model=model,
judge=judge,
)
def _evaluate_issue(*, repo: str, number: int, model: str, judge: Any = None) -> dict:
"""Run the future issue rubric (triage kind='issue') in dry-run."""
return triage(
repo=repo,
kind="issue",
number=number,
close=False,
model=model,
judge=judge,
)
def _would_be_closed(kind: str, result: dict) -> bool:
"""True if the future triage would auto-close this PR/issue based on the
rubric (regardless of grace-period gating).
For PRs we trust ``review_gate``'s ``passing`` field — it combines the LLM
verdict and the Greptile score. For issues we read the LLM verdict
directly. Both fields are ``None``/missing on skip paths
(skip-internal-author, skip-llm-error, etc.) where the future bot would
NOT close the item — those return False.
"""
if kind == "pr":
passing = result.get("passing")
if passing is None:
return False # skipped — nothing for the heads-up to warn about
return passing is False
verdict = result.get("verdict") or {}
return (verdict.get("verdict") or "").lower() == "fail"
def _process_one(
*,
repo: str,
kind: str,
number: int,
model: str,
cutoff: dt.date,
dry_run: bool,
judge: Any = None,
skip_marker_check: bool = False,
allowlist: frozenset[str] = ALLOWLIST_LOGINS,
) -> dict:
"""Evaluate one PR/issue and post a heads-up if it would be auto-closed.
Returns a per-item dict for the summary table.
"""
base = {"kind": kind, "number": number}
fetcher = fetch_pr if kind == "pr" else fetch_issue
item = fetcher(repo, number)
if (item.get("state") or "") != "open":
return {**base, "action": "skip-not-open"}
if allowlist:
login = (item.get("user") or {}).get("login") or ""
if login.lower() not in allowlist:
return {**base, "action": "skip-not-allowlisted"}
elif is_internal_contributor(item):
return {**base, "action": "skip-internal-author"}
if not skip_marker_check and _has_heads_up_marker(item):
return {**base, "action": "skip-already-marked-in-body"}
if not skip_marker_check and _comments_have_marker(repo, number):
return {**base, "action": "skip-already-notified"}
if kind == "pr":
result = _evaluate_pr(repo=repo, number=number, model=model, judge=judge)
else:
result = _evaluate_issue(repo=repo, number=number, model=model, judge=judge)
if not _would_be_closed(kind, result):
return {**base, "action": "skip-passing", "evaluator": result.get("action")}
verdict = result.get("verdict") or {}
greptile_score = result.get("greptile_score") if kind == "pr" else None
comment = format_heads_up_comment(
kind=kind, verdict=verdict, greptile_score=greptile_score, cutoff=cutoff
)
maybe_post_comment(repo, number, comment, dry_run=dry_run)
return {
**base,
"action": "heads-up-posted" if not dry_run else "would-post-heads-up",
"verdict": (verdict.get("verdict") or "").lower(),
"greptile_score": greptile_score,
}
def _print_summary(results: list[dict]) -> None:
"""Tally per-action counts so a dry-run preview tells you at a glance how
many comments the real run would post."""
counts: dict[str, int] = {}
for r in results:
counts[r["action"]] = counts.get(r["action"], 0) + 1
print("\n=== rollout heads-up summary ===")
for action in sorted(counts):
print(f" {action:35s} {counts[action]}")
print(f" total {len(results)}")
def run(
*,
repo: str,
close: bool,
cutoff: dt.date,
model: str,
kinds: tuple[str, ...] = ("pr", "issue"),
judge: Any = None,
only_numbers: dict[str, list[int]] | None = None,
skip_marker_check: bool = False,
) -> list[dict]:
"""Sweep ``repo`` and post heads-up comments. Returns the per-item results."""
dry_run = not close
if dry_run:
print(
f"[DRY RUN] sweeping {repo}; --close not passed, no comments will be posted."
)
else:
print(f"[REAL RUN] sweeping {repo}; comments WILL be posted.")
print(f"Cutoff date in comment body: {cutoff.isoformat()}")
results: list[dict] = []
for kind in kinds:
if only_numbers and kind in only_numbers:
numbers = list(only_numbers[kind])
else:
numbers = _list_open_numbers(repo, kind)
print(f"\n--- {kind}s: {len(numbers)} open ---")
for n in numbers:
try:
result = _process_one(
repo=repo,
kind=kind,
number=n,
model=model,
cutoff=cutoff,
dry_run=dry_run,
judge=judge,
skip_marker_check=skip_marker_check,
)
except (
Exception
) as exc: # noqa: BLE001 - per-item errors don't abort the sweep
result = {
"kind": kind,
"number": n,
"action": "error",
"error": str(exc),
}
print(f"!! {kind}#{n}: {exc}", file=sys.stderr)
print(f" {kind}#{n}: {result['action']}")
results.append(result)
_print_summary(results)
return results
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--repo", required=True, help="owner/repo")
parser.add_argument(
"--close",
action="store_true",
help=(
"Actually post comments. Without this flag the script is in "
"dry-run mode and only logs what it would do."
),
)
parser.add_argument(
"--close-on",
type=dt.date.fromisoformat,
default=None,
help=(
"Cutoff date shown in the heads-up comment as the rollout date "
f"(default: today + {DEFAULT_GRACE_DAYS} days)."
),
)
parser.add_argument(
"--model",
default=os.environ.get("TRIAGE_MODEL") or DEFAULT_MODEL,
help=f"Model for the rubric LLM judge (default: {DEFAULT_MODEL}).",
)
parser.add_argument(
"--kind",
choices=("pr", "issue", "both"),
default="both",
help="Restrict the sweep to PRs or issues only (default: both).",
)
parser.add_argument(
"--only-pr",
type=int,
action="append",
default=[],
help="Limit the PR sweep to these PR numbers (repeat for several).",
)
parser.add_argument(
"--only-issue",
type=int,
action="append",
default=[],
help="Limit the issue sweep to these issue numbers (repeat for several).",
)
parser.add_argument(
"--ignore-existing-marker",
action="store_true",
help=(
"Re-post on PRs/issues that already carry the heads-up marker. "
"Useful for testing the comment wording on a known PR."
),
)
args = parser.parse_args()
cutoff = args.close_on or (
dt.datetime.now(dt.timezone.utc).date() + dt.timedelta(days=DEFAULT_GRACE_DAYS)
)
kinds: tuple[str, ...]
if args.kind == "pr":
kinds = ("pr",)
elif args.kind == "issue":
kinds = ("issue",)
else:
kinds = ("pr", "issue")
only: dict[str, list[int]] = {}
if args.only_pr:
only["pr"] = args.only_pr
if args.only_issue:
only["issue"] = args.only_issue
# The script must NOT hit the LLM in dry-run if no key is set — we still
# want a useful preview that says "skip-no-llm-key" for items that would
# have been judged. Production runs require OPENAI_API_KEY.
if args.close and not os.environ.get("OPENAI_API_KEY"):
parser.error("OPENAI_API_KEY must be set for --close (real-run) mode.")
run(
repo=args.repo,
close=args.close,
cutoff=cutoff,
model=args.model,
kinds=kinds,
only_numbers=only or None,
skip_marker_check=args.ignore_existing_marker,
)
return 0
if __name__ == "__main__":
sys.exit(main())

View file

@ -60,6 +60,9 @@ jobs:
name: Run tests
runs-on: ubuntu-latest
timeout-minutes: ${{ inputs.job-timeout-minutes }}
permissions:
contents: read
pull-requests: read
outputs:
decision: ${{ steps.changes.outputs.decision }}
@ -69,24 +72,27 @@ jobs:
with:
persist-credentials: false
- name: Detect backend-relevant changes
- name: Detect relevant changes
id: changes
timeout-minutes: 2
uses: ./.github/actions/detect-backend-changes
uses: ./.github/actions/detect-changes
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 5
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:

View file

@ -23,7 +23,7 @@ jobs:
version: "0.10.9"
- name: Update JSON Data
run: |
uv run --frozen --with 'aiohttp==3.13.3' python ".github/workflows/auto_update_price_and_context_window_file.py"
uv run --frozen --with 'aiohttp==3.13.3' python ".github/scripts/auto_update_price_and_context_window_file.py"
- name: Regenerate JSON Schema
run: |
uv run --frozen python ci_cd/generate_model_prices_schema.py

View file

@ -40,3 +40,9 @@ jobs:
run: |
python -m pip install "pyyaml==6.0.3"
python .github/scripts/assert_ci_coverage.py
# The census asks whether a job names a file; this asks whether that job's -k
# then throws it back out. A file both globbed and deselected everywhere runs
# nowhere while counting as covered, which is how the caching suite went unrun.
- name: Assert no -k expression deselects a file from every job that globs it
run: python .github/scripts/assert_ci_coverage.py --slices

View file

@ -24,6 +24,7 @@ jobs:
# re-running basedpyright over the merge-base tree.
permissions:
contents: read
pull-requests: read
actions: read
steps:
@ -37,7 +38,12 @@ jobs:
clean: true
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
- name: Fetch gate base (merge-base with target branch)
if: steps.changes.outputs.decision != 'skip'
env:
GH_TOKEN: ${{ github.token }}
BASE_SHA: ${{ github.event.pull_request.base.sha }}
@ -50,39 +56,47 @@ jobs:
echo "GATE_BASE_SHA=$MERGE_BASE" >> "$GITHUB_ENV"
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Clean Python cache
if: steps.changes.outputs.decision != 'skip'
run: |
find . -type d -name "__pycache__" -exec rm -rf {} + || true
find . -name "*.pyc" -delete || true
- name: Check uv.lock is up to date
if: steps.changes.outputs.decision != 'skip'
run: |
uv lock --check || (echo "❌ uv.lock is out of sync with pyproject.toml. Run 'uv lock' locally and commit the result." && exit 1)
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: |
uv sync --frozen --group proxy-dev --group e2e-dev
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/cache-prisma-binaries
# basedpyright resolves Prisma's generated client (litellm/proxy/schema.prisma)
# only after `prisma generate` writes prisma/client.py et al. Without this the
# DB wrappers typed against the generated client would degrade to Unknown.
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Check ruff format
if: steps.changes.outputs.decision != 'skip'
run: |
git diff --name-only --diff-filter=ACMR "$GATE_BASE_SHA" HEAD -- 'litellm/**/*.py' | grep -v '^litellm/enterprise/' > "$RUNNER_TEMP/ruff_format_files.txt" || true
if [ ! -s "$RUNNER_TEMP/ruff_format_files.txt" ]; then
@ -92,6 +106,7 @@ jobs:
xargs uv run --no-sync ruff format --check --exclude '/enterprise/' < "$RUNNER_TEMP/ruff_format_files.txt"
- name: Debug - Check file state
if: steps.changes.outputs.decision != 'skip'
run: |
echo "Current branch:"
git branch --show-current
@ -101,30 +116,41 @@ jobs:
head -50 litellm/litellm_core_utils/custom_logger_registry.py | tail -10
- name: Run Ruff linting
if: steps.changes.outputs.decision != 'skip'
run: |
cd litellm
uv run --no-sync ruff check .
cd ..
- name: Check strict-rule budget (delta vs base)
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python scripts/ruff_strict_gate.py --base "$GATE_BASE_SHA"
- name: Check type-discipline budget (mutable collections / casts / type guards / kwargs / unexplained suppressions, delta vs base)
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python scripts/type_discipline_gate.py --base "$GATE_BASE_SHA"
- name: Check test-quality budget (zero-assert / mock-echo tests, sys.path.insert, raw env writes, litellm global mutation, credential-gated skips, delta vs base)
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python scripts/test_quality_gate.py --base "$GATE_BASE_SHA"
- name: Print OpenAI version
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python -c "import openai; print(f'OpenAI version: {openai.__version__}')"
- name: Check basedpyright budget (delta vs base)
if: steps.changes.outputs.decision != 'skip'
env:
GH_TOKEN: ${{ github.token }}
run: |
uv run --no-sync python scripts/type_check_gate.py --base "$GATE_BASE_SHA"
- name: Check tests/e2e basedpyright (zero errors)
if: steps.changes.outputs.decision != 'skip'
run: |
if git diff --name-only --diff-filter=ACMRD "$GATE_BASE_SHA" HEAD -- 'tests/e2e/**/*.py' | grep -q .; then
uv run --no-sync basedpyright tests/e2e
@ -133,12 +159,14 @@ jobs:
fi
- name: Check for circular imports
if: steps.changes.outputs.decision != 'skip'
run: |
cd litellm
uv run --no-sync python ../tests/documentation_tests/test_circular_imports.py
cd ..
- name: Check import safety
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)
@ -200,7 +228,7 @@ jobs:
- name: Run secret scan test
run: |
uv run --no-project --with 'pytest==9.0.2' pytest tests/litellm/test_no_hardcoded_secrets.py -v
uv run --no-project --with 'pytest==9.0.2' pytest tests/code_coverage_tests/test_no_hardcoded_secrets.py -v
- name: Run ggshield secret scan
env:

View file

@ -1,6 +1,7 @@
name: UI Build Check
permissions:
contents: read
pull-requests: read
on:
pull_request:
@ -28,7 +29,14 @@ jobs:
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
with:
category: ui
- name: Setup Node.js
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
@ -36,7 +44,9 @@ jobs:
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: npm ci
- name: Build
if: steps.changes.outputs.decision != 'skip'
run: npm run build

View file

@ -1,6 +1,7 @@
name: UI Unit Tests
permissions:
contents: read
pull-requests: read
on:
pull_request:
@ -32,7 +33,14 @@ jobs:
fetch-depth: 1
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
with:
category: ui
- name: Setup Node.js
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
@ -40,36 +48,50 @@ jobs:
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: npm ci
- name: Run UI type tests (Vitest)
if: steps.changes.outputs.decision != 'skip'
env:
CI: "true"
run: npm run test:types
- name: Run UI unit tests (Vitest)
if: steps.changes.outputs.decision != 'skip'
env:
CI: "true"
GH_TOKEN: ${{ github.token }}
BASE_SHA: ${{ github.event.pull_request.base.sha }}
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
run: |
if [ -n "$BASE_SHA" ]; then
merge_base=$(gh api "repos/${{ github.repository }}/compare/${BASE_SHA}...${HEAD_SHA}?per_page=1" --jq '.merge_base_commit.sha')
test -n "$merge_base"
git fetch --no-tags --depth=1 origin "$merge_base" "$HEAD_SHA"
changed_files=()
while IFS= read -r f; do
changed_files+=("$f")
done < <(git diff --name-only --relative "$merge_base" "$HEAD_SHA" -- .)
if [ ${#changed_files[@]} -eq 0 ]; then
echo "No UI files changed in this PR; skipping unit tests."
exit 0
fi
echo "Pull request: running tests related to ${#changed_files[@]} changed UI files"
npm run test -- related "${changed_files[@]}" --run --passWithNoTests \
--pool forks --poolOptions.forks.maxForks=14
else
full_suite() { npm run test -- --run --pool forks --poolOptions.forks.maxForks=14; }
if [ -z "$BASE_SHA" ]; then
echo "Push to $GITHUB_REF_NAME: running the full suite"
npm run test -- --run --pool forks --poolOptions.forks.maxForks=14
full_suite
exit 0
fi
merge_base=$(gh api "repos/${{ github.repository }}/compare/${BASE_SHA}...${HEAD_SHA}?per_page=1" --jq '.merge_base_commit.sha')
test -n "$merge_base"
git fetch --no-tags --depth=1 origin "$merge_base" "$HEAD_SHA"
changed_files=()
while IFS= read -r f; do
changed_files+=("$f")
done < <(git diff --name-only --relative "$merge_base" "$HEAD_SHA" -- .)
if [ ${#changed_files[@]} -eq 0 ]; then
echo "No UI files changed in this PR; skipping unit tests."
exit 0
fi
scope=$(printf '%s\n' "${changed_files[@]}" | bash "$GITHUB_WORKSPACE/.github/scripts/select_ui_test_scope.sh")
if [ "$scope" != related ]; then
echo "Pull request: ${#changed_files[@]} changed UI files reach outside src/, so related would miss their dependents; running the full suite"
full_suite
exit 0
fi
echo "Pull request: running tests related to ${#changed_files[@]} changed UI files"
npm run test -- related "${changed_files[@]}" --run --passWithNoTests \
--pool forks --poolOptions.forks.maxForks=14

View file

@ -10,6 +10,7 @@ on:
permissions:
contents: read
pull-requests: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
@ -25,26 +26,34 @@ jobs:
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
- name: Thank You Message
run: |
echo "### 🙏 Thank you for contributing to LiteLLM!" >> $GITHUB_STEP_SUMMARY
echo "Your PR is being tested now. We appreciate your help in making LiteLLM better!" >> $GITHUB_STEP_SUMMARY
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: |
uv lock --check
.github/scripts/uv_sync_with_retries.sh --frozen --group proxy-dev --extra proxy --extra semantic-router
- name: Run MCP tests
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync pytest tests/mcp_tests -x -vv -n 4 --cov=./litellm --cov-report=xml --durations=5

View file

@ -1,31 +0,0 @@
name: "Unit Tests: Core Utilities"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
core-utils:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: "tests/test_litellm/litellm_core_utils"
workers: 2
reruns: 1
artifact-name: core-utils

View file

@ -23,34 +23,41 @@ jobs:
documentation:
runs-on: ubuntu-latest
timeout-minutes: 10
permissions:
contents: read
pull-requests: read
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
- name: Checkout litellm-docs into docs/my-website (for documentation_tests)
if: steps.changes.outputs.decision != 'skip'
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
repository: BerriAI/litellm-docs
path: docs/my-website
persist-credentials: false
- name: Detect backend-relevant changes
id: changes
uses: ./.github/actions/detect-backend-changes
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip'
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |

View file

@ -1,35 +0,0 @@
name: "Unit Tests: Enterprise, Google GenAI & Routing"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
enterprise-routing:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: >-
tests/test_litellm/enterprise
tests/test_litellm/google_genai
tests/test_litellm/router_utils
tests/test_litellm/router_strategy
workers: 2
reruns: 2
artifact-name: enterprise-routing

View file

@ -1,31 +0,0 @@
name: "Unit Tests: Integrations (Callbacks & Logging)"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
integrations:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: "tests/test_litellm/integrations"
workers: 2
reruns: 3
artifact-name: integrations

View file

@ -1,47 +0,0 @@
name: "Unit Tests: LLM Provider Transformations"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
vertex-ai:
name: Vertex AI
permissions:
contents: read
id-token: write
pull-requests: write
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: "tests/test_litellm/llms/vertex_ai"
workers: 1
reruns: 2
artifact-name: llm-vertex-ai
other-providers:
name: All Other Providers
permissions:
contents: read
id-token: write
pull-requests: write
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: "tests/test_litellm/llms --ignore=tests/test_litellm/llms/vertex_ai"
workers: 2
reruns: 2
artifact-name: llm-other-providers

View file

@ -1,53 +0,0 @@
name: "Unit Tests: MCP, Secrets, Containers & Misc"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
misc:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: >-
tests/test_litellm/batches
tests/test_litellm/secret_managers
tests/test_litellm/a2a_protocol
tests/test_litellm/anthropic_interface
tests/test_litellm/completion_extras
tests/test_litellm/compression
tests/test_litellm/containers
tests/test_litellm/experimental_mcp_client
tests/test_litellm/models
tests/test_litellm/repositories
tests/test_litellm/images
tests/test_litellm/interactions
tests/test_litellm/ocr
tests/test_litellm/passthrough
tests/test_litellm/rag
tests/test_litellm/realtime_api
tests/test_litellm/rerank_api
tests/test_litellm/sandbox
tests/test_litellm/test_router
tests/test_litellm/vector_stores
tests/test_litellm/videos
tests/test_litellm/test_*.py
workers: 2
reruns: 2
artifact-name: misc

View file

@ -1,31 +0,0 @@
name: "Unit Tests: Proxy Auth & Key Management"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
proxy-auth:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: "tests/test_litellm/proxy/auth tests/test_litellm/proxy/hooks tests/test_litellm/proxy/policy_engine tests/test_litellm/proxy/client"
workers: 2
reruns: 2
artifact-name: proxy-auth

View file

@ -42,11 +42,10 @@ concurrency:
# pinning the whole file to one worker (the default --dist=loadscope
# behavior for single-file targets).
jobs:
# Fast guard — fails the workflow if a test_*.py file under
# tests/proxy_unit_tests/ is not referenced by any matrix entry below.
# The semantic-shard design (no catch-all "remaining" bucket) relies on
# every test file being explicitly assigned; this guard prevents a new
# file from silently dropping out of CI.
# Fast guard — fails the workflow when a test directory or file inside a sharded
# tree is claimed by no shard. The semantic-shard design has no catch-all bucket,
# so an unassigned child runs nowhere; assert_ci_coverage.py holds the tree list
# and reads the same test-path keys the coverage census does.
assert-shard-coverage:
runs-on: ubuntu-latest
timeout-minutes: 2
@ -56,31 +55,8 @@ jobs:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Assert every test_*.py is in a matrix shard
run: |
python3 - <<'PY'
import pathlib, sys, yaml
wf = yaml.safe_load(open(".github/workflows/test-unit-proxy-db.yml"))
matrix = wf["jobs"]["proxy-db"]["strategy"]["matrix"]["include"]
referenced = set()
for entry in matrix:
for token in entry["test-path"].split():
if token.startswith("tests/proxy_unit_tests/"):
referenced.add(pathlib.PurePosixPath(token).name)
actual = {p.name for p in pathlib.Path("tests/proxy_unit_tests").iterdir()
if p.name.startswith("test_") and (p.suffix == ".py" or p.is_dir())
and p.name != "test_configs"}
orphans = sorted(actual - referenced)
if orphans:
print("ERROR: the following files/dirs under tests/proxy_unit_tests/")
print(" are not assigned to any shard in test-unit-proxy-db.yml:")
for o in orphans:
print(f" - {o}")
print()
print("Add each to whichever semantic shard it belongs to.")
sys.exit(1)
print(f"OK: all {len(actual)} files assigned to a shard.")
PY
- name: Assert every test directory and file is claimed by a shard
run: python3 .github/scripts/assert_ci_coverage.py --shards
proxy-db:
needs: assert-shard-coverage

View file

@ -1,80 +0,0 @@
name: "Unit Tests: Proxy API Endpoints"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
workflow_dispatch:
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
proxy-endpoints:
permissions:
contents: read
id-token: write
pull-requests: write
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: >-
tests/test_litellm/proxy/analytics_endpoints
tests/test_litellm/proxy/management_endpoints
tests/test_litellm/proxy/memory
tests/test_litellm/proxy/guardrails
tests/test_litellm/proxy/management_helpers
tests/test_litellm/proxy/anthropic_endpoints
tests/test_litellm/proxy/google_endpoints
tests/test_litellm/proxy/openai_files_endpoint
tests/test_litellm/proxy/batches_endpoints
tests/test_litellm/proxy/fine_tuning_endpoints
tests/test_litellm/proxy/vector_store_files_endpoints
tests/test_litellm/proxy/video_endpoints
tests/test_litellm/proxy/response_api_endpoints
tests/test_litellm/proxy/image_endpoints
tests/test_litellm/proxy/vector_store_endpoints
tests/test_litellm/proxy/agent_endpoints
tests/test_litellm/proxy/a2a
tests/test_litellm/proxy/credential_endpoints
tests/test_litellm/proxy/discovery_endpoints
tests/test_litellm/proxy/health_endpoints
tests/test_litellm/proxy/shutdown
tests/test_litellm/proxy/public_endpoints
tests/test_litellm/proxy/prompts
tests/test_litellm/proxy/rag_endpoints
tests/test_litellm/proxy/realtime_endpoints
tests/test_litellm/proxy/ui_crud_endpoints
tests/test_litellm/proxy/config_resolvers
tests/test_litellm/proxy/utils
workers: 2
reruns: 2
artifact-name: proxy-endpoints
# Behavior-pinning tests for litellm/proxy/proxy_server.py. Owns its
# own job (not a path on the proxy-endpoints job above) so its budget
# is independent and its coverage artifact is uploaded separately.
# See: https://www.notion.so/36c43b8acdab81ee845fd5365128a2fc
proxy-server:
permissions:
contents: read
id-token: write
pull-requests: write
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: tests/test_litellm/proxy/proxy_server
workers: 4
reruns: 2
timeout-minutes: 60
job-timeout-minutes: 95
artifact-name: proxy-server

View file

@ -1,42 +0,0 @@
name: "Unit Tests: Proxy Infrastructure"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
proxy-infra:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: >-
tests/test_litellm/proxy/db
tests/test_litellm/proxy/middleware
tests/test_litellm/proxy/spend_tracking
tests/test_litellm/proxy/pass_through_endpoints
tests/test_litellm/proxy/_experimental
tests/test_litellm/proxy/experimental
tests/test_litellm/proxy/common_utils
tests/test_litellm/proxy/enterprise_billing
tests/test_litellm/proxy/types_utils
tests/test_litellm/proxy/logging_endpoints
tests/test_litellm/proxy/test_*.py
workers: 2
reruns: 2
artifact-name: proxy-infra

View file

@ -1,31 +0,0 @@
name: "Unit Tests: Responses, Caching & Types"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
responses-caching-types:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: "tests/test_litellm/responses tests/test_litellm/caching tests/test_litellm/types"
workers: 2
reruns: 2
artifact-name: responses-caching-types

219
.github/workflows/test-unit.yml vendored Normal file
View file

@ -0,0 +1,219 @@
name: "Unit Tests"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
push:
branches:
- main
- litellm_internal_staging
workflow_dispatch:
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
# One caller for every tests/test_litellm shard, replacing the nine thin workflow
# files that each wrapped a single call to _test-unit-base.yml. Adding a shard is
# now one matrix entry rather than a new file.
#
# `name` is the shard id and nothing else, so each check reports as
# "<shard> / Run tests" exactly as it did when the shard had its own file. Those
# strings are the branch ruleset's required contexts, so they are load-bearing:
# renaming an entry renames a required check and the ruleset stops matching it.
#
# Every entry states its timeouts even when they equal the base workflow's
# defaults. An absent matrix key renders as an empty string, which is not a
# number, so a partially-specified entry would fail the call rather than fall
# back to the default.
#
# tests/proxy_unit_tests keeps its own caller (test-unit-proxy-db.yml): it is
# already a matrix and carries a shard-coverage guard that reads that file by
# name. Folding it in here is a follow-up, together with generalising that guard
# into assert_ci_coverage.py.
jobs:
unit:
name: ${{ matrix.shard }}
permissions:
contents: read
id-token: write
pull-requests: write
strategy:
fail-fast: false
matrix:
include:
- shard: core-utils
artifact-name: core-utils
test-path: "tests/test_litellm/litellm_core_utils"
workers: 2
reruns: 1
timeout-minutes: 20
job-timeout-minutes: 55
- shard: enterprise-routing
artifact-name: enterprise-routing
test-path: >-
tests/test_litellm/enterprise
tests/test_litellm/google_genai
tests/test_litellm/router_utils
tests/test_litellm/router_strategy
workers: 2
reruns: 2
timeout-minutes: 20
job-timeout-minutes: 55
- shard: integrations
artifact-name: integrations
test-path: "tests/test_litellm/integrations"
workers: 2
reruns: 3
timeout-minutes: 20
job-timeout-minutes: 55
- shard: Vertex AI
artifact-name: llm-vertex-ai
test-path: "tests/test_litellm/llms/vertex_ai"
workers: 1
reruns: 2
timeout-minutes: 20
job-timeout-minutes: 55
- shard: All Other Providers
artifact-name: llm-other-providers
test-path: "tests/test_litellm/llms --ignore=tests/test_litellm/llms/vertex_ai"
workers: 2
reruns: 2
timeout-minutes: 20
job-timeout-minutes: 55
- shard: misc
artifact-name: misc
test-path: >-
tests/test_litellm/batches
tests/test_litellm/secret_managers
tests/test_litellm/a2a_protocol
tests/test_litellm/anthropic_interface
tests/test_litellm/completion_extras
tests/test_litellm/compression
tests/test_litellm/containers
tests/test_litellm/experimental_mcp_client
tests/test_litellm/models
tests/test_litellm/repositories
tests/test_litellm/images
tests/test_litellm/interactions
tests/test_litellm/ocr
tests/test_litellm/passthrough
tests/test_litellm/rag
tests/test_litellm/realtime_api
tests/test_litellm/rerank_api
tests/test_litellm/sandbox
tests/test_litellm/test_router
tests/test_litellm/vector_stores
tests/test_litellm/videos
tests/test_litellm/test_*.py
workers: 2
reruns: 2
timeout-minutes: 20
job-timeout-minutes: 55
- shard: proxy-auth
artifact-name: proxy-auth
test-path: >-
tests/test_litellm/proxy/auth
tests/test_litellm/proxy/hooks
tests/test_litellm/proxy/policy_engine
tests/test_litellm/proxy/client
workers: 2
reruns: 2
timeout-minutes: 20
job-timeout-minutes: 55
- shard: proxy-endpoints
artifact-name: proxy-endpoints
test-path: >-
tests/test_litellm/proxy/analytics_endpoints
tests/test_litellm/proxy/management_endpoints
tests/test_litellm/proxy/memory
tests/test_litellm/proxy/guardrails
tests/test_litellm/proxy/management_helpers
tests/test_litellm/proxy/anthropic_endpoints
tests/test_litellm/proxy/google_endpoints
tests/test_litellm/proxy/openai_files_endpoint
tests/test_litellm/proxy/batches_endpoints
tests/test_litellm/proxy/fine_tuning_endpoints
tests/test_litellm/proxy/vector_store_files_endpoints
tests/test_litellm/proxy/video_endpoints
tests/test_litellm/proxy/response_api_endpoints
tests/test_litellm/proxy/image_endpoints
tests/test_litellm/proxy/ocr_endpoints
tests/test_litellm/proxy/vector_store_endpoints
tests/test_litellm/proxy/agent_endpoints
tests/test_litellm/proxy/a2a
tests/test_litellm/proxy/credential_endpoints
tests/test_litellm/proxy/discovery_endpoints
tests/test_litellm/proxy/health_endpoints
tests/test_litellm/proxy/shutdown
tests/test_litellm/proxy/public_endpoints
tests/test_litellm/proxy/prompts
tests/test_litellm/proxy/rag_endpoints
tests/test_litellm/proxy/realtime_endpoints
tests/test_litellm/proxy/ui_crud_endpoints
tests/test_litellm/proxy/config_resolvers
tests/test_litellm/proxy/utils
workers: 2
reruns: 2
timeout-minutes: 20
job-timeout-minutes: 55
- shard: proxy-server
artifact-name: proxy-server
test-path: "tests/test_litellm/proxy/proxy_server"
workers: 4
reruns: 2
timeout-minutes: 60
job-timeout-minutes: 95
- shard: proxy-infra
artifact-name: proxy-infra
test-path: >-
tests/test_litellm/proxy/db
tests/test_litellm/proxy/middleware
tests/test_litellm/proxy/spend_tracking
tests/test_litellm/proxy/pass_through_endpoints
tests/test_litellm/proxy/_experimental
tests/test_litellm/proxy/experimental
tests/test_litellm/proxy/common_utils
tests/test_litellm/proxy/enterprise_billing
tests/test_litellm/proxy/types_utils
tests/test_litellm/proxy/logging_endpoints
tests/test_litellm/proxy/test_*.py
workers: 2
reruns: 2
timeout-minutes: 20
job-timeout-minutes: 55
- shard: responses-caching-types
artifact-name: responses-caching-types
test-path: >-
tests/test_litellm/responses
tests/test_litellm/caching
tests/test_litellm/types
workers: 2
reruns: 2
timeout-minutes: 20
job-timeout-minutes: 55
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
workers: ${{ matrix.workers }}
reruns: ${{ matrix.reruns }}
timeout-minutes: ${{ matrix.timeout-minutes }}
job-timeout-minutes: ${{ matrix.job-timeout-minutes }}
artifact-name: ${{ matrix.artifact-name }}

View file

@ -1,92 +0,0 @@
name: Agent Shin — rollout heads-up (one-shot)
# Fires the 7-day heads-up comment on every open external PR/issue that the
# new triage bot would auto-close. The real sweep is a deliberate one-shot:
# trigger it at rollout via a manual `workflow_dispatch` with `dry_run=false`.
# The script is idempotent (skips items that already carry the
# `<!-- agent-shin:rollout-heads-up -->` marker), so a re-run is harmless.
#
# The automatic push trigger runs DRY-RUN only, so merging the script to
# `litellm_internal_staging` never posts a comment; it just confirms the
# workflow is wired up. Posting real comments requires the manual dispatch,
# which is also the only trigger that exposes `OPENAI_API_KEY`. The heads-up
# is intentionally NOT gated on `AGENT_SHIN_ENABLED`: it has to warn
# contributors while that flag is still off, ahead of the flip that turns on
# auto-closing.
#
# The workflow is a thin shell over `.github/scripts/triage_rollout_heads_up.py`.
# Dry-run vs. real run differ in EXACTLY one CLI flag (`--close`), added only
# on a manual dispatch with `dry_run=false`.
on:
push:
branches:
- litellm_internal_staging
paths:
# The presence of this script on staging IS the rollout merge marker.
# Editing the file later would re-fire the workflow; that's safe because
# the script skips PRs/issues that already have the heads-up marker.
- ".github/scripts/triage_rollout_heads_up.py"
workflow_dispatch:
inputs:
dry_run:
description: "Dry run (true = preview only, false = actually post comments)."
required: false
default: "true"
type: choice
options:
- "true"
- "false"
permissions:
contents: read
issues: write
pull-requests: write
jobs:
heads-up:
if: github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
steps:
- name: Checkout triage scripts
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
sparse-checkout: .github/scripts
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Install LLM client
run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt
- name: Run heads-up sweep
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
# Only the manual dispatch (the real-run trigger) needs the LLM key.
# The automatic push trigger runs dry-run and never posts, so it gets
# no key. Mirrors the sibling triage workflows, which expose the key
# only on an enabled/dispatched run rather than unconditionally.
OPENAI_API_KEY: ${{ github.event_name == 'workflow_dispatch' && secrets.OPENAI_API_KEY || '' }}
OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }}
TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }}
# The real run is a deliberate manual dispatch with dry_run=false.
# Use the EXACT "false" comparison so any unexpected input value
# fail-closes to dry-run (mirrors the AGENT_SHIN_ENABLED pattern in
# the sibling workflows). The automatic push trigger always stays
# dry-run, so merging the script never posts.
DRY_RUN_INPUT: ${{ github.event.inputs.dry_run }}
run: |
set -euo pipefail
ARGS=(--repo "${{ github.repository }}")
if [ "${GITHUB_EVENT_NAME:-}" = "workflow_dispatch" ] && [ "${DRY_RUN_INPUT:-true}" = "false" ]; then
ARGS+=(--close)
echo "::notice::Manual rollout dispatch with dry_run=false -> heads-up comments WILL be posted."
elif [ "${GITHUB_EVENT_NAME:-}" = "workflow_dispatch" ]; then
echo "::notice::Manual dispatch in dry-run mode -> previewing only, no comments will be posted."
else
echo "::notice::Automatic push trigger -> dry-run preview only. Fire the real rollout sweep with a manual workflow_dispatch (dry_run=false)."
fi
python3 .github/scripts/triage_rollout_heads_up.py "${ARGS[@]}"

2
.gitignore vendored
View file

@ -1,9 +1,11 @@
.python-version
.venv
tests/e2e/.fixtures/
.venv-typecheck
.venv_policy_test
.env
.claude
CLAUDE.local.md
.newenv
newenv/*
litellm/proxy/myenv/*

View file

@ -13,8 +13,8 @@ Here are the core requirements for any PR submitted to LiteLLM:
- [ ] **Add testing** - Adding at least 1 test is a hard requirement - [see details](#adding-testing)
- [ ] **Ensure your PR passes all checks**:
- [ ] [Unit Tests](#running-unit-tests) - `make test-unit`
- [ ] [Linting / Formatting](#running-linting-and-formatting-checks) - `make lint`
- [ ] [The tests covering your change](#running-unit-tests) pass, e.g. `uv run pytest tests/test_litellm/<your_test_file>.py -v`. CI runs the full unit test matrix, so you don't need to run the whole suite locally
#### UI PRs
@ -71,8 +71,8 @@ make format
# Run all linting checks (matches CI exactly)
make lint
# Run unit tests to ensure nothing is broken
make test-unit
# Run the tests covering your change (CI runs the full suite)
uv run pytest tests/test_litellm/<your_test_file>.py -v
# Commit your changes (must follow Conventional Commits — see above)
git add .
@ -123,12 +123,13 @@ def test_your_feature():
### Running Unit Tests
Run all unit tests (uses parallel execution for speed):
Run the tests covering your change:
```bash
make test-unit
uv run pytest tests/test_litellm/test_your_file.py -v
```
`tests/test_litellm` holds thousands of tests, so running all of it locally takes a long time. CI runs it as a parallel matrix (`make test-unit-llms`, `make test-unit-proxy-core`, and the other `test-unit-*` targets) on beefier boxes, so if, for whatever reason, you must run the whole suite, it's better to rely on CI to do that.
If you're running broader test suites, proxy tests, or anything that touches PostgreSQL-backed fixtures/plugins, install the full local test environment first:
```bash
@ -137,11 +138,6 @@ make install-test-deps
This syncs the locked test environment used across the repo, including `psycopg` v3 plus `psycopg-binary` (used by `pytest-postgresql`), `psycopg2-binary` (used by some proxy E2E tests), and a generated Prisma client for DB-backed proxy tests, so pytest startup matches CI without manual package installs.
Run specific test files:
```bash
uv run pytest tests/test_litellm/test_your_file.py -v
```
### Running Linting and Formatting Checks
Run all linting checks (matches CI exactly):

View file

@ -7,6 +7,7 @@
info lint lint-inner lint-dev lint-checks format \
lint-basedpyright lint-e2e-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \
lint-test-quality lint-test-quality-budget-update \
install-dev install-proxy-dev install-test-deps install-hooks \
install-helm-unittest check-circular-imports check-import-safety check check-inner pre-commit \
lint-install lint-fetch-base bootstrap
@ -35,7 +36,8 @@ help:
@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-ruff-budget-update - Ratchet ruff-strict-budget.json limits down by what this branch fixed"
@echo " make lint-budget-update - Ratchet all budgets down (ruff + type-discipline + basedpyright)"
@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)"
@echo " make check-circular-imports - Check for circular imports"
@echo " make check-import-safety - Check import safety"
@echo " make test - Run all tests"
@ -200,6 +202,12 @@ lint-e2e-basedpyright: $(LINT_E2E_DEP_INSTALL)
lint-type-discipline: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
$(UV_RUN) python scripts/type_discipline_gate.py --base origin/litellm_internal_staging
# Test-quality budget (zero-assert / mock-echo tests, sys.path.insert, raw env writes,
# litellm module-global mutation, credential-gated skips), 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
# --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
@ -221,8 +229,11 @@ lint-ruff-budget-update: install-dev lint-fetch-base
lint-type-discipline-budget-update: install-dev lint-fetch-base
$(UV_RUN) python scripts/type_discipline_gate.py --update
# Ratchet all budgets in one shot (ruff strict + type-discipline + basedpyright)
lint-budget-update: lint-ruff-budget-update lint-type-discipline-budget-update lint-basedpyright-budget-update
lint-test-quality-budget-update: install-dev lint-fetch-base
$(UV_RUN) python scripts/test_quality_gate.py --update
# 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
check-circular-imports: $(LINT_DEP_INSTALL)
cd litellm && $(UV_RUN) python ../tests/documentation_tests/test_circular_imports.py && cd ..
@ -244,7 +255,7 @@ lint:
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-checks: lint-format-check-changed lint-ruff lint-gate lint-type-discipline lint-basedpyright lint-e2e-basedpyright check-circular-imports check-import-safety
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
# Faster linting for local development (only checks changed code)
lint-dev: lint-format-changed check-circular-imports check-import-safety
@ -314,7 +325,7 @@ test-unit-helm: install-helm-unittest
# LLM Translation testing targets
test-llm-translation: install-test-deps
@echo "Running LLM translation tests..."
@python .github/workflows/run_llm_translation_tests.py
@python .github/scripts/run_llm_translation_tests.py
test-llm-translation-single: install-test-deps
@echo "Running single LLM translation test file..."

View file

@ -35,6 +35,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
# Models & routing config
"/model/",
"/v1/model/info",
"/v1/model/deprecations",
"/v2/model/",
"/model_group",
"/model_access_group/",

View file

@ -1,12 +1,12 @@
{
"reportAny": {
"limit": 22945
"limit": 19955
},
"reportArgumentType": {
"limit": 2579
"limit": 2566
},
"reportAssignmentType": {
"limit": 323
"limit": 320
},
"reportAttributeAccessIssue": {
"limit": 488
@ -24,13 +24,13 @@
"limit": 19
},
"reportExplicitAny": {
"limit": 7311
"limit": 6049
},
"reportFunctionMemberAccess": {
"limit": 7
},
"reportGeneralTypeIssues": {
"limit": 157
"limit": 154
},
"reportIncompatibleMethodOverride": {
"limit": 56
@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5707
"limit": 5663
},
"reportMissingTypeArgument": {
"limit": 15640
"limit": 15555
},
"reportMissingTypeStubs": {
"limit": 40
@ -72,7 +72,7 @@
"limit": 0
},
"reportOptionalMemberAccess": {
"limit": 1069
"limit": 1061
},
"reportOptionalOperand": {
"limit": 0
@ -84,7 +84,7 @@
"limit": 56
},
"reportPrivateUsage": {
"limit": 1824
"limit": 1823
},
"reportRedeclaration": {
"limit": 8
@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44776
"limit": 44655
},
"reportUnknownLambdaType": {
"limit": 113
"limit": 109
},
"reportUnknownMemberType": {
"limit": 39237
"limit": 39017
},
"reportUnknownParameterType": {
"limit": 19967
"limit": 19885
},
"reportUnknownVariableType": {
"limit": 30881
"limit": 30572
},
"reportUnnecessaryCast": {
"limit": 117
@ -123,7 +123,7 @@
"limit": 5
},
"reportUnnecessaryIsInstance": {
"limit": 853
"limit": 836
},
"reportUntypedBaseClass": {
"limit": 0
@ -132,7 +132,7 @@
"limit": 27
},
"reportUnusedClass": {
"limit": 23
"limit": 21
},
"reportUnusedFunction": {
"limit": 139

View file

@ -41,6 +41,11 @@ OBJECT_KEYS: dict[str, JsonSchema] = {
},
"additionalProperties": False,
},
"guardrail_cost_per_unit": {
"type": "object",
"description": "USD cost per billable guardrail unit, keyed by the provider's usage counter name (e.g. Bedrock's contentPolicyUnits).",
"additionalProperties": NONNEG_NUMBER,
},
"metadata": {
"type": "object",
"description": "Free-form notes about the entry (e.g. pricing derivation).",
@ -140,6 +145,11 @@ NUMBER_KEYS: dict[str, JsonSchema] = {
"minimum": 1,
"description": "Multiplier applied to all token costs for US data residency (e.g. 1.10 = +10%).",
},
"regional_endpoint_uplift_multiplier": {
"type": "number",
"minimum": 1,
"description": "Multiplier applied to all token costs when served from a non-global Vertex AI endpoint (e.g. 1.10 = +10%).",
},
}
COST_DESCRIPTIONS: dict[str, str] = {

View file

@ -24,12 +24,16 @@ if TYPE_CHECKING:
CHECK_BATCH_COST_USER_AGENT = "LiteLLM Proxy/CheckBatchCost"
TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = (
PROVIDER_TERMINAL_BATCH_STATUSES: Final[Tuple[str, ...]] = (
"completed",
"complete",
"failed",
"expired",
"cancelled",
)
TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = (
*PROVIDER_TERMINAL_BATCH_STATUSES,
"stale_expired",
)
@ -286,6 +290,57 @@ class CheckBatchCost:
404 must not retire the row; the staleness sweep bounds it instead."""
return self.llm_router.get_deployment(model_id=model_id) is not None
@staticmethod
def _is_output_file_gone_at_provider(error: Exception, output_file_id: Optional[str]) -> bool:
"""A 404 naming the output file means there is nothing to fetch on this or any
later poll: providers like Vertex AI advertise an output path for every batch,
including terminal ones that never wrote it. Any other failure may be
transient, so it keeps retrying until the staleness sweep bounds it."""
import openai
from litellm.exceptions import NotFoundError
if not output_file_id:
return False
return isinstance(error, (NotFoundError, openai.NotFoundError)) and output_file_id in str(error)
async def _finalize_unbilled_terminal_job(
self, job: "LiteLLM_ManagedObjectTable", response: "LiteLLMBatch"
) -> None:
"""Persist a terminal batch that has nothing billable, converting any raw
provider file ids to managed ids, and take it out of the poll page."""
try:
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
ensure_batch_response_managed_file_ids,
)
response.id = job.unified_object_id
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=self.proxy_logging_obj.get_proxy_hook("managed_files"),
prisma_client=self.prisma_client,
verbose_proxy_logger=verbose_proxy_logger,
db_batch_object=job,
unified_batch_id=_is_base64_encoded_unified_file_id(job.unified_object_id),
)
update_data: Final[dict] = {
"status": response.status,
"file_object": response.model_dump_json(),
**({"batch_processed": True} if self._has_batch_processed_column else {}),
}
await self.prisma_client.db.litellm_managedobjecttable.update(
where={"id": job.id},
data=update_data,
)
verbose_proxy_logger.info(
f"CheckBatchCost: marked job {job.id} as {response.status} in DB"
)
except Exception as db_err:
verbose_proxy_logger.error(
f"CheckBatchCost: failed to mark job {job.id} as {response.status} in DB: {db_err}"
)
@staticmethod
def _record_error(
prom_logger: Optional["PrometheusLogger"], error_type: str
@ -528,6 +583,7 @@ class CheckBatchCost:
from litellm.files.main import afile_content
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.litellm_logging import deployment_pricing_model_info
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
@ -648,15 +704,20 @@ class CheckBatchCost:
f"{_file_attr}={_raw_file_id!r}: {_e}"
)
# Pass deployment model_info so custom batch pricing
# (input_cost_per_token_batches etc.) is used for cost calc
deployment_model_info = deployment_info.model_info.model_dump() if deployment_info.model_info else {}
# Pass the deployment's router-registered pricing (litellm_params custom
# rates merged with the model's published rates) so custom batch pricing
# (input_cost_per_token_batches etc.) is used for cost calc, exactly as
# the inline retrieve path does.
deployment_model_info = deployment_pricing_model_info(
model_id=model_id,
deployment_model=litellm_model_name,
)
batch_cost, batch_usage, batch_models = (
await calculate_batch_cost_and_usage(
file_content_dictionary=file_content_as_dict,
custom_llm_provider=llm_provider, # type: ignore
model_name=model_name,
model_info=deployment_model_info, # type: ignore[arg-type]
model_info=deployment_model_info,
)
)
logging_obj = LiteLLMLogging(
@ -796,7 +857,7 @@ class CheckBatchCost:
## RETRIEVE THE BATCH JOB OUTPUT FILE
if (
response.status in ("completed", "complete", "expired")
response.status in PROVIDER_TERMINAL_BATCH_STATUSES
and response.output_file_id is not None
):
try:
@ -808,6 +869,15 @@ class CheckBatchCost:
prom_logger=prom_logger,
)
except Exception as tracking_err:
if self._is_output_file_gone_at_provider(
tracking_err, response.output_file_id
) and self._batch_deployment_exists(model_id):
verbose_proxy_logger.warning(
f"CheckBatchCost: output file {response.output_file_id} of batch {batch_id} "
f"does not exist at the provider; retiring job {job.id} unbilled"
)
await self._finalize_unbilled_terminal_job(job, response)
continue
verbose_proxy_logger.error(
f"CheckBatchCost: failed to track cost for batch {batch_id} "
f"(job {job.id}); leaving it unprocessed so the next poll retries: {tracking_err}"
@ -837,45 +907,8 @@ class CheckBatchCost:
f"CheckBatchCost: failed to mark job {job.id} complete in DB: {db_err}"
)
elif response.status in (
"completed",
"complete",
"failed",
"expired",
"cancelled",
):
try:
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
ensure_batch_response_managed_file_ids,
)
response.id = job.unified_object_id
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=self.proxy_logging_obj.get_proxy_hook("managed_files"),
prisma_client=self.prisma_client,
verbose_proxy_logger=verbose_proxy_logger,
db_batch_object=job,
unified_batch_id=_is_base64_encoded_unified_file_id(job.unified_object_id),
)
update_data = {
"status": response.status,
"file_object": response.model_dump_json(),
}
if self._has_batch_processed_column:
update_data["batch_processed"] = True
await self.prisma_client.db.litellm_managedobjecttable.update(
where={"id": job.id},
data=update_data,
)
verbose_proxy_logger.info(
f"CheckBatchCost: marked job {job.id} as {response.status} in DB"
)
except Exception as db_err:
verbose_proxy_logger.error(
f"CheckBatchCost: failed to mark job {job.id} as {response.status} in DB: {db_err}"
)
elif response.status in PROVIDER_TERMINAL_BATCH_STATUSES:
await self._finalize_unbilled_terminal_job(job, response)
# Record polling run metrics (always, even if nothing was processed)
if prom_logger:

View file

@ -41,6 +41,7 @@ from litellm.proxy._types import (
CallTypes,
LiteLLM_ManagedFileTable,
LiteLLM_ManagedObjectTable,
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.openai_files_endpoints.common_utils import (
@ -423,13 +424,26 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# This is because the encoded object ids stored in the managed objects table do not contain the provider information
# To support provider filtering, we would need to store the provider information in the encoded object ids
if provider:
raise Exception("Filtering by 'provider' is not supported when using managed batches.")
raise ProxyException(
message="Filtering by 'provider' is not supported when using managed batches.",
type="invalid_request_error",
param="provider",
code=400,
)
# Model name filtering is not supported for managed batches
# This is because the encoded object ids stored in the managed objects table do not contain the model name
# A hash of the model name + litellm_params for the model name is encoded as the model id. This is not sufficient to reliably map the target model names to the model ids.
if target_model_names:
raise Exception("Filtering by 'target_model_names' is not supported when using managed batches.")
raise ProxyException(
message="Filtering by 'target_model_names' is not supported when using managed batches.",
type="invalid_request_error",
param="target_model_names",
code=400,
)
if limit == 0:
return build_list_page([])
owner_filter = build_owner_filter(user_api_key_dict)
if owner_filter is None:

View file

@ -11,7 +11,7 @@ Endpoints for /project operations
#### PROJECT MANAGEMENT ####
import json
from collections.abc import Mapping, Sequence
from collections.abc import Sequence
from typing import TYPE_CHECKING
from fastapi import APIRouter, Depends, HTTPException, Request
@ -29,7 +29,11 @@ from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy
if TYPE_CHECKING:
from prisma import models as prisma_models
from prisma.actions import LiteLLM_TeamTableActions
from prisma.actions import (
LiteLLM_ProjectTableActions,
LiteLLM_TeamTableActions,
LiteLLM_VerificationTokenActions,
)
router = APIRouter()
@ -39,6 +43,27 @@ def _team_table(prisma_client: PrismaClient) -> "LiteLLM_TeamTableActions[prisma
return team_table
def _project_table(prisma_client: PrismaClient) -> "LiteLLM_ProjectTableActions[prisma_models.LiteLLM_ProjectTable]":
project_table: LiteLLM_ProjectTableActions[prisma_models.LiteLLM_ProjectTable] = (
prisma_client.db.litellm_projecttable
)
return project_table
def _verification_token_table(
prisma_client: PrismaClient,
) -> "LiteLLM_VerificationTokenActions[prisma_models.LiteLLM_VerificationToken]":
verification_token_table: LiteLLM_VerificationTokenActions[prisma_models.LiteLLM_VerificationToken] = (
prisma_client.db.litellm_verificationtoken
)
return verification_token_table
def _jsonified(prisma_client: PrismaClient, payload: dict[str, object]) -> dict[str, object]:
jsonified: dict[str, object] = prisma_client.jsonify_object(payload)
return jsonified
async def _check_user_permission_for_project(
user_api_key_dict: UserAPIKeyAuth,
team_id: str | None,
@ -137,7 +162,7 @@ def _check_team_project_limits(
# --- Validate project models are a subset of team models ---
project_models = data.models
team_models = team_object.models or []
team_models: list[str] = team_object.models or []
if project_models and len(team_models) > 0:
# If team has 'all-proxy-models', skip validation as it allows all models
if SpecialModelNames.all_proxy_models.value not in team_models:
@ -188,11 +213,11 @@ async def _create_budget_for_project(
) -> str:
"""Create a budget for the project and return budget_id."""
budget_params = LiteLLM_BudgetTable.model_fields.keys()
_json_data: Mapping[str, object] = data.json(exclude_none=True)
_json_data: dict[str, object] = data.model_dump(exclude_none=True)
_budget_data = {k: v for k, v in _json_data.items() if k in budget_params}
budget_row = LiteLLM_BudgetTable.model_validate(_budget_data)
new_budget = prisma_client.jsonify_object(budget_row.json(exclude_none=True))
new_budget = _jsonified(prisma_client, budget_row.model_dump(exclude_none=True))
_budget: prisma_models.LiteLLM_BudgetTable = await prisma_client.db.litellm_budgettable.create(
data={
@ -227,7 +252,7 @@ async def _set_project_object_permission(
return None
def _remove_budget_fields_from_project_data(project_data: dict) -> dict:
def _remove_budget_fields_from_project_data(project_data: dict[str, object]) -> dict[str, object]:
"""
Remove budget fields from project data.
Budget fields belong to LiteLLM_BudgetTable, not LiteLLM_ProjectTable.
@ -396,9 +421,7 @@ async def new_project(
data.project_id = str(uuid.uuid4())
else:
# Check if project_id already exists
existing_project = await prisma_client.db.litellm_projecttable.find_unique(
where={"project_id": data.project_id}
)
existing_project = await _project_table(prisma_client).find_unique(where={"project_id": data.project_id})
if existing_project is not None:
raise ProxyException(
message=f"Project id = {data.project_id} already exists. Please use a different project id.",
@ -423,11 +446,14 @@ async def new_project(
)
# Create project row (following organization_endpoints.py pattern)
project_row = LiteLLM_ProjectTable(
**data.json(exclude_none=True),
object_permission_id=object_permission_id,
created_by=user_api_key_dict.user_id or litellm_proxy_admin_name,
updated_by=user_api_key_dict.user_id or litellm_proxy_admin_name,
project_row_payload: dict[str, object] = data.model_dump(exclude_none=True)
project_row = LiteLLM_ProjectTable.model_validate(
{
**project_row_payload,
"object_permission_id": object_permission_id,
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
"updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
}
)
for field in LiteLLM_ManagementEndpoint_MetadataFields:
@ -438,7 +464,7 @@ async def new_project(
value=getattr(data, field),
)
new_project_row = prisma_client.jsonify_object(project_row.json(exclude_none=True))
new_project_row = _jsonified(prisma_client, project_row.model_dump(exclude_none=True))
# Remove budget fields (following organization_endpoints.py pattern)
new_project_row = _remove_budget_fields_from_project_data(new_project_row)
@ -560,7 +586,7 @@ async def update_project(
# Fetch existing project
existing_project: (
prisma_models.LiteLLM_ProjectTable | None
) = await prisma_client.db.litellm_projecttable.find_unique(where={"project_id": data.project_id})
) = await _project_table(prisma_client).find_unique(where={"project_id": data.project_id})
if existing_project is None:
raise ProxyException(
@ -617,8 +643,7 @@ async def update_project(
)
# Prepare update data
update_data = data.json(exclude_none=True, exclude={"project_id"})
update_data = prisma_client.jsonify_object(update_data)
update_data = _jsonified(prisma_client, data.model_dump(exclude_none=True, exclude={"project_id"}))
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
# Handle budget updates
@ -660,9 +685,10 @@ async def update_project(
# Handle metadata fields
for field in LiteLLM_ManagementEndpoint_MetadataFields:
if field in update_data:
if update_data.get("metadata") is None:
update_data["metadata"] = {}
update_data["metadata"][field] = update_data.pop(field)
existing_metadata = update_data.get("metadata")
metadata_dict: dict[str, object] = existing_metadata if isinstance(existing_metadata, dict) else {}
metadata_dict[field] = update_data.pop(field)
update_data["metadata"] = metadata_dict
# Remove budget fields (following organization_endpoints.py pattern)
update_data = _remove_budget_fields_from_project_data(update_data)
@ -748,11 +774,11 @@ async def delete_project(
detail={"error": "Only admins can delete projects"},
)
deleted_projects = []
deleted_projects: list[prisma_models.LiteLLM_ProjectTable | None] = []
for project_id in data.project_ids:
# Check if project exists
existing_project = await prisma_client.db.litellm_projecttable.find_unique(where={"project_id": project_id})
existing_project = await _project_table(prisma_client).find_unique(where={"project_id": project_id})
if existing_project is None:
raise ProxyException(
@ -765,7 +791,7 @@ async def delete_project(
# Check if there are any keys associated with this project
associated_keys: Sequence[
prisma_models.LiteLLM_VerificationToken
] = await prisma_client.db.litellm_verificationtoken.find_many(where={"project_id": project_id})
] = await _verification_token_table(prisma_client).find_many(where={"project_id": project_id})
if len(associated_keys) > 0:
raise ProxyException(
@ -778,7 +804,7 @@ async def delete_project(
# Delete the project
deleted_project: (
prisma_models.LiteLLM_ProjectTable | None
) = await prisma_client.db.litellm_projecttable.delete(where={"project_id": project_id})
) = await _project_table(prisma_client).delete(where={"project_id": project_id})
await delete_cached_project_object(
project_id=project_id,
@ -829,7 +855,7 @@ async def project_info(
)
# Fetch project
project: prisma_models.LiteLLM_ProjectTable | None = await prisma_client.db.litellm_projecttable.find_unique(
project: prisma_models.LiteLLM_ProjectTable | None = await _project_table(prisma_client).find_unique(
where={"project_id": project_id},
include={"litellm_budget_table": True, "object_permission": True},
)
@ -901,7 +927,7 @@ async def list_projects(
if user_api_key_has_admin_view(user_api_key_dict):
projects: Sequence[
prisma_models.LiteLLM_ProjectTable
] = await prisma_client.db.litellm_projecttable.find_many(
] = await _project_table(prisma_client).find_many(
include={"litellm_budget_table": True, "object_permission": True}
)
else:
@ -911,9 +937,9 @@ async def list_projects(
user_record: prisma_models.LiteLLM_UserTable | None = await prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_api_key_dict.user_id},
)
user_team_ids: Sequence[str] = user_record.teams if user_record is not None and user_record.teams else []
user_team_ids: list[str] = user_record.teams if user_record is not None and user_record.teams else []
projects = await prisma_client.db.litellm_projecttable.find_many(
projects = await _project_table(prisma_client).find_many(
where={"team_id": {"in": user_team_ids}},
include={"litellm_budget_table": True, "object_permission": True},
)

View file

@ -10,7 +10,8 @@ All /vector_store management endpoints
import copy
import json
from typing import List, Optional
from collections.abc import Mapping
from typing import TYPE_CHECKING, Final, List, Optional, Protocol
from fastapi import APIRouter, Depends, HTTPException
@ -32,9 +33,35 @@ from litellm.types.vector_stores import (
)
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
router = APIRouter()
class ManagedVectorStoreRow(Protocol):
"""A ``litellm_managedvectorstorestable`` row as returned by Prisma."""
def model_dump(self) -> LiteLLM_ManagedVectorStore: ...
class ManagedVectorStoreTable(Protocol):
"""The Prisma actions namespace for ``litellm_managedvectorstorestable``."""
async def find_unique(self, where: Mapping[str, str | None]) -> ManagedVectorStoreRow | None: ...
async def create(self, data: Mapping[str, object]) -> ManagedVectorStoreRow: ...
async def delete(self, where: Mapping[str, str | None]) -> ManagedVectorStoreRow | None: ...
async def update(self, where: Mapping[str, str | None], data: Mapping[str, object]) -> ManagedVectorStoreRow: ...
def managed_vector_store_table(prisma_client: "PrismaClient") -> ManagedVectorStoreTable:
"""The Prisma table actions for managed vector stores, behind a typed surface."""
return prisma_client.db.litellm_managedvectorstorestable
########################################################
# Management Endpoints
########################################################
@ -66,7 +93,7 @@ async def new_vector_store(
try:
# Check if vector store already exists
existing_vector_store = (
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
await managed_vector_store_table(prisma_client).find_unique(
where={"vector_store_id": vector_store.get("vector_store_id")}
)
)
@ -92,7 +119,7 @@ async def new_vector_store(
del vector_store["litellm_params"]
_new_vector_store = (
await prisma_client.db.litellm_managedvectorstorestable.create(
await managed_vector_store_table(prisma_client).create(
data={
**vector_store,
"litellm_params": litellm_params_json,
@ -213,7 +240,7 @@ async def delete_vector_store(
try:
# Check if vector store exists
existing_vector_store = (
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
await managed_vector_store_table(prisma_client).find_unique(
where={"vector_store_id": data.vector_store_id}
)
)
@ -224,7 +251,7 @@ async def delete_vector_store(
)
# Delete vector store
await prisma_client.db.litellm_managedvectorstorestable.delete(
await managed_vector_store_table(prisma_client).delete(
where={"vector_store_id": data.vector_store_id}
)
@ -288,7 +315,7 @@ async def get_vector_store_info(
return {"vector_store": vector_store_pydantic_obj}
vector_store = (
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
await managed_vector_store_table(prisma_client).find_unique(
where={"vector_store_id": data.vector_store_id}
)
)
@ -298,7 +325,7 @@ async def get_vector_store_info(
detail=f"Vector store with ID {data.vector_store_id} not found",
)
vector_store_dict = vector_store.model_dump() # type: ignore[attr-defined]
vector_store_dict = vector_store.model_dump()
return {"vector_store": vector_store_dict}
except Exception as e:
verbose_proxy_logger.exception(f"Error getting vector store info: {str(e)}")
@ -322,13 +349,13 @@ async def update_vector_store(
try:
update_data = data.model_dump(exclude_unset=True)
vector_store_id = update_data.pop("vector_store_id")
vector_store_id: Final[str] = update_data.pop("vector_store_id")
if update_data.get("vector_store_metadata") is not None:
update_data["vector_store_metadata"] = safe_dumps(
update_data["vector_store_metadata"]
)
updated = await prisma_client.db.litellm_managedvectorstorestable.update(
updated = await managed_vector_store_table(prisma_client).update(
where={"vector_store_id": vector_store_id},
data=update_data,
)

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.56"
version = "0.1.57"
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.56"
version = "0.1.57"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -83,6 +83,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
"/azure_ai/",
"/aws/",
"/bedrock/",
"/comprehendmedical",
"/cohere/",
"/gemini/",
"/google/",

View file

@ -18,7 +18,7 @@ type: application
# This is the chart version. This version number should be incremented each time you make changes
# to the chart and its templates, including the app version.
# Versions are expected to follow Semantic Versioning (https://semver.org/)
version: 1.1.1
version: 1.1.2
# This is the version number of the application being deployed. This version number should be
# incremented each time you make changes to the application. Versions are not expected to

View file

@ -29,7 +29,7 @@ If `db.useStackgresOperator` is used (not yet implemented):
| `masterkey` | The Master API Key for LiteLLM. If not specified, a random key in the `sk-...` format is generated. | N/A |
| `environmentSecrets` | An optional array of Secret object names. The keys and values in these secrets will be presented to the LiteLLM proxy pod as environment variables. See below for an example Secret object. | `[]` |
| `environmentConfigMaps` | An optional array of ConfigMap object names. The keys and values in these configmaps will be presented to the LiteLLM proxy pod as environment variables. See below for an example Secret object. | `[]` |
| `image.repository` | LiteLLM Proxy image repository | `docker.litellm.ai/berriai/litellm` |
| `image.repository` | LiteLLM Proxy image repository | `ghcr.io/berriai/litellm` |
| `image.pullPolicy` | LiteLLM Proxy image pull policy | `IfNotPresent` |
| `image.tag` | Overrides the image tag whose default the latest version of LiteLLM at the time this chart was published. | `""` |
| `imagePullSecrets` | Registry credentials for the LiteLLM and initContainer images. | `[]` |

View file

@ -119,4 +119,7 @@ spec:
{{- end }}
ttlSecondsAfterFinished: {{ .Values.migrationJob.ttlSecondsAfterFinished }}
backoffLimit: {{ .Values.migrationJob.backoffLimit }}
{{- with .Values.migrationJob.activeDeadlineSeconds }}
activeDeadlineSeconds: {{ . }}
{{- end }}
{{- end }}

View file

@ -15,7 +15,7 @@ tests:
pattern: -litellm$
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-database:test
value: ghcr.io/berriai/litellm:test
- it: should work with tolerations
template: deployment.yaml
set:
@ -337,7 +337,7 @@ tests:
template: deployment.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
repository: ghcr.io/berriai/litellm
tag: test
extraInitContainers:
- name: init-tpl
@ -348,7 +348,7 @@ tests:
path: spec.template.spec.initContainers
content:
name: init-tpl
image: "ghcr.io/berriai/litellm-database:test"
image: "ghcr.io/berriai/litellm:test"
command: ["echo", "hello"]
- it: should work with extraContainers
template: deployment.yaml
@ -366,7 +366,7 @@ tests:
template: deployment.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
repository: ghcr.io/berriai/litellm
tag: test
extraContainers:
- name: sidecar-tpl
@ -376,12 +376,12 @@ tests:
path: spec.template.spec.containers
content:
name: sidecar-tpl
image: "ghcr.io/berriai/litellm-database:test"
image: "ghcr.io/berriai/litellm:test"
- it: should support tpl in podAnnotations
template: deployment.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
repository: ghcr.io/berriai/litellm
tag: test
# Mirrors the real-world scenario this feature unblocks:
# user disables the built-in ConfigMap (and its built-in checksum/config
@ -398,7 +398,7 @@ tests:
value: "test"
- equal:
path: spec.template.metadata.annotations["example.com/some-key"]
value: "ghcr.io/berriai/litellm-database"
value: "ghcr.io/berriai/litellm"
- equal:
path: spec.template.metadata.annotations["example.com/literal"]
value: "plain-string-value"

View file

@ -208,7 +208,7 @@ tests:
template: migrations-job.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
repository: ghcr.io/berriai/litellm
tag: test
migrationJob:
enabled: true
@ -221,7 +221,7 @@ tests:
path: spec.template.spec.initContainers
content:
name: init-tpl
image: "ghcr.io/berriai/litellm-database:test"
image: "ghcr.io/berriai/litellm:test"
command: ["echo", "hello"]
- it: should work with extraContainers
template: migrations-job.yaml
@ -241,7 +241,7 @@ tests:
template: migrations-job.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
repository: ghcr.io/berriai/litellm
tag: test
migrationJob:
enabled: true
@ -253,7 +253,7 @@ tests:
path: spec.template.spec.containers
content:
name: sidecar-tpl
image: "ghcr.io/berriai/litellm-database:test"
image: "ghcr.io/berriai/litellm:test"
- it: should render the pod-level securityContext from podSecurityContext
template: migrations-job.yaml
set:
@ -314,3 +314,31 @@ tests:
operator: Equal
value: litellm-e2e
effect: NoSchedule
- it: bounds the Job with a deadline by default, so a blocked migration cannot stall the release forever
set:
migrationJob:
enabled: true
asserts:
- equal:
path: spec.activeDeadlineSeconds
value: 1800
- it: honours an operator-supplied deadline
set:
migrationJob:
enabled: true
activeDeadlineSeconds: 600
asserts:
- equal:
path: spec.activeDeadlineSeconds
value: 600
- it: omits the deadline entirely when it is nulled out, restoring the unbounded behaviour
set:
migrationJob:
enabled: true
activeDeadlineSeconds: null
asserts:
- notExists:
path: spec.activeDeadlineSeconds

View file

@ -6,8 +6,9 @@ replicaCount: 1
# numWorkers: 2
image:
# Use "ghcr.io/berriai/litellm-database" for optimized image with database
repository: ghcr.io/berriai/litellm-database
# Bundles the prisma CLI and engines, which is what lets the migrations job
# and the proxy's own schema check run without network access.
repository: ghcr.io/berriai/litellm
pullPolicy: Always
# Overrides the image tag whose default is the chart appVersion.
# tag: "latest"
@ -427,6 +428,13 @@ migrationJob:
enabled: true # Enable or disable the schema migration Job
retries: 3 # Number of retries for the Job in case of failure
backoffLimit: 4 # Backoff limit for Job restarts
# Wall-clock budget for the whole Job, shared across every `backoffLimit`
# retry rather than granted per attempt. Without it a migration that blocks
# on the database never fails, and when the Helm hook is enabled the release
# waits on it forever: `helm upgrade` and any GitOps controller driving it
# stop reconciling the whole chart until someone deletes the Job by hand.
# Set to null to opt out and restore the unbounded behaviour.
activeDeadlineSeconds: 1800
disableSchemaUpdate: false # Skip schema migrations for specific environments. When True, the job will exit with code 0.
# Optional service account for the migration job.
# Only used when migrationJob.hooks.helm.enabled=true and serviceAccount.create=true.

View file

@ -24,7 +24,7 @@
"/v1/video" "/v1/videos" "/video" "/videos" "/v1/search" "/search"
"/v1/containers" "/containers" "/v1/evals" "/v1/memory" "/queue/chat"
"/v1beta" "/interactions"
"/anthropic" "/azure" "/azure_ai" "/aws" "/bedrock" "/cohere" "/gemini" "/google"
"/anthropic" "/azure" "/azure_ai" "/aws" "/bedrock" "/comprehendmedical" "/cohere" "/gemini" "/google"
"/vertex_ai" "/vertex-ai" "/assemblyai" "/eu.assemblyai" "/langfuse" "/vllm"
"/mistral" "/groq" "/voyage" "/cursor" "/milvus" "/openai_passthrough"
"/toolset"

View file

@ -21,6 +21,9 @@ metadata:
spec:
backoffLimit: {{ .Values.migrationJob.backoffLimit }}
ttlSecondsAfterFinished: {{ .Values.migrationJob.ttlSecondsAfterFinished }}
{{- with .Values.migrationJob.activeDeadlineSeconds }}
activeDeadlineSeconds: {{ . }}
{{- end }}
template:
metadata:
{{- /* The Job's selector is generated by the controller rather than

View file

@ -167,3 +167,24 @@ tests:
- equal:
path: spec.template.metadata.labels['app.kubernetes.io/component']
value: batch-migrations
- it: bounds the Job with a deadline by default, so a blocked migration cannot stall the release forever
asserts:
- equal:
path: spec.activeDeadlineSeconds
value: 1800
- it: honours an operator-supplied deadline
set:
migrationJob.activeDeadlineSeconds: 600
asserts:
- equal:
path: spec.activeDeadlineSeconds
value: 600
- it: omits the deadline entirely when it is nulled out, restoring the unbounded behaviour
set:
migrationJob.activeDeadlineSeconds: null
asserts:
- notExists:
path: spec.activeDeadlineSeconds

View file

@ -56,6 +56,15 @@ migrationJob:
enabled: true
backoffLimit: 4
ttlSecondsAfterFinished: 120
# Wall-clock budget for the whole Job, shared across every `backoffLimit`
# retry rather than granted per attempt. Without it a migration that blocks
# on the database never fails, and because this is a pre-upgrade hook the
# release waits on it forever: `helm upgrade` and any GitOps controller
# driving it stop reconciling the whole chart until someone deletes the Job
# by hand. A migration that has exhausted its retries is not going to
# succeed on the next one, so failing is strictly better than hanging.
# Set to null to opt out and restore the unbounded behaviour.
activeDeadlineSeconds: 1800
resources: {}
# ServiceAccount for the Job pod only.
#

View file

@ -0,0 +1,9 @@
-- CreateTable
CREATE TABLE "LiteLLM_ProxyWorkerHeartbeat" (
"worker_id" TEXT NOT NULL,
"hostname" TEXT NOT NULL,
"started_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"last_heartbeat_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
CONSTRAINT "LiteLLM_ProxyWorkerHeartbeat_pkey" PRIMARY KEY ("worker_id")
);

View file

@ -0,0 +1,7 @@
ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN IF NOT EXISTS "group_id" TEXT;
UPDATE "LiteLLM_ShadowEvalJob" SET "group_id" = "id" WHERE "group_id" IS NULL;
ALTER TABLE "LiteLLM_ShadowEvalJob" ALTER COLUMN "group_id" SET NOT NULL;
CREATE INDEX IF NOT EXISTS "LiteLLM_ShadowEvalJob_group_id_idx" ON "LiteLLM_ShadowEvalJob"("group_id");

View file

@ -0,0 +1,16 @@
-- CreateTable
CREATE TABLE "LiteLLM_DailyGuardrailUsageUnits" (
"guardrail_id" TEXT NOT NULL,
"date" TEXT NOT NULL,
"team_id" TEXT NOT NULL,
"api_key" TEXT NOT NULL,
"usage_unit" TEXT NOT NULL,
"units" BIGINT NOT NULL DEFAULT 0,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL,
CONSTRAINT "LiteLLM_DailyGuardrailUsageUnits_pkey" PRIMARY KEY ("guardrail_id","date","team_id","api_key","usage_unit")
);
-- CreateIndex
CREATE INDEX "LiteLLM_DailyGuardrailUsageUnits_date_idx" ON "LiteLLM_DailyGuardrailUsageUnits"("date");

View file

@ -0,0 +1,3 @@
ALTER TABLE "LiteLLM_SpendLogs"
ADD COLUMN IF NOT EXISTS "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
ADD COLUMN IF NOT EXISTS "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP;

View file

@ -0,0 +1,5 @@
-- AlterTable
ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN "stopped_by" TEXT;
UPDATE "LiteLLM_ShadowEvalJob" SET stopped_by = 'unknown'
WHERE stopped_at IS NOT NULL AND ends_at > (NOW() AT TIME ZONE 'utc');

View file

@ -0,0 +1,4 @@
UPDATE "LiteLLM_SpendLogs"
SET "created_at" = "endTime",
"updated_at" = "endTime"
WHERE "created_at" > "endTime" + interval '1 hour';

View file

@ -641,6 +641,8 @@ model LiteLLM_SpendLogs {
mcp_namespaced_tool_name String?
agent_id String?
proxy_server_request Json? @default("{}")
created_at DateTime @default(now()) @map("created_at")
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
@@index([startTime])
@@index([startTime, request_id])
@@index([end_user])
@ -945,6 +947,17 @@ model LiteLLM_DailyTagSpend {
}
// One row per live proxy worker process. Workers upsert their row on a fixed
// heartbeat; counting rows with a recent heartbeat tells how many workers share
// this database, which lets the Admin UI hide its "no Redis" warning for
// deployments that are provably a single worker.
model LiteLLM_ProxyWorkerHeartbeat {
worker_id String @id
hostname String
started_at DateTime @default(now())
last_heartbeat_at DateTime @default(now())
}
// Track the status of cron jobs running. Only allow one pod to run the job at a time
model LiteLLM_CronJob {
cronjob_id String @id @default(cuid()) // Unique ID for the record
@ -1069,6 +1082,21 @@ model LiteLLM_DailyGuardrailMetrics {
@@index([guardrail_id])
}
// Daily guardrail billable usage units (one row per guardrail/day/team/key/unit type)
model LiteLLM_DailyGuardrailUsageUnits {
guardrail_id String
date String // YYYY-MM-DD
team_id String // empty string when the request had no team
api_key String // hashed virtual key; empty string when unknown
usage_unit String // provider counter name, e.g. Bedrock's contentPolicyUnits
units BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@id([guardrail_id, date, team_id, api_key, usage_unit])
@@index([date])
}
// Daily policy metrics for usage dashboard (one row per policy per day)
model LiteLLM_DailyPolicyMetrics {
policy_id String
@ -1450,28 +1478,38 @@ model LiteLLM_AutoRouterSession {
@@index([last_turn_at], map: "idx_autorouter_session_last_turn")
}
// Shadow eval: evaluation of an auto-router against a key's live traffic, in either
// direction. forward duplicates the requests the key did not route through the router
// through it, answering whether the key should adopt it; reverse duplicates the requests
// the router did serve against a fixed baseline model, answering whether a key already on
// it still benefits. Either way a sampled slice runs in a detached task and an LLM judge
// compares real vs shadow responses blind. The job row is immutable config plus
// stopped_at; every count, status, and spend figure is derived from the append-only
// attempt rows, so nothing can disagree across pods or stop races.
// Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in
// either direction. forward duplicates the requests the keys did not route through the
// router through it, answering whether they should adopt it; reverse duplicates the
// requests the router did serve against a fixed baseline model, answering whether a key
// already on it still benefits. Either way a sampled slice runs in a detached task and an
// LLM judge compares real vs shadow responses blind. Each row is ONE key's leg of a job:
// immutable config plus that key's own turn budget and stop state, so one key exhausting
// its budget never ends a sibling's sampling. A job is the set of legs sharing group_id
// (the id the API reports), written together by one atomic create_many with identical
// config; single-key jobs predating group_id were backfilled group_id = id. "One active
// job per (key, direction)" is a partial unique index on (api_key_id, direction) WHERE
// stopped_at IS NULL, expressed only in the migration because schema.prisma cannot state
// partial indexes; it is what makes a concurrent start on another pod race-safe rather
// than read-then-create. Every count, status, and spend figure is derived from the
// append-only attempt rows, so nothing can disagree across pods or stop races.
model LiteLLM_ShadowEvalJob {
id String @id @default(cuid())
api_key_id String // hashed virtual key whose traffic is shadowed
group_id String // legs of one job share this; the API's job id
api_key_id String // hashed virtual key whose traffic this leg shadows
router_name String // the auto-router under evaluation, in either direction
direction String @default("forward") // forward | reverse
baseline_model String? // reverse only: the fixed model the router is judged against
judge_model String
shadow_percentage Float
max_turns Int // sample budget: judge at most this many turns
max_turns Int // this key's sample budget: judge at most this many turns
created_at DateTime @default(now())
created_by String?
ends_at DateTime
stopped_at DateTime?
stopped_by String? // operator who stopped it early; null when it ended on its own
@@index([group_id])
@@index([api_key_id])
@@index([created_at])
}

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.86"
version = "0.4.87"
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.86"
version = "0.4.87"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -26,6 +26,7 @@ def _dev_env_hot_reload_enabled() -> bool:
if os.getenv("LITELLM_MODE", "DEV") == "DEV":
_dotenv.load_dotenv(override=_dev_env_hot_reload_enabled())
from collections.abc import Sequence
from typing import (
Any,
Callable,
@ -217,7 +218,10 @@ add_user_information_to_llm_headers: Optional[bool] = (
overwrite_user_with_key_hash: bool = (
False # force the outgoing `user` param to the hashed api key, so providers see a stable, tamper-proof id
)
store_audit_logs = False # Enterprise feature, allow users to see audit logs
bedrock_request_metadata_fields: Optional[Sequence[str]] = (
None # allow-list of `user_api_key_*` fields (+ `spend_logs_metadata`) sent as Bedrock `requestMetadata`
)
store_audit_logs: bool | None = None
skip_system_message_in_guardrail: bool = False
skip_tool_message_in_guardrail: bool = False
### end of callbacks #############
@ -788,6 +792,8 @@ def _populate_provider_model_sets(model_cost_map: Dict) -> None:
nlp_cloud_models.add(key)
elif value.get("litellm_provider") == "aleph_alpha":
aleph_alpha_models.add(key)
elif value.get("litellm_provider") == "bedrock" and value.get("mode") == "guardrail":
pass
elif value.get("litellm_provider") == "bedrock" and not is_bedrock_pricing_only_model(key):
bedrock_models.add(key)
elif value.get("litellm_provider") == "bedrock_converse":

View file

@ -10,7 +10,7 @@ from typing import Any, Final
import litellm
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.secret_redaction import redact_string
from litellm.litellm_core_utils.secret_redaction import redact_string, redact_structured_value
set_verbose = False
@ -59,6 +59,12 @@ def _redact_string(value: str) -> str:
return redact_string(value)
def _redact_structured_value(key: str | None, value: str) -> str:
if not _ENABLE_SECRET_REDACTION:
return value
return redact_structured_value(key, value)
def redact_secrets(value: str) -> str:
"""Public API: redact known secret/credential patterns from an arbitrary string.
@ -265,7 +271,7 @@ class JsonFormatter(Formatter):
if record.exc_info:
json_record["stacktrace"] = record.exc_text or self.formatException(record.exc_info)
return safe_dumps(json_record)
return safe_dumps(json_record, value_transform=_redact_structured_value)
class CorrelationPlainFormatter(logging.Formatter):
@ -276,7 +282,7 @@ class CorrelationPlainFormatter(logging.Formatter):
"""
def format(self, record: logging.LogRecord) -> str:
formatted: Final = super().format(record)
formatted: Final = _redact_string(super().format(record))
trace_id: Final = getattr(record, "trace_id", None)
session_id: Final = getattr(record, "session_id", None)
if not trace_id and not session_id:

View file

@ -7,9 +7,10 @@ import hashlib
import json
import time
from collections.abc import AsyncIterator
from typing import Any, Final, NamedTuple, cast
from typing import Any, Final, NamedTuple, Protocol
import httpx
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._logging import verbose_logger
from litellm.a2a_protocol.providers.watsonx_orchestrate.transformation import (
@ -38,11 +39,59 @@ class WXORequestParams(NamedTuple):
thread_id: str | None
class WXOLitellmParams(TypedDict, total=False):
"""litellm_params keys read when routing an A2A request to watsonx Orchestrate."""
cp4d_host: ReadOnly[str]
instance_id: ReadOnly[str]
wxo_agent_id: ReadOnly[str]
api_key: ReadOnly[str]
username: ReadOnly[str | None]
auth_mode: ReadOnly[str]
thread_id: ReadOnly[str | None]
class _IBMCloudTokenBody(TypedDict):
"""Fields read from the IBM Cloud IAM token response."""
access_token: ReadOnly[str]
expires_in: ReadOnly[NotRequired[int]]
class _CP4DTokenBody(TypedDict):
"""Fields read from the CP4D authorize response."""
token: ReadOnly[str]
expiration: ReadOnly[NotRequired[float]]
class _WXORun(TypedDict, total=False):
"""Fields the handler reads from a WXO run object or run event."""
status: ReadOnly[str]
run_id: ReadOnly[str]
id: ReadOnly[str]
class _SSELineSource(Protocol):
def aiter_lines(self) -> AsyncIterator[str]: ...
class _WXOView(TypedDict, total=False):
"""Typed reads of otherwise untyped watsonx Orchestrate and httpx values."""
ibm_cloud_token: ReadOnly[_IBMCloudTokenBody]
cp4d_token: ReadOnly[_CP4DTokenBody]
run: ReadOnly[_WXORun]
content_type: ReadOnly[str]
sse_source: ReadOnly[_SSELineSource]
class WatsonxOrchestrateHandler:
@staticmethod
def _http_client(timeout: float = 90.0) -> AsyncHTTPHandler:
return get_async_httpx_client(
llm_provider=cast(Any, httpxSpecialProvider.A2AProvider),
llm_provider=httpxSpecialProvider.A2AProvider,
params={"timeout": timeout},
)
@ -57,7 +106,7 @@ class WatsonxOrchestrateHandler:
return hashlib.sha256(material.encode()).hexdigest()
@staticmethod
def _cp4d_token_ttl_seconds(expiration: Any, now_wall: float | None = None) -> int:
def _cp4d_token_ttl_seconds(expiration: float, now_wall: float | None = None) -> int:
# CP4D returns expiration as absolute Unix epoch seconds, not a duration.
expires_at: Final = int(expiration)
wall: Final = now_wall if now_wall is not None else time.time()
@ -90,9 +139,9 @@ class WatsonxOrchestrateHandler:
headers={"Content-Type": "application/x-www-form-urlencoded"},
)
response.raise_for_status()
payload = response.json()
token = str(payload["access_token"])
ttl_s = int(payload.get("expires_in", 3600))
iam_payload: Final[_WXOView] = {"ibm_cloud_token": response.json()}
token = str(iam_payload["ibm_cloud_token"]["access_token"])
ttl_s = int(iam_payload["ibm_cloud_token"].get("expires_in", 3600))
else:
if not username:
raise ValueError("'username' is required in litellm_params when auth_mode='cp4d'")
@ -103,9 +152,9 @@ class WatsonxOrchestrateHandler:
headers={"Content-Type": "application/json"},
)
response.raise_for_status()
payload = response.json()
token = str(payload["token"])
expiration: Final = payload.get("expiration")
cp4d_payload: Final[_WXOView] = {"cp4d_token": response.json()}
token = str(cp4d_payload["cp4d_token"]["token"])
expiration: Final = cp4d_payload["cp4d_token"].get("expiration")
if expiration is None:
ttl_s = 3600
else:
@ -118,6 +167,16 @@ class WatsonxOrchestrateHandler:
del _token_cache[stale_key]
return token
@staticmethod
def _run_body(response: httpx.Response) -> _WXORun:
view: Final[_WXOView] = {"run": response.json()}
return view["run"]
@staticmethod
def _decode_run_event(payload: str | bytes) -> _WXORun:
view: Final[_WXOView] = {"run": json.loads(payload)}
return view["run"]
@staticmethod
async def _poll_run(
base_url: str,
@ -126,14 +185,14 @@ class WatsonxOrchestrateHandler:
client: AsyncHTTPHandler,
max_attempts: int = _MAX_POLL_ATTEMPTS,
interval_s: float = _POLL_INTERVAL_S,
) -> dict[str, Any]:
) -> _WXORun:
url: Final = f"{base_url}/v1/orchestrate/runs/{run_id}"
for attempt in range(max_attempts):
await asyncio.sleep(interval_s)
response = await client.get(url, headers=auth_headers)
response.raise_for_status()
result: dict[str, Any] = response.json()
result = WatsonxOrchestrateHandler._run_body(response)
status = result.get("status", "")
verbose_logger.debug("WXO: Poll %s/%s run='%s' status='%s'", attempt + 1, max_attempts, run_id, status)
if status in WatsonxOrchestrateTransformation.TERMINAL_STATES:
@ -145,11 +204,11 @@ class WatsonxOrchestrateHandler:
@staticmethod
async def _get_successful_run_data(
run_data: dict[str, Any],
run_data: _WXORun,
base_url: str,
auth_headers: dict[str, str],
client: AsyncHTTPHandler,
) -> dict[str, Any]:
) -> _WXORun:
status = run_data.get("status", "")
if status not in WatsonxOrchestrateTransformation.TERMINAL_STATES:
run_id: Final = run_data.get("run_id") or run_data.get("id") or ""
@ -170,15 +229,16 @@ class WatsonxOrchestrateHandler:
@staticmethod
async def _accumulate_wxo_sse_text(response: Any) -> str:
source: Final[_WXOView] = {"sse_source": response}
accumulated_text = ""
async for line in response.aiter_lines():
async for line in source["sse_source"].aiter_lines():
if not line.startswith("data:"):
continue
data_str = line[5:].strip()
if not data_str or data_str == "[DONE]":
continue
try:
event = json.loads(data_str)
event = WatsonxOrchestrateHandler._decode_run_event(data_str)
except json.JSONDecodeError:
continue
chunk_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(event)
@ -187,7 +247,7 @@ class WatsonxOrchestrateHandler:
return accumulated_text
@staticmethod
def _extract_litellm_params(litellm_params: dict[str, Any]) -> WXORequestParams:
def _extract_litellm_params(litellm_params: WXOLitellmParams) -> WXORequestParams:
cp4d_host: Final = litellm_params.get("cp4d_host") or ""
instance_id: Final = litellm_params.get("instance_id") or ""
wxo_agent_id: Final = litellm_params.get("wxo_agent_id") or ""
@ -215,9 +275,9 @@ class WatsonxOrchestrateHandler:
@staticmethod
async def handle_non_streaming(
request_id: str,
params: dict[str, Any],
litellm_params: dict[str, Any],
) -> dict[str, Any]:
params: dict[str, object],
litellm_params: WXOLitellmParams,
) -> dict[str, object]:
wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
client: Final = WatsonxOrchestrateHandler._http_client(timeout=90.0)
@ -246,7 +306,8 @@ class WatsonxOrchestrateHandler:
headers=auth_headers,
)
run_response.raise_for_status()
run_data: dict[str, Any] = run_response.json()
started: Final[_WXOView] = {"run": run_response.json()}
run_data: _WXORun = started["run"]
run_data = await WatsonxOrchestrateHandler._get_successful_run_data(
run_data=run_data,
@ -261,11 +322,11 @@ class WatsonxOrchestrateHandler:
@staticmethod
async def handle_streaming(
request_id: str,
params: dict[str, Any],
litellm_params: dict[str, Any],
params: dict[str, object],
litellm_params: WXOLitellmParams,
chunk_size: int = 50,
delay_ms: int = 10,
) -> AsyncIterator[dict[str, Any]]:
) -> AsyncIterator[dict[str, object]]:
wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
client: Final = WatsonxOrchestrateHandler._http_client(timeout=120.0)
@ -316,10 +377,11 @@ class WatsonxOrchestrateHandler:
yield chunk
return
content_type: Final = response.headers.get("content-type", "").lower()
header_view: Final[_WXOView] = {"content_type": response.headers.get("content-type", "")}
content_type: Final = header_view["content_type"].lower()
if "text/event-stream" not in content_type:
response_body: Final = await response.aread()
result = json.loads(response_body)
result = WatsonxOrchestrateHandler._decode_run_event(response_body)
result = await WatsonxOrchestrateHandler._get_successful_run_data(
run_data=result,
base_url=base_url,

View file

@ -1,5 +1,5 @@
import json
from collections.abc import Iterable, Iterator
from collections.abc import Iterable, Iterator, Mapping
from dataclasses import dataclass
from typing import Any, Final, Literal
@ -48,6 +48,7 @@ async def _handle_completed_batch(
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],
model_name: str | None = None,
litellm_params: dict | None = None,
model_info: ModelInfo | None = None,
) -> tuple[float, Usage, list[str]]:
"""Fetch a completed batch's output file and aggregate its cost, usage, and
models in a single pass over the JSONL lines, so the parsed file content is
@ -58,7 +59,21 @@ async def _handle_completed_batch(
custom_llm_provider: The LLM provider
model_name: Optional model name
litellm_params: Optional litellm parameters containing credentials (api_key, api_base, etc.)
model_info: Optional deployment-level model info with custom pricing,
threaded through so a deployment's configured rates win over the
global cost map.
"""
# A completed batch whose request lines all failed has no output file - the
# results are written to a separate error_file_id and output_file_id is None.
# There is nothing to price or measure, so report an empty result set instead
# of calling _fetch_batch_output_file_content, which raises on a missing
# output file. Without this guard the logging worker crashes on every
# aretrieve_batch poll and the completed batch's zero-cost accounting is lost.
# The generic retrieval helper keeps raising for callers that explicitly ask
# for a missing output file.
if batch.output_file_id is None:
return 0.0, Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0), []
file_content = await _fetch_batch_output_file_content(batch, custom_llm_provider, litellm_params=litellm_params)
if (
@ -72,9 +87,10 @@ async def _handle_completed_batch(
return batch_cost, batch_usage, [model_name]
return _aggregate_batch_cost_usage_models(
entries=_iter_batch_input_entries(file_content),
entries=_iter_batch_output_entries(file_content),
custom_llm_provider=custom_llm_provider,
model_name=model_name,
model_info=model_info,
)
@ -95,43 +111,91 @@ def _iter_successful_output_line_stats(
model_name: str | None,
model_info: ModelInfo | None,
) -> Iterator[_BatchOutputLineStats]:
for entry in entries:
stats = _safe_output_line_stats(entry, custom_llm_provider, model_name, model_info)
if stats is not None:
yield stats
def _safe_output_line_stats(
entry: Mapping[str, Any],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
model_name: str | None,
model_info: ModelInfo | None,
) -> _BatchOutputLineStats | None:
"""Return the stats for one batch output line, or None for a line that is
unsuccessful or cannot be costed, so a single bad line never aborts the
whole batch's cost accounting."""
custom_id: Final = entry.get("custom_id") if isinstance(entry, dict) else None
try:
if not _batch_response_was_successful(entry, custom_llm_provider):
return None
return _compute_output_line_stats(entry, custom_llm_provider, model_name, model_info)
except Exception as e: # noqa: BLE001 # any single line's costing failure must not abort the whole batch
verbose_logger.warning(
"batch output line could not be costed, so it is billed at $0 and the rest of the batch "
"is still billed. custom_id=%s error=%s",
custom_id,
str(e),
)
return None
def _compute_output_line_stats(
entry: Mapping[str, Any],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
model_name: str | None,
model_info: ModelInfo | None,
) -> _BatchOutputLineStats:
response_body: Final = _get_response_from_batch_job_output_file(entry, custom_llm_provider)
usage: Final = _get_batch_job_usage_from_response_body(response_body, custom_llm_provider)
prompt_details: Final = parse_prompt_tokens_details(usage)
raw_model: Final = response_body.get("model")
response_model: Final = raw_model if isinstance(raw_model, str) and raw_model else None
return _BatchOutputLineStats(
cost=_output_line_cost(
response_body=response_body,
usage=usage,
custom_llm_provider=custom_llm_provider,
model_name=model_name,
response_model=response_model,
model_info=model_info,
),
prompt_tokens=usage.prompt_tokens,
completion_tokens=usage.completion_tokens,
total_tokens=usage.total_tokens,
cache_read_tokens=prompt_details["cache_hit_tokens"],
cache_creation_tokens=prompt_details["cache_creation_tokens"],
model=response_model,
)
def _output_line_cost(
response_body: Mapping[str, Any],
usage: Usage,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
model_name: str | None,
response_model: str | None,
model_info: ModelInfo | None,
) -> float:
from litellm.cost_calculator import batch_cost_calculator
for entry in entries:
if not _batch_response_was_successful(entry, custom_llm_provider):
continue
response_body = _get_response_from_batch_job_output_file(entry, custom_llm_provider)
usage = _get_batch_job_usage_from_response_body(response_body, custom_llm_provider)
prompt_details = parse_prompt_tokens_details(usage)
raw_model = response_body.get("model")
response_model = raw_model if isinstance(raw_model, str) and raw_model else None
if model_info is not None or custom_llm_provider in ("anthropic", "bedrock"):
if custom_llm_provider == "bedrock" and model_name:
cost_model = model_name
else:
cost_model = response_model or model_name or ""
prompt_cost, completion_cost = batch_cost_calculator(
usage=usage,
model=cost_model,
custom_llm_provider=custom_llm_provider,
model_info=model_info,
)
line_cost = prompt_cost + completion_cost
else:
line_cost = litellm.completion_cost(
completion_response=response_body,
custom_llm_provider=custom_llm_provider,
call_type=CallTypes.aretrieve_batch.value,
)
yield _BatchOutputLineStats(
cost=line_cost,
prompt_tokens=usage.prompt_tokens,
completion_tokens=usage.completion_tokens,
total_tokens=usage.total_tokens,
cache_read_tokens=prompt_details["cache_hit_tokens"],
cache_creation_tokens=prompt_details["cache_creation_tokens"],
model=response_model,
if model_info is None and custom_llm_provider not in ("anthropic", "bedrock"):
return litellm.completion_cost(
completion_response=response_body,
custom_llm_provider=custom_llm_provider,
call_type=CallTypes.aretrieve_batch.value,
)
cost_model: Final = (
model_name if custom_llm_provider == "bedrock" and model_name else response_model or model_name or ""
)
prompt_cost, completion_cost = batch_cost_calculator(
usage=usage,
model=cost_model,
custom_llm_provider=custom_llm_provider,
model_info=model_info,
)
return prompt_cost + completion_cost
def _aggregate_batch_cost_usage_models(
@ -322,9 +386,10 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict:
def _get_file_content_as_dictionary(file_content: bytes) -> list[dict]:
"""
Get the file content as a list of dictionaries from JSON Lines format
Get the file content as a list of dictionaries from JSON Lines format,
skipping malformed lines
"""
return list(_iter_batch_input_entries(file_content))
return list(_iter_batch_output_entries(file_content))
def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
@ -345,15 +410,29 @@ def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
yield line
def _iter_batch_input_entries(file_content: bytes) -> Iterator[dict]:
def _iter_batch_output_entries(file_content: bytes) -> Iterator[dict]:
"""
Yield parsed batch input JSONL entries one at a time without materializing the
whole file as a list, so peak memory stays bounded. Raises on a malformed line;
callers that must survive bad rows should iterate ``_iter_batch_input_lines``
and parse per-row instead.
Yield parsed batch output JSONL entries one at a time without materializing
the whole file as a list, so peak memory stays bounded. A malformed or
non-object line is skipped with a warning so one bad line never aborts the
whole batch's cost accounting.
"""
for line in _iter_batch_input_lines(file_content):
yield json.loads(line)
entry = _parse_batch_output_line(line)
if entry is not None:
yield entry
def _parse_batch_output_line(line: bytes) -> dict | None:
try:
parsed: Final = json.loads(line)
except ValueError as e:
verbose_logger.warning("skipping malformed batch output line: %s", str(e))
return None
if isinstance(parsed, dict):
return parsed
verbose_logger.warning("skipping non-object batch output line of type %s", type(parsed).__name__)
return None
# A batch request's input tokens scale roughly with its serialized size, so this
@ -424,17 +503,31 @@ def _count_prompt_or_input_tokens(model: str, value: Any) -> int:
return 0
def _get_batch_job_usage_from_response_body(response_body: dict, custom_llm_provider: str = "openai") -> Usage:
def _get_batch_job_usage_from_response_body(
response_body: Mapping[str, Any], custom_llm_provider: str = "openai"
) -> Usage:
"""
Get the tokens of a batch job from the response body
"""
if custom_llm_provider in ("anthropic", "bedrock"):
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
return AnthropicConfig().calculate_usage(
usage_object=response_body.get("usage", None) or {},
usage_object: Final = response_body.get("usage", None) or {}
if custom_llm_provider == "bedrock" and AmazonConverseConfig.is_converse_usage_shape(usage_object):
return AmazonConverseConfig().usage_from_batch_output(usage_object)
anthropic_usage: Final = AnthropicConfig().calculate_usage(
usage_object=usage_object,
reasoning_content=None,
)
if usage_object and anthropic_usage.total_tokens == 0:
verbose_logger.warning(
"batch output line reported usage this parser does not understand, so it will be billed at $0. "
"provider=%s usage_keys=%s",
custom_llm_provider,
sorted(usage_object.keys()),
)
return anthropic_usage
from litellm.responses.utils import ResponseAPILoggingUtils
_usage_dict: Final = response_body.get("usage", None) or {}
@ -444,7 +537,7 @@ def _get_batch_job_usage_from_response_body(response_body: dict, custom_llm_prov
return usage
def _get_anthropic_result_from_batch_results_line(batch_results_line: dict) -> dict:
def _get_anthropic_result_from_batch_results_line(batch_results_line: Mapping[str, Any]) -> dict:
"""
Get the ``result`` object from a line of an Anthropic message batch results JSONL file.
@ -454,7 +547,9 @@ def _get_anthropic_result_from_batch_results_line(batch_results_line: dict) -> d
return batch_results_line.get("result", None) or {}
def _get_response_from_batch_job_output_file(batch_job_output_file: dict, custom_llm_provider: str = "openai") -> Any:
def _get_response_from_batch_job_output_file(
batch_job_output_file: Mapping[str, Any], custom_llm_provider: str = "openai"
) -> Any:
"""
Get the response from the batch job output file
"""
@ -467,7 +562,9 @@ def _get_response_from_batch_job_output_file(batch_job_output_file: dict, custom
return _response_body
def _batch_response_was_successful(batch_job_output_file: dict, custom_llm_provider: str = "openai") -> bool:
def _batch_response_was_successful(
batch_job_output_file: Mapping[str, Any], custom_llm_provider: str = "openai"
) -> bool:
"""
Check if the batch job response was successful

View file

@ -107,7 +107,7 @@ async def acreate_batch(
completion_window: Literal["24h"],
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"],
input_file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy"] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -157,7 +157,7 @@ def create_batch(
completion_window: Literal["24h"],
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"],
input_file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy"] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -339,7 +339,9 @@ def create_batch(
@client
async def aretrieve_batch(
batch_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai",
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic"
] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -385,7 +387,9 @@ def _handle_retrieve_batch_providers_without_provider_config(
litellm_params: dict,
_retrieve_batch_request: RetrieveBatchRequest,
_is_async: bool,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai",
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic"
] = "openai",
logging_obj: Any | None = None,
):
api_base: str | None = None
@ -508,7 +512,9 @@ def _handle_retrieve_batch_providers_without_provider_config(
@client
def retrieve_batch(
batch_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic"] = "openai",
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic"
] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -826,7 +832,7 @@ def list_batches(
async def acancel_batch(
batch_id: str,
model: str | None = None,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy"] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -872,7 +878,7 @@ async def acancel_batch(
def cancel_batch(
batch_id: str,
model: str | None = None,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] | str = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy"] | str = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -993,9 +999,14 @@ def cancel_batch(
timeout=timeout,
max_retries=optional_params.max_retries,
)
elif custom_llm_provider == "bedrock":
response = BedrockBatchesHandler.cancel_batch(
batch_id=batch_id,
**kwargs,
)
else:
raise litellm.exceptions.BadRequestError(
message=f"LiteLLM doesn't support {custom_llm_provider} for 'cancel_batch'. Only 'openai', 'azure', and 'vertex_ai' are supported.",
message=f"LiteLLM doesn't support {custom_llm_provider} for 'cancel_batch'. Only 'openai', 'azure', 'vertex_ai', and 'bedrock' are supported.",
model="n/a",
llm_provider=custom_llm_provider,
response=httpx.Response(

View file

@ -12,8 +12,11 @@ This module is dependency-injected: callers pass the proxy ``llm_router`` and
from __future__ import annotations
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any, Final
import litellm
if TYPE_CHECKING:
from litellm.router import Router
@ -41,3 +44,28 @@ def build_router_embedding_metadata(
metadata: Final[dict[str, Any]] = dict(request_metadata or {})
metadata["semantic-cache-embedding"] = True
return metadata
def resolve_embedding_max_input_tokens(
configured_max_input_tokens: int | None,
embedding_model: str,
router: Router | None,
) -> int | None:
"""Explicit cache setting first, else the Router deployment's configured ``max_input_tokens``."""
if configured_max_input_tokens is not None:
return configured_max_input_tokens
if router is None:
return None
deployment_max_input_tokens, _ = router.get_configured_token_limits(embedding_model)
return deployment_max_input_tokens
def truncate_embedding_input(prompt: str, embedding_model: str, max_input_tokens: int | None) -> str:
"""Keep only the first ``max_input_tokens`` tokens of ``prompt`` for the embedding call."""
if max_input_tokens is None:
return prompt
tokens: Final[Sequence[int]] = litellm.encode(model=embedding_model, text=prompt)
if len(tokens) <= max_input_tokens:
return prompt
truncated: Final[str] = litellm.decode(model=embedding_model, tokens=tokens[:max_input_tokens])
return truncated

View file

@ -97,6 +97,7 @@ class Cache:
qdrant_quantization_config: str | None = None,
qdrant_semantic_cache_embedding_model: str = "text-embedding-ada-002",
qdrant_semantic_cache_vector_size: int | None = None,
semantic_cache_embedding_max_input_tokens: int | None = None,
# GCP IAM authentication parameters
gcp_service_account: str | None = None,
gcp_ssl_ca_certs: str | None = None,
@ -122,6 +123,7 @@ class Cache:
qdrant_api_key (str, optional): The api_key for the local or cloud qdrant cluster.
qdrant_collection_name (str, optional): The name for your qdrant collection. Required if type is "qdrant-semantic".
similarity_threshold (float, optional): The similarity threshold for semantic-caching, Required if type is "redis-semantic" or "qdrant-semantic".
semantic_cache_embedding_max_input_tokens (int, optional): Truncate prompts to this many tokens before embedding them for semantic caching. Defaults to the embedding deployment's configured max_input_tokens.
# Disk Cache Args
disk_cache_dir (str, optional): The directory for the disk cache. Defaults to None.
@ -192,6 +194,7 @@ class Cache:
similarity_threshold=similarity_threshold,
embedding_model=redis_semantic_cache_embedding_model,
index_name=redis_semantic_cache_index_name,
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
**kwargs,
)
elif type == LiteLLMCacheType.VALKEY_SEMANTIC:
@ -207,6 +210,7 @@ class Cache:
embedding_model=valkey_semantic_cache_embedding_model,
index_name=valkey_semantic_cache_index_name,
startup_nodes=redis_startup_nodes,
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
**kwargs,
)
elif type == LiteLLMCacheType.QDRANT_SEMANTIC:
@ -218,6 +222,7 @@ class Cache:
quantization_config=qdrant_quantization_config,
embedding_model=qdrant_semantic_cache_embedding_model,
vector_size=qdrant_semantic_cache_vector_size,
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
)
elif type == LiteLLMCacheType.LOCAL:
self.cache = InMemoryCache()

View file

@ -18,7 +18,7 @@ import asyncio
import datetime
import inspect
import time
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Mapping
from typing import TYPE_CHECKING, Any, Final, Optional, TypeVar
from pydantic import BaseModel
@ -106,7 +106,7 @@ def _is_chat_completion_cached_dict(cached_result: dict) -> bool:
return "choices" in cached_result
def _should_defer_streaming_cache_hit_callbacks(*, kwargs: dict[str, Any]) -> bool:
def _should_defer_streaming_cache_hit_callbacks(*, kwargs: dict[str, object]) -> bool:
"""
When stream=True, do not run success callbacks at cache-hit time.
@ -119,11 +119,21 @@ def _should_defer_streaming_cache_hit_callbacks(*, kwargs: dict[str, Any]) -> bo
return kwargs.get("stream", False) is True
def _prompt_tokens_details_as_mapping(details: "PromptTokensDetailsWrapper") -> Mapping[str, object]:
"""Dump prompt token details to an opaque field mapping, tolerating non-pydantic stand-ins."""
return details.model_dump(exclude_none=True) if hasattr(details, "model_dump") else {}
def _request_cache_key(request_kwargs: Mapping[str, Any]) -> str | None:
"""Read the caller-supplied ``cache_key`` off the request kwargs."""
return request_kwargs.get("cache_key", None)
class LLMCachingHandler:
def __init__(
self,
original_function: Callable,
request_kwargs: dict[str, Any],
request_kwargs: dict[str, object],
start_time: datetime.datetime,
):
from litellm.caching import DualCache, RedisCache
@ -150,7 +160,7 @@ class LLMCachingHandler:
start_time: datetime.datetime,
call_type: str,
kwargs: dict[str, Any],
args: tuple[Any, ...] | None = None,
args: tuple[object, ...] | None = None,
) -> CachingHandlerResponse | None:
"""
Internal method to get from the cache.
@ -289,7 +299,7 @@ class LLMCachingHandler:
start_time: datetime.datetime,
call_type: str,
kwargs: dict[str, Any],
args: tuple[Any, ...] | None = None,
args: tuple[object, ...] | None = None,
) -> CachingHandlerResponse:
cached_result: Any | None = None
@ -366,7 +376,7 @@ class LLMCachingHandler:
return CachingHandlerResponse(cached_result=cached_result)
return CachingHandlerResponse(cached_result=cached_result)
def handle_kwargs_input_list_or_str(self, kwargs: dict[str, Any]) -> list[str]:
def handle_kwargs_input_list_or_str(self, kwargs: dict[str, object]) -> list[str]:
"""
Handles the input of kwargs['input'] being a list or a string
"""
@ -548,8 +558,8 @@ class LLMCachingHandler:
if details2 is None:
return details1
dict1: Final = details1.model_dump(exclude_none=True) if hasattr(details1, "model_dump") else {}
dict2: Final = details2.model_dump(exclude_none=True) if hasattr(details2, "model_dump") else {}
dict1: Final = _prompt_tokens_details_as_mapping(details1)
dict2: Final = _prompt_tokens_details_as_mapping(details2)
merged: Final[dict] = {}
for key in set(dict1.keys()) | set(dict2.keys()):
@ -671,7 +681,9 @@ class LLMCachingHandler:
cache_hit=cache_hit,
)
async def _retrieve_from_cache(self, call_type: str, kwargs: dict[str, Any], args: tuple[Any, ...]) -> Any | None:
async def _retrieve_from_cache(
self, call_type: str, kwargs: dict[str, object], args: tuple[object, ...]
) -> Any | None:
"""
Internal method to
- get cache key
@ -727,7 +739,8 @@ class LLMCachingHandler:
cached_result = None
else:
request_kwargs: Final = new_kwargs.copy()
request_cache_key: Final = request_kwargs.pop("cache_key", None)
request_cache_key: Final = _request_cache_key(request_kwargs)
request_kwargs.pop("cache_key", None)
if litellm.cache._supports_async() is True:
## check if dual cache is supported ##
self.preset_cache_key = request_cache_key or litellm.cache.get_cache_key(**request_kwargs)
@ -749,10 +762,10 @@ class LLMCachingHandler:
self,
cached_result: Any,
call_type: str,
kwargs: dict[str, Any],
kwargs: dict[str, object],
logging_obj: LiteLLMLoggingObj,
model: str,
args: tuple[Any, ...],
args: tuple[object, ...],
custom_llm_provider: str | None = None,
) -> (
ModelResponse
@ -948,7 +961,7 @@ class LLMCachingHandler:
result: Any,
original_function: Callable,
kwargs: dict[str, Any],
args: tuple[Any, ...] | None = None,
args: tuple[object, ...] | None = None,
):
"""
Internal method to check the type of the result & cache used and adds the result to the cache accordingly
@ -1013,8 +1026,8 @@ class LLMCachingHandler:
def sync_set_cache(
self,
result: Any,
kwargs: dict[str, Any],
args: tuple[Any, ...] | None = None,
kwargs: dict[str, object],
args: tuple[object, ...] | None = None,
):
"""
Sync internal method to add the result to the cache
@ -1204,8 +1217,8 @@ class LLMCachingHandler:
def convert_args_to_kwargs(
original_function: Callable,
args: tuple[Any, ...] | None = None,
) -> dict[str, Any]:
args: tuple[object, ...] | None = None,
) -> dict[str, object]:
# Get the signature of the original function
signature: Final = inspect.signature(original_function)

View file

@ -12,7 +12,7 @@ import ast
import asyncio
import json
import os
from typing import Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, cast
import litellm
from litellm._logging import print_verbose
@ -22,12 +22,21 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
)
from litellm.types.utils import EmbeddingResponse
from ._embedding_router import build_router_embedding_metadata, resolve_embedding_router
from ._embedding_router import (
build_router_embedding_metadata,
resolve_embedding_max_input_tokens,
resolve_embedding_router,
truncate_embedding_input,
)
from .base_cache import BaseCache
if TYPE_CHECKING:
from litellm.router import Router
class QdrantSemanticCache(BaseCache):
CACHE_KEY_FIELD_NAME = "litellm_cache_key"
embedding_max_input_tokens: int | None = None
def __init__(
self,
@ -39,6 +48,7 @@ class QdrantSemanticCache(BaseCache):
embedding_model="text-embedding-ada-002",
host_type=None,
vector_size=None,
embedding_max_input_tokens: int | None = None,
):
from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
@ -57,6 +67,7 @@ class QdrantSemanticCache(BaseCache):
raise Exception("similarity_threshold must be provided, passed None")
self.similarity_threshold = similarity_threshold
self.embedding_model = embedding_model
self.embedding_max_input_tokens = embedding_max_input_tokens
self.vector_size = vector_size if vector_size is not None else QDRANT_VECTOR_SIZE
headers = {}
@ -188,6 +199,13 @@ class QdrantSemanticCache(BaseCache):
cached_key: Final = payload.get(self.CACHE_KEY_FIELD_NAME)
return cached_key is not None and str(cached_key) == str(key)
def _embedding_input(self, prompt: str, router: "Router | None") -> str:
return truncate_embedding_input(
prompt,
self.embedding_model,
resolve_embedding_max_input_tokens(self.embedding_max_input_tokens, self.embedding_model, router),
)
def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse:
"""Embed via the proxy Router when it serves the model, else direct."""
try:
@ -197,16 +215,17 @@ class QdrantSemanticCache(BaseCache):
llm_router = None
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
embedding_input: Final = self._embedding_input(prompt, router)
if router is not None:
return router.embedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
metadata=build_router_embedding_metadata(metadata),
)
return litellm.embedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
)
@ -218,17 +237,18 @@ class QdrantSemanticCache(BaseCache):
llm_router = None
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
embedding_input: Final = self._embedding_input(prompt, router)
if router is not None:
return await router.aembedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
metadata=build_router_embedding_metadata(metadata),
)
return await litellm.aembedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
)

View file

@ -49,7 +49,7 @@ if TYPE_CHECKING:
cluster_pipeline = ClusterPipeline
async_redis_client = Redis
async_redis_cluster_client = RedisCluster
Span = _Span | Any
Span = _Span
else:
pipeline = Any
cluster_pipeline = Any
@ -625,7 +625,11 @@ class RedisCache(BaseCache):
f"{self.namespace}-{hashlib.sha256(script.encode()).hexdigest()[:16]}"
)
async def run_script(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any:
async def run_script(
keys: Sequence[str],
args: Sequence[str | bytes | int | float],
client: object = None,
) -> object:
async def execute() -> object:
executor: Callable[..., Awaitable[Any]] | None = litellm.in_memory_llm_clients_cache.get_cache(
key=script_cache_key
@ -650,7 +654,11 @@ class RedisCache(BaseCache):
if hasattr(_redis_client, "register_script"):
registered_script: Final = _redis_client.register_script(script)
async def standalone_executor(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any:
async def standalone_executor(
keys: Sequence[str],
args: Sequence[str | bytes | int | float],
client: object = None,
) -> object:
namespaced_keys: Final = tuple(self.check_and_fix_namespace(key=key) for key in keys)
return await registered_script(keys=namespaced_keys, args=args, client=client)
@ -659,7 +667,11 @@ class RedisCache(BaseCache):
if hasattr(_redis_client, "script_load"):
script_sha: Final = _redis_client.script_load(script)
async def cluster_executor(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any:
async def cluster_executor(
keys: Sequence[str],
args: Sequence[str | bytes | int | float],
client: object = None,
) -> object:
namespaced_keys: Final = tuple(self.check_and_fix_namespace(key=key) for key in keys)
return await _redis_client.evalsha(script_sha, len(namespaced_keys), *namespaced_keys, *args)
@ -757,7 +769,7 @@ class RedisCache(BaseCache):
async def _pipeline_helper(
self,
pipe: pipeline | cluster_pipeline,
cache_list: list[tuple[Any, Any]],
cache_list: Sequence[tuple[str, object]],
ttl: float | None,
) -> list:
"""
@ -783,7 +795,9 @@ class RedisCache(BaseCache):
return results
@_redis_circuit_breaker_guard
async def async_set_cache_pipeline(self, cache_list: list[tuple[Any, Any]], ttl: float | None = None, **kwargs):
async def async_set_cache_pipeline(
self, cache_list: Sequence[tuple[str, object]], ttl: float | None = None, **kwargs
):
"""
Use Redis Pipelines for bulk write operations
"""
@ -795,7 +809,7 @@ class RedisCache(BaseCache):
start_time: Final = time.time()
print_verbose(f"Set Async Redis Cache: key list: {cache_list}\nttl={ttl}, redis_version={self.redis_version}")
cache_value: Final[Any] = None
cache_value: Final = None
try:
async with _redis_client.pipeline(transaction=False) as pipe:
results: Final = await self._pipeline_helper(pipe, cache_list, ttl)
@ -1074,7 +1088,7 @@ class RedisCache(BaseCache):
# NON blocking - notify users Redis is throwing an exception
verbose_logger.error("litellm.caching.caching: get() - Got exception from REDIS: ", e)
def _run_redis_mget_operation(self, keys: list[str]) -> list[Any]:
def _run_redis_mget_operation(self, keys: list[str]) -> Sequence[bytes | str | None]:
"""
Wrapper to call `mget` on the redis client
@ -1082,7 +1096,7 @@ class RedisCache(BaseCache):
"""
return self.redis_client.mget(keys=keys)
async def _async_run_redis_mget_operation(self, keys: list[str]) -> list[Any]:
async def _async_run_redis_mget_operation(self, keys: list[str]) -> Sequence[bytes | str | None]:
"""
Wrapper to call `mget` on the redis client
@ -1115,7 +1129,7 @@ class RedisCache(BaseCache):
cache_key = self.check_and_fix_namespace(key=cache_key or "")
_keys.append(cache_key)
start_time: Final = time.time()
results: Final[list] = self._run_redis_mget_operation(keys=_keys)
results: Final = self._run_redis_mget_operation(keys=_keys)
end_time: Final = time.time()
_duration: Final = end_time - start_time
self.service_logger_obj.service_success_hook(
@ -1522,7 +1536,7 @@ class RedisCache(BaseCache):
async def async_rpush(
self,
key: str,
values: list[Any],
values: Sequence[str | bytes | int | float],
parent_otel_span: Span | None = None,
**kwargs,
) -> int:

View file

@ -14,7 +14,7 @@ import asyncio
import json
import os
from collections.abc import Callable, Mapping
from typing import Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, cast
import litellm
from litellm._logging import print_verbose, verbose_logger
@ -23,9 +23,17 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
)
from litellm.types.utils import EmbeddingResponse
from ._embedding_router import build_router_embedding_metadata, resolve_embedding_router
from ._embedding_router import (
build_router_embedding_metadata,
resolve_embedding_max_input_tokens,
resolve_embedding_router,
truncate_embedding_input,
)
from .base_cache import BaseCache
if TYPE_CHECKING:
from litellm.router import Router
class RedisSemanticCache(BaseCache):
"""
@ -38,6 +46,7 @@ class RedisSemanticCache(BaseCache):
DEFAULT_REDIS_INDEX_NAME: str = "litellm_semantic_cache_index"
CACHE_KEY_FIELD_NAME: str = "litellm_cache_key"
embedding_max_input_tokens: int | None = None
def __init__(
self,
@ -48,6 +57,7 @@ class RedisSemanticCache(BaseCache):
similarity_threshold: float | None = None,
embedding_model: str = "text-embedding-ada-002",
index_name: str | None = None,
embedding_max_input_tokens: int | None = None,
**kwargs: object,
):
"""
@ -62,6 +72,8 @@ class RedisSemanticCache(BaseCache):
where 1.0 requires exact matches and 0.0 accepts any match
embedding_model: Model to use for generating embeddings
index_name: Name for the Redis index
embedding_max_input_tokens: Truncate prompts to this many tokens before
embedding; defaults to the Router deployment's configured max_input_tokens
ttl: Default time-to-live for cache entries in seconds
**kwargs: Additional arguments passed to the Redis client
@ -86,6 +98,7 @@ class RedisSemanticCache(BaseCache):
# While similarity: 1 = most similar, 0 = least similar
self.distance_threshold = 1 - similarity_threshold
self.embedding_model = embedding_model
self.embedding_max_input_tokens = embedding_max_input_tokens
# Set up Redis connection
if redis_url is None:
@ -307,6 +320,13 @@ class RedisSemanticCache(BaseCache):
return dict_method()
return value
def _embedding_input(self, prompt: str, router: "Router | None") -> str:
return truncate_embedding_input(
prompt,
self.embedding_model,
resolve_embedding_max_input_tokens(self.embedding_max_input_tokens, self.embedding_model, router),
)
def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> list[float]:
"""
Routes through the proxy Router when the embedding model is a Router
@ -320,12 +340,13 @@ class RedisSemanticCache(BaseCache):
llm_router = None
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
embedding_input: Final = self._embedding_input(prompt, router)
if router is not None:
embedding_response = cast(
EmbeddingResponse,
router.embedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
metadata=build_router_embedding_metadata(metadata),
),
@ -335,7 +356,7 @@ class RedisSemanticCache(BaseCache):
EmbeddingResponse,
litellm.embedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
),
)
@ -490,18 +511,19 @@ class RedisSemanticCache(BaseCache):
llm_router = None
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
embedding_input: Final = self._embedding_input(prompt, router)
try:
if router is not None:
embedding_response = await router.aembedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
metadata=build_router_embedding_metadata(metadata),
)
else:
embedding_response = await litellm.aembedding(
model=self.embedding_model,
input=prompt,
input=embedding_input,
cache={"no-store": True, "no-cache": True},
)
return embedding_response["data"][0]["embedding"]

View file

@ -17,7 +17,6 @@ RedisSemanticCache since those are backend agnostic.
import asyncio
import hashlib
import os
import struct
from dataclasses import dataclass
from typing import Any, Final
@ -29,6 +28,7 @@ from redis.commands.search.query import Query
from litellm._logging import print_verbose
from litellm._uuid import uuid
from litellm.llms.valkey.common_utils import build_valkey_url, pack_vector
from .redis_semantic_cache import RedisSemanticCache
@ -61,6 +61,7 @@ class ValkeySemanticCache(RedisSemanticCache):
startup_nodes: list | None = None,
sync_client: Redis | None = None,
async_client: AsyncRedis | None = None,
embedding_max_input_tokens: int | None = None,
**kwargs: Any,
):
if similarity_threshold is None:
@ -78,6 +79,7 @@ class ValkeySemanticCache(RedisSemanticCache):
self.similarity_threshold = similarity_threshold
self.embedding_model = embedding_model
self.embedding_max_input_tokens = embedding_max_input_tokens
self.index_name = index_name or self.DEFAULT_VALKEY_INDEX_NAME
self.key_prefix = f"{self.index_name}:"
self._index_dim: int | None = None
@ -92,19 +94,17 @@ class ValkeySemanticCache(RedisSemanticCache):
@staticmethod
def _build_valkey_url(host: str | None, port: str | None, password: str | None, ssl: bool = False) -> str:
host = host or os.environ.get("VALKEY_HOST") or os.environ.get("REDIS_HOST")
port = port or os.environ.get("VALKEY_PORT") or os.environ.get("REDIS_PORT")
password = password or os.environ.get("VALKEY_PASSWORD") or os.environ.get("REDIS_PASSWORD")
resolved_host: Final = host or os.environ.get("VALKEY_HOST") or os.environ.get("REDIS_HOST")
resolved_port: Final = port or os.environ.get("VALKEY_PORT") or os.environ.get("REDIS_PORT")
resolved_password: Final = password or os.environ.get("VALKEY_PASSWORD") or os.environ.get("REDIS_PASSWORD")
if not host or not port:
if not resolved_host or not resolved_port:
raise ValueError(
"Missing required Valkey configuration. Provide host and port "
"(or VALKEY_HOST/VALKEY_PORT), or pass redis_url."
)
credentials: Final = f":{password}@" if password else ""
scheme: Final = "rediss" if ssl else "redis"
return f"{scheme}://{credentials}{host}:{port}"
return build_valkey_url(host=resolved_host, port=resolved_port, password=resolved_password, ssl=ssl)
@classmethod
def _scope_tag(cls, key: str) -> str:
@ -116,7 +116,7 @@ class ValkeySemanticCache(RedisSemanticCache):
@staticmethod
def _embedding_to_bytes(embedding: list[float]) -> bytes:
return struct.pack(f"<{len(embedding)}f", *embedding)
return pack_vector(embedding)
def _index_schema(self, dim: int) -> tuple[TagField, VectorField]:
return (

View file

@ -185,6 +185,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
if not isinstance(tool_choice, dict):
return tool_choice
choice_type: Final = tool_choice.get("type")
if isinstance(choice_type, str) and choice_type in ("auto", "none", "required"):
return choice_type
if choice_type not in ("function", "custom"):
return tool_choice
if isinstance(tool_choice.get("name"), str) and tool_choice.get("name"):

View file

@ -1,5 +1,6 @@
import os
import sys
from types import MappingProxyType
from typing import Final, Literal
from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_or_none
@ -242,6 +243,12 @@ AIOHTTP_NEEDS_CLEANUP_CLOSED: Final = (3, 13, 0) <= sys.version_info < (
# https://github.com/openai/openai-agents-python/blob/cf1b933660e44fd37b4350c41febab8221801409/src/agents/realtime/openai_realtime.py#L235
_max_size_env: Final = os.getenv("REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES")
REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES: Final = int(_max_size_env) if _max_size_env is not None else None
REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS: Final = float(
os.getenv("REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS", "20.0")
)
# RFC 6455 caps the close frame payload at 125 bytes, 2 of which carry the status code
WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123
# SSL/TLS cipher configuration for faster handshakes
# Strategy: Strongly prefer fast modern ciphers, but allow fallback to commonly supported ones
@ -1487,6 +1494,7 @@ WEEKLY_SPEND_REPORT_JOB_ID: Final = "weekly_spend_report_job"
MONTHLY_SPEND_REPORT_JOB_ID: Final = "monthly_spend_report_job"
PROMETHEUS_FALLBACK_STATS_JOB_ID: Final = "prometheus_fallback_stats_job"
SLACK_DAILY_REPORT_LOCK_ID: Final = "slack_daily_report"
SLACK_MODEL_DEPRECATION_LOCK_ID: Final = "slack_model_deprecation_warning"
SPEND_LOG_RUN_LOOPS: Final = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500))
SPEND_LOG_CLEANUP_BATCH_SIZE: Final = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000))
SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(os.getenv("SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3))
@ -1593,6 +1601,13 @@ DEFAULT_MCP_ACCESS_GROUP_NEGATIVE_CACHE_TTL: Final = 10
# in a single ``/{name1,name2,...}/mcp`` URL. Bounds the per-request DB / cache
# fan-out an authenticated caller can trigger by stuffing the path with tokens.
DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS: Final = 16
# Ceilings on the cached auth registries; larger tables fall back to per-row lookups
# instead of holding an unbounded id set in every worker.
TAG_REGISTRY_MAX_SIZE: Final = 5000
END_USER_RESTRICTED_REGISTRY_MAX_SIZE: Final = 5000
# How long a failed registry load is remembered as "unusable", so a degraded Postgres
# is not re-scanned on every request on top of the per-id lookups it falls back to.
REGISTRY_ERROR_NEGATIVE_CACHE_TTL: Final = 30
# Sentry Scrubbing Configuration
SENTRY_DENYLIST: Final = [
@ -1756,3 +1771,18 @@ PTU_LAPSED_ALERT_LIMIT: Final[int] = 10
# one run delete a charge another just wrote. A stale row is hours old and a concurrent
# one is seconds old, so a few minutes separates them.
PTU_PRUNE_SKEW_GRACE_SECONDS: Final[int] = 300
# How long enqueued-token reservations for batches live without a refund. Providers
# complete or expire batches within their completion window (24h for OpenAI), so a
# reservation still unrefunded after 8 days belongs to a batch whose terminal state
# was never observed (e.g. proxy restart); expiry returns the tokens to the caller.
BATCH_ENQUEUED_TOKEN_TTL_SECONDS: Final[int] = 8 * 24 * 60 * 60
# Key/team metadata field that opts batches into enqueued-token limiting. Only proxy
# admins may write it: when present it replaces the standard RPM/TPM checks for
# batch submissions.
BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY: Final = "batch_enqueued_token_limit"
# Shared read-only empty mapping, for defaulting optional Mapping parameters without
# constructing a fresh mutable dict at each call site.
EMPTY_MAPPING: Final = MappingProxyType({})

View file

@ -1,9 +1,11 @@
import asyncio
import contextvars
import json
from collections.abc import Coroutine
from collections.abc import Coroutine, Mapping
from functools import partial
from typing import Any, Final, Literal, overload
from typing import Final, Literal, overload
import httpx
import litellm
from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT
@ -48,16 +50,16 @@ __all__ = [
@client
async def acreate_container(
name: str,
expires_after: dict[str, Any] | None = None,
expires_after: Mapping[str, object] | None = None,
file_ids: list[str] | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
# LiteLLM specific params,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerObject:
"""Asynchronously calls the `create_container` function with the given arguments and keyword arguments.
@ -120,9 +122,9 @@ async def acreate_container(
@overload
def create_container(
name: str,
expires_after: dict[str, Any] | None = None,
expires_after: Mapping[str, object] | None = None,
file_ids: list[str] | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -130,16 +132,16 @@ def create_container(
*,
acreate_container: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, ContainerObject]:
) -> Coroutine[object, object, ContainerObject]:
...
@overload
def create_container(
name: str,
expires_after: dict[str, Any] | None = None,
expires_after: Mapping[str, object] | None = None,
file_ids: list[str] | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -156,20 +158,20 @@ def create_container(
@client
def create_container(
name: str,
expires_after: dict[str, Any] | None = None,
expires_after: Mapping[str, object] | None = None,
file_ids: list[str] | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerObject | Coroutine[Any, Any, ContainerObject]:
) -> ContainerObject | Coroutine[object, object, ContainerObject]:
"""Create a container using the OpenAI Container API.
Currently supports OpenAI
@ -281,13 +283,13 @@ async def alist_containers(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerListResponse:
"""Asynchronously list containers.
@ -351,7 +353,7 @@ def list_containers(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -359,7 +361,7 @@ def list_containers(
*,
alist_containers: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, ContainerListResponse]:
) -> Coroutine[object, object, ContainerListResponse]:
...
@ -368,7 +370,7 @@ def list_containers(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -387,18 +389,18 @@ def list_containers(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerListResponse | Coroutine[Any, Any, ContainerListResponse]:
) -> ContainerListResponse | Coroutine[object, object, ContainerListResponse]:
"""List containers using the OpenAI Container API.
Currently supports OpenAI
@ -481,13 +483,13 @@ def list_containers(
@client
async def aretrieve_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerObject:
"""Asynchronously retrieve a container.
@ -545,7 +547,7 @@ async def aretrieve_container(
@overload
def retrieve_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -553,14 +555,14 @@ def retrieve_container(
*,
aretrieve_container: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, ContainerObject]:
) -> Coroutine[object, object, ContainerObject]:
...
@overload
def retrieve_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -577,18 +579,18 @@ def retrieve_container(
@client
def retrieve_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerObject | Coroutine[Any, Any, ContainerObject]:
) -> ContainerObject | Coroutine[object, object, ContainerObject]:
"""Retrieve a container using the OpenAI Container API.
Currently supports OpenAI
@ -696,13 +698,13 @@ def retrieve_container(
@client
async def adelete_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> DeleteContainerResult:
"""Asynchronously delete a container.
@ -760,7 +762,7 @@ async def adelete_container(
@overload
def delete_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -768,14 +770,14 @@ def delete_container(
*,
adelete_container: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, DeleteContainerResult]:
) -> Coroutine[object, object, DeleteContainerResult]:
...
@overload
def delete_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -792,18 +794,18 @@ def delete_container(
@client
def delete_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> DeleteContainerResult | Coroutine[Any, Any, DeleteContainerResult]:
) -> DeleteContainerResult | Coroutine[object, object, DeleteContainerResult]:
"""Delete a container using the OpenAI Container API.
Currently supports OpenAI
@ -914,11 +916,11 @@ async def alist_container_files(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerFileListResponse:
"""Asynchronously list files in a container.
@ -985,7 +987,7 @@ def list_container_files(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600,
timeout: float | httpx.Timeout = 600,
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -993,7 +995,7 @@ def list_container_files(
*,
alist_container_files: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, ContainerFileListResponse]:
) -> Coroutine[object, object, ContainerFileListResponse]:
...
@ -1003,7 +1005,7 @@ def list_container_files(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600,
timeout: float | httpx.Timeout = 600,
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -1023,16 +1025,16 @@ def list_container_files(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerFileListResponse | Coroutine[Any, Any, ContainerFileListResponse]:
) -> ContainerFileListResponse | Coroutine[object, object, ContainerFileListResponse]:
"""List files in a container using the OpenAI Container API.
Currently supports OpenAI
@ -1125,11 +1127,11 @@ def list_container_files(
async def aupload_container_file(
container_id: str,
file: FileTypes,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerFileObject:
"""Asynchronously upload a file to a container.
@ -1211,7 +1213,7 @@ async def aupload_container_file(
def upload_container_file(
container_id: str,
file: FileTypes,
timeout=600,
timeout: float | httpx.Timeout = 600,
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -1219,7 +1221,7 @@ def upload_container_file(
*,
aupload_container_file: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, ContainerFileObject]:
) -> Coroutine[object, object, ContainerFileObject]:
...
@ -1227,7 +1229,7 @@ def upload_container_file(
def upload_container_file(
container_id: str,
file: FileTypes,
timeout=600,
timeout: float | httpx.Timeout = 600,
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -1245,16 +1247,16 @@ def upload_container_file(
def upload_container_file(
container_id: str,
file: FileTypes,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerFileObject | Coroutine[Any, Any, ContainerFileObject]:
) -> ContainerFileObject | Coroutine[object, object, ContainerFileObject]:
"""Upload a file to a container using the OpenAI Container API.
This endpoint allows uploading files directly to a container session,

View file

@ -102,6 +102,7 @@ from litellm.types.utils import (
LlmProviders,
LlmProvidersSet,
ModelInfo,
PromptTokensDetailsWrapper,
ServiceTier,
StandardBuiltInToolsParams,
TranscriptionUsageDurationObject,
@ -286,7 +287,7 @@ def _transcription_usage_has_token_details(
prompt_tokens_val: Final = getattr(usage_block, "prompt_tokens", 0) or 0
completion_tokens_val: Final = getattr(usage_block, "completion_tokens", 0) or 0
prompt_details: Final = getattr(usage_block, "prompt_tokens_details", None)
prompt_details: Final[PromptTokensDetailsWrapper | None] = getattr(usage_block, "prompt_tokens_details", None)
if prompt_details is not None:
audio_token_count: Final = getattr(prompt_details, "audio_tokens", 0) or 0
@ -326,6 +327,8 @@ def cost_per_token(
service_tier: str | None = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: str | None = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
### VERTEX LOCATION ###
vertex_location: str | None = None, # for Vertex AI regional-endpoint uplift (e.g. "us-east5", "global")
response: Any | None = None,
### REQUEST MODEL ###
request_model: str | None = None, # original request model for router detection
@ -375,7 +378,7 @@ def cost_per_token(
_is_anthropic_style = False
if usage_object is not None:
_pt_details: Final = getattr(usage_object, "prompt_tokens_details", None)
_pt_details: Final[PromptTokensDetailsWrapper | None] = getattr(usage_object, "prompt_tokens_details", None)
if _pt_details is not None:
_cache_read_tokens = float(getattr(_pt_details, "cached_tokens", 0) or 0)
# OpenAI-compatible providers report cache-write tokens under
@ -385,8 +388,8 @@ def cost_per_token(
getattr(_pt_details, "cache_write_tokens", 0) or getattr(_pt_details, "cache_creation_tokens", 0) or 0
)
_anthropic_read: Final = getattr(usage_object, "cache_read_input_tokens", None)
_anthropic_create: Final = getattr(usage_object, "cache_creation_input_tokens", None)
_anthropic_read: Final[int | None] = getattr(usage_object, "cache_read_input_tokens", None)
_anthropic_create: Final[int | None] = getattr(usage_object, "cache_creation_input_tokens", None)
if _anthropic_read is not None or _anthropic_create is not None:
_is_anthropic_style = True
if _anthropic_read is not None:
@ -586,6 +589,7 @@ def cost_per_token(
prompt_characters=prompt_characters,
completion_characters=completion_characters,
usage=usage_block,
vertex_location=vertex_location,
)
elif cost_router == "cost_per_token":
return google_cost_per_token(
@ -593,6 +597,7 @@ def cost_per_token(
custom_llm_provider=custom_llm_provider,
usage=usage_block,
service_tier=service_tier,
vertex_location=vertex_location,
)
elif custom_llm_provider == "anthropic":
return anthropic_cost_per_token(model=model, usage=usage_block, service_tier=service_tier)
@ -703,7 +708,7 @@ def get_replicate_completion_pricing(completion_response: dict, total_time=0.0):
return a100_80gb_price_per_second_public * total_time / 1000
def has_hidden_params(obj: Any) -> bool:
def has_hidden_params(obj: object) -> bool:
return hasattr(obj, "_hidden_params")
@ -728,7 +733,7 @@ def _get_provider_for_cost_calc(
def _select_model_name_for_cost_calc(
model: str | None,
completion_response: Any | None,
completion_response: object | None,
base_model: str | None = None,
custom_pricing: bool | None = None,
custom_llm_provider: str | None = None,
@ -804,7 +809,7 @@ def _model_contains_known_llm_provider(model: str) -> bool:
return _provider_prefix in LlmProvidersSet
def _get_response_model(completion_response: Any) -> str | None:
def _get_response_model(completion_response: object) -> str | None:
"""
Extract the model name from a completion response object.
@ -866,8 +871,18 @@ def _normalize_service_tier(service_tier: object) -> str | None:
return service_tier
def _extract_service_tier(source: object) -> str | None:
"""Read a raw ``service_tier`` off a response body or usage object, dict or pydantic model alike."""
if isinstance(source, BaseModel):
return getattr(source, "service_tier", None)
elif isinstance(source, dict):
return source.get("service_tier")
return None
def _get_usage_object(
completion_response: Any,
completion_response: object,
) -> Usage | None:
usage_obj: Final = cast(
Usage | ResponseAPIUsage | dict | BaseModel,
@ -1060,6 +1075,7 @@ def _store_cost_breakdown_in_logging_obj(
reasoning_cost: float | None = None,
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
) -> None:
"""
Helper function to store cost breakdown in the logging object.
@ -1079,6 +1095,7 @@ def _store_cost_breakdown_in_logging_obj(
margin_total_amount: Total margin added in USD
service_tier: Tier the costs above were priced on, already resolved
data_residency: Region uplift the costs above were priced on, already resolved
vertex_location: Vertex AI location the costs above were priced on, already resolved
"""
if litellm_logging_obj is None:
return
@ -1102,6 +1119,7 @@ def _store_cost_breakdown_in_logging_obj(
reasoning_cost=reasoning_cost,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
)
except Exception as breakdown_error:
@ -1110,7 +1128,7 @@ def _store_cost_breakdown_in_logging_obj(
def completion_cost(
completion_response=None,
completion_response: object | None = None,
model: str | None = None,
prompt="",
messages: list = [],
@ -1138,6 +1156,8 @@ def completion_cost(
service_tier: str | None = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: str | None = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
### VERTEX LOCATION ###
vertex_location: str | None = None, # for Vertex AI regional-endpoint uplift (e.g. "us-east5", "global")
) -> float:
"""
Calculate the cost of a given completion call fot GPT-3.5-turbo, llama2, any litellm supported llm.
@ -1197,19 +1217,13 @@ def completion_cost(
# Extract service_tier from completion_response if not provided
if service_tier is None and completion_response is not None:
if isinstance(completion_response, BaseModel):
service_tier = getattr(completion_response, "service_tier", None)
elif isinstance(completion_response, dict):
service_tier = completion_response.get("service_tier")
service_tier = _extract_service_tier(completion_response)
service_tier = _normalize_service_tier(service_tier)
# Extract service_tier from usage object if not provided
if service_tier is None and cost_per_token_usage_object is not None:
if isinstance(cost_per_token_usage_object, BaseModel):
service_tier = getattr(cost_per_token_usage_object, "service_tier", None)
elif isinstance(cost_per_token_usage_object, dict):
service_tier = cost_per_token_usage_object.get("service_tier")
service_tier = _extract_service_tier(cost_per_token_usage_object)
service_tier = _normalize_service_tier(service_tier)
@ -1412,7 +1426,7 @@ def completion_cost(
if completion_response is not None and isinstance(completion_response, RerankResponse):
meta_obj = completion_response.meta
if meta_obj is not None:
billed_units = meta_obj.get("billed_units", {}) or {}
billed_units: RerankBilledUnits = meta_obj.get("billed_units") or {}
else:
billed_units = {}
@ -1572,6 +1586,7 @@ def completion_cost(
rerank_billed_units=rerank_billed_units,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
response=completion_response,
request_model=request_model_for_cost,
)
@ -1659,6 +1674,7 @@ def completion_cost(
usage=cost_per_token_usage_object,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
)
_reasoning_cost = _token_type_breakdown.reasoning_cost
_cache_read_cost = _token_type_breakdown.cache_read_cost
@ -1681,6 +1697,7 @@ def completion_cost(
reasoning_cost=_reasoning_cost,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
)
return _final_cost
@ -1760,6 +1777,8 @@ def response_cost_calculator(
service_tier: str | None = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: str | None = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
### VERTEX LOCATION ###
vertex_location: str | None = None, # for Vertex AI regional-endpoint uplift (e.g. "us-east5", "global")
) -> float:
"""
Returns
@ -1792,6 +1811,7 @@ def response_cost_calculator(
litellm_logging_obj=litellm_logging_obj,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
)
return response_cost
except Exception as e:
@ -1801,7 +1821,7 @@ def response_cost_calculator(
def ocr_cost(
model: str,
custom_llm_provider: str | None,
response: Any | None = None,
response: object | None = None,
) -> tuple[float, float]:
"""
Args:
@ -2160,7 +2180,7 @@ def batch_cost_calculator(
output_cost_per_token: Final = model_info.get("output_cost_per_token")
total_prompt_cost = 0.0
total_completion_cost = 0.0
if input_cost_per_token_batches:
if input_cost_per_token_batches is not None:
total_prompt_cost = usage.prompt_tokens * input_cost_per_token_batches
elif input_cost_per_token:
details: Final = parse_prompt_tokens_details(usage)
@ -2180,7 +2200,7 @@ def batch_cost_calculator(
cache_creation_cost: Final = model_info.get("cache_creation_input_token_cost") or input_cost_per_token
total_prompt_cost += cache_creation_tokens * cache_creation_cost / 2
if output_cost_per_token_batches:
if output_cost_per_token_batches is not None:
total_completion_cost = usage.completion_tokens * output_cost_per_token_batches
elif output_cost_per_token:
total_completion_cost = (

View file

@ -23,12 +23,15 @@ FileCreateProvider = Literal[
"vertex_ai",
"bedrock",
"hosted_vllm",
"litellm_proxy",
"manus",
"anthropic",
]
FileRetrieveProvider = Literal["openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "manus", "anthropic"]
FileDeleteProvider = Literal["openai", "azure", "gemini", "manus", "anthropic"]
FileListProvider = Literal["openai", "azure", "manus", "anthropic"]
FileRetrieveProvider = Literal[
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic"
]
FileDeleteProvider = Literal["openai", "azure", "gemini", "litellm_proxy", "manus", "anthropic"]
FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic"]
import litellm
from litellm import get_secret_str
from litellm.files.streaming import FileContentStreamingResponse

View file

@ -1,7 +1,9 @@
from collections.abc import AsyncIterator, Iterator
from typing import Literal, NamedTuple
FileContentProvider = Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic", "manus"]
FileContentProvider = Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "manus"
]
class FileContentStreamingResult(NamedTuple):

View file

@ -1,5 +1,5 @@
import json
from collections.abc import AsyncIterator, Iterator
from collections.abc import AsyncIterator, Iterator, Sequence
from typing import Any, Final, TypedDict, cast
from typing_extensions import ReadOnly
@ -27,6 +27,7 @@ from litellm.types.utils import (
ModelResponse,
ModelResponseStream,
StreamingChoices,
Usage,
)
@ -43,6 +44,29 @@ class _GenAIPart(TypedDict, total=False):
functionCall: ReadOnly[dict[str, object]]
class _GenAIFunctionDeclaration(TypedDict, total=False):
name: ReadOnly[str]
description: ReadOnly[str]
parametersJsonSchema: ReadOnly[dict[str, object]]
class _GenAITool(TypedDict, total=False):
functionDeclarations: ReadOnly[list[_GenAIFunctionDeclaration]]
class _GenAIFunctionCallingConfig(TypedDict, total=False):
mode: ReadOnly[str]
class _GenAIToolConfig(TypedDict, total=False):
functionCallingConfig: ReadOnly[_GenAIFunctionCallingConfig]
def _decode_tool_call_arguments(raw_arguments: str) -> object:
"""Decode a tool call's JSON-encoded arguments into the value Google GenAI expects."""
return json.loads(raw_arguments)
class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
"""
Wrapper for streaming Google GenAI generate_content responses.
@ -51,7 +75,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
sent_first_chunk: bool = False
# State tracking for accumulating partial tool calls
accumulated_tool_calls: dict[str, dict[str, str]]
accumulated_tool_calls: dict[int, dict[str, str]]
def __init__(self, completion_stream: object):
self.sent_first_chunk = False
@ -108,7 +132,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
try:
# For tool calls with no arguments, accumulated_args will be "", which is not valid JSON.
# We default to an empty JSON object in this case.
parsed_args = json.loads(tool_call_data["arguments"] or "{}")
parsed_args = _decode_tool_call_arguments(tool_call_data["arguments"] or "{}")
function_call_part: _GenAIPart = {
"functionCall": {
"name": tool_call_data["name"] or "undefined_tool_name",
@ -319,7 +343,7 @@ class GoogleGenAIAdapter:
def _transform_google_genai_tools_to_openai(
self,
tools: list[dict[str, Any]],
tools: Sequence[_GenAITool],
) -> list[ChatCompletionToolParam]:
"""Transform Google GenAI tools to OpenAI tools format"""
openai_tools: Final[list[dict[str, object]]] = []
@ -346,7 +370,7 @@ class GoogleGenAIAdapter:
def _transform_google_genai_tool_config_to_openai(
self,
tool_config: dict[str, Any],
tool_config: _GenAIToolConfig,
) -> ChatCompletionToolChoiceValues | None:
"""Transform Google GenAI tool_config to OpenAI tool_choice"""
function_calling_config: Final = tool_config.get("functionCallingConfig", {})
@ -563,7 +587,7 @@ class GoogleGenAIAdapter:
parts = self._transform_openai_delta_to_google_genai_parts_with_accumulation(choice.delta, wrapper)
else:
parts = []
finish_reason = getattr(choice, "finish_reason", None)
finish_reason: str | None = getattr(choice, "finish_reason", None)
else:
# Fallback for generic choice objects
message_content: Final = getattr(choice, "delta", {}).get("content", "")
@ -625,7 +649,11 @@ class GoogleGenAIAdapter:
for tool_call in message.tool_calls:
if hasattr(tool_call, "function") and tool_call.function:
try:
args = json.loads(tool_call.function.arguments) if tool_call.function.arguments else {}
args = (
_decode_tool_call_arguments(tool_call.function.arguments)
if tool_call.function.arguments
else {}
)
except json.JSONDecodeError:
args = {}
@ -661,7 +689,7 @@ class GoogleGenAIAdapter:
continue
# 3. Use `index` as the primary key for accumulation
tool_call_index = getattr(tool_call, "index", None)
tool_call_index: int | None = getattr(tool_call, "index", None)
if tool_call_index is None:
continue # Index is essential for tracking streaming tool calls
@ -695,7 +723,7 @@ class GoogleGenAIAdapter:
# 5. Attempt to parse arguments even if name hasn't arrived.
try:
# Attempt to parse the accumulated arguments string
parsed_args = json.loads(accumulated_args)
parsed_args = _decode_tool_call_arguments(accumulated_args)
# If parsing succeeds, but we don't have a name yet, wait.
# The part will be created by a later chunk that brings the name.
@ -729,7 +757,7 @@ class GoogleGenAIAdapter:
return mapping.get(finish_reason, "STOP")
def _map_usage(self, usage: Any) -> dict[str, int]:
def _map_usage(self, usage: Usage | None) -> dict[str, int]:
"""Map OpenAI usage to Google GenAI usage format"""
return {
"promptTokenCount": getattr(usage, "prompt_tokens", 0) or 0,

View file

@ -5,6 +5,7 @@ import datetime
import os
import random
import time
from collections.abc import Callable
from datetime import timedelta
from typing import TYPE_CHECKING, Any, Final, Literal
@ -17,7 +18,11 @@ import litellm.litellm_core_utils.litellm_logging
import litellm.types
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.constants import HOURS_IN_A_DAY, SLACK_DAILY_REPORT_LOCK_ID
from litellm.constants import (
HOURS_IN_A_DAY,
SLACK_DAILY_REPORT_LOCK_ID,
SLACK_MODEL_DEPRECATION_LOCK_ID,
)
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.integrations.SlackAlerting.budget_alert_types import get_budget_alert_type
from litellm.integrations.SlackAlerting.hanging_request_check import (
@ -45,6 +50,10 @@ from litellm.repositories.table_repositories import InvitationLinkRepository
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.user_repository import UserRepository
from litellm.types.integrations.slack_alerting import *
from litellm.types.proxy.model_deprecation import (
DEFAULT_DEPRECATION_CHECK_INTERVAL_SECONDS,
DEPRECATION_IDLE_POLL_SECONDS,
)
from ..email_templates.templates import *
from .batching_handler import send_to_webhook, squash_payloads
@ -59,6 +68,12 @@ else:
Router = Any
def _proxy_llm_router() -> Router | None:
from litellm.proxy.proxy_server import llm_router
return llm_router
class SlackAlerting(CustomBatchLogger):
"""
Class for sending Slack Alerts
@ -1044,6 +1059,99 @@ Model Info:
async def model_removed_alert(self, model_name: str):
pass
def _deprecation_alerts_enabled(self) -> bool:
return self.alerting is not None and AlertType.model_deprecation_warnings in self.alert_types
async def send_model_deprecation_alert(
self,
llm_router: Router | None = None,
pod_lock_manager: "PodLockManager | None" = None,
) -> bool:
"""Alert on the router's deprecated and imminent models, True when one was sent
The daily lock is claimed only once there is something to say, so an empty pass never blocks a
later real one, and a sent alert is stamped in the shared cache for a day so sibling pods stop asking
"""
if not self._deprecation_alerts_enabled():
return False
from litellm.proxy.common_utils.model_deprecation import (
collect_model_deprecations,
format_deprecation_alert_message,
)
snapshot: Final = collect_model_deprecations(llm_router=llm_router)
message: Final = format_deprecation_alert_message(snapshot)
if message is None:
return False
if not await self._claimed_deprecation_alert_window(pod_lock_manager):
return False
level: Final[Literal["Low", "Medium", "High"]] = "High" if snapshot.deprecated else "Medium"
await self.send_alert(
message=message,
level=level,
alert_type=AlertType.model_deprecation_warnings,
alerting_metadata={ # mutable-ok: send_alert takes a dict payload
"deprecated_count": len(snapshot.deprecated),
"imminent_count": len(snapshot.imminent),
"upcoming_count": len(snapshot.upcoming),
},
)
await self.internal_usage_cache.async_set_cache(
key=SlackAlertingCacheKeys.deprecation_alert_sent_key.value,
value=time.time(),
ttl=DEFAULT_DEPRECATION_CHECK_INTERVAL_SECONDS,
)
return True
async def _claimed_deprecation_alert_window(self, pod_lock_manager: "PodLockManager | None") -> bool:
"""Without a redis backed lock there is no fleet to coordinate, so a lone pod always alerts"""
if pod_lock_manager is None:
return True
return (
await pod_lock_manager.acquire_lock(
cronjob_id=SLACK_MODEL_DEPRECATION_LOCK_ID,
ttl=DEFAULT_DEPRECATION_CHECK_INTERVAL_SECONDS,
allow_reentrant=False,
)
) is not False
async def _deprecation_alert_sent_within_a_day(self) -> bool:
return (
await self.internal_usage_cache.async_get_cache(key=SlackAlertingCacheKeys.deprecation_alert_sent_key.value)
) is not None
async def _run_deprecation_alert_pass(
self, llm_router: Router | None, pod_lock_manager: "PodLockManager | None"
) -> bool:
if llm_router is None or not self._deprecation_alerts_enabled():
return False
if await self._deprecation_alert_sent_within_a_day():
return False
return await self.send_model_deprecation_alert(llm_router=llm_router, pod_lock_manager=pod_lock_manager)
async def run_scheduled_deprecation_check(
self,
get_llm_router: Callable[[], Router | None] = _proxy_llm_router,
pod_lock_manager: "PodLockManager | None" = None,
) -> None:
"""Poll every pass for a loaded router, the alert being on, and no alert in the last day, then alert
A pass that could not alert (no router yet, alert type off, a sibling pod holds the daily lock, or a
redis blip at claim time) is retried on the next poll instead of costing a day, while a pass that
raised (a missing webhook, say) backs off a full day so a misconfiguration logs once, not every poll
"""
while True:
try:
await self._run_deprecation_alert_pass(get_llm_router(), pod_lock_manager)
except Exception as e: # noqa: BLE001 # a failed alert must not kill the loop
verbose_proxy_logger.exception("Error in model deprecation alert loop: %s", e)
await asyncio.sleep(DEFAULT_DEPRECATION_CHECK_INTERVAL_SECONDS)
continue
await asyncio.sleep(DEPRECATION_IDLE_POLL_SECONDS)
async def send_webhook_alert(self, webhook_event: WebhookEvent) -> bool:
"""
Sends structured alert to webhook, if set.

View file

@ -10,17 +10,31 @@ Supported for both `v1/chat/completions` (via the prompt-management hook) and
"""
import copy
import os
import re
from collections.abc import Iterable, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, cast
from urllib.parse import urlparse
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.custom_prompt_management import CustomPromptManagement
from litellm.integrations.prompt_management_base import PromptManagementClient
from litellm.litellm_core_utils.prompt_templates.common_utils import (
with_prompt_cache_breakpoint,
)
from litellm.types.integrations.anthropic_cache_control_hook import (
CacheControlInjectionPoint,
CacheControlMessageInjectionPoint,
)
from litellm.types.llms.openai import AllMessageValues, ChatCompletionCachedContent
from litellm.types.llms.anthropic import AnthropicSystemMessageContent
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionCachedContent,
ChatCompletionTextObject,
PromptCacheBreakpoint,
PromptCacheOptions,
)
from litellm.types.prompts.init_prompts import PromptSpec
from litellm.types.utils import StandardCallbackDynamicParams
@ -34,6 +48,55 @@ else:
# breakpoints: "A maximum of 4 blocks with cache_control may be provided."
MAX_CACHE_CONTROL_BLOCKS: Final = 4
CACHE_BREAKPOINT_KEYS: Final = ("cache_control", "prompt_cache_breakpoint")
OPENAI_PROMPT_CACHE_BREAKPOINT_MIN_GPT_VERSION: Final = (5, 6)
_GPT_VERSION_PATTERN: Final = re.compile(r"^gpt-(\d+)(?:\.(\d+))?")
OPENAI_PROMPT_CACHE_BREAKPOINT_BLOCK_TYPES: Final = frozenset(
{"text", "image", "image_url", "file", "input_audio", "input_text", "input_image", "input_file"}
)
OPENAI_API_HOST: Final = "api.openai.com"
OPENAI_API_BASE_ENV_VARS: Final = ("OPENAI_BASE_URL", "OPENAI_API_BASE")
def supports_openai_prompt_cache_breakpoint(model: str) -> bool:
model_map_flag: Final = _model_map_prompt_cache_breakpoint_flag(model)
if model_map_flag is not None:
return model_map_flag
version_match: Final = _GPT_VERSION_PATTERN.match(model.rsplit("/", 1)[-1].lower())
if version_match is None:
return False
version: Final = (int(version_match.group(1)), int(version_match.group(2) or 0))
return version >= OPENAI_PROMPT_CACHE_BREAKPOINT_MIN_GPT_VERSION
def _model_map_prompt_cache_breakpoint_flag(model: str) -> bool | None:
import litellm
entries: Final = (litellm.model_cost.get(key) for key in (model, model.rsplit("/", 1)[-1]))
flags: Final = (entry.get("supports_prompt_cache_breakpoint") for entry in entries if isinstance(entry, dict))
return next((bool(flag) for flag in flags if flag is not None), None)
def targets_openai_api(api_base: object) -> bool:
import litellm
resolved: Final = next(
(value for value in (api_base, litellm.api_base, *map(os.getenv, OPENAI_API_BASE_ENV_VARS)) if value),
None,
)
if not isinstance(resolved, str):
return True
host: Final = urlparse(resolved).hostname
return host is not None and (host == OPENAI_API_HOST or host.endswith(f".{OPENAI_API_HOST}"))
def _carries_cache_breakpoint(block: object) -> bool:
return isinstance(block, dict) and any(block.get(key) is not None for key in CACHE_BREAKPOINT_KEYS)
def _accepts_prompt_cache_breakpoint(block: object) -> bool:
return isinstance(block, dict) and block.get("type") in OPENAI_PROMPT_CACHE_BREAKPOINT_BLOCK_TYPES
class AnthropicCacheControlHook(CustomPromptManagement):
def get_chat_completion_prompt(
@ -81,13 +144,32 @@ class AnthropicCacheControlHook(CustomPromptManagement):
# provider transform, where each tool_config point appends at most one
# cachePoint to the tools. That block also counts toward Anthropic's
# limit, so reserve a slot for it here to leave room.
reserved_blocks: Final = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0
stamped_dialect: Final = injection_points[0].get("_litellm_openai_dialect")
openai_dialect: Final = (
stamped_dialect
if isinstance(stamped_dialect, bool)
else AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(
model,
non_default_params.get("custom_llm_provider"),
non_default_params.get("api_base") or non_default_params.get("base_url"),
non_default_params.get("prompt_cache_options"),
)
)
reserved_blocks: Final = (
1 if not openai_dialect and any(p.get("location") == "tool_config" for p in remaining_points) else 0
)
breakpoints_before: Final = AnthropicCacheControlHook._count_request_cache_breakpoints(processed_messages)
processed_messages = self._apply_message_injections(
points=message_points,
messages=processed_messages,
max_blocks=MAX_CACHE_CONTROL_BLOCKS - reserved_blocks,
openai_dialect=openai_dialect,
)
if (
openai_dialect
and AnthropicCacheControlHook._count_request_cache_breakpoints(processed_messages) > breakpoints_before
):
non_default_params.setdefault("prompt_cache_options", PromptCacheOptions(mode="explicit"))
# Pass through non-message injection points for provider-specific handling
if remaining_points:
@ -97,11 +179,43 @@ class AnthropicCacheControlHook(CustomPromptManagement):
return model, processed_messages, non_default_params
@staticmethod
def _targets_openai_prompt_cache_breakpoint(
model: str | None,
custom_llm_provider: str | None,
api_base: object = None,
prompt_cache_options: object = None,
) -> bool:
if model is None or not supports_openai_prompt_cache_breakpoint(model):
return False
if (custom_llm_provider or AnthropicCacheControlHook._resolve_provider(model)) != "openai":
return False
return prompt_cache_options is not None or targets_openai_api(api_base)
@staticmethod
def _resolve_provider(model: str) -> str | None:
from litellm.exceptions import BadRequestError
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
try:
_, provider, _, _ = get_llm_provider(model=model)
except BadRequestError:
return None
return provider
@staticmethod
def _count_request_cache_breakpoints(messages: Iterable[object], system: object = None) -> int:
system_blocks: Final = (
sum(1 for block in system if _carries_cache_breakpoint(block)) if isinstance(system, list) else 0
)
return system_blocks + sum(AnthropicCacheControlHook._count_cache_control_blocks(msg) for msg in messages)
@staticmethod
def _apply_message_injections(
points: list[CacheControlMessageInjectionPoint],
messages: list[AllMessageValues],
max_blocks: int,
openai_dialect: bool = False,
) -> list[AllMessageValues]:
"""Apply message-level cache control injection points in order.
@ -112,7 +226,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
``max_blocks`` is reached. Injection points are honored in config order,
so earlier points win when slots are scarce.
"""
used_blocks = sum(AnthropicCacheControlHook._count_cache_control_blocks(msg) for msg in messages)
used_blocks = AnthropicCacheControlHook._count_request_cache_breakpoints(messages)
limit_reached = False
for point in points:
@ -134,16 +248,17 @@ class AnthropicCacheControlHook(CustomPromptManagement):
continue
messages[target_index] = AnthropicCacheControlHook._safe_insert_cache_control_in_message(
messages[target_index], control
messages[target_index], control, openai_dialect
)
used_blocks += 1
if AnthropicCacheControlHook._message_has_cache_control(messages[target_index]):
used_blocks += 1
if limit_reached:
break
if limit_reached:
verbose_logger.warning(
"AnthropicCacheControlHook: Reached the Anthropic limit of %s cache_control blocks. Skipping further injection.",
"AnthropicCacheControlHook: Reached the provider limit of %s cache breakpoints. Skipping further injection.",
MAX_CACHE_CONTROL_BLOCKS,
)
@ -189,16 +304,13 @@ class AnthropicCacheControlHook(CustomPromptManagement):
return []
@staticmethod
def _count_cache_control_blocks(message: AllMessageValues) -> int:
"""Count cache_control breakpoints on a message (message + content level)."""
count = 0
if message.get("cache_control") is not None:
count += 1
def _count_cache_control_blocks(message: object) -> int:
if not isinstance(message, dict):
return 0
count = 1 if _carries_cache_breakpoint(message) else 0
content: Final = message.get("content")
if isinstance(content, list):
for block in content:
if isinstance(block, dict) and block.get("cache_control") is not None:
count += 1
count += sum(1 for block in content if _carries_cache_breakpoint(block))
return count
@staticmethod
@ -208,7 +320,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
@staticmethod
def _safe_insert_cache_control_in_message(
message: AllMessageValues, control: ChatCompletionCachedContent
message: AllMessageValues, control: ChatCompletionCachedContent, openai_dialect: bool = False
) -> AllMessageValues:
"""
Safe way to insert cache control in a message
@ -221,6 +333,9 @@ class AnthropicCacheControlHook(CustomPromptManagement):
Per Anthropic's API specification, when using multiple content blocks,
only the last content block can have cache_control.
"""
if openai_dialect:
return AnthropicCacheControlHook._insert_prompt_cache_breakpoint_in_message(message)
message_content: Final = message.get("content", None)
# 1. if string, insert cache control in the message
@ -232,11 +347,51 @@ class AnthropicCacheControlHook(CustomPromptManagement):
message_content[-1]["cache_control"] = control
return message
@staticmethod
def _insert_prompt_cache_breakpoint_in_message(message: AllMessageValues) -> AllMessageValues:
if message.get("role") == "assistant":
return message
message_content: Final = message.get("content", None)
if isinstance(message_content, str):
marked: Final = copy.copy(message)
marked["content"] = [
with_prompt_cache_breakpoint(
ChatCompletionTextObject(type="text", text=message_content), PromptCacheBreakpoint(mode="explicit")
)
]
return marked
if isinstance(message_content, list):
target_index: Final = next(
(
index
for index in range(len(message_content) - 1, -1, -1)
if _accepts_prompt_cache_breakpoint(message_content[index])
),
None,
)
if target_index is not None:
message_content[target_index] = with_prompt_cache_breakpoint(
message_content[target_index], PromptCacheBreakpoint(mode="explicit")
)
return message
@staticmethod
def _system_block_with_breakpoint(
block: Mapping[str, object], control: ChatCompletionCachedContent, openai_dialect: bool
) -> Mapping[str, object]:
marker: Final = (
("prompt_cache_breakpoint", PromptCacheBreakpoint(mode="explicit"))
if openai_dialect
else ("cache_control", control)
)
return {**block, marker[0]: marker[1]}
@staticmethod
def apply_to_anthropic_messages_request(
messages: list[dict],
system: str | list | None,
injection_points: list[CacheControlInjectionPoint],
openai_dialect: bool = False,
) -> tuple[list[dict], str | list | None, list[CacheControlInjectionPoint]]:
"""Apply cache control injection for the Anthropic-native v1/messages endpoint.
@ -262,30 +417,32 @@ class AnthropicCacheControlHook(CustomPromptManagement):
else:
remaining_points.append(point)
reserved_blocks: Final = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0
reserved_blocks: Final = (
1 if not openai_dialect and any(p.get("location") == "tool_config" for p in remaining_points) else 0
)
max_blocks: Final = MAX_CACHE_CONTROL_BLOCKS - reserved_blocks
used_blocks = sum(
AnthropicCacheControlHook._count_cache_control_blocks(cast(AllMessageValues, msg))
for msg in processed_messages
)
if isinstance(processed_system, list):
used_blocks += sum(
1 for b in processed_system if isinstance(b, dict) and b.get("cache_control") is not None
)
message_blocks: Final = AnthropicCacheControlHook._count_request_cache_breakpoints(processed_messages)
system_blocks = AnthropicCacheControlHook._count_request_cache_breakpoints((), processed_system)
if system_points and processed_system is not None and used_blocks < max_blocks:
if system_points and processed_system is not None and message_blocks + system_blocks < max_blocks:
system_already_has_cc: Final = isinstance(processed_system, list) and any(
isinstance(b, dict) and b.get("cache_control") is not None for b in processed_system
_carries_cache_breakpoint(b) for b in processed_system
)
if not system_already_has_cc:
control: Final = system_points[0].get("control") or ChatCompletionCachedContent(type="ephemeral")
if isinstance(processed_system, str):
processed_system = [{"type": "text", "text": processed_system, "cache_control": control}]
used_blocks += 1
processed_system = [
AnthropicCacheControlHook._system_block_with_breakpoint(
AnthropicSystemMessageContent(type="text", text=processed_system), control, openai_dialect
)
]
system_blocks += 1
elif len(processed_system) > 0 and isinstance(processed_system[-1], dict):
processed_system[-1] = {**processed_system[-1], "cache_control": control}
used_blocks += 1
processed_system[-1] = AnthropicCacheControlHook._system_block_with_breakpoint(
processed_system[-1], control, openai_dialect
)
system_blocks += 1
for i, msg in enumerate(processed_messages):
content = msg.get("content")
@ -295,7 +452,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
processed_messages = AnthropicCacheControlHook._apply_message_injections(
points=message_points,
messages=cast(list[AllMessageValues], processed_messages),
max_blocks=max_blocks - used_blocks,
max_blocks=max_blocks - system_blocks,
openai_dialect=openai_dialect,
)
return processed_messages, processed_system, remaining_points
@ -315,17 +473,57 @@ class AnthropicCacheControlHook(CustomPromptManagement):
return ChatCompletionCachedContent(type="ephemeral")
@staticmethod
def _stamped_as_judged(points: list[CacheControlInjectionPoint]) -> list[dict[str, object]]:
def _stamped_as_judged(points: Sequence[CacheControlInjectionPoint]) -> Sequence[Mapping[str, object]]:
"""Mark written-back points as having passed the client cache_control judgment.
Builds copies because config-owned point dicts are shared across
requests; mutating them would leak the stamp into future requests.
"""
return [{**point, "_litellm_judged": True} for point in points]
return AnthropicCacheControlHook._stamped(points, "_litellm_judged", True)
@staticmethod
def _judged_configured_points(
points: Sequence[CacheControlInjectionPoint],
messages: list[AllMessageValues],
tools: list[object] | None,
model: str,
custom_llm_provider: str | None,
api_base: object,
prompt_cache_options: object,
) -> Sequence[Mapping[str, object]] | None:
if AnthropicCacheControlHook._should_stand_down(points, messages, None, tools):
return None
return AnthropicCacheControlHook._stamped_with_dialect(
points, model, custom_llm_provider, api_base, prompt_cache_options
)
@staticmethod
def _stamped_with_dialect(
points: Sequence[CacheControlInjectionPoint],
model: str,
custom_llm_provider: str | None,
api_base: object,
prompt_cache_options: object,
) -> Sequence[Mapping[str, object]]:
if not supports_openai_prompt_cache_breakpoint(model):
return points
return AnthropicCacheControlHook._stamped(
points,
"_litellm_openai_dialect",
AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(
model, custom_llm_provider, api_base, prompt_cache_options
),
)
@staticmethod
def _stamped(
points: Sequence[CacheControlInjectionPoint], key: str, value: object
) -> Sequence[Mapping[str, object]]:
return [{**point, key: value} for point in points]
@staticmethod
def _should_stand_down(
points: list[CacheControlInjectionPoint],
points: Sequence[CacheControlInjectionPoint],
messages: list[AllMessageValues],
system: str | list | None,
tools: list | None,
@ -359,11 +557,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
carry the mark either at the top level (Anthropic shape) or nested under
``function`` (OpenAI shape); the Anthropic chat transform accepts both.
"""
if any(AnthropicCacheControlHook._count_cache_control_blocks(msg) for msg in messages):
if AnthropicCacheControlHook._count_request_cache_breakpoints(messages, system) > 0:
return True
if isinstance(system, list):
if any(isinstance(block, dict) and block.get("cache_control") is not None for block in system):
return True
if tools is not None:
return any(
isinstance(tool, dict)
@ -438,6 +633,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
custom_llm_provider: str | None,
tools: list | None = None,
enable_prompt_caching: bool | None = None,
api_base: object = None,
) -> None:
"""For /chat/completions: resolve the injection points the request should carry.
@ -452,10 +648,19 @@ class AnthropicCacheControlHook(CustomPromptManagement):
unchanged.
"""
if non_default_params.get("cache_control_injection_points"):
if AnthropicCacheControlHook._should_stand_down(
non_default_params["cache_control_injection_points"], messages, None, tools
):
judged: Final = AnthropicCacheControlHook._judged_configured_points(
non_default_params["cache_control_injection_points"],
messages,
tools,
model,
custom_llm_provider,
api_base,
non_default_params.get("prompt_cache_options"),
)
if judged is None:
non_default_params.pop("cache_control_injection_points")
else:
non_default_params["cache_control_injection_points"] = judged
return
points: Final = AnthropicCacheControlHook.get_default_injection_points(
messages=messages,
@ -476,6 +681,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
model: str | None = None,
custom_llm_provider: str | None = None,
tools: list[dict] | None = None,
api_base: str | None = None,
) -> tuple[list[dict], str | list | None]:
"""Extract cache_control_injection_points from kwargs and apply if present.
@ -513,11 +719,21 @@ class AnthropicCacheControlHook(CustomPromptManagement):
if not injection_points:
return messages, system
openai_dialect: Final = AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(
model, custom_llm_provider, api_base, kwargs.get("prompt_cache_options")
)
breakpoints_before: Final = AnthropicCacheControlHook._count_request_cache_breakpoints(messages, system)
messages, system, remaining = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=system,
injection_points=injection_points,
openai_dialect=openai_dialect,
)
if (
openai_dialect
and AnthropicCacheControlHook._count_request_cache_breakpoints(messages, system) > breakpoints_before
):
kwargs.setdefault("prompt_cache_options", PromptCacheOptions(mode="explicit"))
if remaining:
kwargs["cache_control_injection_points"] = AnthropicCacheControlHook._stamped_as_judged(remaining)
return messages, system

View file

@ -359,10 +359,10 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
"""
Determine if prompt management should run based on the prompt_id.
For Arize Phoenix, we always return True and handle the prompt loading
in the _compile_prompt_helper method.
Arize Phoenix needs a prompt_id to compile, so it declines requests without one;
prompt loading itself happens in the _compile_prompt_helper method.
"""
return True
return prompt_id is not None
def _compile_prompt_helper(
self,

View file

@ -102,6 +102,9 @@ class CustomGuardrail(CustomLogger):
# If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path.
use_native_during_call_hook: ClassVar[bool] = False
# If True, every proxy lifecycle event runs this guardrail's own hooks, not apply_guardrail.
use_native_lifecycle_hooks: ClassVar[bool] = False
records_own_guardrail_information: ClassVar[bool] = False
def __init__(
@ -632,7 +635,7 @@ class CustomGuardrail(CustomLogger):
return type(self).apply_guardrail is not CustomGuardrail.apply_guardrail
def _deployment_pre_call_target(self) -> "CustomLogger":
if not self.uses_apply_guardrail_interface():
if not self.uses_apply_guardrail_interface() or self.use_native_lifecycle_hooks:
return self
try:
from litellm.proxy.utils import unified_guardrail

View file

@ -60,13 +60,13 @@ class LLMResponse(BaseModel):
default=None,
description="Total cost of the LLM call in USD as computed by LiteLLM.",
)
output_logprobs: dict[str, Any] | None = Field(
output_logprobs: dict[str, object] | None = Field(
default=None,
description="Optional. When available, logprobs are used to compute Uncertainty.",
)
created_at: str = Field(..., description='timestamp constructed in "%Y-%m-%dT%H:%M:%S" format')
tags: list[str] | None = None
user_metadata: dict[str, Any] | None = None
user_metadata: dict[str, object] | None = None
class GalileoObserve(CustomLogger):
@ -238,13 +238,13 @@ class GalileoObserve(CustomLogger):
return created_at
@staticmethod
def _token_metrics_from_record(record: Mapping[str, Any]) -> dict[str, Any]:
def _token_metrics_from_record(record: Mapping[str, Any]) -> dict[str, object]:
num_input_tokens: Final = int(record.get("num_input_tokens") or 0)
num_output_tokens: Final = int(record.get("num_output_tokens") or 0)
num_total_tokens = int(record.get("num_total_tokens") or 0)
if num_total_tokens == 0 and (num_input_tokens or num_output_tokens):
num_total_tokens = num_input_tokens + num_output_tokens
metrics: Final[dict[str, Any]] = {
metrics: Final[dict[str, object]] = {
"num_input_tokens": num_input_tokens,
"num_output_tokens": num_output_tokens,
"num_total_tokens": num_total_tokens,
@ -260,10 +260,10 @@ class GalileoObserve(CustomLogger):
*,
trace_id: str,
span_id: str,
) -> dict[str, Any]:
) -> dict[str, object]:
created_at: Final = GalileoObserve._normalize_created_at(record.get("created_at", ""))
span: Final[dict[str, Any]] = {
span: Final[dict[str, object]] = {
"type": "llm",
"id": span_id,
"trace_id": trace_id,
@ -287,7 +287,7 @@ class GalileoObserve(CustomLogger):
return span
@staticmethod
def _record_to_v2_trace(record: Mapping[str, Any]) -> dict[str, Any]:
def _record_to_v2_trace(record: Mapping[str, Any]) -> dict[str, object]:
trace_id: Final = str(uuid.uuid4())
span_id: Final = str(uuid.uuid4())
created_at: Final = GalileoObserve._normalize_created_at(record.get("created_at", ""))
@ -307,8 +307,8 @@ class GalileoObserve(CustomLogger):
"spans": [GalileoObserve._record_to_v2_span(record, trace_id=trace_id, span_id=span_id)],
}
def _build_traces_payload(self, records: Sequence[Mapping[str, Any]]) -> dict[str, Any]:
payload: Final[dict[str, Any]] = {
def _build_traces_payload(self, records: Sequence[Mapping[str, object]]) -> dict[str, object]:
payload: Final[dict[str, object]] = {
"traces": [self._record_to_v2_trace(record) for record in records],
"logging_method": "api_direct",
"reliable": False,
@ -318,7 +318,7 @@ class GalileoObserve(CustomLogger):
payload["log_stream_id"] = self.log_stream_id
return payload
def _get_ingest_request(self) -> tuple[str, dict[str, Any]] | None:
def _get_ingest_request(self) -> tuple[str, dict[str, object]] | None:
if not self.base_url or not self.project_id:
return None
@ -427,9 +427,9 @@ class GalileoObserve(CustomLogger):
pass
@staticmethod
def _build_prompt(kwargs: Mapping[str, Any]) -> dict[str, Any]:
def _build_prompt(kwargs: Mapping[str, Any]) -> dict[str, object]:
optional_params: Final[Mapping[str, object]] = kwargs.get("optional_params", {}) or {}
prompt: Final[dict[str, Any]] = {"messages": kwargs.get("messages")}
prompt: Final[dict[str, object]] = {"messages": kwargs.get("messages")}
if optional_params.get("functions") is not None:
prompt["functions"] = optional_params["functions"]
if optional_params.get("tools") is not None:
@ -451,7 +451,7 @@ class GalileoObserve(CustomLogger):
return json.dumps(value, default=_json_default)
@staticmethod
def _prompt_to_input_text(prompt: Mapping[str, Any]) -> str:
def _prompt_to_input_text(prompt: Mapping[str, object]) -> str:
messages: Final[object] = prompt.get("messages")
if messages is not None:
text: Final = GalileoObserve._input_text_from_messages(messages)
@ -464,7 +464,7 @@ class GalileoObserve(CustomLogger):
if response_obj.choices and len(response_obj.choices) > 0:
message: Final = response_obj["choices"][0]["message"]
if hasattr(message, "json"):
message_json: Final = message.json()
message_json: Final[object] = message.json()
if isinstance(message_json, str):
return json.loads(message_json)
return message_json
@ -488,7 +488,7 @@ class GalileoObserve(CustomLogger):
return None
@staticmethod
def _langfuse_style_rerank_prompt(kwargs: Mapping[str, object]) -> dict[str, Any]:
def _langfuse_style_rerank_prompt(kwargs: Mapping[str, object]) -> dict[str, object]:
"""Match Langfuse rerank input: prompt = {"messages": kwargs.get("messages")}."""
return {"messages": kwargs.get("messages")}

View file

@ -5,7 +5,7 @@ import traceback
from collections.abc import Callable, Iterable, Mapping
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast
from packaging.version import Version
@ -49,10 +49,21 @@ else:
_DENIED_STEERING_KEYS: Final = frozenset({"headers", "endpoint", "caching_groups", "previous_models"})
_NO_METADATA: Final[Mapping[str, Any]] = MappingProxyType({})
_NO_METADATA: Final[Mapping[str, object]] = MappingProxyType({})
_REDACTED_PROXY_HEADERS: Final[frozenset[str]] = frozenset({"authorization", "cookie", "referer"})
def _object_mapping(value: object) -> Mapping[str, object] | None:
"""Return ``value`` as an opaque mapping when it is a dict."""
return value if isinstance(value, dict) else None
class _UsageObject(Protocol):
"""Token-count surface the Langfuse logger reads off a response usage payload."""
def get(self, key: Literal["cache_creation_input_tokens", "cache_read_input_tokens"], /) -> int | None: ...
def _extract_cache_read_input_tokens(usage_obj) -> int:
"""
Extract cache_read_input_tokens from usage object.
@ -82,6 +93,11 @@ def _extract_cache_read_input_tokens(usage_obj) -> int:
return cache_read_input_tokens
def _logging_id(start_time: datetime | None, response_obj: object) -> str | None:
"""Typed view of the timestamped response id Langfuse uses as the generation id."""
return litellm.utils.get_logging_id(start_time, response_obj)
def _as_steering_flag(value: object) -> bool:
"""A string ``str_to_bool`` does not recognise falls back to its truthiness."""
if isinstance(value, str):
@ -222,7 +238,7 @@ class LangFuseLogger:
return langfuse_client
@staticmethod
def add_metadata_from_header(litellm_params: dict, metadata: dict) -> dict:
def add_metadata_from_header(litellm_params: dict, metadata: dict) -> dict[str, object]:
"""
Adds metadata from proxy request headers to Langfuse logging if keys start with "langfuse_"
and overwrites litellm_params.metadata if already included.
@ -494,7 +510,7 @@ class LangFuseLogger:
def _log_langfuse_v2(
self,
user_id: str | None,
metadata: dict,
metadata: dict[str, object],
litellm_params: dict,
output: str | dict | list | None,
start_time: datetime | None,
@ -519,7 +535,7 @@ class LangFuseLogger:
else []
)
allowlisted_metadata: Final[StandardLoggingMetadata | dict[str, Any]] = (
allowlisted_metadata: Final[StandardLoggingMetadata | Mapping[str, object]] = (
standard_logging_object["metadata"] if standard_logging_object is not None else _NO_METADATA
)
end_user_id: Final = allowlisted_metadata.get("user_api_key_end_user_id", None)
@ -531,11 +547,12 @@ class LangFuseLogger:
# Clean Metadata before logging - never log raw metadata
# the raw metadata can contain circular references which leads to infinite recursion
# we clean out all extra litellm metadata params before logging
clean_metadata: dict[str, Any] = {}
clean_metadata: dict[str, object] = {}
if prompt_management_metadata is not None:
clean_metadata["prompt_management_metadata"] = prompt_management_metadata
if isinstance(metadata, dict):
for key, value in metadata.items():
metadata_entries: Final = _object_mapping(metadata)
if metadata_entries is not None:
for key, value in metadata_entries.items():
# generate langfuse tags - Default Tags sent to Langfuse from LiteLLM Proxy
if (
litellm.langfuse_default_tags is not None
@ -705,8 +722,8 @@ class LangFuseLogger:
usage_details = None
if response_obj is not None:
if hasattr(response_obj, "id") and response_obj.get("id", None) is not None:
generation_id = litellm.utils.get_logging_id(start_time, response_obj)
_usage_obj: Final = getattr(response_obj, "usage", None)
generation_id = _logging_id(start_time, response_obj)
_usage_obj: Final[_UsageObject | None] = getattr(response_obj, "usage", None)
if _usage_obj:
# Safely get usage values, defaulting None to 0 for Langfuse compatibility.
@ -811,7 +828,7 @@ class LangFuseLogger:
@staticmethod
def _get_chat_content_for_langfuse(
response_obj: ModelResponse,
):
) -> str | None:
"""
Get the chat content for Langfuse logging
"""
@ -1078,7 +1095,7 @@ def log_provider_specific_information_as_span(
None
"""
_hidden_params: Final = clean_metadata.get("hidden_params", None)
_hidden_params: Final[Mapping[str, object] | None] = clean_metadata.get("hidden_params", None)
if _hidden_params is None:
return

View file

@ -1,5 +1,8 @@
import os
from collections.abc import Mapping
import threading
from collections import OrderedDict
from collections.abc import Callable, Mapping
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
@ -17,6 +20,7 @@ from litellm.integrations.opentelemetry_utils.gen_ai_semconv import (
OTELSemconvCategory,
parse_semconv_opt_in,
)
from litellm.integrations.otel.model.db_endpoint import db_span_attributes
from litellm.integrations.otel.model.semconv import Metric
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.secret_redaction import redact_string
@ -38,6 +42,7 @@ from litellm.types.utils import (
# OpenTelemetry imports moved to individual functions to avoid import errors when not installed
if TYPE_CHECKING:
from opentelemetry.sdk.resources import Resource as _Resource
from opentelemetry.sdk.trace import TracerProvider as _SDKTracerProvider
from opentelemetry.sdk.trace.export import SpanExporter as _SpanExporter
from opentelemetry.trace import Context as _Context
@ -83,6 +88,12 @@ class _ResponseWithUsageView(TypedDict, total=False):
usage: "_UsageCompletionTokensView | None"
# Cap on credential-scoped providers held at once; each one owns an exporter thread.
_MAX_DYNAMIC_TRACER_PROVIDERS: Final = 256
# Dedicated so a slow exporter shutdown cannot starve the shared logging executor.
_PROVIDER_SHUTDOWN_EXECUTOR: Final = ThreadPoolExecutor(max_workers=4, thread_name_prefix="OtelProviderShutdown")
LITELLM_TRACER_NAME: Final = os.getenv("OTEL_TRACER_NAME", "litellm")
LITELLM_METER_NAME: Final = os.getenv("LITELLM_METER_NAME", "litellm")
LITELLM_LOGGER_NAME: Final = os.getenv("LITELLM_LOGGER_NAME", "litellm")
@ -227,6 +238,34 @@ def _freeze_for_dedupe(value: object, _depth: int = 0) -> HashableScope:
return repr(value)
def _shutdown_tracer_provider(provider: "_SDKTracerProvider") -> None:
"""Flush and stop a dropped provider so its exporter thread is reclaimed."""
try:
provider.shutdown()
except Exception as e: # noqa: BLE001 # exporter shutdown must not fail the request that dropped it
verbose_logger.debug("OpenTelemetry: error shutting down dropped tracer provider: %s", e)
@dataclass(frozen=True, slots=True)
class _CachedTracerProvider:
"""A cached credential-scoped provider plus whether it may be shut down when dropped."""
provider: "_SDKTracerProvider"
owns_exporter: bool
def _provider_owns_exporter(exporter: "str | _SpanExporter") -> bool:
"""Whether a provider built for ``exporter`` may be shut down when it is dropped.
``_get_span_processor`` builds a fresh exporter for a named kind, but wraps a
caller-supplied ``SpanExporter`` instance as-is, and that instance is shared with the
logger's own provider. Shutting a dropped provider down would then stop exporting for
the whole process. The shared case also uses ``SimpleSpanProcessor``, so it owns no
thread and there is nothing to reclaim.
"""
return not hasattr(exporter, "export")
@dataclass
class OpenTelemetryConfig:
exporter: str | SpanExporter = "console"
@ -322,6 +361,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
tracer_provider: object | None = None,
logger_provider: object | None = None,
meter_provider: object | None = None,
max_dynamic_tracer_providers: int = _MAX_DYNAMIC_TRACER_PROVIDERS,
**kwargs,
):
team_metadata_keys_override: Final = kwargs.pop("baggage_team_metadata_keys", None)
@ -347,7 +387,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
self.OTEL_EXPORTER = self.config.exporter
self.OTEL_ENDPOINT = self.config.endpoint
self.OTEL_HEADERS = self.config.headers
self._tracer_provider_cache: dict[str, _SDKTracerProvider] = {}
self._tracer_provider_cache: OrderedDict[str, _CachedTracerProvider] = OrderedDict()
self._tracer_provider_cache_lock: Final = threading.Lock()
self._max_dynamic_tracer_providers: Final = max(1, max_dynamic_tracer_providers)
self._litellm_resource_memo: _Resource | None = None
self._init_tracing(tracer_provider)
_debug_otel: Final = str(os.getenv("DEBUG_OTEL", "False")).lower()
@ -373,7 +416,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
self._init_otel_logger_on_litellm_proxy()
@staticmethod
def _get_litellm_resource(config: OpenTelemetryConfig):
def _get_litellm_resource(config: OpenTelemetryConfig) -> "_Resource":
"""Create an OpenTelemetry Resource using config-driven defaults."""
from opentelemetry.sdk.resources import OTELResourceDetector, Resource
@ -388,6 +431,21 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
env_resource: Final = otel_resource_detector.detect()
return base_resource.merge(env_resource)
def _litellm_resource(self) -> "_Resource":
"""The Resource every provider on this logger is built with, frozen at first use.
``Resource.create`` scans every installed distribution's entry points, roughly 3ms and
200 file opens, and the dynamic providers reach it from the async logging path. Freezing
also keeps them consistent with whatever this logger built at startup. ``cached_property``
locks class-wide before 3.12, which this file still supports.
"""
memo: Final = self._litellm_resource_memo
if memo is not None:
return memo
built: Final = self._get_litellm_resource(self.config)
self._litellm_resource_memo = built
return built
def _init_otel_logger_on_litellm_proxy(self):
"""
Initializes OpenTelemetry for litellm proxy server
@ -555,7 +613,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
from opentelemetry.trace import SpanKind
def create_tracer_provider():
provider: Final = TracerProvider(resource=self._get_litellm_resource(self.config))
provider: Final = TracerProvider(resource=self._litellm_resource())
provider.add_span_processor(self._get_span_processor())
return provider
@ -593,7 +651,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
metric_reader: Final = self._get_metric_reader()
return MeterProvider(
metric_readers=[metric_reader],
resource=self._get_litellm_resource(self.config),
resource=self._litellm_resource(),
)
meter_provider = self._get_or_create_provider(
@ -651,7 +709,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
from opentelemetry.sdk._logs.export import BatchLogRecordProcessor
def create_logger_provider():
provider: Final = OTLoggerProvider(resource=self._get_litellm_resource(self.config))
provider: Final = OTLoggerProvider(resource=self._litellm_resource())
log_exporter: Final = self._get_log_exporter()
provider.add_log_record_processor(BatchLogRecordProcessor(log_exporter))
return provider
@ -678,6 +736,28 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
self._handle_failure(kwargs, response_obj, start_time, end_time)
def _start_service_span(self, payload: ServiceLoggerPayload, parent_otel_span: Span, start_time_ns: int) -> Span:
"""Open a service span, named and classified by what the service is.
A datastore call is an outbound CLIENT span carrying ``db.*`` semconv.
Without those a Postgres span says only ``service=postgres``, so the
backend falls back to the transport peer, which for Prisma is the local
query engine on loopback. Everything else stays an INTERNAL span.
"""
from opentelemetry import trace
from opentelemetry.trace import SpanKind
attributes: Final = db_span_attributes(payload.service.value, payload.call_type)
span: Final = self.tracer.start_span(
name=payload.service,
context=trace.set_span_in_context(parent_otel_span),
start_time=start_time_ns,
kind=SpanKind.CLIENT if attributes else SpanKind.INTERNAL,
)
for key, value in attributes.items():
self.safe_set_attribute(span=span, key=key, value=value)
return span
async def async_service_success_hook(
self,
payload: ServiceLoggerPayload,
@ -686,7 +766,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
end_time: datetime | float | None = None,
event_metadata: dict | None = None,
):
from opentelemetry import trace
from opentelemetry.trace import Status, StatusCode
_start_time_ns = 0
@ -703,12 +782,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
_end_time_ns = self._to_ns(end_time)
if parent_otel_span is not None:
_span_name: Final = payload.service
service_logging_span: Final = self.tracer.start_span(
name=_span_name,
context=trace.set_span_in_context(parent_otel_span),
start_time=_start_time_ns,
)
service_logging_span: Final = self._start_service_span(payload, parent_otel_span, _start_time_ns)
self.safe_set_attribute(
span=service_logging_span,
key="call_type",
@ -746,7 +820,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
end_time: float | datetime | None = None,
event_metadata: dict | None = None,
):
from opentelemetry import trace
from opentelemetry.trace import Status, StatusCode
_start_time_ns = 0
@ -763,12 +836,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
_end_time_ns = self._to_ns(end_time)
if parent_otel_span is not None:
_span_name: Final = payload.service
service_logging_span: Final = self.tracer.start_span(
name=_span_name,
context=trace.set_span_in_context(parent_otel_span),
start_time=_start_time_ns,
)
service_logging_span: Final = self._start_service_span(payload, parent_otel_span, _start_time_ns)
self.safe_set_attribute(
span=service_logging_span,
key="call_type",
@ -1027,38 +1095,94 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
return self.construct_dynamic_otel_config(standard_callback_dynamic_params=standard_callback_dynamic_params)
def _get_tracer_with_dynamic_config(self, dynamic_config: OpenTelemetryConfig):
def _insert_or_drop(
self, cache_key: str, built: "_CachedTracerProvider"
) -> "tuple[_CachedTracerProvider, _CachedTracerProvider | None]":
"""Cache ``built`` under ``cache_key``, returning the entry to use and what to drop.
Caller holds ``_tracer_provider_cache_lock``. The drop is either the loser of a
concurrent build for this key or the LRU victim its insertion pushed out.
"""
raced: Final = self._tracer_provider_cache.get(cache_key)
if raced is not None:
self._tracer_provider_cache.move_to_end(cache_key)
return raced, built
self._tracer_provider_cache[cache_key] = built
if len(self._tracer_provider_cache) > self._max_dynamic_tracer_providers:
return built, self._tracer_provider_cache.popitem(last=False)[1]
return built, None
def _cached_dynamic_tracer(
self,
cache_key: str,
build: Callable[[], "_SDKTracerProvider"],
owns_exporter: bool,
) -> "_Tracer":
"""Return the tracer for ``cache_key``, building and caching a provider on miss.
A provider that owns its exporter also owns a ``BatchSpanProcessor`` worker thread
that only stops on ``shutdown()``, so the cache is a bounded LRU and whatever it
drops is shut down. Without both, a proxy serving key-scoped credentials accumulates
one live thread per credential set for the life of the process.
``owns_exporter`` also decides ``shutdown_on_exit`` at build time: a provider we may
never shut down must not hold an interpreter-exit hook, which would both pin it in
memory for the life of the process and stop the shared exporter at exit. Those
providers use ``SimpleSpanProcessor``, which buffers nothing, so the hook costs them
no flush.
``owns_exporter`` describes the provider being built, and is cached with it, because
the two dynamic entry points share this cache and can disagree: whether the LRU
victim may be shut down is a property of the victim, never of the request that
happened to evict it.
"""
with self._tracer_provider_cache_lock:
cached: Final = self._tracer_provider_cache.get(cache_key)
if cached is not None:
self._tracer_provider_cache.move_to_end(cache_key)
return cached.provider.get_tracer(LITELLM_TRACER_NAME)
# Built outside the lock: exporter construction can block on DNS/TLS.
built: Final = _CachedTracerProvider(provider=build(), owns_exporter=owns_exporter)
with self._tracer_provider_cache_lock:
winner, dropped = self._insert_or_drop(cache_key, built)
if dropped is not None and dropped.owns_exporter:
# Off the caller's thread: shutdown joins the exporter worker.
_PROVIDER_SHUTDOWN_EXECUTOR.submit(_shutdown_tracer_provider, dropped.provider)
return winner.provider.get_tracer(LITELLM_TRACER_NAME)
def _get_tracer_with_dynamic_config(self, dynamic_config: OpenTelemetryConfig) -> "_Tracer":
"""Create (or reuse) a tracer whose exporter target comes from a per-request config."""
from opentelemetry.sdk.trace import TracerProvider
cache_key = f"dynamic_config:{dynamic_config.exporter}:{dynamic_config.endpoint}:{dynamic_config.headers}"
if cache_key in self._tracer_provider_cache:
return self._tracer_provider_cache[cache_key].get_tracer(LITELLM_TRACER_NAME)
owns_exporter: Final = _provider_owns_exporter(dynamic_config.exporter)
temp_provider: Final = TracerProvider(resource=self._get_litellm_resource(self.config))
temp_provider.add_span_processor(self._get_span_processor(config_override=dynamic_config))
def _build() -> "_SDKTracerProvider":
provider: Final = TracerProvider(resource=self._litellm_resource(), shutdown_on_exit=owns_exporter)
provider.add_span_processor(self._get_span_processor(config_override=dynamic_config))
return provider
self._tracer_provider_cache[cache_key] = temp_provider
cache_key: Final = (
f"dynamic_config:{dynamic_config.exporter}:{dynamic_config.endpoint}:{dynamic_config.headers}"
)
return self._cached_dynamic_tracer(cache_key, _build, owns_exporter)
return temp_provider.get_tracer(LITELLM_TRACER_NAME)
def _get_tracer_with_dynamic_headers(self, dynamic_headers: dict):
"""Create a temporary tracer with dynamic headers for this request only."""
def _get_tracer_with_dynamic_headers(self, dynamic_headers: Mapping[str, str]) -> "_Tracer":
"""Create (or reuse) a tracer whose OTLP headers come from a per-request credential set."""
from opentelemetry.sdk.trace import TracerProvider
# Prevents thread exhaustion by reusing providers for the same credential sets (e.g. per-team keys)
owns_exporter: Final = _provider_owns_exporter(self.OTEL_EXPORTER)
def _build() -> "_SDKTracerProvider":
provider: Final = TracerProvider(resource=self._litellm_resource(), shutdown_on_exit=owns_exporter)
provider.add_span_processor(self._get_span_processor(dynamic_headers=dynamic_headers))
return provider
cache_key: Final = str(sorted(dynamic_headers.items()))
if cache_key in self._tracer_provider_cache:
return self._tracer_provider_cache[cache_key].get_tracer(LITELLM_TRACER_NAME)
# Create a temporary tracer provider with dynamic headers
temp_provider: Final = TracerProvider(resource=self._get_litellm_resource(self.config))
temp_provider.add_span_processor(self._get_span_processor(dynamic_headers=dynamic_headers))
# Store in cache for reuse
self._tracer_provider_cache[cache_key] = temp_provider
return temp_provider.get_tracer(LITELLM_TRACER_NAME)
return self._cached_dynamic_tracer(cache_key, _build, owns_exporter)
def construct_dynamic_otel_headers(
self, standard_callback_dynamic_params: StandardCallbackDynamicParams
@ -2832,7 +2956,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
def _get_span_processor(
self,
dynamic_headers: dict | None = None,
dynamic_headers: Mapping[str, str] | None = None,
config_override: OpenTelemetryConfig | None = None,
):
from opentelemetry.sdk.trace.export import (
@ -3144,7 +3268,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
@staticmethod
def _get_headers_dictionary(
headers: str | dict | None,
headers: "str | Mapping[str, str] | None",
) -> dict[str, str]:
"""
Convert a string or dictionary of headers into a dictionary of headers.
@ -3158,8 +3282,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
for part in parts:
key, value = part.split("=", 1)
_split_otel_headers[key] = value
elif isinstance(headers, dict):
_split_otel_headers = headers
elif isinstance(headers, Mapping):
_split_otel_headers.update(headers)
return _split_otel_headers
async def async_management_endpoint_success_hook(

View file

@ -62,8 +62,13 @@ from litellm.integrations.otel.plumbing.providers import (
from litellm.integrations.otel.plumbing.routing import TenantTracerCache
if TYPE_CHECKING:
from opentelemetry.metrics import MeterProvider
from litellm.caching.dual_cache import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.services import ServiceLoggerPayload
from litellm.types.utils import (
CallTypesLiteral,
StandardLoggingGuardrailInformation,
StandardLoggingPayload,
)
@ -140,7 +145,7 @@ class OpenTelemetryV2(CustomLogger):
callback_name: str | None = None,
tracer_provider: TracerProvider | None = None,
logger_provider: LoggerProvider | None = None,
meter_provider: Any | None = None,
meter_provider: "MeterProvider | None" = None,
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
@ -162,7 +167,7 @@ class OpenTelemetryV2(CustomLogger):
self._open_llm_calls: OrderedDict[str, _LLMCallSpan] = OrderedDict()
self._init_otel_logger_on_litellm_proxy()
def _init_metrics(self, meter_provider: Any | None) -> "GenAIMetricRecorder | None":
def _init_metrics(self, meter_provider: "MeterProvider | None") -> "GenAIMetricRecorder | None":
"""Create the six GenAI histograms when metrics are enabled, else ``None``.
``meter_provider`` is an explicit override (tests inject one); otherwise the
@ -340,7 +345,7 @@ class OpenTelemetryV2(CustomLogger):
def _emit_mcp_tool_call(
self,
kwargs: Mapping[str, Any],
kwargs: Mapping[str, object],
start_time: datetime | float | None,
end_time: datetime | float | None,
) -> bool:
@ -417,7 +422,7 @@ class OpenTelemetryV2(CustomLogger):
def _close_llm_call(
self,
kwargs: Mapping[str, Any],
kwargs: Mapping[str, object],
start_time: datetime | float | None,
end_time: datetime | float | None,
) -> Span | None:
@ -474,7 +479,7 @@ class OpenTelemetryV2(CustomLogger):
async def async_service_success_hook(
self,
payload: Any,
payload: "ServiceLoggerPayload",
parent_otel_span: Span | None = None,
start_time: datetime | float | None = None,
end_time: datetime | float | None = None,
@ -491,7 +496,7 @@ class OpenTelemetryV2(CustomLogger):
async def async_service_failure_hook(
self,
payload: Any,
payload: "ServiceLoggerPayload",
error: str | None = "",
parent_otel_span: Span | None = None,
start_time: datetime | float | None = None,
@ -509,7 +514,7 @@ class OpenTelemetryV2(CustomLogger):
def _emit_service(
self,
payload: Any,
payload: "ServiceLoggerPayload",
*,
parent_otel_span: Span | None,
start_time: datetime | float | None,
@ -559,7 +564,7 @@ class OpenTelemetryV2(CustomLogger):
# / errors are the FastAPI instrumentor's job, so we don't touch it here.
# ====================================================================== #
def seed_request_identity(self, user_api_key_dict: Any, model: Any = None) -> None:
def seed_request_identity(self, user_api_key_dict: object, model: str | None = None) -> None:
"""Attach request-identity Baggage to the current context + server span.
Seeding identity into Baggage makes **every** span emitted afterwards for
@ -615,10 +620,10 @@ class OpenTelemetryV2(CustomLogger):
async def async_pre_call_hook(
self,
user_api_key_dict: Any,
cache: Any,
user_api_key_dict: "UserAPIKeyAuth",
cache: "DualCache",
data: dict,
call_type: Any,
call_type: "CallTypesLiteral",
) -> dict:
self.seed_request_identity(
user_api_key_dict,
@ -790,7 +795,7 @@ def emit_guardrail_span(entry: "StandardLoggingGuardrailInformation") -> None:
pass
def seed_request_identity(user_api_key_dict: Any, model: Any = None) -> None:
def seed_request_identity(user_api_key_dict: object, model: str | None = None) -> None:
logger: Final = _registered_v2_logger()
if logger is not None:
logger.seed_request_identity(user_api_key_dict, model=model)

View file

@ -18,6 +18,7 @@ from litellm.integrations.otel.mappers.utils import (
serialize_messages,
tool_definition_attrs,
)
from litellm.integrations.otel.model.db_endpoint import db_span_attributes
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
LLMCallSpanData,
@ -27,7 +28,6 @@ from litellm.integrations.otel.model.payloads import (
ToolDefinition,
)
from litellm.integrations.otel.model.semconv import (
DB,
MCP,
Error,
GenAI,
@ -36,7 +36,6 @@ from litellm.integrations.otel.model.semconv import (
RpcSystem,
Server,
)
from litellm.integrations.otel.model.spans import db_system
class GenAIMapper:
@ -182,12 +181,8 @@ class GenAIMapper:
def _service(cls, data: ServiceSpanData) -> AttributeMap:
attrs: Final = collect(cls._SERVICE_ATTRS, data)
# An outbound datastore call (DB_CALL / CLIENT span) also carries db.*
# semconv. Internal services (router, budget jobs, …) have no db.system,
# so they get only the litellm.service.* keys above.
system: Final = db_system(data.service_name)
if system is not None:
attrs[DB.SYSTEM_NAME] = system
if data.call_type:
attrs[DB.OPERATION_NAME] = data.call_type
# semconv naming the server it reached. Internal services (router, budget
# jobs, …) have no db.system, so they get only the litellm.service.* keys.
attrs.update(db_span_attributes(data.service_name, data.call_type))
attrs.update({f"{LiteLLM.METADATA_PREFIX}{key}": value for key, value in data.event_metadata.items()})
return attrs

View file

@ -0,0 +1,164 @@
"""OTel ``db.*`` / ``server.*`` attributes naming the database litellm talks to.
Prisma reaches PostgreSQL through a query engine listening on loopback, so
transport-level instrumentation attributes the work to ``localhost`` and an
operator cannot tell it is a PostgreSQL call or correlate it with the database's
own metrics. These attributes name the real server on litellm's DB spans.
Only the host, port, database and schema of the DSN are read, so no credential
can reach an exporter.
"""
from __future__ import annotations
import os
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
from urllib.parse import ParseResult, parse_qs, unquote, urlparse
from litellm.integrations.otel.model.semconv import DB, Server
from litellm.integrations.otel.model.spans import POSTGRESQL, db_system
_DATABASE_URL_ENV: Final = "DATABASE_URL"
_READ_REPLICA_ENV: Final = "DATABASE_URL_READ_REPLICA"
_DEFAULT_POSTGRES_PORT: Final = 5432
_DEFAULT_POSTGRES_SCHEMA: Final = "public"
_POSTGRES_SCHEMES: Final = frozenset({"postgres", "postgresql"})
_EMPTY_ATTRIBUTES: Final[Mapping[str, str | int]] = MappingProxyType({})
@dataclass(frozen=True, slots=True)
class DatabaseEndpoint:
"""The non-sensitive identity of a PostgreSQL server, parsed from a DSN."""
address: str | None
port: int | None
namespace: str | None
def parse_database_endpoint(url: str | None) -> DatabaseEndpoint | None:
"""Parse a PostgreSQL DSN into its exportable endpoint identity.
Returns ``None`` for an absent, malformed or non-PostgreSQL URL rather than
raising: an unparseable DSN must degrade to a span without endpoint
attributes, never break the request that emitted it.
"""
if not url:
return None
try:
parsed: Final = urlparse(url)
if parsed.scheme not in _POSTGRES_SCHEMES:
return None
query: Final = parse_qs(parsed.query)
raw_database: Final = (parsed.path or "").lstrip("/")
if _is_misparsed_authority(parsed, url, raw_database):
return None
# ``host=`` beats the netloc: it is how libpq names a Unix socket
# directory and how the Cloud SQL connector sits behind a localhost
# netloc, where the netloc is the very answer this module replaces.
address: Final = _first(query.get("host")) or parsed.hostname
# ``port=`` accompanies ``host=`` in a libpq URI, so honour it the same way.
port: Final = _port(_first(query.get("port")), parsed.port) if address else None
namespace: Final = _namespace(unquote(raw_database), _first(query.get("schema")))
except ValueError:
return None
if address is None and namespace is None:
return None
return DatabaseEndpoint(address=address, port=port, namespace=namespace)
def _is_misparsed_authority(parsed: ParseResult, url: str, raw_database: str) -> bool:
"""Whether the URL authority may have been truncated by an unencoded character.
``/``, ``#`` or ``?`` in a password ends the netloc early, so urlparse hands
back the username as the host, the leading digits of the password as the
port, and the rest of the credential as the path, query or fragment. The
stranded userinfo ``@`` is the only surviving evidence.
A database name cannot hold an unencoded slash either, so a second path
segment is the same evidence.
A DSN that carries the at-sign in a query parameter instead, such as
``?application_name=svc@prod``, is indistinguishable from a mis-split by any
property of the parse: both leave no userinfo, a host, a port and a path.
Since guessing wrong publishes a credential fragment to a tracing backend,
that ambiguity resolves to refusing the endpoint. Such a DSN loses
``server.address`` and ``db.namespace`` and keeps the rest of the span,
which is the cheaper error of the two. Percent-encode the at-sign to keep
them.
"""
if "/" in raw_database:
return True
return "@" in url and "@" not in parsed.netloc
def _first(values: Sequence[str] | None) -> str:
return values[0] if values else ""
def _port(from_query: str, from_netloc: int | None) -> int:
return int(from_query) if from_query.isdigit() else (from_netloc or _DEFAULT_POSTGRES_PORT)
def _namespace(database: str, schema: str) -> str | None:
"""``{database}|{schema}`` per the PostgreSQL semconv, dropping absent halves.
Only Prisma's literal default schema stays implicit. The match is
case-sensitive because Prisma quotes the name, so ``?schema=PUBLIC`` builds
a second schema alongside ``public`` and the two must not collapse to one
namespace.
"""
qualifier: Final = "" if schema == _DEFAULT_POSTGRES_SCHEMA else schema
return "|".join(part for part in (database, qualifier) if part) or None
def postgres_endpoint() -> DatabaseEndpoint | None:
"""The PostgreSQL endpoint the process is currently connected to.
Read from ``os.environ`` on every span, deliberately, on both counts.
The environment is what Prisma itself connects with, so the span cannot
disagree with the connection; ``get_secret_str`` would consult a configured
secret manager first and could name a different server than the one serving
the query. And the value is not static: the RDS IAM refresh rebuilds the URL
from ``DATABASE_HOST``/``PORT``/``NAME``/``SCHEMA`` every rotation, the
reconnect path re-reads ``DATABASE_URL``, and the DB-backed
``environment_variables`` config overlay can rewrite any of them after
startup, so a value cached for the process lifetime goes stale against a
connection that has genuinely moved. Nothing is memoized either: a cache
keyed on the URL would hold a rotated credential past its rotation, and the
parse is a single ``urlparse`` on a short string.
A configured read replica yields ``None``: ``RoutingPrismaWrapper`` picks
reader or writer per Prisma call, underneath the span, so naming the writer
would attribute replica reads to the primary.
"""
if os.environ.get(_READ_REPLICA_ENV):
return None
return parse_database_endpoint(os.environ.get(_DATABASE_URL_ENV, ""))
def db_span_attributes(service_name: str, call_type: str | None = None) -> Mapping[str, str | int]:
"""The ``db.*``/``server.*`` attributes for a datastore service call.
Empty for services that are not outbound datastore calls. Endpoint
attributes are PostgreSQL-only: ``DATABASE_URL`` says nothing about where
the redis-backed services point. ``db.system`` rides alongside the current
``db.system.name`` because Datadog's OTLP intake still types a database span
from the older key.
"""
system: Final = db_system(service_name)
if system is None:
return _EMPTY_ATTRIBUTES
endpoint: Final = postgres_endpoint() if system == POSTGRESQL else None
pairs: Final[tuple[tuple[str, str | int | None], ...]] = (
(DB.SYSTEM_NAME, system),
(DB.SYSTEM_LEGACY, system),
(DB.OPERATION_NAME, call_type),
(Server.ADDRESS, endpoint.address if endpoint is not None else None),
(Server.PORT, endpoint.port if endpoint is not None else None),
(DB.NAMESPACE, endpoint.namespace if endpoint is not None else None),
)
return MappingProxyType({key: value for key, value in pairs if value})

View file

@ -238,7 +238,11 @@ class DB:
"""
SYSTEM_NAME: Final = "db.system.name"
# Superseded by SYSTEM_NAME, dual-emitted because Datadog's OTLP intake
# still infers a span's database type from this key.
SYSTEM_LEGACY: Final = "db.system"
OPERATION_NAME: Final = "db.operation.name"
NAMESPACE: Final = "db.namespace"
class HTTP:

View file

@ -115,10 +115,12 @@ SPAN_REGISTRY: Final[dict[SpanRole, SpanSpec]] = {
# redis-backed spend queues. Any service not mapped here is litellm-internal work
# and stays an INTERNAL ``SERVICE`` span. This table is the single source of
# datastore knowledge — both the role classifier and the mapper read it.
POSTGRESQL: Final = "postgresql"
_DB_SYSTEM_BY_SERVICE: Final[dict[str, str]] = {
"redis": "redis",
"postgres": "postgresql",
"batch_write_to_db": "postgresql",
"postgres": POSTGRESQL,
"batch_write_to_db": POSTGRESQL,
}

View file

@ -9,7 +9,9 @@ import os
import sys
from collections.abc import Awaitable, Callable, Mapping, Sequence
from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any, Final, Literal, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, cast
from pydantic import BaseModel
import litellm
from litellm._logging import print_verbose, verbose_logger
@ -38,6 +40,7 @@ from litellm.proxy._types import (
LiteLLM_UserTable,
UserAPIKeyAuth,
)
from litellm.repositories.base_repository import BaseRepository
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.user_repository import UserRepository
@ -58,6 +61,9 @@ if TYPE_CHECKING:
else:
AsyncIOScheduler = Any
_BudgetRowT: Final = TypeVar("_BudgetRowT")
_TableRowT: Final = TypeVar("_TableRowT", bound=BaseModel)
_DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT: Final = 5.0
_NON_ENUM_METRIC_LABELS: Final[frozenset[str]] = frozenset(
@ -73,6 +79,36 @@ _NON_ENUM_METRIC_LABELS: Final[frozenset[str]] = frozenset(
)
class _PaginatedPrismaTable(Protocol[_TableRowT]):
"""The slice of a prisma table action surface used for budget-metric pagination."""
async def find_many(
self,
*,
skip: int,
take: int,
order: Mapping[str, str],
include: Mapping[str, bool] | None = None,
) -> list[_TableRowT]: ...
async def count(self) -> int: ...
def _paginated_table(repository: BaseRepository[_TableRowT]) -> _PaginatedPrismaTable[_TableRowT]:
"""View a repository's prisma table through the pagination surface budget metrics need."""
return repository.table
class _OrgBudgetRow(Protocol):
"""The budget columns joined onto an organization row."""
@property
def max_budget(self) -> float | None: ...
@property
def budget_reset_at(self) -> datetime | None: ...
class _ExcludedLabelMetric:
"""Proxies a prometheus metric whose declared ``labelnames`` had globally
excluded labels removed, dropping those labels from every ``labels(...)``
@ -1531,7 +1567,7 @@ class PrometheusLogger(CustomLogger):
cache_creation_detail_tokens: Final = PrometheusLogger._resolve_cache_write_tokens(prompt_details)
detail_metrics: Final[list[tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]]] = [
detail_metrics: Final[list[tuple[Any, DEFINED_PROMETHEUS_METRICS, object]]] = [
(
self.litellm_input_cached_tokens_metric,
"litellm_input_cached_tokens_metric",
@ -1584,7 +1620,7 @@ class PrometheusLogger(CustomLogger):
if not isinstance(usage_object, dict):
return
media_metrics: Final[list[tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]]] = [
media_metrics: Final[list[tuple[Any, DEFINED_PROMETHEUS_METRICS, object]]] = [
(
self.litellm_video_duration_seconds_metric,
"litellm_video_duration_seconds_metric",
@ -1606,7 +1642,7 @@ class PrometheusLogger(CustomLogger):
def _inc_sparse_usage_counters(
self,
counters_with_values: list[tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]],
counters_with_values: Sequence[tuple[Any, DEFINED_PROMETHEUS_METRICS, object]],
enum_values: UserAPIKeyLabelValues,
label_context: PrometheusLabelFactoryContext | None = None,
) -> None:
@ -2133,7 +2169,7 @@ class PrometheusLogger(CustomLogger):
def _extract_status_code(
self,
kwargs: dict | None = None,
enum_values: Any | None = None,
enum_values: UserAPIKeyLabelValues | None = None,
exception: Exception | None = None,
) -> int | None:
"""
@ -2151,7 +2187,7 @@ class PrometheusLogger(CustomLogger):
Returns:
Status code as integer if found, None otherwise
"""
status_code = None
status_code: int | None = None
# Try from enum_values first (most common in our callbacks)
if enum_values and hasattr(enum_values, "status_code") and enum_values.status_code:
@ -2225,8 +2261,8 @@ class PrometheusLogger(CustomLogger):
def _should_skip_metrics_for_invalid_key(
self,
kwargs: dict | None = None,
user_api_key_dict: Any | None = None,
enum_values: Any | None = None,
user_api_key_dict: UserAPIKeyAuth | None = None,
enum_values: UserAPIKeyLabelValues | None = None,
standard_logging_payload: dict | StandardLoggingPayload | None = None,
exception: Exception | None = None,
) -> bool:
@ -2391,7 +2427,7 @@ class PrometheusLogger(CustomLogger):
for all successful requests (both streaming and non-streaming).
"""
def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any:
def _safe_get(self, obj: Any, key: str, default: object = None) -> Any:
"""Get value from dict or Pydantic model."""
if obj is None:
return default
@ -3273,8 +3309,8 @@ class PrometheusLogger(CustomLogger):
async def _initialize_budget_metrics(
self,
data_fetch_function: Callable[..., Awaitable[tuple[list[Any], int | None]]],
set_metrics_function: Callable[[list[Any]], Awaitable[None]],
data_fetch_function: Callable[..., Awaitable[tuple[list[_BudgetRowT], int | None]]],
set_metrics_function: Callable[[list[_BudgetRowT]], Awaitable[None]],
data_type: Literal["teams", "keys", "users", "orgs"],
):
"""
@ -3393,12 +3429,12 @@ class PrometheusLogger(CustomLogger):
async def fetch_users(page_size: int, page: int) -> tuple[list[LiteLLM_UserTable], int | None]:
skip: Final = (page - 1) * page_size
users: Final = await UserRepository(prisma_client).table.find_many(
users: Final = await _paginated_table(UserRepository(prisma_client)).find_many(
skip=skip,
take=page_size,
order={"created_at": "desc"},
)
total_count: Final = await UserRepository(prisma_client).table.count()
total_count: Final = await _paginated_table(UserRepository(prisma_client)).count()
return users, total_count
await self._initialize_budget_metrics(
@ -3419,13 +3455,13 @@ class PrometheusLogger(CustomLogger):
async def fetch_orgs(page_size: int, page: int) -> tuple[list, int | None]:
skip: Final = (page - 1) * page_size
orgs: Final = await OrganizationRepository(prisma_client).table.find_many(
orgs: Final = await _paginated_table(OrganizationRepository(prisma_client)).find_many(
skip=skip,
take=page_size,
order={"created_at": "desc"},
include={"litellm_budget_table": True},
)
total_count: Final = await OrganizationRepository(prisma_client).table.count()
total_count: Final = await _paginated_table(OrganizationRepository(prisma_client)).count()
return orgs, total_count
await self._initialize_budget_metrics(
@ -3488,7 +3524,7 @@ class PrometheusLogger(CustomLogger):
try:
# Get total user count
total_users: Final = await UserRepository(prisma_client).table.count()
total_users: Final = await _paginated_table(UserRepository(prisma_client)).count()
self.litellm_total_users_metric.set(total_users)
verbose_logger.debug("Prometheus: set litellm_total_users to %s", total_users)
@ -3497,13 +3533,13 @@ class PrometheusLogger(CustomLogger):
verbose_logger.debug("Prometheus: set litellm_active_users to %s", billable_users)
# Get total team count
total_teams: Final = await TeamRepository(prisma_client).table.count()
total_teams: Final = await _paginated_table(TeamRepository(prisma_client)).count()
self.litellm_teams_count_metric.set(total_teams)
verbose_logger.debug("Prometheus: set litellm_teams_count to %s", total_teams)
except Exception as e:
verbose_logger.exception("Error initializing user/team count metrics: %s", e)
async def _set_key_list_budget_metrics(self, keys: list[str | UserAPIKeyAuth]):
async def _set_key_list_budget_metrics(self, keys: list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken]):
"""Helper function to set budget metrics for a list of keys"""
for key in keys:
if isinstance(key, UserAPIKeyAuth):
@ -3522,7 +3558,7 @@ class PrometheusLogger(CustomLogger):
async def _set_org_list_budget_metrics(self, orgs: list):
"""Helper function to set budget metrics for a list of orgs"""
for org in orgs:
budget_table = getattr(org, "litellm_budget_table", None)
budget_table: _OrgBudgetRow | None = getattr(org, "litellm_budget_table", None)
self._set_org_budget_metrics(
org_id=org.organization_id or "",
org_alias=org.organization_alias or "",
@ -4051,6 +4087,11 @@ class PrometheusLogger(CustomLogger):
verbose_proxy_logger.debug("Starting Prometheus Metrics on /metrics (no authentication)")
def _label_source(enum_values: UserAPIKeyLabelValues) -> Mapping[str, object]:
"""Flatten the label values into the opaque name/value mapping the label filters read."""
return enum_values.model_dump()
def _prometheus_labels_from_context(
supported_enum_labels: list[str],
ctx: PrometheusLabelFactoryContext,
@ -4098,7 +4139,7 @@ def prometheus_label_factory(
return _prometheus_labels_from_context(supported_enum_labels, label_context)
# Extract dictionary from Pydantic object
enum_dict: Final = enum_values.model_dump()
enum_dict: Final = _label_source(enum_values)
# Filter supported labels and sanitize values to prevent breaking
# the Prometheus text format (e.g. U+2028 Line Separator in label values)
@ -4154,7 +4195,7 @@ def get_custom_labels_from_metadata(metadata: dict) -> dict[str, str]:
keys_parts = key.split(".")
# Traverse through the dictionary using the parts
value: Any = metadata
value: object = metadata
for part in keys_parts:
if isinstance(value, dict):
value = value.get(part, None) # Get the value, return None if not found
@ -4171,7 +4212,7 @@ def get_custom_labels_from_metadata(metadata: dict) -> dict[str, str]:
def _get_combined_custom_metadata_from_standard_logging_payload(
standard_logging_payload: dict | None,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Combine the metadata sources that can supply custom Prometheus labels.

View file

@ -165,7 +165,7 @@ class PromptManagementBase(ABC):
ignore_prompt_manager_optional_params: bool | None = False,
) -> tuple[str, list[AllMessageValues], dict]:
if prompt_id is None:
raise ValueError("prompt_id is required for Prompt Management Base class")
return model, messages, non_default_params
if not self.should_run_prompt_management(
prompt_id=prompt_id,
prompt_spec=prompt_spec,

View file

@ -9,13 +9,14 @@ across pods or stop races; the hook reads active jobs through a short-TTL cache.
import asyncio
import hashlib
import random
import traceback
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from itertools import groupby
from operator import itemgetter
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from typing import TYPE_CHECKING, Final, Literal
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError, field_validator, model_validator
@ -32,6 +33,7 @@ from litellm.litellm_core_utils.llm_judge import (
parse_json_verdict,
)
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.types.management_endpoints.auto_router_endpoints import ShadowEvalDirection
from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN
@ -55,7 +57,7 @@ _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 = 500
JUDGE_MAX_OUTPUT_TOKENS: Final = 1500
_MAX_ERROR_CHARS: Final = 500
@ -305,16 +307,21 @@ Criteria: correctness, completeness, clarity, conciseness.
Return ONLY valid JSON in this exact format, no other text:
{
"preference": "A" | "B" | "tie",
"confidence": <0.0 to 1.0>,
"reasoning": "<one sentence>"
"confidence": <0.0 to 1.0>
}"""
class PairwiseVerdict(BaseModel):
"""The judge's blind A/B verdict, validated at the parse boundary."""
"""The judge's blind A/B verdict: the response_format schema sent with the judge call
and the validation contract on its reply. Both fields are required and preference is
closed over the prompt's labels, so a malformed or truncated reply is an
unparseable-verdict error row, never a defaulted or fabricated verdict."""
preference: str = "tie"
confidence: float = 0.0
preference: Literal["A", "B", "tie"]
confidence: float
PAIRWISE_JUDGE_RESPONSE_FORMAT: Final = type_to_response_format_param(PairwiseVerdict)
def _sample_hits(request_id: str, job_id: str, percentage: float) -> bool:
@ -325,6 +332,14 @@ def _sample_hits(request_id: str, job_id: str, percentage: float) -> bool:
return bucket * 100.0 < percentage
def _failure_detail(e: BaseException) -> str:
"""Exception class, message, and the raising frame, so an attempt's error row names
the faulty code path without needing debug logs on the pod."""
frames: Final = traceback.extract_tb(e.__traceback__)
location: Final = f" at {frames[-1].filename.rsplit('/', 1)[-1]}:{frames[-1].lineno}" if frames else ""
return f"{type(e).__name__}{location}: {e}"
def _judge_call_cost(response: object) -> float:
"""Price a judge call, treating an unmapped judge model as free rather than fatal."""
import litellm
@ -764,7 +779,9 @@ class ShadowEvalLogger(CustomLogger):
try:
response: Final = await router.acompletion(
model=target_model,
messages=messages, # pyright: ignore[reportArgumentType] # snapshot of the SDK's own message dicts
messages=[ # mutable-ok: provider transforms rewrite messages in place, so the router gets its own copy
dict(m) for m in messages
], # pyright: ignore[reportArgumentType] # snapshot of the SDK's own message dicts
metadata=shadow_metadata,
num_retries=0,
fallbacks=[], # mutable-ok: SDK kwarg; a failed shadow is a recorded error, never a spend multiplier
@ -772,7 +789,7 @@ class ShadowEvalLogger(CustomLogger):
)
except Exception as e: # noqa: BLE001 # provider errors become error rows, not crashes
verbose_logger.debug("shadow_eval: router call failed: %s", e)
return _CallFailure(f"shadow router call failed: {e}")
return _CallFailure(f"shadow router call failed: {_failure_detail(e)}")
text: Final = _chat_final_text(response)
if not text:
return _CallFailure("shadow router returned an empty response")
@ -815,6 +832,7 @@ class ShadowEvalLogger(CustomLogger):
judge_messages, # pyright: ignore[reportArgumentType] # plain SDK message dicts
temperature=0,
max_tokens=JUDGE_MAX_OUTPUT_TOKENS,
response_format=PAIRWISE_JUDGE_RESPONSE_FORMAT,
metadata=judge_metadata,
)
except Exception as e: # noqa: BLE001 # judge outages become error rows, not crashes

View file

@ -12,6 +12,8 @@ import uuid
from collections.abc import AsyncIterator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
from typing_extensions import ReadOnly
import litellm
from litellm._logging import verbose_logger
from litellm.anthropic_interface import messages as anthropic_messages
@ -90,6 +92,20 @@ class _SearchToolConfig(TypedDict, total=False):
litellm_params: Mapping[str, object] | None
class _DeploymentKwargsView(TypedDict):
"""Typed reads of the untyped request kwargs seen by the deployment hook."""
custom_llm_provider: ReadOnly[str]
litellm_params: ReadOnly[Mapping[str, object]]
model: ReadOnly[str]
class _UserAuthView(TypedDict):
"""Typed read of the optional team attached to the caller's auth object."""
team_id: ReadOnly[str | None]
class WebSearchInterceptionLogger(CustomLogger):
"""
CustomLogger that intercepts WebSearch tool calls for models that don't
@ -265,7 +281,9 @@ class WebSearchInterceptionLogger(CustomLogger):
)
return response
async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None:
async def async_pre_call_deployment_hook(
self, kwargs: dict[str, Any], call_type: CallTypes | None
) -> dict[str, object] | None:
"""
Pre-call hook to convert native Anthropic web_search tools to regular tools.
@ -275,12 +293,17 @@ class WebSearchInterceptionLogger(CustomLogger):
"""
# Check if this is for an enabled provider
# Try top-level kwargs first, then nested litellm_params, then derive from model name
custom_llm_provider = kwargs.get("custom_llm_provider", "") or kwargs.get("litellm_params", {}).get(
kwargs_view: Final[_DeploymentKwargsView] = {
"custom_llm_provider": kwargs.get("custom_llm_provider", ""),
"litellm_params": kwargs.get("litellm_params", {}),
"model": kwargs.get("model", ""),
}
custom_llm_provider = kwargs_view["custom_llm_provider"] or kwargs_view["litellm_params"].get(
"custom_llm_provider", ""
)
if not custom_llm_provider:
try:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=kwargs.get("model", ""))
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=kwargs_view["model"])
except Exception:
custom_llm_provider = ""
if custom_llm_provider not in self.enabled_providers:
@ -1422,7 +1445,8 @@ class WebSearchInterceptionLogger(CustomLogger):
valid_token=user_api_key_auth,
)
team_id: Final = getattr(user_api_key_auth, "team_id", None)
auth_view: Final[_UserAuthView] = {"team_id": getattr(user_api_key_auth, "team_id", None)}
team_id: Final = auth_view["team_id"]
if team_id:
from litellm.proxy.proxy_server import (
prisma_client,

View file

@ -12,6 +12,8 @@ from collections.abc import Mapping
from pathlib import Path
from typing import Final
CLI_TOKEN_FRESHNESS_BUFFER_SECONDS: Final = 360
def get_cli_token_file_path() -> str:
"""Get the path to the CLI token file"""
@ -72,13 +74,18 @@ def get_litellm_gateway_api_key(
return token_data["key"]
def is_cli_token_fresh(token_data: Mapping[str, object], buffer_hours: float = 0.1) -> bool:
def is_cli_token_fresh(
token_data: Mapping[str, object], buffer_hours: float = CLI_TOKEN_FRESHNESS_BUFFER_SECONDS / 3600
) -> bool:
"""Check whether a cached CLI token (as stored in token.json) is still
within its expiration window. Used by `lite auth print-token` to fail
fast, without a network round trip, once the cached token is past
`LITELLM_CLI_JWT_EXPIRATION_HOURS`."""
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
expires_at: Final = token_data.get("expires_at")
if isinstance(expires_at, (int, float)):
return time.time() < expires_at - buffer_hours * 3600
timestamp: Final = token_data.get("timestamp")
if not isinstance(timestamp, (int, float)):
return False

View file

@ -13,6 +13,7 @@ import traceback
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime as dt_object
from functools import lru_cache
from types import MappingProxyType, TracebackType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
from httpx import Response
@ -63,6 +64,10 @@ from litellm.integrations.mlflow import MlflowLogger
from litellm.integrations.sqs import SQSLogger
from litellm.litellm_core_utils.core_helpers import reconstruct_model_name
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import (
cost_breakdown_with_guardrail,
guardrail_information_cost,
)
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
StandardBuiltInToolCostTracking,
)
@ -107,6 +112,7 @@ from litellm.types.utils import (
LiteLLMBatch,
LiteLLMLoggingBaseClass,
LiteLLMRealtimeStreamLoggingObject,
ModelInfo,
ModelResponse,
ModelResponseStream,
RawRequestTypedDict,
@ -306,6 +312,95 @@ def _get_cached_prometheus_logger():
return _PrometheusLogger
_DEPLOYMENT_PRICING_KEYS: Final = (
"input_cost_per_token",
"output_cost_per_token",
"input_cost_per_token_batches",
"output_cost_per_token_batches",
)
def deployment_pricing_model_info(model_id: str | None, deployment_model: str | None) -> ModelInfo | None:
"""Pricing the router registered under this deployment's model_info.id.
Returns None when the deployment declares no pricing of its own, so the
caller falls back to the global cost map. The raw registration is what
decides that: the router registers an entry for every deployment, and
get_model_info fills absent costs with 0, so asking it directly cannot
tell "configured as free" apart from "no pricing configured". A deployment
may declare only one side of its pricing, so the side it leaves out keeps
the model's published rates instead of billing as zero. Ownership is per
token direction: declaring either rate for a direction takes that whole
direction, so a published batch rate can never displace a standard rate
the deployment configured itself.
"""
if model_id is None:
return None
registered: Final = litellm.model_cost.get(model_id)
if not isinstance(registered, dict) or not any(registered.get(key) is not None for key in _DEPLOYMENT_PRICING_KEYS):
return None
try:
merged: Final = litellm.get_model_info(model=model_id).copy()
except Exception: # noqa: BLE001 # get_model_info raises for ids it cannot resolve a provider for
return None
published: Final = _published_pricing(deployment_model)
if published is None:
return merged
declares_input: Final = (
registered.get("input_cost_per_token") is not None or registered.get("input_cost_per_token_batches") is not None
)
declares_output: Final = (
registered.get("output_cost_per_token") is not None
or registered.get("output_cost_per_token_batches") is not None
)
if not declares_input:
merged["input_cost_per_token"] = published.get("input_cost_per_token")
merged["input_cost_per_token_batches"] = published.get("input_cost_per_token_batches")
if not declares_output:
merged["output_cost_per_token"] = published.get("output_cost_per_token")
merged["output_cost_per_token_batches"] = published.get("output_cost_per_token_batches")
return merged
def _published_pricing(deployment_model: str | None) -> ModelInfo | None:
"""The cost map's own entry for the deployment's model, when it resolves."""
if deployment_model is None:
return None
try:
return litellm.get_model_info(model=deployment_model)
except Exception: # noqa: BLE001 # no published entry to layer the declared rates over
return None
def _resolve_vertex_location_for_cost(
custom_llm_provider: str | None,
litellm_params: Mapping[str, object] | None,
optional_params: Mapping[str, object] | None,
model: str,
) -> str | None:
"""
The Vertex AI location a request was served from, resolved the same way
dispatch resolves it, so regional deployments price with the
regional-endpoint uplift. None for non-Vertex providers.
Chat dispatch reads the location from request kwargs, which reach this
logging object through optional_params: on the proxy the logging object is
created before the router picks a deployment, so the deployment's location
never lands in litellm_params.
"""
if custom_llm_provider is None or not custom_llm_provider.startswith("vertex_ai"):
return None
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
empty: Final[Mapping[str, object]] = MappingProxyType({})
configured_location: Final = (
VertexBase.explicit_vertex_ai_location(optional_params or empty)
or VertexBase.explicit_vertex_ai_location(litellm_params or empty)
or VertexBase.safe_get_vertex_ai_location(empty)
)
return VertexBase.get_vertex_region(configured_location, model)
class Logging(LiteLLMLoggingBaseClass):
global \
supabaseClient, \
@ -578,6 +673,28 @@ class Logging(LiteLLMLoggingBaseClass):
return model_id
return None
def get_deployment_model_for_cost(self) -> str | None:
"""The provider-qualified model to price against.
On a batch retrieve both self.model and litellm_params["model"] can be
unset, and self.model can otherwise carry the router's model_group alias,
which no cost map resolves. model_call_details holds the deployment's own
provider-qualified model, so it is preferred.
"""
candidates: Final = (
(self.model_call_details or {}).get("model") if hasattr(self, "model_call_details") else None,
self.litellm_params.get("model") if hasattr(self, "litellm_params") else None,
self.model,
)
return next((candidate for candidate in candidates if isinstance(candidate, str) and candidate), None)
def get_router_deployment_model_info(self) -> ModelInfo | None:
"""See deployment_pricing_model_info; None means fall back to the global cost map."""
return deployment_pricing_model_info(
model_id=self.get_router_model_id(),
deployment_model=self.get_deployment_model_for_cost(),
)
def update_environment_variables(
self,
litellm_params: dict,
@ -715,8 +832,8 @@ class Logging(LiteLLMLoggingBaseClass):
eg. AnthropicCacheControlHook and BedrockKnowledgeBaseHook both don't require a `prompt_id` to be passed in, they are triggered by dynamic params
"""
for param in non_default_params:
if param in DynamicPromptManagementParamLiteral.list_all_params():
for param in DynamicPromptManagementParamLiteral.list_all_params():
if non_default_params.get(param):
return True
#############################################################################
@ -849,6 +966,23 @@ class Logging(LiteLLMLoggingBaseClass):
return None
@staticmethod
def _prompt_manager_runs_without_prompt_id(
logger: CustomLogger,
prompt_spec: PromptSpec | None,
dynamic_callback_params: StandardCallbackDynamicParams | None,
) -> bool:
if not isinstance(logger, CustomPromptManagement):
return False
try:
return logger.should_run_prompt_management(
prompt_id=None,
prompt_spec=prompt_spec,
dynamic_callback_params=dynamic_callback_params or StandardCallbackDynamicParams(),
)
except Exception:
return False
def get_custom_logger_for_prompt_management(
self,
model: str,
@ -899,8 +1033,13 @@ class Logging(LiteLLMLoggingBaseClass):
callback_type=CustomPromptManagement
)
if prompt_management_loggers:
logger: Final = prompt_management_loggers[0]
for logger in prompt_management_loggers:
if prompt_id is None and not self._prompt_manager_runs_without_prompt_id(
logger=logger,
prompt_spec=prompt_spec,
dynamic_callback_params=dynamic_callback_params,
):
continue
self.model_call_details["prompt_integration"] = logger.__class__.__name__
return logger
@ -1006,10 +1145,10 @@ class Logging(LiteLLMLoggingBaseClass):
data=additional_args.get("complete_input_dict", {}),
)
_metadata["raw_request"] = str(curl_command)
_metadata["raw_request"] = _redact_string(str(curl_command))
# split up, so it's easier to parse in the UI
self.model_call_details["raw_request_typed_dict"] = RawRequestTypedDict(
raw_request_api_base=str(additional_args.get("api_base") or ""),
raw_request_api_base=self._get_masked_api_base(str(additional_args.get("api_base") or "")),
raw_request_body=self._get_raw_request_body(additional_args.get("complete_input_dict", {})),
# NOTE: setting ignore_sensitive_headers to True will cause
# the Authorization header to be leaked when calls to the health
@ -1023,8 +1162,10 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["raw_request_typed_dict"] = RawRequestTypedDict(
error=str(e),
)
_metadata["raw_request"] = f"Unable to Log \
_metadata["raw_request"] = _redact_string(
f"Unable to Log \
raw request: {e}"
)
if getattr(self, "logger_fn", None) and callable(self.logger_fn):
try:
self.logger_fn(
@ -1118,15 +1259,16 @@ class Logging(LiteLLMLoggingBaseClass):
if _is_debugging_on() or self.litellm_request_debug:
if json_logs:
masked_headers: Final = self._get_masked_headers(headers)
masked_api_base: Final = self._get_masked_api_base(str(api_base or ""))
if self.litellm_request_debug:
verbose_logger.warning( # .warning ensures this shows up in all environments
"POST Request Sent from LiteLLM",
extra={"api_base": {api_base}, **masked_headers},
extra={"api_base": {masked_api_base}, **masked_headers},
)
else:
verbose_logger.debug(
"POST Request Sent from LiteLLM",
extra={"api_base": {api_base}, **masked_headers},
extra={"api_base": {masked_api_base}, **masked_headers},
)
else:
headers = additional_args.get("headers", {})
@ -1166,8 +1308,6 @@ class Logging(LiteLLMLoggingBaseClass):
curl_command = "\nRequest Sent from LiteLLM:\n"
request_str: Final = additional_args.get("request_str", "")
curl_command += request_str
elif api_base == "":
curl_command = str(self.model_call_details)
return curl_command
def _get_masked_headers(self, headers: dict, ignore_sensitive_headers: bool = False) -> dict:
@ -1189,6 +1329,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["additional_args"] = additional_args
self.model_call_details["log_event_type"] = "post_api_call"
attr: Literal["warning", "debug"]
if self.litellm_request_debug:
attr = "warning"
else:
@ -1342,6 +1483,7 @@ class Logging(LiteLLMLoggingBaseClass):
reasoning_cost: float | None = None,
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
) -> None:
"""
Helper method to store cost breakdown in the logging object.
@ -1360,6 +1502,7 @@ class Logging(LiteLLMLoggingBaseClass):
margin_total_amount: Total margin added in USD
service_tier: Tier the costs above were priced on, already resolved
data_residency: Region uplift the costs above were priced on, already resolved
vertex_location: Vertex AI location the costs above were priced on, already resolved
"""
self.cost_breakdown = CostBreakdown(
@ -1369,6 +1512,7 @@ class Logging(LiteLLMLoggingBaseClass):
tool_usage_cost=cost_for_built_in_tools_cost_usd_dollar,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
)
if cache_read_cost is not None and cache_read_cost > 0:
self.cost_breakdown["cache_read_cost"] = cache_read_cost
@ -1484,6 +1628,12 @@ class Logging(LiteLLMLoggingBaseClass):
if hasattr(self, "litellm_params") and self.litellm_params
else None
),
"vertex_location": _resolve_vertex_location_for_cost(
custom_llm_provider=self.model_call_details.get("custom_llm_provider", None),
litellm_params=(self.litellm_params if hasattr(self, "litellm_params") else None),
optional_params=self.optional_params,
model=litellm_model_name or self.model,
),
}
except Exception as e: # error creating kwargs for cost calculation
debug_info = StandardLoggingModelCostFailureDebugInformation(
@ -1802,7 +1952,7 @@ class Logging(LiteLLMLoggingBaseClass):
if self.model_call_details.get("litellm_params") is None:
return
metadata_hidden_params: Final = hidden_params.copy()
response_cost: Final = self.model_call_details.get("response_cost")
response_cost: Final[object] = self.model_call_details.get("response_cost")
if metadata_hidden_params.get("response_cost") is None and response_cost is not None:
metadata_hidden_params["response_cost"] = response_cost
@ -1844,7 +1994,10 @@ class Logging(LiteLLMLoggingBaseClass):
logging_result, start_time, end_time
)
if (standard_logging_payload := self.model_call_details.get("standard_logging_object")) is not None:
standard_logging_payload: Final[StandardLoggingPayload | None] = self.model_call_details.get(
"standard_logging_object"
)
if standard_logging_payload is not None:
emit_standard_logging_payload(standard_logging_payload)
def _build_standard_logging_payload(
@ -2109,7 +2262,7 @@ class Logging(LiteLLMLoggingBaseClass):
def _success_handler_body(
self,
result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml)
result: object = None,
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
cache_hit: bool | None = None,
@ -2150,7 +2303,10 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["standard_logging_object"] = self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
)
if (standard_logging_payload := self.model_call_details.get("standard_logging_object")) is not None:
standard_logging_payload: Final[StandardLoggingPayload | None] = self.model_call_details.get(
"standard_logging_object"
)
if standard_logging_payload is not None:
# Only emit for sync requests (async_success_handler handles async)
if is_sync_request:
emit_standard_logging_payload(standard_logging_payload)
@ -2592,7 +2748,9 @@ class Logging(LiteLLMLoggingBaseClass):
) = await _handle_completed_batch(
batch=result,
custom_llm_provider=self.custom_llm_provider,
model_name=self.get_deployment_model_for_cost(),
litellm_params=self.litellm_params,
model_info=self.get_router_deployment_model_info(),
)
result._hidden_params["response_cost"] = response_cost
@ -2981,7 +3139,7 @@ class Logging(LiteLLMLoggingBaseClass):
global_callbacks=litellm.failure_callback,
)
result = None # result sent to all loggers, init this to None incase it's not created
result: object = None # result sent to all loggers, init this to None incase it's not created
result = redact_message_input_output_from_logging(
model_call_details=(self.model_call_details if hasattr(self, "model_call_details") else {}),
@ -3395,11 +3553,11 @@ class Logging(LiteLLMLoggingBaseClass):
def _get_assembled_streaming_response(
self,
result: ModelResponse | TextCompletionResponse | ModelResponseStream | ResponseCompletedEvent | Any,
result: ModelResponse | TextCompletionResponse | ModelResponseStream | ResponseCompletedEvent | object,
start_time: datetime.datetime,
end_time: datetime.datetime,
is_async: bool,
streaming_chunks: list[Any],
streaming_chunks: list[object],
) -> ModelResponse | TextCompletionResponse | ResponsesAPIResponse | None:
if self.stream is not True:
return None
@ -3677,9 +3835,7 @@ def set_callbacks(callback_list, function_id=None):
from sentry_sdk.scrubber import EventScrubber
sentry_sdk_instance = sentry_sdk
sentry_trace_rate = (
os.environ.get("SENTRY_API_TRACE_RATE") if "SENTRY_API_TRACE_RATE" in os.environ else "1.0"
)
sentry_trace_rate = os.environ.get("SENTRY_API_TRACE_RATE", "1.0")
sentry_sample_rate = (
os.environ.get("SENTRY_API_SAMPLE_RATE") if "SENTRY_API_SAMPLE_RATE" in os.environ else "1.0"
)
@ -5150,13 +5306,13 @@ class StandardLoggingPayloadSetup:
# ProxyException uses .code, LiteLLM exceptions use .status_code,
# httpx.HTTPStatusError exposes status only as .response.status_code.
# Stringified for Prisma JSON compatibility.
error_code_attr: Final = getattr(original_exception, "code", None)
error_code_attr: Final[object] = getattr(original_exception, "code", None)
if error_code_attr is not None and str(error_code_attr) not in ("", "None"):
error_status: str = str(error_code_attr)
else:
status_code_attr = getattr(original_exception, "status_code", None)
status_code_attr: object = getattr(original_exception, "status_code", None)
if status_code_attr is None:
response_attr: Final = getattr(original_exception, "response", None)
response_attr: Final[object] = getattr(original_exception, "response", None)
status_code_attr = getattr(response_attr, "status_code", None)
error_status = str(status_code_attr) if status_code_attr is not None else ""
error_class: Final[str] = str(original_exception.__class__.__name__) if original_exception else ""
@ -5165,7 +5321,7 @@ class StandardLoggingPayloadSetup:
# Get traceback information (first 100 lines)
traceback_info = traceback_str or ""
if original_exception:
tb: Final = getattr(original_exception, "__traceback__", None)
tb: Final[TracebackType | None] = getattr(original_exception, "__traceback__", None)
if tb:
tb_lines: Final = traceback.format_tb(tb)
traceback_info += "".join(tb_lines[:MAXIMUM_TRACEBACK_LINES_TO_LOG]) # Limit to first 100 lines
@ -5276,11 +5432,11 @@ class StandardLoggingPayloadSetup:
"""
dynamic_litellm_session_id: Final = litellm_params.get("litellm_session_id")
dynamic_litellm_trace_id: Final = litellm_params.get("litellm_trace_id")
metadata: Final = litellm_params.get("metadata")
metadata: Final[Mapping[str, object] | None] = litellm_params.get("metadata")
metadata_session_id: Final = metadata.get("session_id") if metadata else None
metadata_trace_id: Final = metadata.get("trace_id") if metadata else None
ordered_candidates: Final[tuple[Any, Any, Any, Any]] = (
ordered_candidates: Final[tuple[object, object, object, object]] = (
(dynamic_litellm_trace_id, dynamic_litellm_session_id, metadata_trace_id, metadata_session_id)
if litellm.request_correlation_in_logs
else (dynamic_litellm_session_id, dynamic_litellm_trace_id, metadata_session_id, metadata_trace_id)
@ -5305,10 +5461,10 @@ class StandardLoggingPayloadSetup:
"""
if not litellm.request_correlation_in_logs:
return ""
dynamic_litellm_session_id: Final = litellm_params.get("litellm_session_id")
dynamic_litellm_session_id: Final[object] = litellm_params.get("litellm_session_id")
if dynamic_litellm_session_id:
return str(dynamic_litellm_session_id)
metadata: Final = litellm_params.get("metadata")
metadata: Final[Mapping[str, object] | None] = litellm_params.get("metadata")
metadata_session_id: Final = metadata.get("session_id") if metadata else None
if metadata_session_id:
return str(metadata_session_id)
@ -5559,12 +5715,14 @@ def get_standard_logging_object_payload(
base_model = metadata.get("deployment")
custom_pricing: Final = use_custom_pricing_for_model(litellm_params=litellm_params)
raw_response_cost: Final = kwargs.get("response_cost")
response_cost: Final[float] = raw_response_cost or 0.0
llm_response_cost: Final[float] = raw_response_cost or 0.0
guardrail_cost: Final = guardrail_information_cost(metadata.get("standard_logging_guardrail_information"))
response_cost: Final[float] = llm_response_cost + guardrail_cost
# clean up litellm hidden params
clean_hidden_params: Final = StandardLoggingPayloadSetup.get_hidden_params(hidden_params)
if clean_hidden_params["response_cost"] is None and raw_response_cost is not None:
clean_hidden_params["response_cost"] = response_cost
clean_hidden_params["response_cost"] = llm_response_cost
model_cost_information: Final = StandardLoggingPayloadSetup.get_model_cost_information(
base_model=base_model,
@ -5644,7 +5802,7 @@ def get_standard_logging_object_payload(
metadata=clean_metadata,
cache_key=clean_hidden_params["cache_key"],
response_cost=response_cost,
cost_breakdown=logging_obj.cost_breakdown,
cost_breakdown=cost_breakdown_with_guardrail(logging_obj.cost_breakdown, guardrail_cost),
total_tokens=usage_dict.get("total_tokens", 0),
prompt_tokens=usage_dict.get("prompt_tokens", 0),
completion_tokens=usage_dict.get("completion_tokens", 0),

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