mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into fireworks_nim_vllm_compat
This commit is contained in:
commit
ffc9b87317
678 changed files with 28971 additions and 9546 deletions
|
|
@ -17,3 +17,24 @@
|
|||
|
||||
# style: unify ruff format width on 120 (#31518)
|
||||
48b5a5a0cc5a694a11219416ee0b6eb6e620e74e
|
||||
|
||||
# refactor(imports): move collections.abc names out of typing (#35495)
|
||||
397e8e4918777e4e60a7f5e88699e0a9a7dabb3d
|
||||
|
||||
# refactor(lint): apply every safe ruff autofix and zero 28 strict-rule budgets (#35495)
|
||||
b604e2b20c6db2099085a2f0e59b7e99e87eed6f
|
||||
|
||||
# refactor(logging): drop redundant !s conversion flags from f-strings (#35546)
|
||||
7b2d3440cba3160277470f7a0180098ae9b87864
|
||||
|
||||
# perf: build log messages lazily so filtered-out log records cost nothing (#35703)
|
||||
c9887a1f94bc1e7e4bdfe64d640f0509a0bc19dd
|
||||
|
||||
# feat(lint): enforce Final on locals and freeze function parameters (#35807)
|
||||
2708620d6a599cc73c1950a942d26ac26a7ed3d4
|
||||
|
||||
# chore(lint): remove litellm/types from the ruff lint exclusion (#35926)
|
||||
4e32a8bf6a1e1af1e04b67c759841ccef44b2235
|
||||
|
||||
# chore(lint): strip inert type: ignore comments and zero LIT009/LIT010/LIT011 headroom (#35928)
|
||||
338e411103ad5d7003e97f34f04fa36bca542dbe
|
||||
|
|
|
|||
159
.github/ci-coverage-allowlist.yml
vendored
Normal file
159
.github/ci-coverage-allowlist.yml
vendored
Normal file
|
|
@ -0,0 +1,159 @@
|
|||
description: >-
|
||||
Paths deliberately outside CI coverage, each with the reason it is exempt.
|
||||
assert_ci_coverage.py fails when a test file or Dockerfile is neither invoked
|
||||
by a job nor listed here, so every entry below is a decision on the record.
|
||||
|
||||
test_paths:
|
||||
- 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
|
||||
paths:
|
||||
- tests/e2e
|
||||
- reason: >-
|
||||
The documentation and code-quality workflows execute four files in this directory by name as
|
||||
scripts and pytest never collects the directory, so these six run nowhere; listed individually
|
||||
so a seventh cannot inherit the exemption
|
||||
paths:
|
||||
- tests/documentation_tests/test_exception_types.py
|
||||
- tests/documentation_tests/test_general_setting_keys.py
|
||||
- tests/documentation_tests/test_optional_params.py
|
||||
- tests/documentation_tests/test_readme_providers.py
|
||||
- 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
|
||||
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
|
||||
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
|
||||
paths:
|
||||
- tests/vector_store_tests/rag/test_rag_bedrock.py
|
||||
- tests/vector_store_tests/rag/test_rag_openai.py
|
||||
- tests/vector_store_tests/rag/test_rag_s3_vectors.py
|
||||
- tests/vector_store_tests/rag/test_rag_vertex_ai.py
|
||||
- tests/vector_store_tests/test_azure_ai_vector_store.py
|
||||
- tests/vector_store_tests/test_azure_vector_store.py
|
||||
- tests/vector_store_tests/test_bedrock_vector_store.py
|
||||
- tests/vector_store_tests/test_gemini_vector_store.py
|
||||
- tests/vector_store_tests/test_milvus_vector_store.py
|
||||
- tests/vector_store_tests/test_openai_vector_store.py
|
||||
- tests/vector_store_tests/test_ragflow_vector_store.py
|
||||
- tests/vector_store_tests/test_s3_vectors_vector_store.py
|
||||
- tests/vector_store_tests/test_vertex_ai_search_api_vector_store.py
|
||||
- tests/vector_store_tests/test_vertex_ai_vector_store.py
|
||||
- reason: >-
|
||||
Throughput and memory-growth measurements whose runtime and variance make them unsuitable for
|
||||
a per-pull-request job
|
||||
paths:
|
||||
- tests/load_tests/test_datadog_load_test.py
|
||||
- tests/load_tests/test_langsmith_load_test.py
|
||||
- tests/load_tests/test_linear_memory_growth.py
|
||||
- tests/load_tests/test_memory_usage.py
|
||||
- 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: >-
|
||||
Third-party integration tests that skip themselves without OCI configuration or sandbox
|
||||
credentials, neither of which a pull request job holds
|
||||
paths:
|
||||
- tests/integration/sandbox/test_e2b_sandbox.py
|
||||
- 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
|
||||
paths:
|
||||
- tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py
|
||||
|
||||
dockerfiles:
|
||||
- reason: >-
|
||||
The componentized images the microservices chart deploys are built by no job; wiring both into
|
||||
the scan workflow costs a full image build each and is deferred to a change that prices the
|
||||
whole set
|
||||
paths:
|
||||
- backend/Dockerfile
|
||||
- gateway/Dockerfile
|
||||
- reason: >-
|
||||
The dashboard container is a static Next.js export served by nginx, and the dashboard build
|
||||
and lint workflows already exercise that output, so building the image adds no signal about it
|
||||
paths:
|
||||
- ui/Dockerfile
|
||||
- reason: >-
|
||||
The Rust gateway ships as its own chart and package with a separate release pipeline, so its
|
||||
image is not part of this repo's Python image set
|
||||
paths:
|
||||
- litellm-rust/crates/ai-gateway/Dockerfile
|
||||
- reason: >-
|
||||
An example image under cookbook/ that is documentation rather than a shipped artifact
|
||||
paths:
|
||||
- cookbook/litellm-ollama-docker-image/Dockerfile
|
||||
262
.github/scripts/assert_ci_coverage.py
vendored
Normal file
262
.github/scripts/assert_ci_coverage.py
vendored
Normal file
|
|
@ -0,0 +1,262 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import pathlib
|
||||
import re
|
||||
import sys
|
||||
from collections.abc import Iterable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
|
||||
import yaml
|
||||
|
||||
REPO_ROOT = pathlib.Path(__file__).resolve().parents[2]
|
||||
WORKFLOW_DIR = REPO_ROOT / ".github" / "workflows"
|
||||
CIRCLECI_CONFIG = REPO_ROOT / ".circleci" / "config.yml"
|
||||
ALLOWLIST_FILE = REPO_ROOT / ".github" / "ci-coverage-allowlist.yml"
|
||||
TESTS_ROOT = REPO_ROOT / "tests"
|
||||
|
||||
ALLOWLIST_KEYS = frozenset({"description", "test_paths", "dockerfiles"})
|
||||
PATH_FILTER_KEYS = frozenset({"paths", "paths-ignore"})
|
||||
TEST_PATH_KEYS = frozenset({"test-path", "test-paths"})
|
||||
DOCKERFILE_INPUT_KEYS = frozenset({"file", "dockerfile"})
|
||||
TEST_RUNNER_RE = re.compile(r"\bpytest\b|\bcircleci tests\b|\bhelm unittest\b|\bplaywright test\b|\bpython[0-9.]*\s")
|
||||
IMAGE_BUILD_RE = re.compile(r"\bdocker\s+(?:buildx\s+)?build\b")
|
||||
TEST_TOKEN_RE = re.compile(r"tests/[A-Za-z0-9_./*?-]+")
|
||||
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("*?")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AllowEntry:
|
||||
paths: tuple[str, ...]
|
||||
reason: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Allowlist:
|
||||
test_paths: tuple[AllowEntry, ...]
|
||||
dockerfiles: tuple[AllowEntry, ...]
|
||||
|
||||
def covers_test(self, relative_path: str) -> bool:
|
||||
return any(_token_covers(path, relative_path) for entry in self.test_paths for path in entry.paths)
|
||||
|
||||
def covers_dockerfile(self, relative_path: str) -> bool:
|
||||
return any(relative_path == path for entry in self.dockerfiles for path in entry.paths)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Scalar:
|
||||
key: str
|
||||
value: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Finding:
|
||||
subject: str
|
||||
detail: str
|
||||
|
||||
|
||||
def _scalars(node: object, key: str) -> tuple[Scalar, ...]:
|
||||
if isinstance(node, str):
|
||||
return (Scalar(key=key, value=node),)
|
||||
if isinstance(node, Mapping):
|
||||
return tuple(
|
||||
scalar
|
||||
for child_key, value in node.items()
|
||||
if child_key not in PATH_FILTER_KEYS
|
||||
for scalar in _scalars(value, str(child_key))
|
||||
)
|
||||
if isinstance(node, Sequence):
|
||||
return tuple(scalar for item in node for scalar in _scalars(item, key))
|
||||
return ()
|
||||
|
||||
|
||||
def _config_files() -> tuple[pathlib.Path, ...]:
|
||||
workflows = tuple(sorted(path for path in WORKFLOW_DIR.iterdir() if path.suffix in (".yml", ".yaml")))
|
||||
circleci = (CIRCLECI_CONFIG,) if CIRCLECI_CONFIG.is_file() else ()
|
||||
return workflows + circleci
|
||||
|
||||
|
||||
def _all_scalars() -> tuple[Scalar, ...]:
|
||||
return tuple(
|
||||
scalar
|
||||
for path in _config_files()
|
||||
for scalar in _scalars(yaml.safe_load(path.read_text(encoding="utf-8")), path.name)
|
||||
)
|
||||
|
||||
|
||||
def _uncommented(value: str) -> str:
|
||||
return COMMENT_RE.sub("", value)
|
||||
|
||||
|
||||
def _invoked_test_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
|
||||
return frozenset(
|
||||
match.group(0).rstrip("/")
|
||||
for scalar in scalars
|
||||
if scalar.key in TEST_PATH_KEYS or TEST_RUNNER_RE.search(scalar.value)
|
||||
for match in TEST_TOKEN_RE.finditer(_uncommented(scalar.value))
|
||||
)
|
||||
|
||||
|
||||
def _built_dockerfile_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
|
||||
return frozenset(
|
||||
match.group(0)
|
||||
for scalar in scalars
|
||||
if scalar.key in DOCKERFILE_INPUT_KEYS or IMAGE_BUILD_RE.search(scalar.value)
|
||||
for match in DOCKERFILE_TOKEN_RE.finditer(_uncommented(scalar.value))
|
||||
)
|
||||
|
||||
|
||||
def _glob_to_regex(token: str) -> re.Pattern[str]:
|
||||
parts = re.split(r"(\*\*/|\*\*|\*|\?)", token)
|
||||
translated = "".join(
|
||||
{"**/": r"(?:.*/)?", "**": r".*", "*": r"[^/]*", "?": r"[^/]"}.get(part, re.escape(part)) for part in parts
|
||||
)
|
||||
return re.compile(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 relative_path == token or relative_path.startswith(f"{token}/")
|
||||
|
||||
|
||||
def _test_files() -> tuple[str, ...]:
|
||||
return tuple(
|
||||
sorted(
|
||||
path.relative_to(REPO_ROOT).as_posix()
|
||||
for path in TESTS_ROOT.rglob("test_*.py")
|
||||
if path.is_file() and "node_modules" not in path.parts
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _dockerfiles() -> tuple[str, ...]:
|
||||
return tuple(
|
||||
sorted(
|
||||
path.relative_to(REPO_ROOT).as_posix()
|
||||
for path in REPO_ROOT.rglob("Dockerfile*")
|
||||
if path.is_file()
|
||||
and ".git" not in path.parts
|
||||
and "node_modules" not in path.parts
|
||||
and not path.name.endswith(".dockerignore")
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _uncovered_tests(allowlist: Allowlist, tokens: frozenset[str]) -> tuple[Finding, ...]:
|
||||
uncovered = tuple(
|
||||
relative_path
|
||||
for relative_path in _test_files()
|
||||
if not any(_token_covers(token, relative_path) for token in tokens) and not allowlist.covers_test(relative_path)
|
||||
)
|
||||
directories = tuple(dict.fromkeys(path.rsplit("/", 1)[0] for path in uncovered))
|
||||
return tuple(
|
||||
Finding(
|
||||
subject=directory,
|
||||
detail=_describe(tuple(p for p in uncovered if p.rsplit("/", 1)[0] == directory)),
|
||||
)
|
||||
for directory in directories
|
||||
)
|
||||
|
||||
|
||||
def _describe(paths: tuple[str, ...]) -> str:
|
||||
names = ", ".join(path.rsplit("/", 1)[1] for path in paths[:3])
|
||||
suffix = f", +{len(paths) - 3} more" if len(paths) > 3 else ""
|
||||
return f"{len(paths)} test file(s) invoked by no job: {names}{suffix}"
|
||||
|
||||
|
||||
def _uncovered_dockerfiles(allowlist: Allowlist, tokens: frozenset[str]) -> tuple[Finding, ...]:
|
||||
return tuple(
|
||||
Finding(subject=relative_path, detail="built by no job")
|
||||
for relative_path in _dockerfiles()
|
||||
if relative_path not in tokens and not allowlist.covers_dockerfile(relative_path)
|
||||
)
|
||||
|
||||
|
||||
def _parse_entry(item: object, section: str) -> AllowEntry:
|
||||
if not isinstance(item, dict):
|
||||
raise SystemExit(f"{ALLOWLIST_FILE.name}: '{section}' entries must be mappings")
|
||||
paths = item.get("paths")
|
||||
reason = item.get("reason")
|
||||
if (
|
||||
not isinstance(paths, list)
|
||||
or not paths
|
||||
or not all(isinstance(path, str) for path in paths)
|
||||
or not isinstance(reason, str)
|
||||
or not reason.strip()
|
||||
):
|
||||
raise SystemExit(
|
||||
f"{ALLOWLIST_FILE.name}: every '{section}' entry needs a non-empty 'paths' "
|
||||
"list of strings and a non-empty 'reason'"
|
||||
)
|
||||
return AllowEntry(paths=tuple(paths), reason=reason)
|
||||
|
||||
|
||||
def _parse_entries(raw: object, section: str) -> tuple[AllowEntry, ...]:
|
||||
if not isinstance(raw, list):
|
||||
raise SystemExit(f"{ALLOWLIST_FILE.name}: '{section}' must be a list")
|
||||
return tuple(_parse_entry(item, section) for item in raw)
|
||||
|
||||
|
||||
def _load_allowlist() -> Allowlist:
|
||||
if not ALLOWLIST_FILE.is_file():
|
||||
return Allowlist(test_paths=(), dockerfiles=())
|
||||
raw = yaml.safe_load(ALLOWLIST_FILE.read_text(encoding="utf-8")) or {}
|
||||
if not isinstance(raw, dict):
|
||||
raise SystemExit(f"{ALLOWLIST_FILE.name}: top level must be a mapping")
|
||||
unknown = sorted(str(key) for key in raw if key not in ALLOWLIST_KEYS)
|
||||
if unknown:
|
||||
raise SystemExit(
|
||||
f"{ALLOWLIST_FILE.name}: unknown top-level key(s) {unknown}; expected only {sorted(ALLOWLIST_KEYS)}"
|
||||
)
|
||||
return Allowlist(
|
||||
test_paths=_parse_entries(raw.get("test_paths", []), "test_paths"),
|
||||
dockerfiles=_parse_entries(raw.get("dockerfiles", []), "dockerfiles"),
|
||||
)
|
||||
|
||||
|
||||
def _write(message: str) -> None:
|
||||
sys.stdout.write(f"{message}\n")
|
||||
|
||||
|
||||
def _report(title: str, findings: tuple[Finding, ...], remedy: str) -> None:
|
||||
_write(f"ERROR: {title}")
|
||||
for finding in findings:
|
||||
_write(f" - {finding.subject}: {finding.detail}")
|
||||
_write("")
|
||||
_write(remedy)
|
||||
_write("")
|
||||
|
||||
|
||||
def main() -> int:
|
||||
allowlist = _load_allowlist()
|
||||
scalars = _all_scalars()
|
||||
|
||||
test_findings = _uncovered_tests(allowlist, _invoked_test_tokens(scalars))
|
||||
dockerfile_findings = _uncovered_dockerfiles(allowlist, _built_dockerfile_tokens(scalars))
|
||||
|
||||
if test_findings:
|
||||
_report(
|
||||
"test files that no CI job invokes",
|
||||
test_findings,
|
||||
"Add each to a job's test path, or list it in .github/ci-coverage-allowlist.yml with a reason.",
|
||||
)
|
||||
if dockerfile_findings:
|
||||
_report(
|
||||
"Dockerfiles that no CI job builds",
|
||||
dockerfile_findings,
|
||||
"Build each in a workflow, or list it in .github/ci-coverage-allowlist.yml with a reason.",
|
||||
)
|
||||
if test_findings or dockerfile_findings:
|
||||
return 1
|
||||
|
||||
_write(
|
||||
f"OK: {len(_test_files())} test files and {len(_dockerfiles())} Dockerfiles are each "
|
||||
"invoked by at least one job or carry an explicit allowlist entry."
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
42
.github/workflows/ci-coverage.yml
vendored
Normal file
42
.github/workflows/ci-coverage.yml
vendored
Normal file
|
|
@ -0,0 +1,42 @@
|
|||
name: "CI Coverage"
|
||||
|
||||
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:
|
||||
assert-ci-coverage:
|
||||
name: assert-ci-coverage
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
permissions:
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Assert every test file and Dockerfile is invoked by a job
|
||||
run: |
|
||||
python -m pip install "pyyaml==6.0.3"
|
||||
python .github/scripts/assert_ci_coverage.py
|
||||
|
|
@ -13,35 +13,16 @@ jobs:
|
|||
contents: write
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
persist-credentials: false
|
||||
|
||||
- name: Create daily oss-agent-shin branch
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
# Configure Git user
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||
|
||||
# Generate branch name with MM_DD_YYYY format
|
||||
BRANCH_NAME="litellm_oss_agent_shin_$(date +'%m_%d_%Y')"
|
||||
echo "Creating branch: $BRANCH_NAME"
|
||||
|
||||
# Fetch all branches
|
||||
git fetch --all
|
||||
|
||||
# Check if the branch already exists
|
||||
if git show-ref --verify --quiet refs/remotes/origin/$BRANCH_NAME; then
|
||||
if gh api "repos/${{ github.repository }}/git/ref/heads/$BRANCH_NAME" --silent 2>/dev/null; then
|
||||
echo "Branch $BRANCH_NAME already exists. Skipping creation."
|
||||
else
|
||||
echo "Creating new branch: $BRANCH_NAME"
|
||||
# Create the new branch from main
|
||||
git checkout -b $BRANCH_NAME origin/main
|
||||
# Push the new branch
|
||||
git push origin $BRANCH_NAME
|
||||
echo "Successfully created and pushed branch: $BRANCH_NAME"
|
||||
exit 0
|
||||
fi
|
||||
MAIN_SHA=$(gh api "repos/${{ github.repository }}/git/ref/heads/main" --jq '.object.sha')
|
||||
gh api "repos/${{ github.repository }}/git/refs" -f ref="refs/heads/$BRANCH_NAME" -f sha="$MAIN_SHA" --silent
|
||||
echo "Successfully created branch: $BRANCH_NAME at $MAIN_SHA"
|
||||
|
|
|
|||
|
|
@ -13,38 +13,19 @@ jobs:
|
|||
contents: write
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
persist-credentials: false
|
||||
|
||||
- name: Create daily staging branch
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
# Configure Git user
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||
|
||||
# Generate branch name with MM_DD_YYYY format
|
||||
BRANCH_NAME="litellm_oss_staging_$(date +'%m_%d_%Y')"
|
||||
echo "Creating branch: $BRANCH_NAME"
|
||||
|
||||
# Fetch all branches
|
||||
git fetch --all
|
||||
|
||||
# Check if the branch already exists
|
||||
if git show-ref --verify --quiet refs/remotes/origin/$BRANCH_NAME; then
|
||||
if gh api "repos/${{ github.repository }}/git/ref/heads/$BRANCH_NAME" --silent 2>/dev/null; then
|
||||
echo "Branch $BRANCH_NAME already exists. Skipping creation."
|
||||
else
|
||||
echo "Creating new branch: $BRANCH_NAME"
|
||||
# Create the new branch from main
|
||||
git checkout -b $BRANCH_NAME origin/main
|
||||
# Push the new branch
|
||||
git push origin $BRANCH_NAME
|
||||
echo "Successfully created and pushed branch: $BRANCH_NAME"
|
||||
exit 0
|
||||
fi
|
||||
MAIN_SHA=$(gh api "repos/${{ github.repository }}/git/ref/heads/main" --jq '.object.sha')
|
||||
gh api "repos/${{ github.repository }}/git/refs" -f ref="refs/heads/$BRANCH_NAME" -f sha="$MAIN_SHA" --silent
|
||||
echo "Successfully created branch: $BRANCH_NAME at $MAIN_SHA"
|
||||
|
||||
create-internal-dev-branch:
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
|
|
@ -53,35 +34,16 @@ jobs:
|
|||
contents: write
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
persist-credentials: false
|
||||
|
||||
- name: Create internal dev branch
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
# Configure Git user
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||
|
||||
# Generate branch name with MM_DD_YYYY format
|
||||
BRANCH_NAME="litellm_internal_dev_$(date +'%m_%d_%Y')"
|
||||
echo "Creating branch: $BRANCH_NAME"
|
||||
|
||||
# Fetch all branches
|
||||
git fetch --all
|
||||
|
||||
# Check if the branch already exists
|
||||
if git show-ref --verify --quiet refs/remotes/origin/$BRANCH_NAME; then
|
||||
if gh api "repos/${{ github.repository }}/git/ref/heads/$BRANCH_NAME" --silent 2>/dev/null; then
|
||||
echo "Branch $BRANCH_NAME already exists. Skipping creation."
|
||||
else
|
||||
echo "Creating new branch: $BRANCH_NAME"
|
||||
# Create the new branch from main
|
||||
git checkout -b $BRANCH_NAME origin/main
|
||||
# Push the new branch
|
||||
git push origin $BRANCH_NAME
|
||||
echo "Successfully created and pushed branch: $BRANCH_NAME"
|
||||
exit 0
|
||||
fi
|
||||
MAIN_SHA=$(gh api "repos/${{ github.repository }}/git/ref/heads/main" --jq '.object.sha')
|
||||
gh api "repos/${{ github.repository }}/git/refs" -f ref="refs/heads/$BRANCH_NAME" -f sha="$MAIN_SHA" --silent
|
||||
echo "Successfully created branch: $BRANCH_NAME at $MAIN_SHA"
|
||||
|
|
|
|||
31
.github/workflows/helm_unit_test.yml
vendored
31
.github/workflows/helm_unit_test.yml
vendored
|
|
@ -23,21 +23,28 @@ jobs:
|
|||
with:
|
||||
version: "3.11.1"
|
||||
|
||||
- name: Download and verify Helm Unit Test Plugin
|
||||
run: |
|
||||
curl -fsSLo "$RUNNER_TEMP/helm-unittest.tgz" https://github.com/helm-unittest/helm-unittest/releases/download/v0.8.2/helm-unittest-linux-amd64-0.8.2.tgz
|
||||
echo "56ab3091e6fa52a7c92ee951def9bed957f295d9ce98483aed404e748d7b3a94 $RUNNER_TEMP/helm-unittest.tgz" | sha256sum -c -
|
||||
|
||||
- name: Install Helm Unit Test Plugin
|
||||
run: |
|
||||
helm plugin install https://github.com/helm-unittest/helm-unittest --version v0.4.4
|
||||
- name: Verify Helm Unit Test Plugin integrity
|
||||
run: |
|
||||
EXPECTED_SHA="e251ba198448629678ff2168e1a469249d998155"
|
||||
PLUGIN_DIR="$(helm env HELM_PLUGINS)/helm-unittest"
|
||||
ACTUAL_SHA="$(git -C "$PLUGIN_DIR" rev-parse HEAD)"
|
||||
if [ "$ACTUAL_SHA" != "$EXPECTED_SHA" ]; then
|
||||
echo "::error::Helm unittest plugin checksum mismatch! Expected $EXPECTED_SHA but got $ACTUAL_SHA"
|
||||
exit 1
|
||||
fi
|
||||
echo "Helm unittest plugin integrity verified: $ACTUAL_SHA"
|
||||
mkdir -p "$PLUGIN_DIR"
|
||||
tar -xzf "$RUNNER_TEMP/helm-unittest.tgz" -C "$PLUGIN_DIR"
|
||||
helm plugin list
|
||||
|
||||
- name: Run unit tests
|
||||
run: |
|
||||
helm unittest -f 'tests/*.yaml' helm/litellm-helm
|
||||
helm unittest -f 'tests/*.yaml' helm/litellm
|
||||
for chart in helm/litellm-helm helm/litellm; do
|
||||
declared="$(grep -h '^suite:' "$chart"/tests/*.yaml | wc -l | tr -d '[:space:]')"
|
||||
output="$(mktemp)"
|
||||
helm unittest -f 'tests/*.yaml' "$chart" | tee "$output"
|
||||
executed="$(sed -n 's/^Test Suites:.*[[:space:]]\([0-9][0-9]*\) total$/\1/p' "$output")"
|
||||
if [ "$declared" != "$executed" ]; then
|
||||
echo "::error::$chart declares $declared test suites but helm-unittest ran $executed. Suites are being skipped silently, so their assertions never execute."
|
||||
exit 1
|
||||
fi
|
||||
echo "$chart: all $declared declared test suites ran"
|
||||
done
|
||||
|
|
|
|||
33
.github/workflows/image-scan.yml
vendored
33
.github/workflows/image-scan.yml
vendored
|
|
@ -8,10 +8,12 @@ on:
|
|||
- litellm_oss_branch
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- Dockerfile
|
||||
- docker/Dockerfile.non_root
|
||||
- migrations/Dockerfile
|
||||
- migrations/run.py
|
||||
- tests/proxy_migration_tests/test_offline_image_migration.py
|
||||
- litellm-proxy-extras/**
|
||||
- tests/proxy_migration_tests/**
|
||||
- uv.lock
|
||||
- ui/litellm-dashboard/package-lock.json
|
||||
- .github/workflows/image-scan.yml
|
||||
|
|
@ -86,6 +88,35 @@ jobs:
|
|||
--fail-on high \
|
||||
--output table
|
||||
|
||||
runtime-image:
|
||||
name: runtime-image
|
||||
runs-on: ubuntu-latest
|
||||
if: >-
|
||||
github.event_name != 'pull_request' ||
|
||||
github.event.pull_request.head.repo.full_name == github.repository
|
||||
timeout-minutes: 30
|
||||
permissions:
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Build runtime image
|
||||
run: docker build -f Dockerfile -t litellm-runtime-scan:${{ github.sha }} .
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Verify offline migration as a non-root uid
|
||||
env:
|
||||
LITELLM_IMAGE: litellm-runtime-scan:${{ github.sha }}
|
||||
run: |
|
||||
python -m pip install "pytest==9.0.3"
|
||||
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py -v
|
||||
|
||||
migrations-image:
|
||||
name: migrations-image
|
||||
runs-on: ubuntu-latest
|
||||
|
|
|
|||
62
.github/workflows/publish-basedpyright-base-counts.yml
vendored
Normal file
62
.github/workflows/publish-basedpyright-base-counts.yml
vendored
Normal file
|
|
@ -0,0 +1,62 @@
|
|||
name: Publish basedpyright base counts
|
||||
|
||||
# Every commit on litellm_internal_staging is some branch's future merge-base.
|
||||
# Publishing its per-rule basedpyright counts as an artifact lets
|
||||
# scripts/type_check_gate.py download them in seconds instead of paying a
|
||||
# 60-110s second basedpyright pass on every fresh worktree or moved merge-base.
|
||||
# No concurrency group on purpose: runs must never cancel each other, because
|
||||
# every sha's artifact matters (any of them can become a merge-base).
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- litellm_internal_staging
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
ref:
|
||||
description: "Ref to compute and publish base counts for"
|
||||
required: false
|
||||
default: litellm_internal_staging
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
publish:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
ref: ${{ inputs.ref || github.sha }}
|
||||
clean: true
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
# The gate provisions its own measurement env (.venv-typecheck: a frozen
|
||||
# uv sync of its canonical dependency groups plus a generated Prisma
|
||||
# client), so no install step here can drift from what local runs measure.
|
||||
- name: Emit basedpyright counts for HEAD
|
||||
env:
|
||||
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
|
||||
run: |
|
||||
python scripts/type_check_gate.py --emit-counts-dir "$RUNNER_TEMP/basedpyright-counts"
|
||||
counts_file=$(ls "$RUNNER_TEMP"/basedpyright-counts/basedpyright-counts-*.json)
|
||||
echo "COUNTS_ARTIFACT_NAME=$(basename "$counts_file" .json)" >> "$GITHUB_ENV"
|
||||
|
||||
- name: Upload counts artifact
|
||||
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
|
||||
with:
|
||||
name: ${{ env.COUNTS_ARTIFACT_NAME }}
|
||||
path: ${{ runner.temp }}/basedpyright-counts/
|
||||
if-no-files-found: error
|
||||
53
.github/workflows/test-linting.yml
vendored
53
.github/workflows/test-linting.yml
vendored
|
|
@ -15,6 +15,12 @@ jobs:
|
|||
lint:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
# actions: read lets scripts/type_check_gate.py download the base-counts
|
||||
# artifact published by publish-basedpyright-base-counts.yml instead of
|
||||
# re-running basedpyright over the merge-base tree.
|
||||
permissions:
|
||||
contents: read
|
||||
actions: read
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
|
|
@ -23,10 +29,21 @@ jobs:
|
|||
# Any-discipline) would otherwise blame on this branch.
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha }}
|
||||
fetch-depth: 0
|
||||
fetch-depth: 1
|
||||
clean: true
|
||||
persist-credentials: false
|
||||
|
||||
- name: Fetch gate base (merge-base with target branch)
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
run: |
|
||||
MERGE_BASE=$(gh api "repos/${{ github.repository }}/compare/${BASE_SHA}...${HEAD_SHA}?per_page=1" --jq '.merge_base_commit.sha')
|
||||
test -n "$MERGE_BASE"
|
||||
git fetch --no-tags --depth=1 origin "$MERGE_BASE"
|
||||
echo "GATE_BASE_SHA=$MERGE_BASE" >> "$GITHUB_ENV"
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
|
|
@ -60,10 +77,8 @@ jobs:
|
|||
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Check ruff format
|
||||
env:
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
run: |
|
||||
git diff --name-only --diff-filter=ACMR "$BASE_SHA"...HEAD -- 'litellm/**/*.py' | grep -v '^litellm/enterprise/' > "$RUNNER_TEMP/ruff_format_files.txt" || true
|
||||
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
|
||||
echo "No changed litellm Python files to check with ruff format."
|
||||
exit 0
|
||||
|
|
@ -86,16 +101,12 @@ jobs:
|
|||
cd ..
|
||||
|
||||
- name: Check strict-rule budget (delta vs base)
|
||||
env:
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
run: |
|
||||
uv run --no-sync python scripts/ruff_strict_gate.py --base "$BASE_SHA"
|
||||
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)
|
||||
env:
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
run: |
|
||||
uv run --no-sync python scripts/type_discipline_gate.py --base "$BASE_SHA"
|
||||
uv run --no-sync python scripts/type_discipline_gate.py --base "$GATE_BASE_SHA"
|
||||
|
||||
- name: Print OpenAI version
|
||||
run: |
|
||||
|
|
@ -103,15 +114,14 @@ jobs:
|
|||
|
||||
- name: Check basedpyright budget (delta vs base)
|
||||
env:
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
|
||||
run: |
|
||||
uv run --no-sync python scripts/type_check_gate.py --base "$BASE_SHA"
|
||||
uv run --no-sync python scripts/type_check_gate.py --base "$GATE_BASE_SHA"
|
||||
|
||||
- name: Check tests/e2e basedpyright (zero errors)
|
||||
env:
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
run: |
|
||||
if git diff --name-only --diff-filter=ACMRD "$BASE_SHA"...HEAD -- 'tests/e2e/**/*.py' | grep -q .; then
|
||||
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
|
||||
else
|
||||
echo "No changed tests/e2e Python files; skipping."
|
||||
|
|
@ -140,9 +150,15 @@ jobs:
|
|||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
fetch-depth: 1
|
||||
persist-credentials: false
|
||||
|
||||
- name: Fetch ratchet base
|
||||
env:
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
run: |
|
||||
git fetch --no-tags --depth=1 origin "$BASE_SHA"
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
|
|
@ -163,7 +179,7 @@ jobs:
|
|||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
fetch-depth: 1
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
|
|
@ -178,13 +194,14 @@ jobs:
|
|||
|
||||
- name: Run secret scan test
|
||||
run: |
|
||||
uv run --frozen --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/litellm/test_no_hardcoded_secrets.py -v
|
||||
|
||||
- name: Run ggshield secret scan
|
||||
env:
|
||||
GITGUARDIAN_API_KEY: ${{ secrets.GITGUARDIAN_API_KEY }}
|
||||
run: |
|
||||
if [ -n "$GITGUARDIAN_API_KEY" ]; then
|
||||
git fetch --no-tags --unshallow origin
|
||||
uv tool run --from 'ggshield==1.48.0' ggshield secret scan repo .
|
||||
else
|
||||
echo "GITGUARDIAN_API_KEY not set, skipping ggshield scan"
|
||||
|
|
|
|||
7
.github/workflows/test-litellm-ui-lint.yml
vendored
7
.github/workflows/test-litellm-ui-lint.yml
vendored
|
|
@ -22,12 +22,13 @@ jobs:
|
|||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
fetch-depth: 1
|
||||
persist-credentials: false
|
||||
|
||||
- name: Collect changed files
|
||||
id: changed
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
run: |
|
||||
|
|
@ -37,7 +38,9 @@ jobs:
|
|||
# landed since, so a PR that touches no UI file still gets linted
|
||||
# against hundreds of other people's files. Diff the PR head against its
|
||||
# own merge base instead, which is exactly what this PR changed.
|
||||
merge_base=$(git merge-base "$BASE_SHA" "$HEAD_SHA")
|
||||
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"
|
||||
: > "$RUNNER_TEMP/prettier_files.txt"
|
||||
: > "$RUNNER_TEMP/eslint_files.txt"
|
||||
while IFS= read -r f; do
|
||||
|
|
|
|||
19
.github/workflows/test-litellm-ui-unit.yml
vendored
19
.github/workflows/test-litellm-ui-unit.yml
vendored
|
|
@ -29,7 +29,7 @@ jobs:
|
|||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
fetch-depth: 1
|
||||
persist-credentials: false
|
||||
|
||||
- name: Setup Node.js
|
||||
|
|
@ -45,11 +45,24 @@ jobs:
|
|||
- name: Run UI unit tests (Vitest)
|
||||
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
|
||||
echo "Pull request: running only tests related to changes since $BASE_SHA"
|
||||
npm run test -- --run --changed "$BASE_SHA" --passWithNoTests \
|
||||
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
|
||||
echo "Push to $GITHUB_REF_NAME: running the full suite"
|
||||
|
|
|
|||
4
.github/workflows/test-unit-misc.yml
vendored
4
.github/workflows/test-unit-misc.yml
vendored
|
|
@ -40,7 +40,11 @@ jobs:
|
|||
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
|
||||
|
|
|
|||
|
|
@ -29,7 +29,9 @@ jobs:
|
|||
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
|
||||
|
|
|
|||
2
.github/workflows/test-unit-proxy-infra.yml
vendored
2
.github/workflows/test-unit-proxy-infra.yml
vendored
|
|
@ -33,6 +33,8 @@ jobs:
|
|||
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
|
||||
|
|
|
|||
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -1,5 +1,6 @@
|
|||
.python-version
|
||||
.venv
|
||||
.venv-typecheck
|
||||
.venv_policy_test
|
||||
.env
|
||||
.claude
|
||||
|
|
|
|||
10
CLAUDE.md
10
CLAUDE.md
|
|
@ -29,7 +29,7 @@ Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We pref
|
|||
|
||||
If you ever make public-facing PR descriptions, comments, issues, commit messages, etc., always follow these guidelines to sound less AI-y:
|
||||
- don't use emojis
|
||||
- don't use "—". Instead, reach for ",", ".", conjunction words, ":", ";", etc. in descending order of preference: vary among them, weighted toward the front of the list, and skip "," where it would cause a comma splice or the sentence is getting long. Overusing any one of them, ";" especially, also feels AI-y
|
||||
- don't use "—". Instead, reach for ",", ".", conjunction words, ":", ";", etc. in descending order of preference: vary among them, weighted toward the front of the list, and skip "," where it would cause a comma splice or the sentence is getting long. Overusing any one of them, ";" especially, also feels AI-y. A word cap does not penalize you for adding more sentences: when writing under tight word budgets, prefer a period split or a conjunction over ";", and keep to at most one ";" per message
|
||||
- don't use the pattern "It's not X, it's Y", "You're not X, you're Y", etc.
|
||||
- don't use bulleted or numbered lists unless it would be nonsensical not to. Instead, prefer prose
|
||||
- don't add a trailing "." at the end of paragraphs (just like this file). That means every paragraph, not just the last one (of the markdown file, PR description, GitHub comment, etc.). Rule of thumb: if you're adding new line(s) before the next sentence, don't add a "."
|
||||
|
|
@ -41,9 +41,11 @@ Python max line length is 120, not 88
|
|||
|
||||
When you fix violations gated by `ruff-strict-budget.json`, `type-discipline-budget.json`, or `basedpyright-code-budget.json`, run `make lint-budget-update` and commit the lowered limits so the ceilings ratchet down instead of leaving stale headroom. It measures the working tree, so it must contain exactly the fixes you're committing
|
||||
|
||||
`make pre-commit` saves its complete output to a log file in .git (overwriting previous pre-commit logs) and prints that path as its first and last output lines. To inspect a run, read or grep that log instead of re-running the multi-minute checks just to see a different slice
|
||||
|
||||
If you're trying to create a new function that relies on untyped stuff, instead of adding more Any's and pushing `reportAny` / `reportExplicitAny` closer to their basedpyright ceilings, just validate it in the caller with Pydantic (a model or `TypeAdapter` that returns the typed thing or raises will do) and then pass the now typed variable in
|
||||
|
||||
If you get an LIT001 or LIT002 fail, refactor the code to follow functional programming best practices rather than introducing mutable data structures. For example, build values in one shot with comprehensions or generators wrapped in `tuple()` / `frozenset()` instead of seeding an empty `list`/`dict`/`set` and mutating it over time. Ideally, `# mutable-ok` is never used; reach for it only as a genuine last resort when an immutable rewrite is truly impossible, and always pair it with a real reason
|
||||
If you get an LIT001 or LIT002 fail, refactor the code to follow functional programming best practices rather than introducing mutable data structures. For example, build values in one shot with comprehensions or generators wrapped in `tuple()` / `MappingProxyType()` / `frozenset()` instead of seeding an empty `list`/`dict`/`set` and mutating it over time. Ideally, `# mutable-ok` is never used; reach for it only as a genuine last resort when an immutable rewrite is truly impossible, and always pair it with a real reason
|
||||
|
||||
Every lint or type suppression must name the exact rule inside brackets and carry a reason comment, e.g. `# pyright: ignore[reportArgumentType] # stubs lack async overload` or `# noqa: TID251 # <reason>`. `# type: ignore` is banned (LIT009): pyrightconfig.json sets `enableTypeIgnoreComments` to false, so it silently does nothing
|
||||
|
||||
|
|
@ -57,7 +59,7 @@ Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages
|
|||
|
||||
When working on a PR, keep the PR description in sync with new commits being made
|
||||
|
||||
Replies/rebuttals to AI PR review bots must be 15-25 word human-readable replies
|
||||
All GitHub comments must be human-readable and 15-25 words max
|
||||
|
||||
Monkeypatching attributes of a class to do testing is an anti-pattern. Prefer dependency-injecting things into classes. That way, at unit test time, you can pass a mocked dependency in
|
||||
|
||||
|
|
@ -70,7 +72,7 @@ Follow these coding conventions for new/updated code (a three-line fix in a lega
|
|||
- Composition over inheritance
|
||||
- Never-nester: early returns over deep nesting
|
||||
- Don't throw; model failures as values (One function (e.g., raise_public) maps error union to existing public exception contracts via exhaustive match + assert_never)
|
||||
- No mutation; don't reassign variables, global or local. Instead of mutable lists and dicts, prefer tuples, frozen dataclasses (with slots=True), etc.
|
||||
- No mutation; don't reassign variables, global or local. Instead of mutable lists and dicts, prefer tuples, frozen dataclasses (with slots=True), `MappingProxyType`, etc.
|
||||
- Annotate every variable with `: Final` (LIT010). Unpacking and walrus targets cannot carry the annotation, so they are implicitly final. Don't rebind them. Never rebind or mutate function parameters (LIT011); `self`/`cls` attribute stores are the exception. If rebinding or in-place mutation is truly unavoidable, suppress with `# rebind-ok: <reason>` explaining why
|
||||
- Use dependency injection
|
||||
- Fully typed; no `Any` or coarse types like `dict[str, Any]` or just `dict`. Every function parameter must be strongly typed
|
||||
|
|
|
|||
|
|
@ -134,7 +134,8 @@ RUN find /app/.venv -type f -path "*/tornado/test/*" -delete && \
|
|||
find /app/.venv -type d -path "*/tornado/test" -delete && \
|
||||
chmod -R a+rX /opt/prisma && \
|
||||
test -x /opt/prisma/binaries/node_modules/.bin/prisma && \
|
||||
test -f /opt/prisma/binaries/node_modules/prisma/build/index.js
|
||||
test -f /opt/prisma/binaries/node_modules/prisma/build/index.js && \
|
||||
python -c "from prisma.client import BINARY_PATHS; paths = list(BINARY_PATHS.query_engine.values()); assert paths and all(p.startswith('/opt/prisma/') for p in paths), paths"
|
||||
|
||||
EXPOSE 4000/tcp
|
||||
|
||||
|
|
|
|||
13
Makefile
13
Makefile
|
|
@ -99,7 +99,10 @@ install-test-deps: install-proxy-dev
|
|||
$(UV_RUN) prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
install-helm-unittest:
|
||||
helm plugin install https://github.com/helm-unittest/helm-unittest --version v0.4.4 || echo "ignore error if plugin exists"
|
||||
@helm plugin list | grep -qE '^unittest[[:space:]]+0\.8\.2([[:space:]]|$$)' || { \
|
||||
helm plugin uninstall unittest >/dev/null 2>&1 || true; \
|
||||
helm plugin install https://github.com/helm-unittest/helm-unittest --version v0.8.2; \
|
||||
}
|
||||
|
||||
# Install git hooks that enforce Conventional Commits and Conventional Branches.
|
||||
# Opt-in: not chained into install-dev.
|
||||
|
|
@ -121,10 +124,10 @@ lint-fetch-base:
|
|||
git fetch origin litellm_internal_staging
|
||||
|
||||
# Mirror test-linting.yml's lint job environment: the proxy-dev group plus a generated
|
||||
# Prisma client, so basedpyright resolves the same modules CI does (without the generated
|
||||
# client the DB wrappers typed against it degrade to Unknown, drifting the budget from
|
||||
# CI's). --inexact tops up the venv instead of pruning the proxy extras gen:api and the
|
||||
# running proxy need.
|
||||
# Prisma client, so `basedpyright tests/e2e` resolves the same modules CI does. The
|
||||
# budget gate itself no longer measures here (scripts/type_check_gate.py provisions its
|
||||
# own .venv-typecheck). --inexact tops up the venv instead of pruning the proxy extras
|
||||
# gen:api and the running proxy need.
|
||||
lint-install:
|
||||
$(UV) sync --inexact --frozen --group proxy-dev --group e2e-dev
|
||||
$(UV_RUN) python scripts/prisma_generate_if_needed.py
|
||||
|
|
|
|||
|
|
@ -59,9 +59,9 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
--extra semantic-router \
|
||||
--python python3
|
||||
|
||||
RUN mkdir -p /home/nonroot && \
|
||||
HOME=/home/nonroot prisma generate --schema=./schema.prisma && \
|
||||
chown -R nonroot:nonroot /home/nonroot/.cache
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
npm_config_cache=/root/.npm \
|
||||
prisma generate --schema=./schema.prisma
|
||||
|
||||
RUN sed -i 's/\r$//' docker/component_entrypoint.sh && chmod +x docker/component_entrypoint.sh
|
||||
|
||||
|
|
@ -83,13 +83,16 @@ ENV HOME=/home/nonroot \
|
|||
PATH="/app/.venv/bin:${PATH}" \
|
||||
PYTHONPATH="/app" \
|
||||
PYTHONDONTWRITEBYTECODE=1 \
|
||||
PYTHONUNBUFFERED=1
|
||||
PYTHONUNBUFFERED=1 \
|
||||
PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries
|
||||
|
||||
COPY --from=builder --chown=nonroot:nonroot /app /app
|
||||
COPY --from=builder --chown=nonroot:nonroot /home/nonroot/.cache /home/nonroot/.cache
|
||||
COPY --from=builder /opt/prisma /opt/prisma
|
||||
|
||||
RUN find /app/.venv -type f -path "*/tornado/test/*" -delete && \
|
||||
find /app/.venv -type d -path "*/tornado/test" -delete
|
||||
find /app/.venv -type d -path "*/tornado/test" -delete && \
|
||||
chmod -R a+rX /opt/prisma && \
|
||||
python -c "from prisma.client import BINARY_PATHS; paths = list(BINARY_PATHS.query_engine.values()); assert paths and all(p.startswith('/opt/prisma/') for p in paths), paths"
|
||||
|
||||
USER nonroot
|
||||
|
||||
|
|
|
|||
|
|
@ -82,6 +82,9 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/user_agent",
|
||||
"/usage/",
|
||||
"/daily/",
|
||||
# Deployment-wide gateway request counts. Scoped to the analytics read rather
|
||||
# than all of /gateway/, which stays free for data-plane routes.
|
||||
"/gateway/daily/",
|
||||
# CloudZero cost-export admin (init / settings / export / dry-run / delete)
|
||||
"/cloudzero/",
|
||||
# Caching admin
|
||||
|
|
|
|||
|
|
@ -1,18 +1,18 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 29809
|
||||
"limit": 28842
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2645
|
||||
"limit": 2634
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 329
|
||||
},
|
||||
"reportAttributeAccessIssue": {
|
||||
"limit": 516
|
||||
"limit": 514
|
||||
},
|
||||
"reportCallIssue": {
|
||||
"limit": 123
|
||||
"limit": 117
|
||||
},
|
||||
"reportConstantRedefinition": {
|
||||
"limit": 40
|
||||
|
|
@ -21,10 +21,10 @@
|
|||
"limit": 215
|
||||
},
|
||||
"reportDuplicateImport": {
|
||||
"limit": 24
|
||||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 9473
|
||||
"limit": 9103
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -54,10 +54,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5855
|
||||
"limit": 5843
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15849
|
||||
"limit": 15816
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -72,7 +72,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportOptionalMemberAccess": {
|
||||
"limit": 1079
|
||||
"limit": 1078
|
||||
},
|
||||
"reportOptionalOperand": {
|
||||
"limit": 0
|
||||
|
|
@ -90,7 +90,7 @@
|
|||
"limit": 8
|
||||
},
|
||||
"reportReturnType": {
|
||||
"limit": 219
|
||||
"limit": 218
|
||||
},
|
||||
"reportTypedDictNotRequiredAccess": {
|
||||
"limit": 27
|
||||
|
|
@ -99,34 +99,34 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 45262
|
||||
"limit": 45110
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 113
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 40452
|
||||
"limit": 39838
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 20309
|
||||
"limit": 20237
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 31978
|
||||
"limit": 31383
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 124
|
||||
"limit": 122
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 703
|
||||
"limit": 701
|
||||
},
|
||||
"reportUnnecessaryContains": {
|
||||
"limit": 5
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 866
|
||||
"limit": 864
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 165
|
||||
"limit": 0
|
||||
},
|
||||
"reportUntypedFunctionDecorator": {
|
||||
"limit": 33
|
||||
|
|
@ -138,9 +138,9 @@
|
|||
"limit": 139
|
||||
},
|
||||
"reportUnusedImport": {
|
||||
"limit": 588
|
||||
"limit": 555
|
||||
},
|
||||
"reportUnusedVariable": {
|
||||
"limit": 147
|
||||
"limit": 146
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -133,7 +133,8 @@ RUN find /app/.venv -type f -path "*/tornado/test/*" -delete && \
|
|||
find /app/.venv -type d -path "*/tornado/test" -delete && \
|
||||
chmod -R a+rX /opt/prisma && \
|
||||
test -x /opt/prisma/binaries/node_modules/.bin/prisma && \
|
||||
test -f /opt/prisma/binaries/node_modules/prisma/build/index.js
|
||||
test -f /opt/prisma/binaries/node_modules/prisma/build/index.js && \
|
||||
python -c "from prisma.client import BINARY_PATHS; paths = list(BINARY_PATHS.query_engine.values()); assert paths and all(p.startswith('/opt/prisma/') for p in paths), paths"
|
||||
|
||||
EXPOSE 4000/tcp
|
||||
|
||||
|
|
|
|||
|
|
@ -185,7 +185,8 @@ RUN mkdir -p /nonexistent /app/.cache /var/lib/litellm/assets /var/lib/litellm/u
|
|||
chmod -R a+rX /opt/prisma && \
|
||||
test -x /opt/prisma/binaries/node_modules/.bin/prisma && \
|
||||
test -f /opt/prisma/binaries/node_modules/prisma/build/index.js && \
|
||||
ls /opt/prisma/binaries/node_modules/@prisma/engines/query-engine-* >/dev/null 2>&1
|
||||
ls /opt/prisma/binaries/node_modules/@prisma/engines/query-engine-* >/dev/null 2>&1 && \
|
||||
python -c "from prisma.client import BINARY_PATHS; paths = list(BINARY_PATHS.query_engine.values()); assert paths and all(p.startswith('/opt/prisma/') for p in paths), paths"
|
||||
|
||||
USER 65534
|
||||
|
||||
|
|
|
|||
|
|
@ -47,7 +47,13 @@ RUN uv venv --python python && \
|
|||
"prisma==0.11.0" \
|
||||
"openai==2.24.0"
|
||||
|
||||
RUN prisma generate --schema=./schema.prisma
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
npm_config_cache=/root/.npm \
|
||||
prisma generate --schema=./schema.prisma && \
|
||||
chmod -R a+rX /opt/prisma && \
|
||||
python -c "import sys; from prisma.client import BINARY_PATHS; bad = sorted(p for group in BINARY_PATHS.model_dump().values() for p in group.values() if not p.startswith('/opt/prisma/')); sys.exit('prisma engines baked outside /opt/prisma: %r' % bad) if bad else None"
|
||||
|
||||
ENV PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries
|
||||
|
||||
EXPOSE 4000/tcp
|
||||
|
||||
|
|
|
|||
|
|
@ -296,17 +296,13 @@ class CheckBatchCost:
|
|||
underlying provider model (e.g. ``gpt-5.5``), which no key is allowed to call.
|
||||
"""
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
convert_b64_uid_to_unified_uid,
|
||||
get_models_from_unified_file_id,
|
||||
resolve_managed_output_file_model_name,
|
||||
)
|
||||
|
||||
input_file_id = cls._get_input_file_id(job)
|
||||
target_model_names = (
|
||||
get_models_from_unified_file_id(convert_b64_uid_to_unified_uid(input_file_id)) if input_file_id else []
|
||||
return resolve_managed_output_file_model_name(
|
||||
unified_input_file_id=cls._get_input_file_id(job),
|
||||
fallback_model_name=deployment_info.model_name or None,
|
||||
)
|
||||
if target_model_names:
|
||||
return ",".join(target_model_names)
|
||||
return deployment_info.model_name or None
|
||||
|
||||
@staticmethod
|
||||
def _get_input_file_id(job: "LiteLLM_ManagedObjectTable") -> Optional[str]:
|
||||
|
|
@ -502,6 +498,7 @@ class CheckBatchCost:
|
|||
},
|
||||
"metadata": {
|
||||
"user_api_key_user_id": creator_user_id,
|
||||
"user_api_key_team_id": getattr(job, "team_id", None),
|
||||
**user_info,
|
||||
},
|
||||
},
|
||||
|
|
@ -660,6 +657,20 @@ class CheckBatchCost:
|
|||
|
||||
elif response.status in ("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(),
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
"""
|
||||
Polls LiteLLM_ManagedObjectTable to check if the response is complete.
|
||||
Cost tracking is handled automatically by litellm.aget_responses().
|
||||
Cost tracking is handled automatically by the get-responses call.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Dict, Optional, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -13,11 +13,15 @@ from litellm.constants import (
|
|||
MAX_OBJECTS_PER_POLL_CYCLE,
|
||||
STALE_OBJECT_CLEANUP_BATCH_SIZE,
|
||||
)
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.router import Router
|
||||
|
||||
TERMINAL_RESPONSE_STATUSES = frozenset({"completed", "failed", "cancelled", "incomplete"})
|
||||
|
||||
|
||||
class CheckResponsesCost:
|
||||
def __init__(
|
||||
|
|
@ -33,6 +37,28 @@ class CheckResponsesCost:
|
|||
self.prisma_client: PrismaClient = prisma_client
|
||||
self.llm_router: Router = llm_router
|
||||
|
||||
async def _get_response(
|
||||
self,
|
||||
response_id: str,
|
||||
litellm_metadata: Dict[str, str],
|
||||
) -> ResponsesAPIResponse:
|
||||
"""Fetch the upstream response, using deployment credentials when available.
|
||||
|
||||
LiteLLM-encoded response IDs carry the ``model_id`` of the deployment that
|
||||
served the original request, so routing through ``llm_router`` applies that
|
||||
deployment's ``api_base`` / ``api_key`` / ``api_version``, exactly like
|
||||
``GET /v1/responses/{id}`` does. ``litellm.aget_responses`` on its own only
|
||||
sees provider env vars, so it fails for every deployment whose credentials
|
||||
live in the config; the row then never leaves ``queued``.
|
||||
"""
|
||||
model_id: Optional[str] = ResponsesAPIRequestUtils.get_model_id_from_response_id(response_id)
|
||||
if model_id is None or self.llm_router.get_deployment(model_id=model_id) is None:
|
||||
return await litellm.aget_responses(response_id=response_id, litellm_metadata=litellm_metadata)
|
||||
router_response = await self.llm_router.aget_responses(
|
||||
response_id=response_id, litellm_metadata=litellm_metadata
|
||||
)
|
||||
return cast(ResponsesAPIResponse, router_response)
|
||||
|
||||
async def _expire_stale_rows(
|
||||
self, cutoff: datetime, batch_size: int
|
||||
) -> int:
|
||||
|
|
@ -87,8 +113,8 @@ class CheckResponsesCost:
|
|||
Check if background responses are complete and track their cost.
|
||||
- Get all status="queued" or "in_progress" and file_purpose="response" jobs
|
||||
- Query the provider to check if response is complete
|
||||
- Cost is automatically tracked by litellm.aget_responses()
|
||||
- Mark completed/failed/cancelled responses as complete in the database
|
||||
- Cost is automatically tracked by the get-responses call
|
||||
- Mark responses in a terminal state as complete in the database
|
||||
"""
|
||||
try:
|
||||
await self._cleanup_stale_managed_objects()
|
||||
|
|
@ -134,7 +160,7 @@ class CheckResponsesCost:
|
|||
litellm_metadata["model"] = model_name
|
||||
litellm_metadata["model_group"] = model_name # Use same value for model_group
|
||||
|
||||
response = await litellm.aget_responses(
|
||||
response = await self._get_response(
|
||||
response_id=responses_id_security,
|
||||
litellm_metadata=litellm_metadata,
|
||||
)
|
||||
|
|
@ -144,21 +170,14 @@ class CheckResponsesCost:
|
|||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.info(
|
||||
verbose_proxy_logger.warning(
|
||||
f"Skipping job {unified_object_id} due to error: {e}"
|
||||
)
|
||||
continue
|
||||
|
||||
# Check if response is in a terminal state
|
||||
if response.status == "completed":
|
||||
if response.status in TERMINAL_RESPONSE_STATUSES:
|
||||
verbose_proxy_logger.info(
|
||||
f"Response {unified_object_id} is complete. Cost automatically tracked by aget_responses."
|
||||
)
|
||||
completed_jobs.append(job)
|
||||
|
||||
elif response.status in ["failed", "cancelled"]:
|
||||
verbose_proxy_logger.info(
|
||||
f"Response {unified_object_id} has status {response.status}, marking as complete"
|
||||
f"Response {unified_object_id} has terminal status {response.status}, marking as complete"
|
||||
)
|
||||
completed_jobs.append(job)
|
||||
|
||||
|
|
|
|||
|
|
@ -4,9 +4,11 @@
|
|||
import base64
|
||||
import json
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Dict, Final, List, Literal, Optional, Union, cast
|
||||
from uuid import NAMESPACE_URL, uuid5
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm import Router, verbose_logger
|
||||
|
|
@ -33,8 +35,8 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
get_batch_id_from_unified_batch_id,
|
||||
get_content_type_from_file_object,
|
||||
get_model_id_from_unified_batch_id,
|
||||
get_models_from_unified_file_id,
|
||||
normalize_mime_type_for_provider,
|
||||
resolve_managed_output_file_model_name,
|
||||
)
|
||||
from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccessIssue]
|
||||
AllMessageValues,
|
||||
|
|
@ -73,6 +75,26 @@ else:
|
|||
PrismaClient = Any
|
||||
|
||||
|
||||
def _parse_managed_file_object(
|
||||
raw_file_object: object, unified_file_id: str
|
||||
) -> Optional[OpenAIFileObject]:
|
||||
if raw_file_object is None:
|
||||
return None
|
||||
try:
|
||||
return OpenAIFileObject.model_validate(raw_file_object)
|
||||
except ValidationError as e:
|
||||
verbose_logger.warning(
|
||||
f"Failed to parse managed file object {unified_file_id}: "
|
||||
f"{e.errors(include_input=False, include_url=False, include_context=False)}"
|
||||
)
|
||||
return None
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
f"Failed to parse managed file object {unified_file_id}: {type(e).__name__}"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
||||
# Class variables or attributes
|
||||
def __init__(
|
||||
|
|
@ -383,9 +405,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
}
|
||||
)
|
||||
return [
|
||||
OpenAIFileObject.model_validate(file_object.file_object)
|
||||
for file_object in file_ids
|
||||
if file_object.file_object is not None
|
||||
parsed_file_object.model_copy(update={"id": row.unified_file_id})
|
||||
for row in file_ids
|
||||
if (
|
||||
parsed_file_object := _parse_managed_file_object(
|
||||
row.file_object, row.unified_file_id
|
||||
)
|
||||
)
|
||||
is not None
|
||||
]
|
||||
|
||||
async def check_managed_file_id_access(
|
||||
|
|
@ -1059,10 +1086,13 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
def get_unified_output_file_id(
|
||||
self, output_file_id: str, model_id: str, model_name: Optional[str]
|
||||
) -> str:
|
||||
deterministic_uuid: Final = uuid5(
|
||||
uuid5(NAMESPACE_URL, model_id), output_file_id
|
||||
)
|
||||
unified_output_file_id = (
|
||||
SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
|
||||
"application/json",
|
||||
str(uuid.uuid4()),
|
||||
str(deterministic_uuid),
|
||||
model_name or "",
|
||||
output_file_id,
|
||||
model_id,
|
||||
|
|
@ -1098,21 +1128,13 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
) # managed batch id
|
||||
model_id = cast(Optional[str], response._hidden_params.get("model_id"))
|
||||
model_name = cast(Optional[str], response._hidden_params.get("model_name"))
|
||||
resolved_model_name = model_name
|
||||
|
||||
# Some providers (e.g. Vertex batch retrieve) do not set model_name on
|
||||
# the response. In that case, recover target_model_names from the input
|
||||
# managed file metadata so unified output IDs preserve routing metadata.
|
||||
if not resolved_model_name and isinstance(unified_file_id, str):
|
||||
decoded_unified_file_id = (
|
||||
_is_base64_encoded_unified_file_id(unified_file_id)
|
||||
or unified_file_id
|
||||
)
|
||||
target_model_names = get_models_from_unified_file_id(
|
||||
decoded_unified_file_id
|
||||
)
|
||||
if target_model_names:
|
||||
resolved_model_name = ",".join(target_model_names)
|
||||
resolved_model_name = resolve_managed_output_file_model_name(
|
||||
unified_input_file_id=unified_file_id
|
||||
if isinstance(unified_file_id, str)
|
||||
else response.input_file_id,
|
||||
fallback_model_name=model_name,
|
||||
)
|
||||
original_response_id = response.id
|
||||
|
||||
if (unified_batch_id or unified_file_id) and model_id:
|
||||
|
|
|
|||
|
|
@ -831,7 +831,7 @@ async def project_info(
|
|||
)
|
||||
|
||||
# Check if user has access to this project (admin or team member)
|
||||
is_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
is_admin = user_api_key_has_admin_view(user_api_key_dict)
|
||||
is_team_member = False
|
||||
|
||||
if project.team_id and user_api_key_dict.user_id:
|
||||
|
|
@ -886,7 +886,7 @@ async def list_projects(
|
|||
)
|
||||
|
||||
# If proxy admin, get all projects
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
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(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.53"
|
||||
version = "0.1.54"
|
||||
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.53"
|
||||
version = "0.1.54"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -61,9 +61,9 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
--extra bedrock-realtime \
|
||||
--python python3
|
||||
|
||||
RUN mkdir -p /home/nonroot && \
|
||||
HOME=/home/nonroot prisma generate --schema=./schema.prisma && \
|
||||
chown -R nonroot:nonroot /home/nonroot/.cache
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
npm_config_cache=/root/.npm \
|
||||
prisma generate --schema=./schema.prisma
|
||||
|
||||
RUN sed -i 's/\r$//' docker/component_entrypoint.sh && chmod +x docker/component_entrypoint.sh
|
||||
|
||||
|
|
@ -85,13 +85,16 @@ ENV HOME=/home/nonroot \
|
|||
PATH="/app/.venv/bin:${PATH}" \
|
||||
PYTHONPATH="/app" \
|
||||
PYTHONDONTWRITEBYTECODE=1 \
|
||||
PYTHONUNBUFFERED=1
|
||||
PYTHONUNBUFFERED=1 \
|
||||
PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries
|
||||
|
||||
COPY --from=builder --chown=nonroot:nonroot /app /app
|
||||
COPY --from=builder --chown=nonroot:nonroot /home/nonroot/.cache /home/nonroot/.cache
|
||||
COPY --from=builder /opt/prisma /opt/prisma
|
||||
|
||||
RUN find /app/.venv -type f -path "*/tornado/test/*" -delete && \
|
||||
find /app/.venv -type d -path "*/tornado/test" -delete
|
||||
find /app/.venv -type d -path "*/tornado/test" -delete && \
|
||||
chmod -R a+rX /opt/prisma && \
|
||||
python -c "from prisma.client import BINARY_PATHS; paths = list(BINARY_PATHS.query_engine.values()); assert paths and all(p.startswith('/opt/prisma/') for p in paths), paths"
|
||||
|
||||
USER nonroot
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,15 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_DailyGatewayRequests" (
|
||||
"date" TEXT NOT NULL,
|
||||
"category" TEXT NOT NULL,
|
||||
"route" TEXT NOT NULL,
|
||||
"successful_requests" BIGINT NOT NULL DEFAULT 0,
|
||||
"failed_requests" BIGINT NOT NULL DEFAULT 0,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"updated_at" TIMESTAMP(3) NOT NULL,
|
||||
|
||||
CONSTRAINT "LiteLLM_DailyGatewayRequests_pkey" PRIMARY KEY ("date","category","route")
|
||||
);
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX IF NOT EXISTS "LiteLLM_DailyGatewayRequests_date_idx" ON "LiteLLM_DailyGatewayRequests"("date");
|
||||
|
|
@ -0,0 +1,31 @@
|
|||
CREATE TABLE IF NOT EXISTS "LiteLLM_AutoRouterSession" (
|
||||
"api_key" TEXT NOT NULL,
|
||||
"session_id" TEXT NOT NULL,
|
||||
"router_name" TEXT NOT NULL,
|
||||
"router_type" TEXT NOT NULL,
|
||||
"first_turn_at" TIMESTAMP(3) NOT NULL,
|
||||
"last_turn_at" TIMESTAMP(3) NOT NULL,
|
||||
"last_model" TEXT NOT NULL,
|
||||
"models" JSONB NOT NULL DEFAULT '{}',
|
||||
"turns" INTEGER NOT NULL DEFAULT 0,
|
||||
"unordered_turns" INTEGER NOT NULL DEFAULT 0,
|
||||
"covered_turns" INTEGER NOT NULL DEFAULT 0,
|
||||
"cache_hits" INTEGER NOT NULL DEFAULT 0,
|
||||
"same_model_turns" INTEGER NOT NULL DEFAULT 0,
|
||||
"same_model_hits" INTEGER NOT NULL DEFAULT 0,
|
||||
"first_visit_turns" INTEGER NOT NULL DEFAULT 0,
|
||||
"first_visit_hits" INTEGER NOT NULL DEFAULT 0,
|
||||
"return_turns" INTEGER NOT NULL DEFAULT 0,
|
||||
"return_hits" INTEGER NOT NULL DEFAULT 0,
|
||||
"return_expired_misses" INTEGER NOT NULL DEFAULT 0,
|
||||
"return_within_ttl_misses" INTEGER NOT NULL DEFAULT 0,
|
||||
"ttl_5m_turns" INTEGER NOT NULL DEFAULT 0,
|
||||
"ttl_1h_turns" INTEGER NOT NULL DEFAULT 0,
|
||||
"total_tokens" BIGINT NOT NULL DEFAULT 0,
|
||||
"spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
|
||||
"saved_spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
|
||||
|
||||
CONSTRAINT "LiteLLM_AutoRouterSession_pkey" PRIMARY KEY ("api_key", "session_id", "router_name")
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS "idx_autorouter_session_last_turn" ON "LiteLLM_AutoRouterSession"("last_turn_at");
|
||||
|
|
@ -121,9 +121,13 @@ def heal_incomplete_nodeenv_cache() -> bool:
|
|||
Prisma invocation reinstalls it instead of failing on a missing binary.
|
||||
"""
|
||||
cache_dir = nodeenv_cache_dir()
|
||||
if cache_dir is None or not cache_dir.is_dir():
|
||||
if cache_dir is None:
|
||||
return False
|
||||
if node_binary_path(cache_dir).exists():
|
||||
try:
|
||||
if not cache_dir.is_dir() or node_binary_path(cache_dir).exists():
|
||||
return False
|
||||
except OSError as e:
|
||||
logger.warning("Could not inspect the Node toolchain at %s: %s", cache_dir, e)
|
||||
return False
|
||||
logger.warning(
|
||||
"Node toolchain at %s has no %s, so a previous install was interrupted. "
|
||||
|
|
|
|||
|
|
@ -1118,6 +1118,26 @@ model LiteLLM_DailyToolSpend {
|
|||
@@id([date, tool_name])
|
||||
}
|
||||
|
||||
// Gateway request counts recorded at the ASGI edge by
|
||||
// BillableRequestMetricsMiddleware. This is the source of truth for SGR
|
||||
// (successful gateway requests): it counts what the proxy actually answered,
|
||||
// independent of whether the request reached litellm's logging callbacks.
|
||||
// The key carries no deployment or caller dimension. Every part of it is
|
||||
// chosen by the proxy and drawn from a closed set, so the table is bounded by
|
||||
// (days x categories x routes) rather than by anything a caller can vary.
|
||||
model LiteLLM_DailyGatewayRequests {
|
||||
date String
|
||||
category String
|
||||
route String
|
||||
successful_requests BigInt @default(0)
|
||||
failed_requests BigInt @default(0)
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
|
||||
@@id([date, category, route])
|
||||
@@index([date])
|
||||
}
|
||||
|
||||
// Prompt table for storing prompt configurations
|
||||
model LiteLLM_PromptTable {
|
||||
id String @id @default(uuid())
|
||||
|
|
@ -1393,6 +1413,37 @@ model LiteLLM_AdaptiveRouterSession {
|
|||
@@index([last_activity_at], map: "idx_adaptive_router_session_activity")
|
||||
}
|
||||
|
||||
model LiteLLM_AutoRouterSession {
|
||||
api_key String
|
||||
session_id String
|
||||
router_name String
|
||||
router_type String
|
||||
first_turn_at DateTime
|
||||
last_turn_at DateTime
|
||||
last_model String
|
||||
models Json @default("{}")
|
||||
turns Int @default(0)
|
||||
unordered_turns Int @default(0)
|
||||
covered_turns Int @default(0)
|
||||
cache_hits Int @default(0)
|
||||
same_model_turns Int @default(0)
|
||||
same_model_hits Int @default(0)
|
||||
first_visit_turns Int @default(0)
|
||||
first_visit_hits Int @default(0)
|
||||
return_turns Int @default(0)
|
||||
return_hits Int @default(0)
|
||||
return_expired_misses Int @default(0)
|
||||
return_within_ttl_misses Int @default(0)
|
||||
ttl_5m_turns Int @default(0)
|
||||
ttl_1h_turns Int @default(0)
|
||||
total_tokens BigInt @default(0)
|
||||
spend Float @default(0)
|
||||
saved_spend Float @default(0)
|
||||
|
||||
@@id([api_key, session_id, router_name])
|
||||
@@index([last_turn_at], map: "idx_autorouter_session_last_turn")
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Workflow Run Tracking
|
||||
//
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.83"
|
||||
version = "0.4.84"
|
||||
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.83"
|
||||
version = "0.4.84"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -244,6 +244,7 @@ use_chat_completions_url_for_anthropic_messages: bool = bool(
|
|||
# Or via `litellm_settings.strip_anthropic_total_tokens: true` in
|
||||
# config.yaml.
|
||||
strip_anthropic_total_tokens: bool = False
|
||||
anthropic_sse_ping_interval_seconds: float = 15.0
|
||||
route_all_chat_openai_to_responses: bool = (
|
||||
os.getenv("LITELLM_ROUTE_ALL_CHAT_OPENAI_TO_RESPONSES", "false").lower() == "true"
|
||||
) # When True, routes all OpenAI /chat/completions requests through the Responses API bridge
|
||||
|
|
@ -704,9 +705,8 @@ def is_openai_finetune_model(key: str) -> bool:
|
|||
return key.startswith("ft:") and not key.count(":") > 1
|
||||
|
||||
|
||||
def add_known_models(model_cost_map: Optional[Dict] = None):
|
||||
_map: Final = model_cost_map if model_cost_map is not None else model_cost
|
||||
for key, value in _map.items():
|
||||
def _populate_provider_model_sets(model_cost_map: Dict) -> None:
|
||||
for key, value in model_cost_map.items():
|
||||
if value.get("litellm_provider") == "openai" and not is_openai_finetune_model(key):
|
||||
open_ai_chat_completion_models.add(key)
|
||||
elif value.get("litellm_provider") == "text-completion-openai":
|
||||
|
|
@ -949,7 +949,16 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
|
|||
bedrock_mantle_models.add(key)
|
||||
|
||||
|
||||
add_known_models()
|
||||
def add_known_models(model_cost_map: Optional[Dict] = None):
|
||||
"""Fold `model_cost_map` (defaults to `litellm.model_cost`) into the per-provider model sets,
|
||||
then refresh `models_by_provider` from those sets so the additions reach wildcard expansion.
|
||||
The refresh updates the dict in place, so references captured before a reload stay live.
|
||||
"""
|
||||
_populate_provider_model_sets(model_cost_map if model_cost_map is not None else model_cost)
|
||||
models_by_provider.update(_build_models_by_provider())
|
||||
|
||||
|
||||
_populate_provider_model_sets(model_cost)
|
||||
# known openai compatible endpoints - we'll eventually move this list to the model_prices_and_context_window.json dictionary
|
||||
|
||||
# this is maintained for Exception Mapping
|
||||
|
|
@ -1071,112 +1080,116 @@ model_list_set = set(model_list)
|
|||
# provider_list is lazy-loaded via __getattr__ to avoid importing LlmProviders at import time
|
||||
|
||||
|
||||
models_by_provider: dict = {
|
||||
"openai": open_ai_chat_completion_models | open_ai_text_completion_models,
|
||||
"text-completion-openai": open_ai_text_completion_models,
|
||||
"cohere": cohere_models | cohere_chat_models,
|
||||
"cohere_chat": cohere_chat_models,
|
||||
"anthropic": anthropic_models,
|
||||
"replicate": replicate_models,
|
||||
"huggingface": huggingface_models,
|
||||
"together_ai": together_ai_models,
|
||||
"baseten": baseten_models,
|
||||
"openrouter": openrouter_models,
|
||||
"vercel_ai_gateway": vercel_ai_gateway_models,
|
||||
"datarobot": datarobot_models,
|
||||
"vertex_ai": vertex_chat_models
|
||||
| vertex_text_models
|
||||
| vertex_anthropic_models
|
||||
| vertex_vision_models
|
||||
| vertex_language_models
|
||||
| vertex_deepseek_models
|
||||
| vertex_minimax_models
|
||||
| vertex_moonshot_models
|
||||
| vertex_zai_models,
|
||||
"ai21": ai21_models,
|
||||
"bedrock": bedrock_models | bedrock_converse_models,
|
||||
"petals": petals_models,
|
||||
"ollama": ollama_models,
|
||||
"ollama_chat": ollama_models,
|
||||
"deepinfra": deepinfra_models,
|
||||
"perplexity": perplexity_models,
|
||||
"maritalk": maritalk_models,
|
||||
"watsonx": watsonx_models,
|
||||
"gemini": gemini_models,
|
||||
"fireworks_ai": fireworks_ai_models | fireworks_ai_embedding_models,
|
||||
"aleph_alpha": aleph_alpha_models,
|
||||
"text-completion-codestral": text_completion_codestral_models,
|
||||
"text-completion-inception": text_completion_inception_models,
|
||||
"xai": xai_models,
|
||||
"zai": zai_models,
|
||||
"fal_ai": fal_ai_models,
|
||||
"deepseek": deepseek_models,
|
||||
"tencent": tencent_models,
|
||||
"runwayml": runwayml_models,
|
||||
"mistral": mistral_chat_models,
|
||||
"azure_ai": azure_ai_models,
|
||||
"voyage": voyage_models,
|
||||
"infinity": infinity_models,
|
||||
"databricks": databricks_models,
|
||||
"cloudflare": cloudflare_models,
|
||||
"codestral": codestral_models,
|
||||
"nlp_cloud": nlp_cloud_models,
|
||||
"friendliai": friendliai_models,
|
||||
"palm": palm_models,
|
||||
"groq": groq_models,
|
||||
"azure": azure_models | azure_text_models,
|
||||
"azure_anthropic": azure_anthropic_models,
|
||||
"azure_text": azure_text_models,
|
||||
"anyscale": anyscale_models,
|
||||
"cerebras": cerebras_models,
|
||||
"galadriel": galadriel_models,
|
||||
"nvidia_nim": nvidia_nim_models,
|
||||
"nvidia_riva": nvidia_riva_models,
|
||||
"soniox": soniox_models,
|
||||
"sambanova": sambanova_models | sambanova_embedding_models,
|
||||
"novita": novita_models,
|
||||
"nebius": nebius_models | nebius_embedding_models,
|
||||
"aiml": aiml_models,
|
||||
"assemblyai": assemblyai_models,
|
||||
"jina_ai": jina_ai_models,
|
||||
"snowflake": snowflake_models,
|
||||
"gradient_ai": gradient_ai_models,
|
||||
"meta_llama": llama_models,
|
||||
"nscale": nscale_models,
|
||||
"featherless_ai": featherless_ai_models,
|
||||
"deepgram": deepgram_models,
|
||||
"elevenlabs": elevenlabs_models,
|
||||
"heroku": heroku_models,
|
||||
"dashscope": dashscope_models,
|
||||
"modelscope": modelscope_models,
|
||||
"moonshot": moonshot_models,
|
||||
"publicai": publicai_models,
|
||||
"darkbloom": darkbloom_models,
|
||||
"v0": v0_models,
|
||||
"morph": morph_models,
|
||||
"lambda_ai": lambda_ai_models,
|
||||
"inception": inception_models,
|
||||
"hyperbolic": hyperbolic_models,
|
||||
"black_forest_labs": black_forest_labs_models,
|
||||
"recraft": recraft_models,
|
||||
"cometapi": cometapi_models,
|
||||
"oci": oci_models,
|
||||
"volcengine": volcengine_models,
|
||||
"wandb": wandb_models,
|
||||
"ovhcloud": ovhcloud_models | ovhcloud_embedding_models,
|
||||
"lemonade": lemonade_models,
|
||||
"clarifai": clarifai_models,
|
||||
"amazon_nova": amazon_nova_models,
|
||||
"stability": stability_models,
|
||||
"github_copilot": github_copilot_models,
|
||||
"chatgpt": chatgpt_models,
|
||||
"minimax": minimax_models,
|
||||
"aws_polly": aws_polly_models,
|
||||
"gigachat": gigachat_models,
|
||||
"llamagate": llamagate_models,
|
||||
"reducto": reducto_models,
|
||||
"bedrock_mantle": bedrock_mantle_models,
|
||||
}
|
||||
def _build_models_by_provider() -> dict:
|
||||
return {
|
||||
"openai": open_ai_chat_completion_models | open_ai_text_completion_models,
|
||||
"text-completion-openai": open_ai_text_completion_models,
|
||||
"cohere": cohere_models | cohere_chat_models,
|
||||
"cohere_chat": cohere_chat_models,
|
||||
"anthropic": anthropic_models,
|
||||
"replicate": replicate_models,
|
||||
"huggingface": huggingface_models,
|
||||
"together_ai": together_ai_models,
|
||||
"baseten": baseten_models,
|
||||
"openrouter": openrouter_models,
|
||||
"vercel_ai_gateway": vercel_ai_gateway_models,
|
||||
"datarobot": datarobot_models,
|
||||
"vertex_ai": vertex_chat_models
|
||||
| vertex_text_models
|
||||
| vertex_anthropic_models
|
||||
| vertex_vision_models
|
||||
| vertex_language_models
|
||||
| vertex_deepseek_models
|
||||
| vertex_minimax_models
|
||||
| vertex_moonshot_models
|
||||
| vertex_zai_models,
|
||||
"ai21": ai21_models,
|
||||
"bedrock": bedrock_models | bedrock_converse_models,
|
||||
"petals": petals_models,
|
||||
"ollama": ollama_models,
|
||||
"ollama_chat": ollama_models,
|
||||
"deepinfra": deepinfra_models,
|
||||
"perplexity": perplexity_models,
|
||||
"maritalk": maritalk_models,
|
||||
"watsonx": watsonx_models,
|
||||
"gemini": gemini_models,
|
||||
"fireworks_ai": fireworks_ai_models | fireworks_ai_embedding_models,
|
||||
"aleph_alpha": aleph_alpha_models,
|
||||
"text-completion-codestral": text_completion_codestral_models,
|
||||
"text-completion-inception": text_completion_inception_models,
|
||||
"xai": xai_models,
|
||||
"zai": zai_models,
|
||||
"fal_ai": fal_ai_models,
|
||||
"deepseek": deepseek_models,
|
||||
"tencent": tencent_models,
|
||||
"runwayml": runwayml_models,
|
||||
"mistral": mistral_chat_models,
|
||||
"azure_ai": azure_ai_models,
|
||||
"voyage": voyage_models,
|
||||
"infinity": infinity_models,
|
||||
"databricks": databricks_models,
|
||||
"cloudflare": cloudflare_models,
|
||||
"codestral": codestral_models,
|
||||
"nlp_cloud": nlp_cloud_models,
|
||||
"friendliai": friendliai_models,
|
||||
"palm": palm_models,
|
||||
"groq": groq_models,
|
||||
"azure": azure_models | azure_text_models,
|
||||
"azure_anthropic": azure_anthropic_models,
|
||||
"azure_text": azure_text_models,
|
||||
"anyscale": anyscale_models,
|
||||
"cerebras": cerebras_models,
|
||||
"galadriel": galadriel_models,
|
||||
"nvidia_nim": nvidia_nim_models,
|
||||
"nvidia_riva": nvidia_riva_models,
|
||||
"soniox": soniox_models,
|
||||
"sambanova": sambanova_models | sambanova_embedding_models,
|
||||
"novita": novita_models,
|
||||
"nebius": nebius_models | nebius_embedding_models,
|
||||
"aiml": aiml_models,
|
||||
"assemblyai": assemblyai_models,
|
||||
"jina_ai": jina_ai_models,
|
||||
"snowflake": snowflake_models,
|
||||
"gradient_ai": gradient_ai_models,
|
||||
"meta_llama": llama_models,
|
||||
"nscale": nscale_models,
|
||||
"featherless_ai": featherless_ai_models,
|
||||
"deepgram": deepgram_models,
|
||||
"elevenlabs": elevenlabs_models,
|
||||
"heroku": heroku_models,
|
||||
"dashscope": dashscope_models,
|
||||
"modelscope": modelscope_models,
|
||||
"moonshot": moonshot_models,
|
||||
"publicai": publicai_models,
|
||||
"darkbloom": darkbloom_models,
|
||||
"v0": v0_models,
|
||||
"morph": morph_models,
|
||||
"lambda_ai": lambda_ai_models,
|
||||
"inception": inception_models,
|
||||
"hyperbolic": hyperbolic_models,
|
||||
"black_forest_labs": black_forest_labs_models,
|
||||
"recraft": recraft_models,
|
||||
"cometapi": cometapi_models,
|
||||
"oci": oci_models,
|
||||
"volcengine": volcengine_models,
|
||||
"wandb": wandb_models,
|
||||
"ovhcloud": ovhcloud_models | ovhcloud_embedding_models,
|
||||
"lemonade": lemonade_models,
|
||||
"clarifai": clarifai_models,
|
||||
"amazon_nova": amazon_nova_models,
|
||||
"stability": stability_models,
|
||||
"github_copilot": github_copilot_models,
|
||||
"chatgpt": chatgpt_models,
|
||||
"minimax": minimax_models,
|
||||
"aws_polly": aws_polly_models,
|
||||
"gigachat": gigachat_models,
|
||||
"llamagate": llamagate_models,
|
||||
"reducto": reducto_models,
|
||||
"bedrock_mantle": bedrock_mantle_models,
|
||||
}
|
||||
|
||||
|
||||
models_by_provider: dict = _build_models_by_provider()
|
||||
|
||||
# mapping for those models which have larger equivalents
|
||||
longer_context_model_fallback_dict: dict = {
|
||||
|
|
|
|||
|
|
@ -1461,32 +1461,30 @@ _UTILS_MODULE_IMPORT_MAP: Final = {
|
|||
|
||||
# Export all name tuples and import maps for use in _lazy_imports.py
|
||||
__all__ = [
|
||||
# Name tuples
|
||||
"COST_CALCULATOR_NAMES",
|
||||
"LITELLM_LOGGING_NAMES",
|
||||
"UTILS_NAMES",
|
||||
"TOKEN_COUNTER_NAMES",
|
||||
"LLM_CLIENT_CACHE_NAMES",
|
||||
"BEDROCK_TYPES_NAMES",
|
||||
"TYPES_UTILS_NAMES",
|
||||
"CACHING_NAMES",
|
||||
"HTTP_HANDLER_NAMES",
|
||||
"COST_CALCULATOR_NAMES",
|
||||
"DOTPROMPT_NAMES",
|
||||
"HTTP_HANDLER_NAMES",
|
||||
"LITELLM_LOGGING_NAMES",
|
||||
"LLM_CLIENT_CACHE_NAMES",
|
||||
"LLM_CONFIG_NAMES",
|
||||
"TYPES_NAMES",
|
||||
"LLM_PROVIDER_LOGIC_NAMES",
|
||||
"TOKEN_COUNTER_NAMES",
|
||||
"TYPES_NAMES",
|
||||
"TYPES_UTILS_NAMES",
|
||||
"UTILS_MODULE_NAMES",
|
||||
# Import maps
|
||||
"_UTILS_IMPORT_MAP",
|
||||
"_COST_CALCULATOR_IMPORT_MAP",
|
||||
"_TYPES_UTILS_IMPORT_MAP",
|
||||
"_TOKEN_COUNTER_IMPORT_MAP",
|
||||
"UTILS_NAMES",
|
||||
"_BEDROCK_TYPES_IMPORT_MAP",
|
||||
"_CACHING_IMPORT_MAP",
|
||||
"_LITELLM_LOGGING_IMPORT_MAP",
|
||||
"_COST_CALCULATOR_IMPORT_MAP",
|
||||
"_DOTPROMPT_IMPORT_MAP",
|
||||
"_TYPES_IMPORT_MAP",
|
||||
"_LITELLM_LOGGING_IMPORT_MAP",
|
||||
"_LLM_CONFIGS_IMPORT_MAP",
|
||||
"_LLM_PROVIDER_LOGIC_IMPORT_MAP",
|
||||
"_TOKEN_COUNTER_IMPORT_MAP",
|
||||
"_TYPES_IMPORT_MAP",
|
||||
"_TYPES_UTILS_IMPORT_MAP",
|
||||
"_UTILS_IMPORT_MAP",
|
||||
"_UTILS_MODULE_IMPORT_MAP",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -273,8 +273,42 @@ def _suppress_loggers():
|
|||
apscheduler_scheduler_logger.setLevel(logging.WARNING)
|
||||
|
||||
|
||||
_REDACTED_THIRD_PARTY_LOGGERS: Final[tuple[str, ...]] = (
|
||||
"apscheduler.executors.default",
|
||||
"apscheduler.scheduler",
|
||||
"asyncio",
|
||||
"backoff",
|
||||
"httpx",
|
||||
"uvicorn.error",
|
||||
)
|
||||
|
||||
|
||||
def _redact_third_party_loggers() -> None:
|
||||
"""Extend secret redaction to records litellm does not emit directly.
|
||||
|
||||
litellm's own loggers are covered by the filter on their shared handler, but a
|
||||
litellm value can also reach a log record through a dependency that logs on its
|
||||
own logger. Those records never pass through a litellm handler.
|
||||
|
||||
The filter is attached to each emitting logger rather than to the root logger or
|
||||
to root's handlers. `Logger.handle` applies the emitting logger's filters before
|
||||
any handler runs, so redaction happens once, at the earliest point in the
|
||||
record's life, and covers every downstream handler regardless of who owns it.
|
||||
The alternatives do not hold: `callHandlers` consults ancestors for handlers but
|
||||
never for filters, so a filter on the root logger never sees these records at
|
||||
all, and a filter on a root handler only covers that one handler, leaving
|
||||
handlers registered earlier or on the emitting logger itself untouched.
|
||||
|
||||
Each name is the exact logger a dependency emits on; a parent name would not
|
||||
cover its children, for the same reason the root logger does not.
|
||||
"""
|
||||
for name in _REDACTED_THIRD_PARTY_LOGGERS:
|
||||
logging.getLogger(name).addFilter(_secret_filter)
|
||||
|
||||
|
||||
# Call the suppression function
|
||||
_suppress_loggers()
|
||||
_redact_third_party_loggers()
|
||||
|
||||
ALL_LOGGERS: Final = [
|
||||
logging.getLogger(),
|
||||
|
|
|
|||
|
|
@ -395,9 +395,7 @@ def _get_redis_client_logic(**env_overrides):
|
|||
if _sentinel_password is not None:
|
||||
redis_kwargs["sentinel_password"] = _sentinel_password
|
||||
|
||||
_service_name: Final[str | None] = redis_kwargs.get("service_name", None) or get_secret(
|
||||
"REDIS_SERVICE_NAME"
|
||||
)
|
||||
_service_name: Final[str | None] = redis_kwargs.get("service_name", None) or get_secret("REDIS_SERVICE_NAME")
|
||||
|
||||
if _service_name is not None:
|
||||
redis_kwargs["service_name"] = _service_name
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import asyncio
|
||||
from datetime import datetime, timedelta
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -16,7 +16,7 @@ if TYPE_CHECKING:
|
|||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
OTELClass = OpenTelemetry
|
||||
else:
|
||||
Span = Any
|
||||
|
|
|
|||
|
|
@ -55,19 +55,15 @@ from litellm.a2a_protocol.main import (
|
|||
from litellm.types.agents import LiteLLMSendMessageResponse
|
||||
|
||||
__all__ = [
|
||||
# Client
|
||||
"A2AClient",
|
||||
# Functions
|
||||
"asend_message",
|
||||
"send_message",
|
||||
"asend_message_streaming",
|
||||
"aget_agent_card",
|
||||
"create_a2a_client",
|
||||
# Response types
|
||||
"LiteLLMSendMessageResponse",
|
||||
# Exceptions
|
||||
"A2AError",
|
||||
"A2AConnectionError",
|
||||
"A2AAgentCardError",
|
||||
"A2AClient",
|
||||
"A2AConnectionError",
|
||||
"A2AError",
|
||||
"A2ALocalhostURLError",
|
||||
"LiteLLMSendMessageResponse",
|
||||
"aget_agent_card",
|
||||
"asend_message",
|
||||
"asend_message_streaming",
|
||||
"create_a2a_client",
|
||||
"send_message",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -202,9 +202,9 @@ async def handle_a2a_localhost_retry(
|
|||
# Fix the agent card URL
|
||||
set_agent_card_url(agent_card, error.base_url)
|
||||
|
||||
# Reuse the httpx client LiteLLM attached at creation. It carries this agent's
|
||||
# trace-id and auth headers, so a fresh client would drop them. Only clients built
|
||||
# by ``create_a2a_client`` have it; an externally-supplied client cannot be retried.
|
||||
# Reuse the httpx client and call context LiteLLM attached at creation, since the
|
||||
# context carries this agent's trace-id/auth headers. Only clients built by
|
||||
# ``create_a2a_client`` have them; an externally-supplied client cannot be retried.
|
||||
httpx_client: Final = getattr(a2a_client, "_litellm_httpx_client", None)
|
||||
if httpx_client is None:
|
||||
raise RuntimeError(
|
||||
|
|
@ -220,5 +220,8 @@ async def handle_a2a_localhost_retry(
|
|||
),
|
||||
)
|
||||
new_client._litellm_httpx_client = httpx_client
|
||||
new_client._litellm_call_context = getattr( # pyright: ignore[reportAttributeAccessIssue] # LiteLLM-owned stash
|
||||
a2a_client, "_litellm_call_context", None
|
||||
)
|
||||
new_client._litellm_agent_card = agent_card
|
||||
return new_client
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ A2A Streaming Events (in order):
|
|||
4. Status update (kind: "status-update") - Final status "completed" with final=true
|
||||
"""
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from collections.abc import AsyncIterator, Mapping
|
||||
from typing import Any, Final
|
||||
|
||||
import litellm
|
||||
|
|
@ -21,6 +21,8 @@ from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
|||
)
|
||||
from litellm.a2a_protocol.providers.config_manager import A2AProviderConfigManager
|
||||
from litellm.interactions.agents.utils import merge_agent_headers
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
# litellm_params key carrying the authenticated principal (hashed virtual key) so
|
||||
# A2A provider configs can scope provider-side state (e.g. LangFlow session memory)
|
||||
|
|
@ -45,47 +47,14 @@ class A2ACompletionBridgeHandler:
|
|||
"""
|
||||
|
||||
@staticmethod
|
||||
async def handle_non_streaming(
|
||||
request_id: str,
|
||||
def _build_completion_params(
|
||||
params: dict[str, Any],
|
||||
litellm_params: dict[str, Any],
|
||||
api_base: str | None = None,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
litellm_params: Mapping[str, Any],
|
||||
api_base: str | None,
|
||||
agent_extra_headers: Mapping[str, str] | None,
|
||||
*,
|
||||
_skip_a2a_provider_routing: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Handle non-streaming A2A request via litellm.acompletion.
|
||||
|
||||
Args:
|
||||
request_id: A2A JSON-RPC request ID
|
||||
params: A2A MessageSendParams containing the message
|
||||
litellm_params: Agent's litellm_params (custom_llm_provider, model, etc.)
|
||||
api_base: API base URL from agent_card_params
|
||||
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
|
||||
admin extra_headers) to forward on the upstream HTTP call.
|
||||
|
||||
Returns:
|
||||
A2A SendMessageResponse dict
|
||||
"""
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
if not _skip_a2a_provider_routing:
|
||||
a2a_provider_config: Final = A2AProviderConfigManager.get_provider_config(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model=litellm_params.get("model"),
|
||||
)
|
||||
|
||||
if a2a_provider_config is not None:
|
||||
verbose_logger.info("A2A: Using provider config for %s", custom_llm_provider)
|
||||
|
||||
return await a2a_provider_config.handle_non_streaming(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
|
||||
stream: bool,
|
||||
) -> Mapping[str, Any]:
|
||||
# Extract message from params
|
||||
message: Final = params.get("message", {})
|
||||
|
||||
|
|
@ -93,7 +62,7 @@ class A2ACompletionBridgeHandler:
|
|||
openai_messages: Final = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
|
||||
# Get completion params
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
|
||||
model: Final = litellm_params.get("model", "agent")
|
||||
|
||||
# Build full model string if provider specified
|
||||
|
|
@ -103,14 +72,17 @@ class A2ACompletionBridgeHandler:
|
|||
else:
|
||||
full_model = model
|
||||
|
||||
verbose_logger.info("A2A completion bridge: model=%s, api_base=%s", full_model, api_base)
|
||||
if stream:
|
||||
verbose_logger.info("A2A completion bridge streaming: model=%s, api_base=%s", full_model, api_base)
|
||||
else:
|
||||
verbose_logger.info("A2A completion bridge: model=%s, api_base=%s", full_model, api_base)
|
||||
|
||||
# Build completion params dict
|
||||
completion_params: Final[dict[str, Any]] = {
|
||||
"model": full_model,
|
||||
"messages": openai_messages,
|
||||
"api_base": api_base,
|
||||
"stream": False,
|
||||
"stream": stream,
|
||||
}
|
||||
# Add litellm_params (contains api_key, client_id, client_secret, tenant_id, etc.)
|
||||
litellm_params_to_add: Final = {
|
||||
|
|
@ -134,8 +106,64 @@ class A2ACompletionBridgeHandler:
|
|||
static_headers=completion_params.get("extra_headers"),
|
||||
)
|
||||
|
||||
return completion_params
|
||||
|
||||
@staticmethod
|
||||
async def _acompletion(completion_params: Mapping[str, Any]) -> ModelResponse | CustomStreamWrapper:
|
||||
return await litellm.acompletion(**completion_params)
|
||||
|
||||
@staticmethod
|
||||
async def handle_non_streaming(
|
||||
request_id: str,
|
||||
params: dict[str, Any],
|
||||
litellm_params: dict[str, Any],
|
||||
api_base: str | None = None,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
*,
|
||||
_skip_a2a_provider_routing: bool = False,
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Handle non-streaming A2A request via litellm.acompletion.
|
||||
|
||||
Args:
|
||||
request_id: A2A JSON-RPC request ID
|
||||
params: A2A MessageSendParams containing the message
|
||||
litellm_params: Agent's litellm_params (custom_llm_provider, model, etc.)
|
||||
api_base: API base URL from agent_card_params
|
||||
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
|
||||
admin extra_headers) to forward on the upstream HTTP call.
|
||||
|
||||
Returns:
|
||||
A2A SendMessageResponse dict
|
||||
"""
|
||||
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
|
||||
if not _skip_a2a_provider_routing:
|
||||
a2a_provider_config: Final = A2AProviderConfigManager.get_provider_config(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model=litellm_params.get("model"),
|
||||
)
|
||||
|
||||
if a2a_provider_config is not None:
|
||||
verbose_logger.info("A2A: Using provider config for %s", custom_llm_provider)
|
||||
|
||||
return await a2a_provider_config.handle_non_streaming(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
|
||||
completion_params: Final = A2ACompletionBridgeHandler._build_completion_params(
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
api_base=api_base,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
stream=False,
|
||||
)
|
||||
|
||||
# Call litellm.acompletion
|
||||
response: Final = await litellm.acompletion(**completion_params)
|
||||
response: Final = await A2ACompletionBridgeHandler._acompletion(completion_params)
|
||||
|
||||
# Transform response to A2A format
|
||||
a2a_response: Final = A2ACompletionBridgeTransformation.openai_response_to_a2a_response(
|
||||
|
|
@ -156,7 +184,7 @@ class A2ACompletionBridgeHandler:
|
|||
agent_extra_headers: dict[str, str] | None = None,
|
||||
*,
|
||||
_skip_a2a_provider_routing: bool = False,
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
) -> AsyncIterator[dict[str, object]]:
|
||||
"""
|
||||
Handle streaming A2A request via litellm.acompletion with stream=True.
|
||||
|
||||
|
|
@ -177,7 +205,7 @@ class A2ACompletionBridgeHandler:
|
|||
Yields:
|
||||
A2A streaming response events
|
||||
"""
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
|
||||
if not _skip_a2a_provider_routing:
|
||||
a2a_provider_config: Final = A2AProviderConfigManager.get_provider_config(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -198,60 +226,20 @@ class A2ACompletionBridgeHandler:
|
|||
|
||||
return
|
||||
|
||||
# Extract message from params
|
||||
message: Final = params.get("message", {})
|
||||
|
||||
# Create streaming context
|
||||
ctx: Final = A2AStreamingContext(
|
||||
request_id=request_id,
|
||||
input_message=message,
|
||||
input_message=params.get("message", {}),
|
||||
)
|
||||
|
||||
# Transform A2A message to OpenAI format
|
||||
openai_messages: Final = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
||||
|
||||
# Get completion params
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
model: Final = litellm_params.get("model", "agent")
|
||||
|
||||
# Build full model string if provider specified
|
||||
# Skip prepending if model already starts with the provider prefix
|
||||
if custom_llm_provider and not model.startswith(f"{custom_llm_provider}/"):
|
||||
full_model = f"{custom_llm_provider}/{model}"
|
||||
else:
|
||||
full_model = model
|
||||
|
||||
verbose_logger.info("A2A completion bridge streaming: model=%s, api_base=%s", full_model, api_base)
|
||||
|
||||
# Build completion params dict
|
||||
completion_params: Final[dict[str, Any]] = {
|
||||
"model": full_model,
|
||||
"messages": openai_messages,
|
||||
"api_base": api_base,
|
||||
"stream": True,
|
||||
}
|
||||
# Add litellm_params (contains api_key, client_id, client_secret, tenant_id, etc.)
|
||||
litellm_params_to_add: Final = {
|
||||
k: v
|
||||
for k, v in litellm_params.items()
|
||||
if k not in ("model", "custom_llm_provider") and k not in _AGENT_ONLY_PARAMS
|
||||
}
|
||||
completion_params.update(litellm_params_to_add)
|
||||
# Apply forward metadata AFTER the litellm_params merge so the helper
|
||||
# sees any agent-owner-configured ``extra_body.metadata`` and can keep
|
||||
# those keys authoritative over the client-supplied A2A metadata.
|
||||
A2ACompletionBridgeTransformation.apply_forward_metadata_to_completion_params(
|
||||
completion_params=completion_params,
|
||||
a2a_message=message,
|
||||
completion_params: Final = A2ACompletionBridgeHandler._build_completion_params(
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
api_base=api_base,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
stream=True,
|
||||
)
|
||||
|
||||
if agent_extra_headers:
|
||||
completion_params["extra_headers"] = merge_agent_headers(
|
||||
dynamic_headers=agent_extra_headers,
|
||||
static_headers=completion_params.get("extra_headers"),
|
||||
)
|
||||
|
||||
# 1. Emit initial task event (kind: "task", status: "submitted")
|
||||
task_event: Final = A2ACompletionBridgeTransformation.create_task_event(ctx)
|
||||
yield task_event
|
||||
|
|
@ -266,7 +254,7 @@ class A2ACompletionBridgeHandler:
|
|||
yield working_event
|
||||
|
||||
# Call litellm.acompletion with streaming
|
||||
response: Final = await litellm.acompletion(**completion_params)
|
||||
response: Final = await A2ACompletionBridgeHandler._acompletion(completion_params)
|
||||
|
||||
# 3. Accumulate content and emit artifact update
|
||||
accumulated_text = ""
|
||||
|
|
@ -312,7 +300,7 @@ async def handle_a2a_completion(
|
|||
litellm_params: dict[str, Any],
|
||||
api_base: str | None = None,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""Convenience function for non-streaming A2A completion."""
|
||||
return await A2ACompletionBridgeHandler.handle_non_streaming(
|
||||
request_id=request_id,
|
||||
|
|
@ -329,7 +317,7 @@ async def handle_a2a_completion_streaming(
|
|||
litellm_params: dict[str, Any],
|
||||
api_base: str | None = None,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
) -> AsyncIterator[dict[str, object]]:
|
||||
"""Convenience function for streaming A2A completion."""
|
||||
async for chunk in A2ACompletionBridgeHandler.handle_streaming(
|
||||
request_id=request_id,
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import asyncio
|
|||
import datetime
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Coroutine
|
||||
from http.cookiejar import DefaultCookiePolicy
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional, cast
|
||||
|
||||
import litellm
|
||||
|
|
@ -30,6 +31,7 @@ from litellm.utils import client
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from a2a.client import Client as A2AClientType
|
||||
from a2a.client import ClientCallContext as A2ACallContextType
|
||||
from a2a.compat.v0_3.types import (
|
||||
AgentCard,
|
||||
Message,
|
||||
|
|
@ -45,7 +47,7 @@ A2A_SDK_AVAILABLE = False
|
|||
_a2a_conversions: Any = None
|
||||
|
||||
try:
|
||||
from a2a.client import Client, ClientConfig, create_client
|
||||
from a2a.client import Client, ClientCallContext, ClientConfig, create_client
|
||||
from a2a.compat.v0_3 import conversions as _a2a_conversions
|
||||
from a2a.compat.v0_3.types import (
|
||||
Message,
|
||||
|
|
@ -60,6 +62,7 @@ try:
|
|||
A2A_SDK_AVAILABLE = True
|
||||
except ImportError:
|
||||
Client = None
|
||||
ClientCallContext = None
|
||||
ClientConfig = None
|
||||
create_client = None
|
||||
|
||||
|
|
@ -77,6 +80,8 @@ from litellm.a2a_protocol.exceptions import A2ALocalhostURLError
|
|||
# Use our custom resolver instead of the default A2A SDK resolver
|
||||
A2ACardResolver: Final = LiteLLMA2ACardResolver
|
||||
|
||||
_BLOCK_ALL_COOKIES: Final = DefaultCookiePolicy(allowed_domains=())
|
||||
|
||||
|
||||
def _set_usage_on_logging_obj(
|
||||
kwargs: dict[str, Any],
|
||||
|
|
@ -218,6 +223,10 @@ async def _send_message_via_completion_bridge(
|
|||
return LiteLLMSendMessageResponse.from_dict(response_dict, request_id=str(request.id))
|
||||
|
||||
|
||||
def _get_a2a_call_context(a2a_client: "A2AClientType") -> Optional["A2ACallContextType"]:
|
||||
return getattr(a2a_client, "_litellm_call_context", None)
|
||||
|
||||
|
||||
async def _send_message(a2a_client: "A2AClientType", request: "SendMessageRequest") -> "SendMessageResponse":
|
||||
"""Send a non-streaming message via a2a-sdk 1.x and return JSON-RPC response."""
|
||||
if _a2a_conversions is None:
|
||||
|
|
@ -227,7 +236,7 @@ async def _send_message(a2a_client: "A2AClientType", request: "SendMessageReques
|
|||
|
||||
pb_request: Final = _a2a_conversions.to_core_send_message_request(request)
|
||||
last_event = None
|
||||
async for event in a2a_client.send_message(pb_request):
|
||||
async for event in a2a_client.send_message(pb_request, context=_get_a2a_call_context(a2a_client)):
|
||||
last_event = event
|
||||
if last_event is None:
|
||||
raise RuntimeError("A2A send_message failed: no response received from agent.")
|
||||
|
|
@ -301,7 +310,7 @@ async def _stream_messages(
|
|||
)
|
||||
|
||||
pb_request: Final = _a2a_conversions.to_core_send_message_request(request)
|
||||
async for event in a2a_client.send_message(pb_request):
|
||||
async for event in a2a_client.send_message(pb_request, context=_get_a2a_call_context(a2a_client)):
|
||||
compat_chunk = _a2a_conversions.to_compat_stream_response(
|
||||
event,
|
||||
request_id=request.id,
|
||||
|
|
@ -756,26 +765,13 @@ async def create_a2a_client(
|
|||
|
||||
verbose_logger.info("Creating A2A client for %s", base_url)
|
||||
|
||||
# Use get_async_httpx_client with per-agent params so that different agents
|
||||
# (with different extra_headers) get separate cached clients. The params
|
||||
# dict is hashed into the cache key, keeping agent auth isolated while
|
||||
# still reusing connections within the same agent.
|
||||
#
|
||||
# Only pass params that AsyncHTTPHandler.__init__ accepts (e.g. timeout).
|
||||
# Use "disable_aiohttp_transport" key for cache-key-only data (it's
|
||||
# filtered out before reaching the constructor).
|
||||
_client_params: Final[dict] = {"timeout": timeout}
|
||||
if extra_headers:
|
||||
# Encode headers into a cache-key-only param so each unique header
|
||||
# set produces a distinct cache key.
|
||||
_client_params["disable_aiohttp_transport"] = str(sorted(extra_headers.items()))
|
||||
_async_handler: Final = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.A2AProvider,
|
||||
params=_client_params,
|
||||
params={"timeout": timeout},
|
||||
)
|
||||
httpx_client: Final = _async_handler.client
|
||||
httpx_client.cookies.jar.set_policy(_BLOCK_ALL_COOKIES)
|
||||
if extra_headers:
|
||||
httpx_client.headers.update(extra_headers)
|
||||
verbose_proxy_logger.debug("A2A client created with extra_headers=%s", list(extra_headers.keys()))
|
||||
|
||||
a2a_client: Final = await create_client( # pyright: ignore[reportOptionalCall]
|
||||
|
|
@ -784,11 +780,17 @@ async def create_a2a_client(
|
|||
httpx_client=httpx_client,
|
||||
streaming=streaming,
|
||||
),
|
||||
resolver_http_kwargs={"headers": extra_headers} if extra_headers else None,
|
||||
)
|
||||
# Stash LiteLLM-owned handles on the client so the localhost-retry path can reuse
|
||||
# the configured httpx client (with this agent's trace-id/auth headers) without
|
||||
# excavating a2a-sdk private internals.
|
||||
# the configured httpx client and this agent's headers without excavating
|
||||
# a2a-sdk private internals.
|
||||
a2a_client._litellm_httpx_client = httpx_client
|
||||
a2a_client._litellm_call_context = ( # pyright: ignore[reportAttributeAccessIssue] # LiteLLM-owned stash
|
||||
ClientCallContext(service_parameters=extra_headers) # pyright: ignore[reportOptionalCall] # SDK checked above
|
||||
if extra_headers
|
||||
else None
|
||||
)
|
||||
agent_card: Final = getattr(a2a_client, "_card", None)
|
||||
if agent_card is not None:
|
||||
a2a_client._litellm_agent_card = agent_card
|
||||
|
|
|
|||
|
|
@ -8,8 +8,8 @@ from ..types.llms.openai import *
|
|||
|
||||
def get_optional_params_add_message(
|
||||
role: str | None,
|
||||
content: str | List[MessageContentTextObject | MessageContentImageFileObject | MessageContentImageURLObject] | None,
|
||||
attachments: List[Attachment] | None,
|
||||
content: str | list[MessageContentTextObject | MessageContentImageFileObject | MessageContentImageURLObject] | None,
|
||||
attachments: list[Attachment] | None,
|
||||
metadata: dict | None,
|
||||
custom_llm_provider: str,
|
||||
**kwargs,
|
||||
|
|
@ -57,7 +57,7 @@ def get_optional_params_add_message(
|
|||
optional_params = litellm.AzureOpenAIAssistantsAPIConfig().map_openai_params_create_message_params(
|
||||
non_default_params=non_default_params, optional_params=optional_params
|
||||
)
|
||||
for k in passed_params.keys():
|
||||
for k in passed_params:
|
||||
if k not in default_params:
|
||||
optional_params[k] = passed_params[k]
|
||||
return optional_params
|
||||
|
|
@ -128,7 +128,7 @@ def get_optional_params_image_gen(
|
|||
if n is not None:
|
||||
optional_params["sampleCount"] = int(n)
|
||||
|
||||
for k in passed_params.keys():
|
||||
for k in passed_params:
|
||||
if k not in default_params:
|
||||
optional_params[k] = passed_params[k]
|
||||
return optional_params
|
||||
|
|
|
|||
|
|
@ -9,12 +9,12 @@ Has 4 methods:
|
|||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
Span = Any
|
||||
|
||||
|
|
|
|||
|
|
@ -1,12 +1,12 @@
|
|||
import json
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from .base_cache import BaseCache
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
Span = Any
|
||||
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ import time
|
|||
import traceback
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from threading import Lock
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
|
|
@ -29,7 +29,7 @@ from .redis_cache import RedisCache
|
|||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
Span = Any
|
||||
|
||||
|
|
|
|||
276
litellm/caching/evicted_client_closer.py
Normal file
276
litellm/caching/evicted_client_closer.py
Normal file
|
|
@ -0,0 +1,276 @@
|
|||
"""
|
||||
Deferred close of HTTP/SDK clients that the LLM client cache has evicted.
|
||||
|
||||
Eviction only drops the cache's reference to a client. Every OpenAI/Azure SDK
|
||||
client is a reference cycle (each resource namespace holds the client back), so
|
||||
an evicted client and its pooled TCP connections survive until a generational
|
||||
collection runs, which under load is thousands of requests later.
|
||||
|
||||
Closing at eviction time is not an option: a request that was handed the client
|
||||
just before it was evicted is still using it, and closing it underneath that
|
||||
request raises ``RuntimeError: Cannot send a request, as the client has been
|
||||
closed.``
|
||||
|
||||
So an evicted client is closed once two conditions hold. A grace window must
|
||||
have passed since its eviction, which covers a request that holds the client
|
||||
but is momentarily not on the wire, and the client must report no connection in
|
||||
flight. The second condition is what keeps the first honest: a request may run
|
||||
for ``litellm.request_timeout`` seconds, 6000 by default, and a streaming
|
||||
response is bounded only by how long the upstream keeps sending, so no deadline
|
||||
on its own can promise that a request has finished.
|
||||
|
||||
Only clients litellm itself created are closed; a client the caller supplied is
|
||||
left alone because litellm does not own its lifecycle.
|
||||
|
||||
A client that closes synchronously is closed from wherever the cache is next
|
||||
used. One whose close is a coroutine needs the event loop it was evicted on, so
|
||||
it waits for a call from that loop rather than having work scheduled onto a loop
|
||||
it does not belong to. Queued clients are therefore bucketed by what it takes to
|
||||
close them, and each bucket is ordered by deadline, so a reap walks the entries
|
||||
that are due rather than the whole queue.
|
||||
|
||||
The queue holds its clients weakly, so waiting out a grace window never keeps
|
||||
alive anything the collector would have reclaimed first.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import inspect
|
||||
import threading
|
||||
import time
|
||||
import weakref
|
||||
from collections import deque
|
||||
from collections.abc import Awaitable, Callable, Iterator
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import (
|
||||
EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS,
|
||||
EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING,
|
||||
)
|
||||
|
||||
_CLOSABLE_ANYWHERE: Final = "closable-anywhere"
|
||||
_CLOSABLE_ON_ANY_LOOP: Final = "closable-on-any-loop"
|
||||
|
||||
_BucketKey = str | int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PendingClose:
|
||||
"""A queued close.
|
||||
|
||||
The client is held weakly, so queueing one never keeps alive anything the
|
||||
collector would otherwise have reclaimed first.
|
||||
|
||||
``needs_loop`` is set for a client whose close is a coroutine; those can only
|
||||
be closed from the event loop they were evicted on, recorded in ``loop_id``.
|
||||
A client that closes synchronously carries neither constraint.
|
||||
"""
|
||||
|
||||
client_ref: "weakref.ref[object]"
|
||||
loop_id: int | None
|
||||
needs_loop: bool
|
||||
close_after: float
|
||||
|
||||
|
||||
def _bucket_key(pending: _PendingClose) -> _BucketKey:
|
||||
"""Which reaps can close this entry: any at all, any running a loop, or one loop's."""
|
||||
if not pending.needs_loop:
|
||||
return _CLOSABLE_ANYWHERE
|
||||
if pending.loop_id is None:
|
||||
return _CLOSABLE_ON_ANY_LOOP
|
||||
return pending.loop_id
|
||||
|
||||
|
||||
def _running_loop_id() -> int | None:
|
||||
try:
|
||||
return id(asyncio.get_running_loop())
|
||||
except RuntimeError:
|
||||
return None
|
||||
|
||||
|
||||
def _close_function(client: object) -> Callable[[], object] | None:
|
||||
close_fn: Final[Callable[[], object] | None] = getattr(client, "aclose", None) or getattr(client, "close", None)
|
||||
return close_fn
|
||||
|
||||
|
||||
def _transport_of(client: object) -> object:
|
||||
"""The httpx transport behind an SDK wrapper, a litellm handler, or a bare client."""
|
||||
for holder in (getattr(client, "_client", None), getattr(client, "client", None), client):
|
||||
transport: object = getattr(holder, "_transport", None)
|
||||
if transport is not None:
|
||||
return transport
|
||||
return None
|
||||
|
||||
|
||||
def _connection_is_idle(connection: object) -> bool:
|
||||
"""A pooled connection is idle unless it is servicing a request."""
|
||||
is_idle: Final[object] = getattr(connection, "is_idle", None)
|
||||
return bool(is_idle()) if callable(is_idle) else True
|
||||
|
||||
|
||||
def _pool_has_busy_connection(transport: object) -> bool | None:
|
||||
"""Whether the httpcore pool behind the transport is servicing a request.
|
||||
|
||||
``None`` when there is no such pool, so the caller can ask the other backend.
|
||||
"""
|
||||
pooled: Final[object] = getattr(getattr(transport, "_pool", None), "connections", None)
|
||||
if not isinstance(pooled, (list, tuple)):
|
||||
return None
|
||||
return any(
|
||||
not _connection_is_idle(connection) # pyright: ignore[reportUnknownArgumentType] # untyped pool list
|
||||
for connection in pooled # pyright: ignore[reportUnknownVariableType] # untyped pool list
|
||||
)
|
||||
|
||||
|
||||
def _has_connection_in_flight(client: object) -> bool:
|
||||
"""Whether the client is servicing a request right now.
|
||||
|
||||
Both connection backends litellm uses already account for the connections
|
||||
they have handed out, so this reads the client's own lease accounting rather
|
||||
than inferring it from elapsed time: httpcore reports a non-idle connection
|
||||
for the whole of a response including a stream, and aiohttp holds the
|
||||
connection in ``_acquired`` over the same span.
|
||||
|
||||
A client that cannot answer is reported as idle, which leaves the grace
|
||||
window as the only guard, exactly as it was before this check existed.
|
||||
"""
|
||||
try:
|
||||
transport: Final = _transport_of(client)
|
||||
pooled_busy: Final = _pool_has_busy_connection(transport)
|
||||
if pooled_busy is not None:
|
||||
return pooled_busy
|
||||
session: Final[object] = getattr(transport, "client", None)
|
||||
return bool(getattr(getattr(session, "connector", None), "_acquired", None))
|
||||
except Exception: # noqa: BLE001 - a client that cannot report its state is treated as idle
|
||||
return False
|
||||
|
||||
|
||||
async def _close_quietly(closing: Awaitable[object]) -> None:
|
||||
with contextlib.suppress(Exception):
|
||||
await closing
|
||||
|
||||
|
||||
class EvictedClientCloser:
|
||||
"""Closes evicted, litellm-owned clients once they are idle and out of grace."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
grace_seconds: float = EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS,
|
||||
max_pending: int = EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING,
|
||||
clock: Callable[[], float] = time.monotonic,
|
||||
) -> None:
|
||||
self._grace_seconds = grace_seconds
|
||||
self._max_pending = max_pending
|
||||
self._clock = clock
|
||||
self._owned: weakref.WeakSet[object] = weakref.WeakSet()
|
||||
self._buckets: dict[_BucketKey, deque[_PendingClose]] = {} # mutable-ok: deadline-ordered queues
|
||||
self._pending_count = 0
|
||||
self._queue_lock = threading.Lock() # the cache is reachable from every worker thread's loop
|
||||
self._close_tasks: set[asyncio.Task[None]] = set() # mutable-ok: strong refs to running closes
|
||||
|
||||
def mark_owned(self, client: object) -> None:
|
||||
"""Record that litellm created this client, so it may be closed on eviction."""
|
||||
try:
|
||||
self._owned.add(client)
|
||||
except TypeError:
|
||||
pass # values that cannot be weak-referenced are never litellm clients
|
||||
|
||||
def _is_owned(self, client: object) -> bool:
|
||||
try:
|
||||
return client in self._owned
|
||||
except TypeError:
|
||||
return False # unhashable values are never litellm clients
|
||||
|
||||
def schedule(self, client: object) -> None:
|
||||
"""Queue an evicted client for closing once it is idle and out of grace.
|
||||
|
||||
Past ``max_pending`` the client is left to the collector instead, so a
|
||||
workload that churns the cache cannot grow this queue without bound.
|
||||
Every queued entry comes due within one grace window, so the capacity it
|
||||
occupies is returned within that window rather than held.
|
||||
"""
|
||||
if client is None or not self._is_owned(client):
|
||||
return
|
||||
close_fn: Final = _close_function(client)
|
||||
if close_fn is None:
|
||||
return
|
||||
if self._pending_count >= self._max_pending:
|
||||
return
|
||||
self._enqueue(
|
||||
_PendingClose(
|
||||
client_ref=weakref.ref(client),
|
||||
loop_id=_running_loop_id(),
|
||||
needs_loop=inspect.iscoroutinefunction(close_fn),
|
||||
close_after=self._clock() + self._grace_seconds,
|
||||
)
|
||||
)
|
||||
|
||||
def reap(self) -> None:
|
||||
"""Close every queued client that is due, idle, and closable from here.
|
||||
|
||||
Called from the cache's read path, so the empty-queue exit comes first and
|
||||
the work done past it is proportional to what is due, not to the queue.
|
||||
"""
|
||||
if not self._pending_count:
|
||||
return
|
||||
now: Final = self._clock()
|
||||
for pending in self._take_due(_running_loop_id(), now):
|
||||
client = pending.client_ref()
|
||||
if client is None:
|
||||
continue
|
||||
if _has_connection_in_flight(client):
|
||||
self._enqueue(replace(pending, close_after=now + self._grace_seconds))
|
||||
continue
|
||||
self._close(client)
|
||||
|
||||
@property
|
||||
def pending_count(self) -> int:
|
||||
return self._pending_count
|
||||
|
||||
def _enqueue(self, pending: _PendingClose) -> None:
|
||||
"""Append to the entry's bucket, dropping any dead entries it queues behind.
|
||||
|
||||
Deadlines only ever move forward, so appending keeps each bucket ordered
|
||||
by deadline, and entries whose client the collector already took sit at
|
||||
the front rather than having to be searched for.
|
||||
"""
|
||||
with self._queue_lock:
|
||||
bucket: Final = self._buckets.setdefault(_bucket_key(pending), deque()) # mutable-ok: FIFO by design
|
||||
while bucket and bucket[0].client_ref() is None:
|
||||
bucket.popleft()
|
||||
self._pending_count -= 1
|
||||
bucket.append(pending)
|
||||
self._pending_count += 1
|
||||
|
||||
def _take_due(self, loop_id: int | None, now: float) -> tuple[_PendingClose, ...]:
|
||||
buckets = (_CLOSABLE_ANYWHERE,) if loop_id is None else (_CLOSABLE_ANYWHERE, _CLOSABLE_ON_ANY_LOOP, loop_id)
|
||||
with self._queue_lock:
|
||||
return tuple(pending for key in buckets for pending in self._drain_locked(key, now))
|
||||
|
||||
def _drain_locked(self, key: _BucketKey, now: float) -> Iterator[_PendingClose]:
|
||||
bucket: Final = self._buckets.get(key)
|
||||
if bucket is None:
|
||||
return
|
||||
while bucket and bucket[0].close_after <= now:
|
||||
self._pending_count -= 1
|
||||
yield bucket.popleft()
|
||||
if not bucket:
|
||||
del self._buckets[key]
|
||||
|
||||
def _close(self, client: object) -> None:
|
||||
close_fn: Final = _close_function(client)
|
||||
if close_fn is None:
|
||||
return
|
||||
try:
|
||||
closing: Final = close_fn()
|
||||
except Exception: # noqa: BLE001 - a discarded client's close must never surface to callers
|
||||
return
|
||||
if not inspect.isawaitable(closing):
|
||||
return
|
||||
task: Final = asyncio.get_running_loop().create_task(_close_quietly(closing))
|
||||
self._close_tasks.add(task)
|
||||
task.add_done_callback(self._close_tasks.discard)
|
||||
|
||||
|
||||
default_evicted_client_closer: Final = EvictedClientCloser()
|
||||
|
|
@ -5,21 +5,44 @@ Add the event loop to the cache key, to prevent event loop closed errors.
|
|||
import asyncio
|
||||
from typing import Final
|
||||
|
||||
from .evicted_client_closer import EvictedClientCloser, default_evicted_client_closer
|
||||
from .in_memory_cache import InMemoryCache
|
||||
|
||||
|
||||
class LLMClientCache(InMemoryCache):
|
||||
"""Cache for LLM HTTP clients (OpenAI, Azure, httpx, etc.).
|
||||
|
||||
IMPORTANT: This cache intentionally does NOT close clients on eviction.
|
||||
Evicted clients may still be in use by in-flight requests. Closing them
|
||||
eagerly causes ``RuntimeError: Cannot send a request, as the client has
|
||||
been closed.`` errors in production after the TTL (1 hour) expires.
|
||||
An evicted client is never closed on the spot: a request handed the client
|
||||
just before eviction is still using it, and closing it there raises
|
||||
``RuntimeError: Cannot send a request, as the client has been closed.``
|
||||
|
||||
Clients that are no longer referenced will be garbage-collected normally.
|
||||
For explicit shutdown cleanup, use ``close_litellm_async_clients()``.
|
||||
Nor can eviction be left to rely on garbage collection. The SDK clients are
|
||||
reference cycles, so an evicted client and its open TCP connections survive
|
||||
until a generational collection runs. Instead a client litellm created is
|
||||
handed to ``EvictedClientCloser``, which closes it once a grace window has
|
||||
passed. Clients the caller supplied are left untouched.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_size_in_memory: int | None = 200,
|
||||
default_ttl: int | None = 600,
|
||||
max_size_per_item: int | None = 1024,
|
||||
evicted_client_closer: EvictedClientCloser | None = None,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
max_size_in_memory=max_size_in_memory,
|
||||
default_ttl=default_ttl,
|
||||
max_size_per_item=max_size_per_item,
|
||||
)
|
||||
self.evicted_client_closer = evicted_client_closer or default_evicted_client_closer
|
||||
|
||||
def _remove_key(self, key: str) -> None:
|
||||
evicted: Final[object] = self.cache_dict.get(key)
|
||||
super()._remove_key(key)
|
||||
self.evicted_client_closer.schedule(evicted)
|
||||
self.evicted_client_closer.reap()
|
||||
|
||||
def update_cache_key_with_event_loop(self, key):
|
||||
"""
|
||||
Add the event loop to the cache key, to prevent event loop closed errors.
|
||||
|
|
@ -32,16 +55,22 @@ class LLMClientCache(InMemoryCache):
|
|||
except RuntimeError: # handle no current running event loop
|
||||
return key
|
||||
|
||||
def set_cache(self, key, value, **kwargs):
|
||||
def set_cache(self, key: str, value: object, litellm_owned_client: bool = False, **kwargs):
|
||||
"""``litellm_owned_client`` marks a client litellm built, so it may be closed once evicted."""
|
||||
if litellm_owned_client:
|
||||
self.evicted_client_closer.mark_owned(value)
|
||||
key = self.update_cache_key_with_event_loop(key)
|
||||
return super().set_cache(key, value, **kwargs)
|
||||
|
||||
async def async_set_cache(self, key, value, **kwargs):
|
||||
async def async_set_cache(self, key: str, value: object, litellm_owned_client: bool = False, **kwargs):
|
||||
if litellm_owned_client:
|
||||
self.evicted_client_closer.mark_owned(value)
|
||||
key = self.update_cache_key_with_event_loop(key)
|
||||
return await super().async_set_cache(key, value, **kwargs)
|
||||
|
||||
def get_cache(self, key, **kwargs):
|
||||
key = self.update_cache_key_with_event_loop(key)
|
||||
self.evicted_client_closer.reap()
|
||||
|
||||
return super().get_cache(key, **kwargs)
|
||||
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ import time
|
|||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from contextvars import ContextVar
|
||||
from datetime import timedelta
|
||||
from typing import TYPE_CHECKING, Any, Final, TypeVar, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, TypeVar, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
|
|
@ -49,7 +49,7 @@ if TYPE_CHECKING:
|
|||
cluster_pipeline = ClusterPipeline
|
||||
async_redis_client = Redis
|
||||
async_redis_cluster_client = RedisCluster
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
pipeline = Any
|
||||
cluster_pipeline = Any
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Key differences:
|
|||
- RedisClient NEEDs to be re-used across requests, adds 3000ms latency if it's re-created
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
|
|
@ -16,7 +16,7 @@ if TYPE_CHECKING:
|
|||
|
||||
pipeline = Pipeline
|
||||
async_redis_client = Redis
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
pipeline = Any
|
||||
async_redis_client = Any
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import ast
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Callable, Mapping
|
||||
from typing import Any, Final, cast
|
||||
|
||||
import litellm
|
||||
|
|
@ -47,7 +48,7 @@ class RedisSemanticCache(BaseCache):
|
|||
similarity_threshold: float | None = None,
|
||||
embedding_model: str = "text-embedding-ada-002",
|
||||
index_name: str | None = None,
|
||||
**kwargs,
|
||||
**kwargs: object,
|
||||
):
|
||||
"""
|
||||
Initialize the Redis Semantic Cache.
|
||||
|
|
@ -150,11 +151,11 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
def _init_semantic_cache(
|
||||
self,
|
||||
semantic_cache_cls: Any,
|
||||
semantic_cache_cls: Callable[..., object],
|
||||
index_name: str,
|
||||
redis_url: str,
|
||||
cache_vectorizer: Any,
|
||||
) -> Any:
|
||||
cache_vectorizer: object,
|
||||
) -> object:
|
||||
def _is_schema_mismatch(exc: ValueError) -> bool:
|
||||
error_message: Final = str(exc).lower()
|
||||
return any(phrase in error_message for phrase in ("schema does not match", "index schema"))
|
||||
|
|
@ -206,12 +207,12 @@ class RedisSemanticCache(BaseCache):
|
|||
def _get_cache_filters(self, key: str) -> dict[str, str]:
|
||||
return {self.CACHE_KEY_FIELD_NAME: str(key)}
|
||||
|
||||
def _get_cache_key_filter_expression(self, key: str) -> Any:
|
||||
def _get_cache_key_filter_expression(self, key: str) -> object:
|
||||
from redisvl.query.filter import Tag
|
||||
|
||||
return Tag(self.CACHE_KEY_FIELD_NAME) == str(key)
|
||||
|
||||
def _cache_hit_matches_key(self, cache_hit: dict[str, Any], key: str) -> bool:
|
||||
def _cache_hit_matches_key(self, cache_hit: Mapping[str, object], key: str) -> bool:
|
||||
# Pre-isolation entries with no ``litellm_cache_key`` field cannot be
|
||||
# safely reassigned to a caller's scope and are treated as misses.
|
||||
cached_key = cache_hit.get(self.CACHE_KEY_FIELD_NAME)
|
||||
|
|
@ -297,7 +298,7 @@ class RedisSemanticCache(BaseCache):
|
|||
return
|
||||
|
||||
@staticmethod
|
||||
def _coerce_response_input_value(value: Any) -> Any:
|
||||
def _coerce_response_input_value(value: object) -> object:
|
||||
model_dump: Final = getattr(value, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
return model_dump()
|
||||
|
|
@ -340,7 +341,7 @@ class RedisSemanticCache(BaseCache):
|
|||
)
|
||||
return embedding_response["data"][0]["embedding"]
|
||||
|
||||
def _get_cache_logic(self, cached_response: Any) -> Any:
|
||||
def _get_cache_logic(self, cached_response: Any) -> object:
|
||||
"""
|
||||
Process the cached response to prepare it for use.
|
||||
|
||||
|
|
@ -369,7 +370,7 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
return cached_response
|
||||
|
||||
def set_cache(self, key: str, value: Any, **kwargs) -> None:
|
||||
def set_cache(self, key: str, value: object, **kwargs) -> None:
|
||||
"""
|
||||
Store a value in the semantic cache.
|
||||
|
||||
|
|
@ -405,7 +406,7 @@ class RedisSemanticCache(BaseCache):
|
|||
except Exception as e:
|
||||
print_verbose(f"Error setting {value_str or value} in the Redis semantic cache: {e}")
|
||||
|
||||
def get_cache(self, key: str, **kwargs) -> Any:
|
||||
def get_cache(self, key: str, **kwargs) -> object:
|
||||
"""
|
||||
Retrieve a semantically similar cached response.
|
||||
|
||||
|
|
@ -428,7 +429,7 @@ class RedisSemanticCache(BaseCache):
|
|||
# Check the cache for semantically similar prompts in this exact
|
||||
# LiteLLM cache-key scope.
|
||||
prompt_embedding: Final = self._get_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
check_kwargs: Final[dict[str, Any]] = {
|
||||
check_kwargs: Final[Mapping[str, object]] = {
|
||||
"prompt": prompt,
|
||||
"vector": prompt_embedding,
|
||||
"filter_expression": self._get_cache_key_filter_expression(key),
|
||||
|
|
@ -508,7 +509,7 @@ class RedisSemanticCache(BaseCache):
|
|||
print_verbose(f"Error generating async embedding: {e}")
|
||||
raise ValueError(f"Failed to generate embedding: {e}") from e
|
||||
|
||||
async def async_set_cache(self, key: str, value: Any, **kwargs) -> None:
|
||||
async def async_set_cache(self, key: str, value: object, **kwargs) -> None:
|
||||
"""
|
||||
Asynchronously store a value in the semantic cache.
|
||||
|
||||
|
|
@ -548,7 +549,7 @@ class RedisSemanticCache(BaseCache):
|
|||
except Exception as e:
|
||||
print_verbose(f"Error in async_set_cache: {e}")
|
||||
|
||||
async def async_get_cache(self, key: str, **kwargs) -> Any:
|
||||
async def async_get_cache(self, key: str, **kwargs) -> object:
|
||||
"""
|
||||
Asynchronously retrieve a semantically similar cached response.
|
||||
|
||||
|
|
@ -573,7 +574,7 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
# Check the cache for semantically similar prompts in this exact
|
||||
# LiteLLM cache-key scope.
|
||||
check_kwargs: Final[dict[str, Any]] = {
|
||||
check_kwargs: Final[Mapping[str, object]] = {
|
||||
"prompt": prompt,
|
||||
"vector": prompt_embedding,
|
||||
"filter_expression": self._get_cache_key_filter_expression(key),
|
||||
|
|
@ -615,7 +616,7 @@ class RedisSemanticCache(BaseCache):
|
|||
print_verbose(f"Error in async_get_cache: {e}")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
|
||||
async def _index_info(self) -> dict[str, Any]:
|
||||
async def _index_info(self) -> Mapping[str, object]:
|
||||
"""
|
||||
Get information about the Redis index.
|
||||
|
||||
|
|
@ -625,7 +626,7 @@ class RedisSemanticCache(BaseCache):
|
|||
aindex: Final = await self.llmcache._get_async_index()
|
||||
return await aindex.info()
|
||||
|
||||
async def async_set_cache_pipeline(self, cache_list: list[tuple[str, Any]], **kwargs) -> None:
|
||||
async def async_set_cache_pipeline(self, cache_list: list[tuple[str, Any]], **kwargs: object) -> None:
|
||||
"""
|
||||
Asynchronously store multiple values in the semantic cache.
|
||||
|
||||
|
|
|
|||
|
|
@ -367,7 +367,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
stream_options = normalize_responses_api_stream_options(value)
|
||||
if stream_options is not None:
|
||||
responses_api_request["stream_options"] = stream_options
|
||||
elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys():
|
||||
elif key in ResponsesAPIOptionalRequestParams.__annotations__:
|
||||
responses_api_request[key] = value
|
||||
elif key == "previous_response_id":
|
||||
responses_api_request["previous_response_id"] = value
|
||||
|
|
|
|||
|
|
@ -91,6 +91,8 @@ DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD: Final = float(
|
|||
)
|
||||
MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH: Final = int(os.getenv("MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH", 150))
|
||||
|
||||
DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS: Final = 2000
|
||||
|
||||
# Semantic Guard Defaults
|
||||
DEFAULT_SEMANTIC_GUARD_EMBEDDING_MODEL: Final = str(
|
||||
os.getenv("DEFAULT_SEMANTIC_GUARD_EMBEDDING_MODEL", "text-embedding-3-small")
|
||||
|
|
@ -197,6 +199,16 @@ RUNWAYML_POLLING_TIMEOUT = int(os.getenv("RUNWAYML_POLLING_TIMEOUT", 600)) # 10
|
|||
########## Networking constants ##############################################################
|
||||
_DEFAULT_TTL_FOR_HTTPX_CLIENTS: Final = 3600 # 1 hour, re-use the same httpx client for 1 hour
|
||||
|
||||
# The earliest an evicted, litellm-created client may be closed. A request handed the
|
||||
# client just before eviction is still using it, so nothing is closed inside this window;
|
||||
# past it, the client is closed once it reports no connection in flight.
|
||||
EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS: Final = 900
|
||||
|
||||
# How many evicted clients may be queued for closing at once. Past this, an evicted client
|
||||
# is left to the collector rather than letting a cache-churning workload grow the queue
|
||||
# without bound. Each queued entry is ~100 bytes and comes due within one grace window.
|
||||
EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING: Final = 10_000
|
||||
|
||||
# Aiohttp connection pooling - prevents memory leaks from unbounded connection growth
|
||||
# Set to 0 for unlimited (not recommended for production)
|
||||
AIOHTTP_CONNECTOR_LIMIT: Final = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT", 1000))
|
||||
|
|
|
|||
|
|
@ -23,22 +23,20 @@ from .main import (
|
|||
)
|
||||
|
||||
__all__ = [
|
||||
# Core container operations
|
||||
"acreate_container",
|
||||
"adelete_container",
|
||||
"alist_containers",
|
||||
"aretrieve_container",
|
||||
"create_container",
|
||||
"delete_container",
|
||||
"list_containers",
|
||||
"retrieve_container",
|
||||
# Container file operations (auto-generated from endpoints.json)
|
||||
"adelete_container_file",
|
||||
"alist_container_files",
|
||||
"alist_containers",
|
||||
"aretrieve_container",
|
||||
"aretrieve_container_file",
|
||||
"aretrieve_container_file_content",
|
||||
"create_container",
|
||||
"delete_container",
|
||||
"delete_container_file",
|
||||
"list_container_files",
|
||||
"list_containers",
|
||||
"retrieve_container",
|
||||
"retrieve_container_file",
|
||||
"retrieve_container_file_content",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -80,7 +80,7 @@ async def acreate_fine_tuning_job(
|
|||
hyperparameters: dict | None = {},
|
||||
suffix: str | None = None,
|
||||
validation_file: str | None = None,
|
||||
integrations: List[str] | None = None,
|
||||
integrations: list[str] | None = None,
|
||||
seed: int | None = None,
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -157,7 +157,7 @@ def create_fine_tuning_job(
|
|||
hyperparameters: dict | None = {},
|
||||
suffix: str | None = None,
|
||||
validation_file: str | None = None,
|
||||
integrations: List[str] | None = None,
|
||||
integrations: list[str] | None = None,
|
||||
seed: int | None = None,
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
from collections.abc import AsyncIterator, Coroutine
|
||||
from typing import Any, Final, cast
|
||||
from collections.abc import AsyncIterator, Coroutine, Mapping
|
||||
from typing import Final, cast
|
||||
|
||||
import litellm
|
||||
from litellm.types.google_genai.adapters import GenerateContentCompletionKwargs
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
|
@ -17,12 +18,12 @@ class GenerateContentToCompletionHandler:
|
|||
@staticmethod
|
||||
def _prepare_completion_kwargs(
|
||||
model: str,
|
||||
contents: list[dict[str, Any]] | dict[str, Any],
|
||||
config: dict[str, Any] | None = None,
|
||||
contents: list[dict[str, object]] | dict[str, object],
|
||||
config: dict[str, object] | None = None,
|
||||
stream: bool = False,
|
||||
litellm_params: GenericLiteLLMParams | None = None,
|
||||
extra_kwargs: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
extra_kwargs: Mapping[str, object] | None = None,
|
||||
) -> GenerateContentCompletionKwargs:
|
||||
"""Prepare kwargs for litellm.completion/acompletion"""
|
||||
|
||||
# Transform generate_content request to completion format
|
||||
|
|
@ -34,7 +35,7 @@ class GenerateContentToCompletionHandler:
|
|||
**(extra_kwargs or {}),
|
||||
)
|
||||
|
||||
completion_kwargs: Final[dict[str, Any]] = dict(completion_request)
|
||||
completion_kwargs: Final = dict(completion_request)
|
||||
|
||||
# Forward extra_kwargs that should be passed to completion call
|
||||
if extra_kwargs is not None:
|
||||
|
|
@ -48,17 +49,17 @@ class GenerateContentToCompletionHandler:
|
|||
if stream:
|
||||
completion_kwargs["stream"] = stream
|
||||
|
||||
return completion_kwargs
|
||||
return GenerateContentCompletionKwargs(**completion_kwargs)
|
||||
|
||||
@staticmethod
|
||||
async def async_generate_content_handler(
|
||||
model: str,
|
||||
contents: list[dict[str, Any]] | dict[str, Any],
|
||||
contents: list[dict[str, object]] | dict[str, object],
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
config: dict[str, Any] | None = None,
|
||||
config: dict[str, object] | None = None,
|
||||
stream: bool = False,
|
||||
**kwargs,
|
||||
) -> dict[str, Any] | AsyncIterator[bytes]:
|
||||
**kwargs: object,
|
||||
) -> dict[str, object] | AsyncIterator[bytes]:
|
||||
"""Handle generate_content call asynchronously using completion adapter"""
|
||||
|
||||
completion_kwargs: Final = GenerateContentToCompletionHandler._prepare_completion_kwargs(
|
||||
|
|
@ -103,13 +104,13 @@ class GenerateContentToCompletionHandler:
|
|||
@staticmethod
|
||||
def generate_content_handler(
|
||||
model: str,
|
||||
contents: list[dict[str, Any]] | dict[str, Any],
|
||||
contents: list[dict[str, object]] | dict[str, object],
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
config: dict[str, Any] | None = None,
|
||||
config: dict[str, object] | None = None,
|
||||
stream: bool = False,
|
||||
_is_async: bool = False,
|
||||
**kwargs,
|
||||
) -> dict[str, Any] | AsyncIterator[bytes] | Coroutine[Any, Any, dict[str, Any] | AsyncIterator[bytes]]:
|
||||
**kwargs: object,
|
||||
) -> dict[str, object] | AsyncIterator[bytes] | Coroutine[None, None, dict[str, object] | AsyncIterator[bytes]]:
|
||||
"""Handle generate_content call using completion adapter"""
|
||||
|
||||
if _is_async:
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ this file has Arize ai specific helper functions
|
|||
|
||||
import os
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm.integrations.arize import _utils
|
||||
from litellm.integrations.arize._utils import ArizeOTELAttributes
|
||||
|
|
@ -21,7 +21,7 @@ if TYPE_CHECKING:
|
|||
from litellm.types.integrations.arize import Protocol as _Protocol
|
||||
|
||||
Protocol = _Protocol
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
Protocol = Any
|
||||
Span = Any
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import os
|
||||
import threading
|
||||
from collections import OrderedDict
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.arize import _utils
|
||||
|
|
@ -22,7 +22,7 @@ if TYPE_CHECKING:
|
|||
|
||||
Protocol = _Protocol
|
||||
OpenTelemetryConfig = _OpenTelemetryConfig
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
OpenTelemetry = _OpenTelemetry
|
||||
LITELLM_TRACER_NAME: str
|
||||
else:
|
||||
|
|
@ -430,7 +430,8 @@ class ArizePhoenixLogger(OpenTelemetry):
|
|||
|
||||
otlp_auth_headers = None
|
||||
if api_key is not None:
|
||||
otlp_auth_headers = f"Authorization=Bearer {api_key}"
|
||||
auth_header_key = "authorization" if protocol == "otlp_grpc" else "Authorization"
|
||||
otlp_auth_headers = f"{auth_header_key}=Bearer {api_key}"
|
||||
elif "app.phoenix.arize.com" in endpoint:
|
||||
raise ValueError("PHOENIX_API_KEY must be set when using Phoenix Cloud (app.phoenix.arize.com).")
|
||||
|
||||
|
|
|
|||
|
|
@ -16,7 +16,10 @@ import asyncio
|
|||
import os
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
|
|
@ -27,6 +30,16 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
)
|
||||
from litellm.types.utils import StandardAuditLogPayload, StandardLoggingPayload
|
||||
|
||||
DEFAULT_AZURE_AUTHORITY_HOST: Final = "https://login.microsoftonline.com"
|
||||
DEFAULT_AZURE_MONITOR_SCOPE: Final = "https://monitor.azure.com/.default"
|
||||
|
||||
MONITOR_SCOPE_BY_AUTHORITY_HOST: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
"login.microsoftonline.com": DEFAULT_AZURE_MONITOR_SCOPE,
|
||||
"login.microsoftonline.us": "https://monitor.azure.us/.default",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class AzureSentinelLogger(CustomBatchLogger):
|
||||
"""
|
||||
|
|
@ -42,6 +55,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
client_id: str | None = None,
|
||||
client_secret: str | None = None,
|
||||
audit_stream_name: str | None = None,
|
||||
authority_host: str | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
|
|
@ -62,6 +76,10 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
If not provided, will use AZURE_SENTINEL_CLIENT_SECRET or AZURE_CLIENT_SECRET env var.
|
||||
audit_stream_name (str, optional): Stream name from DCR for audit logs.
|
||||
If not provided, will use AZURE_SENTINEL_AUDIT_STREAM_NAME env var or the standard stream name.
|
||||
authority_host (str, optional): Microsoft Entra authority host that issues the OAuth2 token,
|
||||
e.g. "https://login.microsoftonline.us" for Azure Government. If not provided, will use
|
||||
AZURE_AUTHORITY_HOST env var or default to the Azure Public Cloud authority. The Azure
|
||||
Monitor audience is derived from it.
|
||||
"""
|
||||
self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
|
||||
|
||||
|
|
@ -76,6 +94,9 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
resolved_client_secret: Final = (
|
||||
client_secret or os.getenv("AZURE_SENTINEL_CLIENT_SECRET") or os.getenv("AZURE_CLIENT_SECRET")
|
||||
)
|
||||
resolved_authority_host: Final = self._normalize_authority_host(
|
||||
authority_host or os.getenv("AZURE_AUTHORITY_HOST") or DEFAULT_AZURE_AUTHORITY_HOST
|
||||
)
|
||||
|
||||
if not resolved_dcr_immutable_id:
|
||||
raise ValueError(
|
||||
|
|
@ -119,7 +140,8 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
)
|
||||
|
||||
# OAuth2 scope for Azure Monitor
|
||||
self.oauth_scope = "https://monitor.azure.com/.default"
|
||||
self.authority_host = resolved_authority_host
|
||||
self.oauth_scope = self._resolve_oauth_scope(authority_host=resolved_authority_host)
|
||||
self.oauth_token: str | None = None
|
||||
self.oauth_token_expires_at: float | None = None
|
||||
|
||||
|
|
@ -129,6 +151,26 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
self.log_queue: list[StandardLoggingPayload] = []
|
||||
self.audit_log_queue: list[StandardAuditLogPayload] = []
|
||||
|
||||
@staticmethod
|
||||
def _normalize_authority_host(authority_host: str) -> str:
|
||||
"""
|
||||
Normalize an authority host into an absolute URL with no trailing slash.
|
||||
|
||||
Accepts the scheme-qualified form litellm documents ("https://login.microsoftonline.us")
|
||||
and the bare-host form the azure-identity AzureAuthorityHosts constants use.
|
||||
"""
|
||||
stripped: Final = authority_host.strip().rstrip("/")
|
||||
return stripped if "://" in stripped else f"https://{stripped}"
|
||||
|
||||
@staticmethod
|
||||
def _resolve_oauth_scope(authority_host: str) -> str:
|
||||
"""
|
||||
Map an authority host to the Azure Monitor Logs Ingestion audience for the same cloud,
|
||||
falling back to the Azure Public Cloud audience for an unrecognized host.
|
||||
"""
|
||||
host: Final = urlparse(authority_host).hostname or ""
|
||||
return MONITOR_SCOPE_BY_AUTHORITY_HOST.get(host, DEFAULT_AZURE_MONITOR_SCOPE)
|
||||
|
||||
@staticmethod
|
||||
def _build_api_endpoint(endpoint: str, dcr_immutable_id: str, stream_name: str) -> str:
|
||||
return f"{endpoint.rstrip('/')}/dataCollectionRules/{dcr_immutable_id}/streams/{stream_name}?api-version=2023-01-01"
|
||||
|
|
@ -150,7 +192,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
assert self.client_id is not None, "client_id is required"
|
||||
assert self.client_secret is not None, "client_secret is required"
|
||||
|
||||
token_url: Final = f"https://login.microsoftonline.com/{self.tenant_id}/oauth2/v2.0/token"
|
||||
token_url: Final = f"{self.authority_host}/{self.tenant_id}/oauth2/v2.0/token"
|
||||
|
||||
token_data: Final = {
|
||||
"client_id": self.client_id,
|
||||
|
|
|
|||
|
|
@ -9,13 +9,18 @@ captured stdout back through the typed agentic loop plan.
|
|||
import json
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Final, Literal, TypedDict, cast
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeAlias, TypedDict, runtime_checkable
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.llms.base_llm.sandbox.transformation import (
|
||||
CodeExecutionResult,
|
||||
ContainerHandle,
|
||||
)
|
||||
from litellm.types.integrations.code_interpreter_interception import (
|
||||
CodeInterpreterInterceptionConfig,
|
||||
)
|
||||
|
|
@ -37,6 +42,9 @@ from litellm.types.utils import (
|
|||
ModelResponse,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
LITELLM_CODE_EXECUTION_TOOL_NAME: Final = "litellm_code_execution"
|
||||
_INTERCEPTION_ACTIVE_KEY: Final = "_code_interpreter_interception_active"
|
||||
_SANDBOX_KEY: Final = "_code_interpreter_interception_sandbox_key"
|
||||
|
|
@ -109,26 +117,94 @@ class ChatCompletionFunctionToolChoice(TypedDict):
|
|||
CodeExecutionFunctionToolChoice = ResponsesFunctionToolChoice | ChatCompletionFunctionToolChoice
|
||||
|
||||
|
||||
def _extract_session_id(kwargs: dict[str, Any]) -> str | None:
|
||||
class SandboxToolParams(TypedDict):
|
||||
sandbox_provider: str
|
||||
api_key: str | None
|
||||
api_base: str | None
|
||||
|
||||
|
||||
class SandboxConfigProtocol(Protocol):
|
||||
async def acreate_sandbox(self) -> ContainerHandle: ...
|
||||
|
||||
async def arun_code(self, *, container: ContainerHandle, code: str) -> CodeExecutionResult: ...
|
||||
|
||||
async def adelete_sandbox(self, *, container: ContainerHandle) -> object: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _SupportsOutput(Protocol):
|
||||
output: object
|
||||
|
||||
|
||||
_CachedContainer: TypeAlias = tuple[ContainerHandle, SandboxToolParams | None, float, str | None]
|
||||
|
||||
|
||||
def _output_item_type(item: object) -> object:
|
||||
if isinstance(item, dict):
|
||||
item_mapping: Final[dict[str, object]] = item
|
||||
return item_mapping.get("type")
|
||||
return getattr(item, "type", None)
|
||||
|
||||
|
||||
def _response_output(response: object) -> object:
|
||||
if isinstance(response, dict):
|
||||
response_mapping: Final[Mapping[str, object]] = response
|
||||
return response_mapping.get("output", [])
|
||||
return getattr(response, "output", []) or []
|
||||
|
||||
|
||||
def _tool_call_arguments(arguments: object) -> str:
|
||||
if isinstance(arguments, str):
|
||||
return arguments
|
||||
return "" if arguments is None else str(arguments)
|
||||
|
||||
|
||||
def _narrow_tool_call(tool_call: Mapping[str, object]) -> CodeExecutionToolCall:
|
||||
tool_call_id: Final = tool_call.get("id")
|
||||
call_id: Final = tool_call.get("call_id")
|
||||
return {
|
||||
"id": tool_call_id if isinstance(tool_call_id, str) else None,
|
||||
"call_id": call_id if isinstance(call_id, str) else None,
|
||||
"type": "function",
|
||||
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
|
||||
"arguments": _tool_call_arguments(tool_call.get("arguments")),
|
||||
}
|
||||
|
||||
|
||||
def _extract_session_id(kwargs: dict[str, object]) -> str | None:
|
||||
for meta_key in ("metadata", "litellm_metadata"):
|
||||
meta = kwargs.get(meta_key)
|
||||
if isinstance(meta, dict):
|
||||
sid = meta.get("session_id")
|
||||
metadata: dict[str, object] = meta
|
||||
sid = metadata.get("session_id")
|
||||
if sid and isinstance(sid, str):
|
||||
return sid
|
||||
return None
|
||||
|
||||
|
||||
def _extract_identity(kwargs: dict[str, Any]) -> str:
|
||||
return kwargs.get("user_api_key_hash") or ""
|
||||
def _extract_identity(kwargs: Mapping[str, object]) -> str:
|
||||
identity: Final = kwargs.get("user_api_key_hash")
|
||||
return identity if isinstance(identity, str) else ""
|
||||
|
||||
|
||||
def _resolve_sandbox_tool(sandbox_tool_name: str | None) -> dict[str, Any] | None:
|
||||
def _resolve_sandbox_tool(sandbox_tool_name: str | None) -> SandboxToolParams | None:
|
||||
if sandbox_tool_name is None:
|
||||
return None
|
||||
try:
|
||||
from litellm.sandbox.sandbox_tools import resolve_sandbox_tool
|
||||
except ImportError:
|
||||
return None
|
||||
return resolve_sandbox_tool(sandbox_tool_name)
|
||||
resolved: Final[dict[str, object] | None] = resolve_sandbox_tool(sandbox_tool_name)
|
||||
if resolved is None:
|
||||
return None
|
||||
provider: Final = resolved.get("sandbox_provider")
|
||||
api_key: Final = resolved.get("api_key")
|
||||
api_base: Final = resolved.get("api_base")
|
||||
return SandboxToolParams(
|
||||
sandbox_provider=provider if isinstance(provider, str) else "",
|
||||
api_key=api_key if isinstance(api_key, str) else None,
|
||||
api_base=api_base if isinstance(api_base, str) else None,
|
||||
)
|
||||
|
||||
|
||||
class CodeInterpreterInterceptionLogger(CustomLogger):
|
||||
|
|
@ -149,14 +225,14 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
enabled: bool = True,
|
||||
enabled_providers: list[str] | None = None,
|
||||
sandbox_tool_name: str | None = None,
|
||||
sandbox_config: Any | None = None,
|
||||
sandbox_config: SandboxConfigProtocol | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.enabled = enabled
|
||||
self.enabled_providers = enabled_providers
|
||||
self.sandbox_tool_name = sandbox_tool_name
|
||||
self.sandbox_config = sandbox_config
|
||||
self._container_cache: dict[str, tuple[Any, dict[str, Any] | None, float, str | None]] = {}
|
||||
self._container_cache: dict[str, _CachedContainer] = {}
|
||||
|
||||
@classmethod
|
||||
def from_config_yaml(cls, config: CodeInterpreterInterceptionConfig) -> "CodeInterpreterInterceptionLogger":
|
||||
|
|
@ -171,19 +247,18 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
litellm_settings: dict[str, Any],
|
||||
callback_specific_params: dict[str, Any],
|
||||
) -> "CodeInterpreterInterceptionLogger":
|
||||
params: CodeInterpreterInterceptionConfig = {}
|
||||
if "code_interpreter_interception_params" in litellm_settings:
|
||||
params = litellm_settings["code_interpreter_interception_params"]
|
||||
elif "code_interpreter_interception" in callback_specific_params and isinstance(
|
||||
callback_specific_params["code_interpreter_interception"], dict
|
||||
):
|
||||
params = cast(
|
||||
CodeInterpreterInterceptionConfig,
|
||||
callback_specific_params["code_interpreter_interception"],
|
||||
)
|
||||
params: Final[CodeInterpreterInterceptionConfig] = (
|
||||
litellm_settings["code_interpreter_interception_params"]
|
||||
if "code_interpreter_interception_params" in litellm_settings
|
||||
else callback_specific_params["code_interpreter_interception"]
|
||||
if isinstance(callback_specific_params.get("code_interpreter_interception"), dict)
|
||||
else {}
|
||||
)
|
||||
return CodeInterpreterInterceptionLogger.from_config_yaml(params)
|
||||
|
||||
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, object], call_type: CallTypes | None
|
||||
) -> dict | None:
|
||||
if not kwargs.get("_agentic_loop_depth"):
|
||||
kwargs.pop(_INTERCEPTION_ACTIVE_KEY, None)
|
||||
kwargs.pop(_SANDBOX_KEY, None)
|
||||
|
|
@ -229,13 +304,14 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
return kwargs
|
||||
|
||||
@staticmethod
|
||||
def _strip_interception_metadata(kwargs: dict[str, Any]) -> None:
|
||||
def _strip_interception_metadata(kwargs: dict[str, object]) -> None:
|
||||
metadata: Final = kwargs.get(_LITELLM_METADATA_KEY)
|
||||
if not isinstance(metadata, dict):
|
||||
return
|
||||
current_metadata: Final[dict[str, object]] = metadata
|
||||
filtered_metadata: Final = {
|
||||
key: value
|
||||
for key, value in metadata.items()
|
||||
for key, value in current_metadata.items()
|
||||
if not is_interception_internal_key(key)
|
||||
and not key.startswith("_agentic_loop")
|
||||
and key != "max_agentic_loops"
|
||||
|
|
@ -247,9 +323,9 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
kwargs.pop(_LITELLM_METADATA_KEY, None)
|
||||
|
||||
@staticmethod
|
||||
def _write_interception_metadata(kwargs: dict[str, Any]) -> None:
|
||||
metadata = kwargs.get(_LITELLM_METADATA_KEY)
|
||||
metadata = dict(metadata) if isinstance(metadata, dict) else {}
|
||||
def _write_interception_metadata(kwargs: dict[str, object]) -> None:
|
||||
existing: Final = kwargs.get(_LITELLM_METADATA_KEY)
|
||||
metadata: Final[dict[str, object]] = dict(existing) if isinstance(existing, dict) else {}
|
||||
for key in (_INTERCEPTION_ACTIVE_KEY, _SANDBOX_KEY, _SESSION_SCOPED_KEY, _CONVERTED_STREAM_KEY):
|
||||
if key in kwargs:
|
||||
metadata[key] = kwargs[key]
|
||||
|
|
@ -296,20 +372,21 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
}
|
||||
|
||||
@staticmethod
|
||||
def _tool_choice_targets_code_interpreter(tool_choice: Any) -> bool:
|
||||
def _tool_choice_targets_code_interpreter(tool_choice: object) -> bool:
|
||||
if not isinstance(tool_choice, dict):
|
||||
return False
|
||||
function: Final = tool_choice.get("function")
|
||||
choice: Final[dict[str, object]] = tool_choice
|
||||
function: Final = choice.get("function")
|
||||
return (
|
||||
tool_choice.get("type") == "code_interpreter"
|
||||
or tool_choice.get("name") == "code_interpreter"
|
||||
or tool_choice.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME
|
||||
choice.get("type") == "code_interpreter"
|
||||
or choice.get("name") == "code_interpreter"
|
||||
or choice.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME
|
||||
or (isinstance(function, dict) and function.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME)
|
||||
)
|
||||
|
||||
def _resolve_provider(self, kwargs: dict[str, Any]) -> str | None:
|
||||
def _resolve_provider(self, kwargs: dict[str, object]) -> str | None:
|
||||
provider: Final = kwargs.get("custom_llm_provider")
|
||||
if provider:
|
||||
if isinstance(provider, str) and provider:
|
||||
return provider
|
||||
model: Final = kwargs.get("model")
|
||||
if not isinstance(model, str):
|
||||
|
|
@ -321,7 +398,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
|
||||
async def async_should_run_agentic_loop(
|
||||
self,
|
||||
response: Any,
|
||||
response: object,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
tools: list[dict] | None,
|
||||
|
|
@ -351,12 +428,12 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
tools: dict,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
response: Any,
|
||||
anthropic_messages_provider_config: Any,
|
||||
anthropic_messages_optional_request_params: dict,
|
||||
logging_obj: Any,
|
||||
response: object,
|
||||
anthropic_messages_provider_config: object,
|
||||
anthropic_messages_optional_request_params: dict[str, object],
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
stream: bool,
|
||||
kwargs: dict,
|
||||
kwargs: dict[str, object],
|
||||
) -> AgenticLoopPlan:
|
||||
if kwargs.get("_agentic_loop_api_surface") == CHAT_COMPLETION_AGENTIC_SURFACE:
|
||||
return await self._build_chat_completion_agentic_loop_plan(
|
||||
|
|
@ -368,14 +445,14 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
)
|
||||
|
||||
await self._prune_expired_cache()
|
||||
tool_calls: Final = cast(list[CodeExecutionToolCall], tools.get("tool_calls", []))
|
||||
sandbox_key: Final = kwargs.get(_SANDBOX_KEY)
|
||||
tool_calls: Final = self._agentic_tool_calls(tools)
|
||||
sandbox_key: Final = self._extract_sandbox_key(kwargs)
|
||||
is_session: Final = bool(kwargs.get(_SESSION_SCOPED_KEY))
|
||||
identity: Final = _extract_identity(kwargs) if is_session else None
|
||||
container, params = await self._get_or_create_container(cache_key=sandbox_key, identity=identity)
|
||||
|
||||
try:
|
||||
container_id: Final = cast(str | None, getattr(container, "id", None))
|
||||
container_id: Final = self._container_id(container)
|
||||
input_list: Final = self._normalize_messages(messages)
|
||||
code_interpreter_calls: Final[list[CodeInterpreterCall]] = []
|
||||
for tool_call in tool_calls:
|
||||
|
|
@ -443,14 +520,14 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
kwargs: dict[str, object],
|
||||
) -> AgenticLoopPlan:
|
||||
await self._prune_expired_cache()
|
||||
tool_calls: Final = cast(list[CodeExecutionToolCall], tools.get("tool_calls", []))
|
||||
sandbox_key: Final = cast(str | None, kwargs.get(_SANDBOX_KEY))
|
||||
tool_calls: Final = self._agentic_tool_calls(tools)
|
||||
sandbox_key: Final = self._extract_sandbox_key(kwargs)
|
||||
is_session: Final = bool(kwargs.get(_SESSION_SCOPED_KEY))
|
||||
identity: Final = _extract_identity(cast(dict[str, Any], kwargs)) if is_session else None
|
||||
identity: Final = _extract_identity(kwargs) if is_session else None
|
||||
container, params = await self._get_or_create_container(cache_key=sandbox_key, identity=identity)
|
||||
|
||||
try:
|
||||
container_id: Final = cast(str | None, getattr(container, "id", None))
|
||||
container_id: Final = self._container_id(container)
|
||||
tool_results: Final = [
|
||||
await self._build_chat_completion_tool_result(
|
||||
container=container,
|
||||
|
|
@ -489,10 +566,28 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _container_id(container: ContainerHandle) -> str | None:
|
||||
container_id: Final[object] = getattr(container, "id", None)
|
||||
return container_id if isinstance(container_id, str) else None
|
||||
|
||||
@staticmethod
|
||||
def _agentic_tool_calls(tools: dict[str, object]) -> list[CodeExecutionToolCall]:
|
||||
tool_calls: Final = tools.get("tool_calls")
|
||||
if not isinstance(tool_calls, list):
|
||||
return []
|
||||
items: Final[list[object]] = tool_calls
|
||||
return [_narrow_tool_call(item) for item in items if isinstance(item, dict)]
|
||||
|
||||
@staticmethod
|
||||
def _extract_sandbox_key(kwargs: dict[str, object]) -> str | None:
|
||||
sandbox_key: Final = kwargs.get(_SANDBOX_KEY)
|
||||
return sandbox_key if isinstance(sandbox_key, str) else None
|
||||
|
||||
async def _build_chat_completion_tool_result(
|
||||
self,
|
||||
container: object,
|
||||
params: dict[str, Any] | None,
|
||||
container: ContainerHandle,
|
||||
params: SandboxToolParams | None,
|
||||
tool_call: CodeExecutionToolCall,
|
||||
container_id: str | None,
|
||||
) -> tuple[ChatCompletionToolMessage, CodeInterpreterCall]:
|
||||
|
|
@ -517,10 +612,15 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
)
|
||||
|
||||
async def async_agentic_loop_cleanup_hook(self, plan: AgenticLoopPlan, kwargs: dict) -> None:
|
||||
metadata: Final = plan.metadata or {} if plan else {}
|
||||
metadata: Final[dict[str, object]] = plan.metadata or {} if plan else {}
|
||||
if metadata.get("is_session_scoped"):
|
||||
return
|
||||
await self._delete_container_for_cache_key(metadata.get("sandbox_key"))
|
||||
await self._delete_container_for_cache_key(self._metadata_sandbox_key(metadata))
|
||||
|
||||
@staticmethod
|
||||
def _metadata_sandbox_key(metadata: Mapping[str, object]) -> str | None:
|
||||
sandbox_key: Final = metadata.get("sandbox_key")
|
||||
return sandbox_key if isinstance(sandbox_key, str) else None
|
||||
|
||||
@staticmethod
|
||||
def _filter_agentic_loop_kwargs(kwargs: dict[str, object]) -> dict[str, object]:
|
||||
|
|
@ -531,12 +631,12 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
and not is_interception_internal_key(k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES)
|
||||
}
|
||||
|
||||
def _get_followup_tools(self, tools: object, call_type: CallTypes | None) -> list[dict[str, Any]] | None:
|
||||
def _get_followup_tools(self, tools: object, call_type: CallTypes | None) -> list[dict[str, object]] | None:
|
||||
if not isinstance(tools, list):
|
||||
return None
|
||||
return [
|
||||
(
|
||||
self._get_function_tool(call_type=call_type)
|
||||
dict(self._get_function_tool(call_type=call_type))
|
||||
if isinstance(tool, dict) and tool.get("type") == "code_interpreter"
|
||||
else tool
|
||||
)
|
||||
|
|
@ -549,34 +649,42 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
k: v for k, v in optional_params.items() if k != "tools" and not (k == "tool_choice" and drop_tool_choice)
|
||||
}
|
||||
|
||||
async def async_post_agentic_loop_response_hook(self, response: Any, plan: AgenticLoopPlan, kwargs: dict) -> Any:
|
||||
metadata: Final = plan.metadata or {} if plan else {}
|
||||
async def async_post_agentic_loop_response_hook(
|
||||
self, response: object, plan: AgenticLoopPlan, kwargs: dict
|
||||
) -> object:
|
||||
metadata: Final[dict[str, object]] = plan.metadata or {} if plan else {}
|
||||
if not metadata.get("is_session_scoped"):
|
||||
await self._delete_container_for_cache_key(metadata.get("sandbox_key"))
|
||||
await self._delete_container_for_cache_key(self._metadata_sandbox_key(metadata))
|
||||
|
||||
calls: Final = metadata.get("code_interpreter_calls")
|
||||
if not calls:
|
||||
if not calls or not isinstance(calls, list):
|
||||
return response
|
||||
|
||||
is_dict: Final = isinstance(response, dict)
|
||||
output: Final = response.get("output") if is_dict else getattr(response, "output", None)
|
||||
if not isinstance(output, list):
|
||||
if isinstance(response, dict):
|
||||
response_mapping: Final[dict[str, object]] = response
|
||||
merged_mapping_output: Final = self._merge_code_interpreter_calls(response_mapping.get("output"), calls)
|
||||
if merged_mapping_output is not None:
|
||||
response_mapping["output"] = merged_mapping_output
|
||||
return response
|
||||
|
||||
def _item_type(item: Any) -> Any:
|
||||
return item.get("type") if isinstance(item, dict) else getattr(item, "type", None)
|
||||
|
||||
insert_at: Final = next(
|
||||
(i for i, item in enumerate(output) if _item_type(item) == "message"),
|
||||
len(output),
|
||||
)
|
||||
new_output: Final = output[:insert_at] + list(calls) + output[insert_at:]
|
||||
if is_dict:
|
||||
response["output"] = new_output
|
||||
else:
|
||||
response.output = new_output
|
||||
if not isinstance(response, _SupportsOutput):
|
||||
return response
|
||||
merged_attr_output: Final = self._merge_code_interpreter_calls(response.output, calls)
|
||||
if merged_attr_output is not None:
|
||||
response.output = merged_attr_output
|
||||
return response
|
||||
|
||||
@staticmethod
|
||||
def _merge_code_interpreter_calls(output: object, calls: Sequence[object]) -> list[object] | None:
|
||||
if not isinstance(output, list):
|
||||
return None
|
||||
items: Final[list[object]] = output
|
||||
insert_at: Final = next(
|
||||
(i for i, item in enumerate(items) if _output_item_type(item) == "message"),
|
||||
len(items),
|
||||
)
|
||||
return items[:insert_at] + list(calls) + items[insert_at:]
|
||||
|
||||
@staticmethod
|
||||
def _parse_code(arguments: str) -> str:
|
||||
try:
|
||||
|
|
@ -584,7 +692,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
except (json.JSONDecodeError, TypeError, AttributeError):
|
||||
return ""
|
||||
|
||||
async def _run_tool_call(self, container: Any, params: dict[str, Any] | None, arguments: str) -> str:
|
||||
async def _run_tool_call(self, container: ContainerHandle, params: SandboxToolParams | None, arguments: str) -> str:
|
||||
try:
|
||||
code: Final = json.loads(arguments).get("code", "") if arguments else ""
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
|
|
@ -601,7 +709,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
self,
|
||||
cache_key: str | None,
|
||||
identity: str | None = None,
|
||||
) -> tuple[Any, dict[str, Any] | None]:
|
||||
) -> tuple[ContainerHandle, SandboxToolParams | None]:
|
||||
if cache_key:
|
||||
cached: Final = self._container_cache.get(cache_key)
|
||||
if cached is not None:
|
||||
|
|
@ -623,7 +731,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
self._container_cache.pop(lru_key, None)
|
||||
await self._delete_container(container=lru_entry[0], params=lru_entry[1])
|
||||
|
||||
async def _create_container(self) -> tuple[Any, dict[str, Any] | None]:
|
||||
async def _create_container(self) -> tuple[ContainerHandle, SandboxToolParams | None]:
|
||||
if self.sandbox_config is not None:
|
||||
return await self.sandbox_config.acreate_sandbox(), None
|
||||
|
||||
|
|
@ -641,7 +749,9 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
)
|
||||
return container, params
|
||||
|
||||
async def _run_code(self, container: Any, params: dict[str, Any] | None, code: str) -> Any:
|
||||
async def _run_code(
|
||||
self, container: ContainerHandle, params: SandboxToolParams | None, code: str
|
||||
) -> CodeExecutionResult:
|
||||
if self.sandbox_config is not None:
|
||||
return await self.sandbox_config.arun_code(container=container, code=code)
|
||||
if params is None:
|
||||
|
|
@ -653,7 +763,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
api_key=params.get("api_key"),
|
||||
)
|
||||
|
||||
async def _delete_container(self, container: Any, params: dict[str, Any] | None) -> None:
|
||||
async def _delete_container(self, container: ContainerHandle, params: SandboxToolParams | None) -> None:
|
||||
try:
|
||||
if self.sandbox_config is not None:
|
||||
await self.sandbox_config.adelete_sandbox(container=container)
|
||||
|
|
@ -677,7 +787,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
return
|
||||
await self._delete_container(container=cached[0], params=cached[1])
|
||||
|
||||
def _normalize_messages(self, messages: Any) -> list[dict[str, Any]]:
|
||||
def _normalize_messages(self, messages: object) -> list[dict[str, object]]:
|
||||
if isinstance(messages, str):
|
||||
return [{"role": "user", "content": messages}]
|
||||
if isinstance(messages, list):
|
||||
|
|
@ -685,10 +795,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
return []
|
||||
|
||||
def _extract_code_execution_tool_calls(self, response: object) -> list[CodeExecutionToolCall]:
|
||||
if isinstance(response, dict):
|
||||
output = response.get("output", [])
|
||||
else:
|
||||
output = getattr(response, "output", []) or []
|
||||
output: Final = _response_output(response)
|
||||
if not isinstance(output, list):
|
||||
return []
|
||||
|
||||
|
|
@ -702,9 +809,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
if self._is_code_execution_call(item)
|
||||
]
|
||||
|
||||
def _extract_chat_completion_code_execution_tool_calls(
|
||||
self, response: ModelResponse | dict[str, Any]
|
||||
) -> list[CodeExecutionToolCall]:
|
||||
def _extract_chat_completion_code_execution_tool_calls(self, response: object) -> list[CodeExecutionToolCall]:
|
||||
model_response: Final = self._to_model_response(response)
|
||||
if model_response is None:
|
||||
return []
|
||||
|
|
@ -743,44 +848,46 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
|
|||
|
||||
@staticmethod
|
||||
def _build_chat_completion_assistant_message(
|
||||
tool_calls: list[CodeExecutionToolCall],
|
||||
tool_calls: Sequence[CodeExecutionToolCall],
|
||||
) -> ChatCompletionAssistantMessage:
|
||||
assistant_tool_calls: Final[list[ChatCompletionAssistantToolCall]] = [
|
||||
{
|
||||
"id": tool_call.get("id"),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
|
||||
"arguments": tool_call.get("arguments", ""),
|
||||
},
|
||||
}
|
||||
for tool_call in tool_calls
|
||||
]
|
||||
return {
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
cast(
|
||||
ChatCompletionAssistantToolCall,
|
||||
{
|
||||
"id": tool_call.get("id"),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
|
||||
"arguments": tool_call.get("arguments", ""),
|
||||
},
|
||||
},
|
||||
)
|
||||
for tool_call in tool_calls
|
||||
],
|
||||
"tool_calls": assistant_tool_calls,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _to_model_response(
|
||||
response: ModelResponse | dict[str, Any],
|
||||
) -> ModelResponse | None:
|
||||
def _to_model_response(response: object) -> ModelResponse | None:
|
||||
if isinstance(response, ModelResponse):
|
||||
return response
|
||||
if not isinstance(response, dict):
|
||||
return None
|
||||
response_fields: Final[dict[str, object]] = response
|
||||
try:
|
||||
return ModelResponse(**response)
|
||||
return ModelResponse(**response_fields)
|
||||
except (TypeError, ValidationError):
|
||||
return None
|
||||
|
||||
def _is_code_execution_call(self, item: Any) -> bool:
|
||||
def _is_code_execution_call(self, item: object) -> bool:
|
||||
if isinstance(item, dict):
|
||||
return item.get("type") == "function_call" and item.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME
|
||||
return (
|
||||
getattr(item, "type", None) == "function_call"
|
||||
and getattr(item, "name", None) == LITELLM_CODE_EXECUTION_TOOL_NAME
|
||||
)
|
||||
item_mapping: Final[dict[str, object]] = item
|
||||
return (
|
||||
item_mapping.get("type") == "function_call"
|
||||
and item_mapping.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME
|
||||
)
|
||||
item_type: Final[object] = getattr(item, "type", None)
|
||||
item_name: Final[object] = getattr(item, "name", None)
|
||||
return item_type == "function_call" and item_name == LITELLM_CODE_EXECUTION_TOOL_NAME
|
||||
|
||||
async def _prune_expired_cache(self) -> None:
|
||||
now: Final = time.time()
|
||||
|
|
|
|||
|
|
@ -714,6 +714,29 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
return result
|
||||
|
||||
def supports_scan_only_tool_results(self) -> bool:
|
||||
"""Whether this guardrail can scan tool-result content.
|
||||
|
||||
Guardrails whose own role filtering only ever scans human-authored
|
||||
messages override this to return False, so configuring them with
|
||||
``scan_only_tool_results`` is rejected at initialization instead of
|
||||
silently scanning nothing on every request.
|
||||
"""
|
||||
return True
|
||||
|
||||
def structured_messages_cover_full_request(self) -> bool:
|
||||
"""Whether returned ``structured_messages`` span the whole request.
|
||||
|
||||
Translation handlers hand guardrails only the in-scope subset of the
|
||||
conversation and merge a returned ``structured_messages`` list back
|
||||
into the full request. A guardrail that already rebuilds the complete
|
||||
conversation itself (like CrowdStrike AIDR with its skip filters
|
||||
active) overrides this to return True so the handler installs the
|
||||
returned list as-is instead of merging it a second time, which would
|
||||
duplicate the out-of-scope messages.
|
||||
"""
|
||||
return False
|
||||
|
||||
def should_run_guardrail(
|
||||
self,
|
||||
data,
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
import re
|
||||
import traceback
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
|
@ -39,7 +39,7 @@ if TYPE_CHECKING:
|
|||
)
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
Span = Any
|
||||
LiteLLMLoggingObj = Any
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ class PromptTemplate:
|
|||
self.output_format = self.metadata.get("output", {}).get("format")
|
||||
self.output_schema = self.metadata.get("output", {}).get("schema", {})
|
||||
self.optional_params = {}
|
||||
for key in self.metadata.keys():
|
||||
for key in self.metadata:
|
||||
if key not in restricted_keys:
|
||||
self.optional_params[key] = self.metadata[key]
|
||||
|
||||
|
|
|
|||
|
|
@ -4,8 +4,9 @@ import json
|
|||
import os
|
||||
import re
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Final, cast
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timezone, tzinfo
|
||||
from typing import Any, Final, TypedDict, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, Field
|
||||
|
|
@ -34,6 +35,17 @@ GALILEO_CLOUD_API_BASE_URL: Final = "https://api.galileo.ai"
|
|||
GALILEO_MAX_IN_MEMORY_RECORDS: Final = 1000
|
||||
|
||||
|
||||
class GalileoStandardLoggingFields(TypedDict, total=False):
|
||||
call_type: str
|
||||
model: str
|
||||
prompt_tokens: int
|
||||
completion_tokens: int
|
||||
total_tokens: int
|
||||
response_cost: float
|
||||
startTime: float
|
||||
endTime: float
|
||||
|
||||
|
||||
class LLMResponse(BaseModel):
|
||||
latency_ms: int
|
||||
status_code: int
|
||||
|
|
@ -59,7 +71,7 @@ class LLMResponse(BaseModel):
|
|||
|
||||
class GalileoObserve(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
self.in_memory_records: list[dict] = []
|
||||
self.in_memory_records: list[Mapping[str, object]] = []
|
||||
self.batch_size = 1
|
||||
self.api_key = os.getenv("GALILEO_API_KEY")
|
||||
self.project_id = os.getenv("GALILEO_PROJECT_ID")
|
||||
|
|
@ -176,7 +188,7 @@ class GalileoObserve(CustomLogger):
|
|||
return False
|
||||
|
||||
@staticmethod
|
||||
def _galileo_input_messages(messages: Any | None, input_text: str) -> list[dict[str, str]]:
|
||||
def _galileo_input_messages(messages: object, input_text: str) -> list[dict[str, str]]:
|
||||
if isinstance(messages, dict):
|
||||
messages = messages.get("messages")
|
||||
if not messages:
|
||||
|
|
@ -203,11 +215,11 @@ class GalileoObserve(CustomLogger):
|
|||
return [{"role": "user", "content": input_text}]
|
||||
|
||||
@staticmethod
|
||||
def _local_timezone():
|
||||
def _local_timezone() -> tzinfo:
|
||||
return datetime.now().astimezone().tzinfo or timezone.utc
|
||||
|
||||
@staticmethod
|
||||
def _format_created_at(dt: datetime | Any) -> str:
|
||||
def _format_created_at(dt: object) -> str:
|
||||
"""Serialize timestamps as UTC ISO-8601 for Galileo."""
|
||||
if not isinstance(dt, datetime):
|
||||
return str(dt)
|
||||
|
|
@ -226,7 +238,7 @@ class GalileoObserve(CustomLogger):
|
|||
return created_at
|
||||
|
||||
@staticmethod
|
||||
def _token_metrics_from_record(record: dict[str, Any]) -> dict[str, Any]:
|
||||
def _token_metrics_from_record(record: Mapping[str, Any]) -> dict[str, Any]:
|
||||
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)
|
||||
|
|
@ -244,7 +256,7 @@ class GalileoObserve(CustomLogger):
|
|||
|
||||
@staticmethod
|
||||
def _record_to_v2_span(
|
||||
record: dict[str, Any],
|
||||
record: Mapping[str, Any],
|
||||
*,
|
||||
trace_id: str,
|
||||
span_id: str,
|
||||
|
|
@ -275,7 +287,7 @@ class GalileoObserve(CustomLogger):
|
|||
return span
|
||||
|
||||
@staticmethod
|
||||
def _record_to_v2_trace(record: dict[str, Any]) -> dict[str, Any]:
|
||||
def _record_to_v2_trace(record: Mapping[str, Any]) -> dict[str, Any]:
|
||||
trace_id: Final = str(uuid.uuid4())
|
||||
span_id: Final = str(uuid.uuid4())
|
||||
created_at: Final = GalileoObserve._normalize_created_at(record.get("created_at", ""))
|
||||
|
|
@ -295,7 +307,7 @@ class GalileoObserve(CustomLogger):
|
|||
"spans": [GalileoObserve._record_to_v2_span(record, trace_id=trace_id, span_id=span_id)],
|
||||
}
|
||||
|
||||
def _build_traces_payload(self, records: list[dict]) -> dict[str, Any]:
|
||||
def _build_traces_payload(self, records: Sequence[Mapping[str, Any]]) -> dict[str, Any]:
|
||||
payload: Final[dict[str, Any]] = {
|
||||
"traces": [self._record_to_v2_trace(record) for record in records],
|
||||
"logging_method": "api_direct",
|
||||
|
|
@ -357,7 +369,7 @@ class GalileoObserve(CustomLogger):
|
|||
@staticmethod
|
||||
def _log_v2_payload_validation(payload: dict[str, Any]) -> None:
|
||||
missing_fields: Final[list[str]] = []
|
||||
traces: Final = payload.get("traces", [])
|
||||
traces: Final[Sequence[object]] = payload.get("traces", [])
|
||||
if not traces:
|
||||
missing_fields.append("traces")
|
||||
|
||||
|
|
@ -385,7 +397,7 @@ class GalileoObserve(CustomLogger):
|
|||
)
|
||||
|
||||
def _log_flush_payload(self, url: str, payload: dict[str, Any]) -> None:
|
||||
traces: Final = payload.get("traces", [])
|
||||
traces: Final[Sequence[object]] = payload.get("traces", [])
|
||||
verbose_logger.debug(
|
||||
"Galileo Logger flush URL: %s trace_count=%s",
|
||||
url,
|
||||
|
|
@ -415,8 +427,8 @@ class GalileoObserve(CustomLogger):
|
|||
pass
|
||||
|
||||
@staticmethod
|
||||
def _build_prompt(kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
optional_params: Final = kwargs.get("optional_params", {}) or {}
|
||||
def _build_prompt(kwargs: Mapping[str, Any]) -> dict[str, Any]:
|
||||
optional_params: Final[Mapping[str, object]] = kwargs.get("optional_params", {}) or {}
|
||||
prompt: Final[dict[str, Any]] = {"messages": kwargs.get("messages")}
|
||||
if optional_params.get("functions") is not None:
|
||||
prompt["functions"] = optional_params["functions"]
|
||||
|
|
@ -425,13 +437,13 @@ class GalileoObserve(CustomLogger):
|
|||
return prompt
|
||||
|
||||
@staticmethod
|
||||
def _serialize_galileo_output(value: Any) -> str:
|
||||
def _serialize_galileo_output(value: object) -> str:
|
||||
if value is None:
|
||||
return ""
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
|
||||
def _json_default(obj: Any) -> Any:
|
||||
def _json_default(obj: Any) -> object:
|
||||
if hasattr(obj, "model_dump"):
|
||||
return obj.model_dump()
|
||||
return str(obj)
|
||||
|
|
@ -439,8 +451,8 @@ class GalileoObserve(CustomLogger):
|
|||
return json.dumps(value, default=_json_default)
|
||||
|
||||
@staticmethod
|
||||
def _prompt_to_input_text(prompt: dict[str, Any]) -> str:
|
||||
messages: Final = prompt.get("messages")
|
||||
def _prompt_to_input_text(prompt: Mapping[str, Any]) -> str:
|
||||
messages: Final[object] = prompt.get("messages")
|
||||
if messages is not None:
|
||||
text: Final = GalileoObserve._input_text_from_messages(messages)
|
||||
if text:
|
||||
|
|
@ -448,7 +460,7 @@ class GalileoObserve(CustomLogger):
|
|||
return json.dumps(prompt, default=str)
|
||||
|
||||
@staticmethod
|
||||
def _get_chat_content_for_galileo(response_obj: litellm.ModelResponse) -> Any:
|
||||
def _get_chat_content_for_galileo(response_obj: litellm.ModelResponse) -> object:
|
||||
if response_obj.choices and len(response_obj.choices) > 0:
|
||||
message: Final = response_obj["choices"][0]["message"]
|
||||
if hasattr(message, "json"):
|
||||
|
|
@ -470,23 +482,23 @@ class GalileoObserve(CustomLogger):
|
|||
@staticmethod
|
||||
def _get_responses_api_content_for_galileo(
|
||||
response_obj: ResponsesAPIResponse,
|
||||
) -> Any:
|
||||
) -> object:
|
||||
if hasattr(response_obj, "output") and response_obj.output:
|
||||
return response_obj.output
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _langfuse_style_rerank_prompt(kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
def _langfuse_style_rerank_prompt(kwargs: Mapping[str, object]) -> dict[str, Any]:
|
||||
"""Match Langfuse rerank input: prompt = {"messages": kwargs.get("messages")}."""
|
||||
return {"messages": kwargs.get("messages")}
|
||||
|
||||
def _get_galileo_input_output_content(
|
||||
self,
|
||||
kwargs: dict[str, Any],
|
||||
response_obj: Any,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
level: str = "DEFAULT",
|
||||
status_message: str | None = None,
|
||||
) -> tuple[str, str, Any]:
|
||||
) -> tuple[str, str, object]:
|
||||
"""
|
||||
Mirror Langfuse _get_langfuse_input_output_content for Galileo ingest.
|
||||
|
||||
|
|
@ -582,12 +594,12 @@ class GalileoObserve(CustomLogger):
|
|||
|
||||
return self._prompt_to_input_text(prompt), "", kwargs.get("messages") or []
|
||||
|
||||
def get_output_str_from_response(self, response_obj: Any, kwargs: dict[str, Any]) -> str:
|
||||
def get_output_str_from_response(self, response_obj: object, kwargs: Mapping[str, object]) -> str:
|
||||
_, output_text, _ = self._get_galileo_input_output_content(kwargs=kwargs, response_obj=response_obj)
|
||||
return output_text
|
||||
|
||||
@staticmethod
|
||||
def _input_text_from_messages(messages: Any) -> str:
|
||||
def _input_text_from_messages(messages: object) -> str:
|
||||
"""Return a plain-string summary of the input suitable for the trace-level input field."""
|
||||
if isinstance(messages, str):
|
||||
return messages
|
||||
|
|
@ -613,7 +625,13 @@ class GalileoObserve(CustomLogger):
|
|||
return str(content)
|
||||
return ""
|
||||
|
||||
async def async_log_success_event(self, kwargs: Any, response_obj: Any, start_time: Any, end_time: Any):
|
||||
async def async_log_success_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: object,
|
||||
end_time: object,
|
||||
) -> None:
|
||||
verbose_logger.debug("On Async Success")
|
||||
try:
|
||||
await self._async_log_success_event_impl(
|
||||
|
|
@ -625,7 +643,13 @@ class GalileoObserve(CustomLogger):
|
|||
except Exception:
|
||||
verbose_logger.exception("Galileo Logger: unexpected error in async_log_success_event")
|
||||
|
||||
async def _async_log_success_event_impl(self, kwargs: Any, response_obj: Any, start_time: Any, end_time: Any):
|
||||
async def _async_log_success_event_impl(
|
||||
self,
|
||||
kwargs: Mapping[str, Any],
|
||||
response_obj: object,
|
||||
start_time: object,
|
||||
end_time: object,
|
||||
) -> None:
|
||||
if not self._is_configured():
|
||||
verbose_logger.debug(
|
||||
"Galileo Logger: skipping — GALILEO_PROJECT_ID=%s GALILEO_API_KEY=%s GALILEO_BASE_URL=%s",
|
||||
|
|
@ -635,7 +659,7 @@ class GalileoObserve(CustomLogger):
|
|||
)
|
||||
return
|
||||
|
||||
slo: Final[dict[str, Any] | None] = kwargs.get("standard_logging_object")
|
||||
slo: Final[GalileoStandardLoggingFields | None] = kwargs.get("standard_logging_object")
|
||||
if slo is None:
|
||||
verbose_logger.debug("Galileo Logger: no standard_logging_object in kwargs, skipping")
|
||||
return
|
||||
|
|
@ -646,8 +670,8 @@ class GalileoObserve(CustomLogger):
|
|||
kwargs=kwargs, response_obj=response_obj
|
||||
)
|
||||
|
||||
raw_start: Final = slo.get("startTime")
|
||||
raw_end: Final = slo.get("endTime")
|
||||
raw_start: Final[float | None] = slo.get("startTime")
|
||||
raw_end: Final[float | None] = slo.get("endTime")
|
||||
if raw_start is None or raw_end is None:
|
||||
verbose_logger.debug(
|
||||
"Galileo Logger: standard_logging_object missing startTime/endTime, "
|
||||
|
|
@ -710,7 +734,7 @@ class GalileoObserve(CustomLogger):
|
|||
if len(self.in_memory_records) >= self.batch_size:
|
||||
await self.flush_in_memory_records()
|
||||
|
||||
async def flush_in_memory_records(self):
|
||||
async def flush_in_memory_records(self) -> None:
|
||||
if not self.in_memory_records:
|
||||
return
|
||||
|
||||
|
|
@ -774,5 +798,11 @@ class GalileoObserve(CustomLogger):
|
|||
if not self.use_v2_api and response.status_code in (401, 403):
|
||||
self.headers = None
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
async def async_log_failure_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: object,
|
||||
end_time: object,
|
||||
) -> None:
|
||||
verbose_logger.debug("On Async Failure")
|
||||
|
|
|
|||
|
|
@ -128,8 +128,8 @@ class LangFuseLogger:
|
|||
self.langfuse_client = create_mock_langfuse_client()
|
||||
self.is_mock_mode = True
|
||||
else:
|
||||
http_client: Final = _get_httpx_client()
|
||||
self.langfuse_client = http_client.client
|
||||
self._http_handler: Final = _get_httpx_client()
|
||||
self.langfuse_client = self._http_handler.client
|
||||
self.is_mock_mode = False
|
||||
|
||||
parameters: Final = {
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import base64
|
|||
import json
|
||||
import os
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.arize import _utils
|
||||
|
|
@ -18,7 +18,7 @@ from litellm.types.utils import StandardCallbackDynamicParams
|
|||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
Span = Any
|
||||
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Call Hook for LiteLLM Proxy which allows Langfuse prompt management.
|
|||
|
||||
import os
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, cast
|
||||
|
||||
from packaging.version import Version
|
||||
|
||||
|
|
@ -30,7 +30,7 @@ if TYPE_CHECKING:
|
|||
|
||||
LangfuseClass: TypeAlias = Langfuse
|
||||
|
||||
PROMPT_CLIENT = Union[TextPromptClient, ChatPromptClient]
|
||||
PROMPT_CLIENT = TextPromptClient | ChatPromptClient
|
||||
else:
|
||||
PROMPT_CLIENT = Any
|
||||
LangfuseClass = Any
|
||||
|
|
|
|||
|
|
@ -1,12 +1,12 @@
|
|||
import json
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm.proxy._types import SpanAttributes
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
Span = Any
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import os
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry
|
||||
|
||||
|
|
@ -13,7 +13,7 @@ if TYPE_CHECKING:
|
|||
|
||||
Protocol = _Protocol
|
||||
OpenTelemetryConfig = _OpenTelemetryConfig
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
Protocol = Any
|
||||
OpenTelemetryConfig = Any
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import os
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -47,12 +47,12 @@ if TYPE_CHECKING:
|
|||
)
|
||||
from litellm.proxy.proxy_server import UserAPIKeyAuth as _UserAPIKeyAuth
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Tracer = Union[_Tracer, Any]
|
||||
Context = Union[_Context, Any]
|
||||
SpanExporter = Union[_SpanExporter, Any]
|
||||
UserAPIKeyAuth = Union[_UserAPIKeyAuth, Any]
|
||||
ManagementEndpointLoggingPayload = Union[_ManagementEndpointLoggingPayload, Any]
|
||||
Span = _Span | Any
|
||||
Tracer = _Tracer | Any
|
||||
Context = _Context | Any
|
||||
SpanExporter = _SpanExporter | Any
|
||||
UserAPIKeyAuth = _UserAPIKeyAuth | Any
|
||||
ManagementEndpointLoggingPayload = _ManagementEndpointLoggingPayload | Any
|
||||
else:
|
||||
Span = Any
|
||||
Tracer = Any
|
||||
|
|
@ -186,16 +186,7 @@ def _normalize_team_metadata_keys(value: Any) -> list[str]:
|
|||
|
||||
_FREEZE_MAX_DEPTH: Final = 16
|
||||
|
||||
HashableScope = Union[
|
||||
str,
|
||||
int,
|
||||
float,
|
||||
bool,
|
||||
bytes,
|
||||
None,
|
||||
tuple["HashableScope", ...],
|
||||
frozenset["HashableScope"],
|
||||
]
|
||||
HashableScope = str | int | float | bool | bytes | None | tuple["HashableScope", ...] | frozenset["HashableScope"]
|
||||
|
||||
|
||||
def _freeze_for_dedupe(value: object, _depth: int = 0) -> HashableScope:
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ Events:
|
|||
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
|
|
@ -40,7 +40,7 @@ if TYPE_CHECKING:
|
|||
|
||||
from litellm.integrations.opentelemetry import OpenTelemetryConfig
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
Span = Any
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
"""Type definitions for Opik payload building."""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Final, Literal, Union
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
@ -42,5 +42,5 @@ class SpanPayload:
|
|||
total_cost: float | None = None
|
||||
|
||||
|
||||
PayloadItem = Union[TracePayload, SpanPayload]
|
||||
PayloadItem = TracePayload | SpanPayload
|
||||
TraceSpanPayloadTuple: Final = tuple[TracePayload | None, SpanPayload]
|
||||
|
|
|
|||
|
|
@ -72,53 +72,49 @@ from litellm.integrations.otel.model.spans import (
|
|||
)
|
||||
|
||||
__all__ = [
|
||||
# config
|
||||
"OTEL_V2_ENV",
|
||||
"OpenTelemetryV2Config",
|
||||
"is_otel_v2_enabled",
|
||||
# semconv
|
||||
"BAGGAGE_PROMOTED_KEYS",
|
||||
"DB",
|
||||
"DEFAULT_BAGGAGE_METADATA_KEYS",
|
||||
"HTTP",
|
||||
"MCP",
|
||||
"OTEL_V2_ENV",
|
||||
"SPAN_REGISTRY",
|
||||
"Client",
|
||||
"Error",
|
||||
"GenAI",
|
||||
"GenAIOperation",
|
||||
"GenAIProvider",
|
||||
"HTTP",
|
||||
"JsonRpc",
|
||||
"LiteLLM",
|
||||
"LiteLLMError",
|
||||
"MCP",
|
||||
"MCPMethod",
|
||||
"Metric",
|
||||
"Network",
|
||||
"NetworkTransport",
|
||||
"Server",
|
||||
"resolve_operation",
|
||||
"resolve_provider",
|
||||
# spans
|
||||
"SPAN_REGISTRY",
|
||||
"LiteLLMSpanKind",
|
||||
"SpanRole",
|
||||
"SpanSpec",
|
||||
"db_system",
|
||||
"span_role_for_service",
|
||||
"validate_registry",
|
||||
# payloads
|
||||
"GuardrailSpanData",
|
||||
"JsonRpc",
|
||||
"LLMCallSpanData",
|
||||
"LLMRequestParams",
|
||||
"LLMUsage",
|
||||
"LiteLLM",
|
||||
"LiteLLMError",
|
||||
"LiteLLMSpanKind",
|
||||
"MCPListToolsSpanData",
|
||||
"MCPMethod",
|
||||
"MCPToolCallSpanData",
|
||||
"Metric",
|
||||
"Network",
|
||||
"NetworkTransport",
|
||||
"OpenTelemetryV2Config",
|
||||
"ProxyRequestSpanData",
|
||||
"RequestContext",
|
||||
"RequestIdentity",
|
||||
"Server",
|
||||
"ServerInfo",
|
||||
"ServiceSpanData",
|
||||
"SpanError",
|
||||
"SpanRole",
|
||||
"SpanSpec",
|
||||
"db_system",
|
||||
"is_mcp_list_tools",
|
||||
"is_mcp_tool_call",
|
||||
"is_otel_v2_enabled",
|
||||
"promoted_baggage",
|
||||
"resolve_operation",
|
||||
"resolve_provider",
|
||||
"span_role_for_service",
|
||||
"validate_registry",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -66,18 +66,13 @@ from litellm.interactions.main import (
|
|||
)
|
||||
|
||||
__all__ = [
|
||||
# Create
|
||||
"create",
|
||||
"acreate",
|
||||
# Get
|
||||
"get",
|
||||
"aget",
|
||||
# Delete
|
||||
"delete",
|
||||
"adelete",
|
||||
# Cancel
|
||||
"cancel",
|
||||
"acancel",
|
||||
# Sub-modules
|
||||
"acreate",
|
||||
"adelete",
|
||||
"agents",
|
||||
"aget",
|
||||
"cancel",
|
||||
"create",
|
||||
"delete",
|
||||
"get",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
## Helper utilities
|
||||
import copy
|
||||
from collections.abc import Iterable
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Union
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -14,7 +14,7 @@ if TYPE_CHECKING:
|
|||
|
||||
from litellm.types.utils import ModelResponseStream
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
Span = Any
|
||||
|
||||
|
|
@ -131,6 +131,9 @@ _FINISH_REASON_MAP: Final[dict[str, OpenAIChatCompletionFinishReason]] = {
|
|||
"content_filter": "content_filter",
|
||||
# Anthropic Sonnet 4
|
||||
"content_filtered": "content_filter",
|
||||
# Generic error passthrough (OpenRouter and other OpenAI-compatible providers
|
||||
# emit lowercase "error" when a provider fails mid-stream)
|
||||
"error": "stop",
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -48,7 +48,7 @@ O(number of rules); callers must only invoke them on a cache miss.
|
|||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Union
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
|
|
@ -100,7 +100,7 @@ class _CapabilityRule:
|
|||
model_info: dict
|
||||
|
||||
|
||||
_CompiledRule = Union[_RoutingRule, _CapabilityRule]
|
||||
_CompiledRule = _RoutingRule | _CapabilityRule
|
||||
|
||||
|
||||
def _compile_rule(rule: object) -> tuple[_CompiledRule, ...]:
|
||||
|
|
|
|||
|
|
@ -115,8 +115,11 @@ def get_litellm_params(
|
|||
litellm_request_debug: bool | None = None,
|
||||
**kwargs,
|
||||
) -> dict:
|
||||
_litellm_metadata_dict: Final = litellm_metadata if isinstance(litellm_metadata, dict) else None
|
||||
resolved_metadata: Final = _litellm_metadata_dict.copy() if not metadata and _litellm_metadata_dict else metadata
|
||||
|
||||
# Derive litellm_session_id / litellm_trace_id from metadata when not provided (call chaining)
|
||||
_meta: Final = metadata or {}
|
||||
_meta: Final = resolved_metadata or {}
|
||||
if litellm_session_id is None:
|
||||
litellm_session_id = _meta.get("session_id") or _meta.get("trace_id")
|
||||
if litellm_trace_id is None:
|
||||
|
|
@ -139,7 +142,7 @@ def get_litellm_params(
|
|||
"model_alias_map": model_alias_map,
|
||||
"completion_call_id": completion_call_id,
|
||||
"aembedding": aembedding,
|
||||
"metadata": metadata,
|
||||
"metadata": resolved_metadata,
|
||||
"model_info": model_info,
|
||||
"proxy_server_request": proxy_server_request,
|
||||
"preset_cache_key": preset_cache_key,
|
||||
|
|
|
|||
|
|
@ -585,8 +585,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
"""
|
||||
base_litellm_params: Final[dict[str, Any]] = {}
|
||||
|
||||
if "metadata" in kwargs:
|
||||
base_litellm_params["metadata"] = kwargs["metadata"]
|
||||
if isinstance(kwargs.get("metadata"), dict):
|
||||
base_litellm_params["metadata"] = kwargs["metadata"].copy()
|
||||
if "litellm_metadata" in kwargs and isinstance(kwargs["litellm_metadata"], dict):
|
||||
base_litellm_params["litellm_metadata"] = kwargs["litellm_metadata"]
|
||||
if "metadata" not in base_litellm_params:
|
||||
|
|
@ -1399,10 +1399,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
litellm_params=(self.litellm_params if hasattr(self, "litellm_params") else None)
|
||||
)
|
||||
|
||||
prompt = "" # use for tts cost calc
|
||||
_input: Final = self.model_call_details.get("input", None)
|
||||
if _input is not None and isinstance(_input, str):
|
||||
prompt = _input
|
||||
prompt = self._prompt_for_cost_calculation()
|
||||
|
||||
if cache_hit is None:
|
||||
cache_hit = self.model_call_details.get("cache_hit", False)
|
||||
|
|
@ -1461,6 +1458,19 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
|
||||
return None
|
||||
|
||||
def _prompt_for_cost_calculation(self) -> str:
|
||||
"""
|
||||
The raw input string is only priced directly for text-to-speech, which bills per character.
|
||||
Every other call type gets its billable units from the response usage object, and call types
|
||||
that carry no usage at all (file content retrieval, and anything else `function_setup` cannot
|
||||
build messages for) only have the ``"default-message-value"`` placeholder here, so passing the
|
||||
input along would token-price that placeholder.
|
||||
"""
|
||||
if self.call_type not in (CallTypes.speech.value, CallTypes.aspeech.value):
|
||||
return ""
|
||||
_input = self.model_call_details.get("input", None)
|
||||
return _input if isinstance(_input, str) else ""
|
||||
|
||||
def _generate_content_result_as_model_response(self, result: object) -> ModelResponse | None:
|
||||
"""
|
||||
Native Google :generateContent bodies report token usage under
|
||||
|
|
@ -4827,7 +4837,7 @@ class StandardLoggingPayloadSetup:
|
|||
|
||||
# Populate well-known typed fields with int/str coercion where needed
|
||||
typed_keys: Final[dict] = {}
|
||||
for key in StandardLoggingAdditionalHeaders.__annotations__.keys():
|
||||
for key in StandardLoggingAdditionalHeaders.__annotations__:
|
||||
_key = key.lower().replace("_", "-")
|
||||
typed_keys[_key] = key
|
||||
if _key in additiona_headers:
|
||||
|
|
@ -4859,7 +4869,7 @@ class StandardLoggingPayloadSetup:
|
|||
usage_object=None,
|
||||
)
|
||||
if hidden_params is not None:
|
||||
for key in StandardLoggingHiddenParams.__annotations__.keys():
|
||||
for key in StandardLoggingHiddenParams.__annotations__:
|
||||
if key in hidden_params:
|
||||
if key == "additional_headers":
|
||||
clean_hidden_params["additional_headers"] = StandardLoggingPayloadSetup.get_additional_headers(
|
||||
|
|
@ -5501,7 +5511,7 @@ def get_standard_logging_metadata(
|
|||
)
|
||||
if isinstance(metadata, dict):
|
||||
# Update the clean_metadata with values from input metadata that match StandardLoggingMetadata fields
|
||||
for key in StandardLoggingMetadata.__annotations__.keys():
|
||||
for key in StandardLoggingMetadata.__annotations__:
|
||||
if key in metadata:
|
||||
clean_metadata[key] = metadata[key]
|
||||
|
||||
|
|
|
|||
|
|
@ -681,6 +681,23 @@ def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str |
|
|||
return 1.0
|
||||
|
||||
|
||||
def _resolve_reasoning_token_cost(
|
||||
model_info: ModelInfo,
|
||||
service_tier: str | None,
|
||||
completion_base_cost: float,
|
||||
) -> float:
|
||||
tier_reasoning_key: Final = _get_service_tier_cost_key("output_cost_per_reasoning_token", service_tier)
|
||||
if model_info.get(tier_reasoning_key) is not None:
|
||||
tier_reasoning_cost: Final = _get_cost_per_unit(model_info, tier_reasoning_key, None)
|
||||
if tier_reasoning_cost is not None:
|
||||
return tier_reasoning_cost
|
||||
tier_output_key: Final = _get_service_tier_cost_key("output_cost_per_token", service_tier)
|
||||
if tier_output_key != "output_cost_per_token" and model_info.get(tier_output_key) is not None:
|
||||
return completion_base_cost
|
||||
standard_reasoning_cost: Final = _get_cost_per_unit(model_info, "output_cost_per_reasoning_token", None)
|
||||
return standard_reasoning_cost if standard_reasoning_cost is not None else completion_base_cost
|
||||
|
||||
|
||||
def generic_cost_per_token(
|
||||
model: str,
|
||||
usage: Usage,
|
||||
|
|
@ -817,9 +834,10 @@ def generic_cost_per_token(
|
|||
|
||||
## REASONING COST
|
||||
if not is_text_tokens_total and reasoning_tokens and reasoning_tokens > 0:
|
||||
_output_cost_per_reasoning_token = _get_cost_per_unit(model_info, "output_cost_per_reasoning_token", None)
|
||||
_output_cost_per_reasoning_token = (
|
||||
_output_cost_per_reasoning_token if _output_cost_per_reasoning_token is not None else completion_base_cost
|
||||
_output_cost_per_reasoning_token = _resolve_reasoning_token_cost(
|
||||
model_info=model_info,
|
||||
service_tier=service_tier,
|
||||
completion_base_cost=completion_base_cost,
|
||||
)
|
||||
completion_cost += float(reasoning_tokens) * _output_cost_per_reasoning_token
|
||||
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import inspect
|
|||
import re
|
||||
import time
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MAX_BASE64_LENGTH_FOR_LOGGING
|
||||
|
|
@ -23,7 +23,7 @@ if TYPE_CHECKING:
|
|||
)
|
||||
|
||||
LiteLLMModelResponse = _ModelResponse
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
else:
|
||||
LiteLLMModelResponse = Any
|
||||
LiteLLMLoggingObject = Any
|
||||
|
|
|
|||
|
|
@ -47,7 +47,7 @@ def is_model_response_stream_empty(model_response: ModelResponseStream) -> bool:
|
|||
|
||||
# Check for any non-base fields that are set
|
||||
# Access model_fields on the class, not the instance, to avoid Pydantic 2.11+ deprecation warnings
|
||||
for model_response_field in type(model_response).model_fields.keys():
|
||||
for model_response_field in type(model_response).model_fields:
|
||||
# Skip base fields that are always set
|
||||
if model_response_field in BASE_FIELDS:
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ import time
|
|||
import traceback
|
||||
from collections.abc import AsyncIterator, Callable, Iterator
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Final, NoReturn, TypeVar, Union, cast
|
||||
from typing import Any, Final, NoReturn, TypeVar, cast
|
||||
|
||||
import anyio
|
||||
import httpx
|
||||
|
|
@ -99,7 +99,7 @@ class _ProviderChunkEarlyReturn:
|
|||
value: Any
|
||||
|
||||
|
||||
_ProviderChunkResult = Union[_ProviderChunkParsed, _ProviderChunkEarlyReturn]
|
||||
_ProviderChunkResult = _ProviderChunkParsed | _ProviderChunkEarlyReturn
|
||||
|
||||
|
||||
class CustomStreamWrapper:
|
||||
|
|
@ -256,9 +256,7 @@ class CustomStreamWrapper:
|
|||
chunk = chunk.strip()
|
||||
self.complete_response = self.complete_response.strip()
|
||||
|
||||
if chunk.startswith(self.complete_response):
|
||||
# Remove last_sent_chunk only if it appears at the start of the new chunk
|
||||
chunk = chunk[len(self.complete_response) :]
|
||||
chunk = chunk.removeprefix(self.complete_response)
|
||||
|
||||
self.complete_response += chunk
|
||||
return chunk
|
||||
|
|
|
|||
|
|
@ -13,8 +13,12 @@ Pattern Overview:
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
|
|
@ -22,10 +26,13 @@ from litellm.llms.anthropic.experimental_pass_through.adapters.transformation im
|
|||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
anthropic_tool_name,
|
||||
effective_scan_only_tool_results_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
openai_messages_without_system,
|
||||
openai_messages_without_tool,
|
||||
merge_guardrailed_scoped_messages,
|
||||
merge_returned_tools_into_request_tools,
|
||||
scoped_structured_message_indices,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
|
|
@ -58,6 +65,50 @@ if TYPE_CHECKING:
|
|||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MessageContentTarget:
|
||||
msg_idx: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ContentBlockTextTarget:
|
||||
msg_idx: int
|
||||
content_idx: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ToolResultStringTarget:
|
||||
msg_idx: int
|
||||
content_idx: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ToolResultBlockTextTarget:
|
||||
msg_idx: int
|
||||
content_idx: int
|
||||
block_idx: int
|
||||
|
||||
|
||||
InputWriteBackTarget = (
|
||||
MessageContentTarget | ContentBlockTextTarget | ToolResultStringTarget | ToolResultBlockTextTarget
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ScannedText:
|
||||
text: str
|
||||
target: InputWriteBackTarget
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ExtractedInput:
|
||||
scanned: tuple[ScannedText, ...]
|
||||
images: tuple[str, ...]
|
||||
|
||||
|
||||
EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
|
||||
|
||||
|
||||
class AnthropicMessagesHandler(BaseTranslation):
|
||||
"""
|
||||
Handler for processing Anthropic messages with guardrails.
|
||||
|
|
@ -278,34 +329,42 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
skip_system: Final = effective_skip_system_message_for_guardrail(guardrail_to_apply)
|
||||
skip_tool: Final = effective_skip_tool_message_for_guardrail(guardrail_to_apply)
|
||||
scan_only_tool_results: Final = effective_scan_only_tool_results_for_guardrail(guardrail_to_apply)
|
||||
|
||||
chat_completion_compatible_request: Final = self._translate_to_openai(data)
|
||||
|
||||
structured_messages = cast(
|
||||
full_structured_messages: Final = cast(
|
||||
list[AllMessageValues],
|
||||
chat_completion_compatible_request.get("messages", []),
|
||||
)
|
||||
if skip_system:
|
||||
structured_messages = openai_messages_without_system(structured_messages)
|
||||
if skip_tool:
|
||||
structured_messages = openai_messages_without_tool(structured_messages)
|
||||
scoped_message_indices: Final = scoped_structured_message_indices(
|
||||
full_structured_messages,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
skip_system=skip_system,
|
||||
skip_tool=skip_tool,
|
||||
)
|
||||
structured_messages: Final = [full_structured_messages[index] for index in scoped_message_indices]
|
||||
|
||||
texts_to_check: Final[list[str]] = []
|
||||
images_to_check: Final[list[str]] = []
|
||||
tools_to_check: Final[list[ChatCompletionToolParam]] = chat_completion_compatible_request.get("tools", [])
|
||||
task_mappings: Final[list[tuple[int, int | None]]] = []
|
||||
tools_to_check: Final[list[ChatCompletionToolParam]] = (
|
||||
[] if scan_only_tool_results else chat_completion_compatible_request.get("tools", [])
|
||||
)
|
||||
|
||||
# Step 1: Extract all text content and images
|
||||
for msg_idx, message in enumerate(messages):
|
||||
extracted: Final = tuple(
|
||||
self._extract_input_text_and_images(
|
||||
message=message,
|
||||
msg_idx=msg_idx,
|
||||
texts_to_check=texts_to_check,
|
||||
images_to_check=images_to_check,
|
||||
task_mappings=task_mappings,
|
||||
skip_system_message=skip_system,
|
||||
skip_tool_message=skip_tool,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
)
|
||||
for msg_idx, message in enumerate(messages)
|
||||
)
|
||||
scanned: Final = tuple(item for one_message in extracted for item in one_message.scanned)
|
||||
texts_to_check: Final = [item.text for item in scanned] # mutable-ok: GenericGuardrailAPIInputs takes list[str]
|
||||
images_to_check: Final = [
|
||||
image for one_message in extracted for image in one_message.images
|
||||
] # mutable-ok: GenericGuardrailAPIInputs takes list[str]
|
||||
|
||||
# Step 2: Apply guardrail to all texts in batch
|
||||
if texts_to_check:
|
||||
|
|
@ -339,20 +398,37 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
if converted_tool is not None:
|
||||
anthropic_tools.append(converted_tool)
|
||||
# Note: MCP servers are handled separately in the main transformation
|
||||
data["tools"] = anthropic_tools
|
||||
data["tools"] = (
|
||||
merge_returned_tools_into_request_tools(
|
||||
request_tools=data.get("tools"),
|
||||
returned_tools=anthropic_tools,
|
||||
tool_name=anthropic_tool_name,
|
||||
)
|
||||
if scan_only_tool_results
|
||||
else anthropic_tools
|
||||
)
|
||||
|
||||
guardrailed_structured_messages: Final = guardrailed_inputs.get("structured_messages")
|
||||
if (
|
||||
guardrailed_structured_messages is not None
|
||||
and guardrailed_structured_messages is not original_structured_messages
|
||||
):
|
||||
self._write_back_structured_messages(data, guardrailed_structured_messages)
|
||||
self._write_back_structured_messages(
|
||||
data,
|
||||
guardrailed_structured_messages
|
||||
if guardrail_to_apply.structured_messages_cover_full_request()
|
||||
else merge_guardrailed_scoped_messages(
|
||||
full_messages=full_structured_messages,
|
||||
scoped_indices=scoped_message_indices,
|
||||
guardrailed_scoped=guardrailed_structured_messages,
|
||||
),
|
||||
)
|
||||
else:
|
||||
# Step 3: Map guardrail responses back to original message structure
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=messages,
|
||||
responses=guardrailed_texts,
|
||||
task_mappings=task_mappings,
|
||||
scanned=scanned,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug("Anthropic Messages: Processed input messages: %s", messages)
|
||||
|
|
@ -405,99 +481,150 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
names.append(str(tool["name"]))
|
||||
return names
|
||||
|
||||
@classmethod
|
||||
def _extract_input_text_and_images(
|
||||
self,
|
||||
cls,
|
||||
message: dict[str, Any],
|
||||
msg_idx: int,
|
||||
texts_to_check: list[str],
|
||||
images_to_check: list[str],
|
||||
task_mappings: list[tuple[int, int | None]],
|
||||
skip_system_message: bool = False,
|
||||
skip_tool_message: bool = False,
|
||||
) -> None:
|
||||
scan_only_tool_results: bool = False,
|
||||
) -> ExtractedInput:
|
||||
"""
|
||||
Extract text content and images from a message.
|
||||
|
||||
Override this method to customize text/image extraction logic.
|
||||
"""
|
||||
role: Final = str(message.get("role") or "").lower()
|
||||
if skip_system_message and role == "system":
|
||||
return
|
||||
if skip_tool_message and role == "tool":
|
||||
return
|
||||
if (skip_system_message and role == "system") or (skip_tool_message and role == "tool"):
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
|
||||
content: Final = message.get("content", None)
|
||||
tools: Final = message.get("tools", None)
|
||||
if content is None and tools is None:
|
||||
return
|
||||
if isinstance(content, str):
|
||||
if scan_only_tool_results:
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
return ExtractedInput(scanned=(ScannedText(content, MessageContentTarget(msg_idx)),), images=())
|
||||
if not isinstance(content, list):
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
|
||||
## CHECK FOR TEXT + IMAGES
|
||||
if content is not None and isinstance(content, str):
|
||||
# Simple string content
|
||||
texts_to_check.append(content)
|
||||
task_mappings.append((msg_idx, None))
|
||||
|
||||
elif content is not None and isinstance(content, list):
|
||||
# List content (e.g., multimodal with text and images)
|
||||
for content_idx, content_item in enumerate(content):
|
||||
# Extract text
|
||||
text_str = content_item.get("text", None)
|
||||
if text_str is not None:
|
||||
texts_to_check.append(text_str)
|
||||
task_mappings.append((msg_idx, int(content_idx)))
|
||||
|
||||
# Extract images
|
||||
if content_item.get("type") == "image":
|
||||
source = content_item.get("source", {})
|
||||
if isinstance(source, dict):
|
||||
# Could be base64 or url
|
||||
data = source.get("data")
|
||||
if data:
|
||||
images_to_check.append(data)
|
||||
|
||||
def _extract_input_tools(
|
||||
self,
|
||||
tools: list[dict[str, Any]],
|
||||
tools_to_check: list[ChatCompletionToolParam],
|
||||
) -> None:
|
||||
"""
|
||||
Extract tools from a message.
|
||||
"""
|
||||
## CHECK FOR TOOLS
|
||||
if tools is not None and isinstance(tools, list):
|
||||
# TRANSFORM ANTHROPIC TOOLS TO OPENAI TOOLS
|
||||
openai_tools: Final = self.adapter.translate_anthropic_tools_to_openai(
|
||||
tools=cast(list[AllAnthropicToolsValues], tools)
|
||||
blocks: Final = tuple(
|
||||
cls._extract_content_block(
|
||||
content_item=content_item,
|
||||
msg_idx=msg_idx,
|
||||
content_idx=content_idx,
|
||||
skip_tool_message=skip_tool_message,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
)
|
||||
tools_to_check.extend(openai_tools)
|
||||
for content_idx, content_item in enumerate(content)
|
||||
if isinstance(content_item, dict)
|
||||
)
|
||||
return ExtractedInput(
|
||||
scanned=tuple(item for block in blocks for item in block.scanned),
|
||||
images=tuple(image for block in blocks for image in block.images),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _extract_content_block(
|
||||
cls,
|
||||
content_item: Mapping[str, Any],
|
||||
msg_idx: int,
|
||||
content_idx: int,
|
||||
skip_tool_message: bool,
|
||||
scan_only_tool_results: bool = False,
|
||||
) -> ExtractedInput:
|
||||
if content_item.get("type") == "tool_result":
|
||||
if skip_tool_message:
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
return cls._extract_tool_result(content_item=content_item, msg_idx=msg_idx, content_idx=content_idx)
|
||||
|
||||
if scan_only_tool_results:
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
|
||||
text_str: Final = content_item.get("text", None)
|
||||
return ExtractedInput(
|
||||
scanned=(
|
||||
() if text_str is None else (ScannedText(text_str, ContentBlockTextTarget(msg_idx, content_idx)),)
|
||||
),
|
||||
images=cls._image_sources(content_item) if content_item.get("type") == "image" else (),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _extract_tool_result(
|
||||
cls,
|
||||
content_item: Mapping[str, Any],
|
||||
msg_idx: int,
|
||||
content_idx: int,
|
||||
) -> ExtractedInput:
|
||||
tool_result_content: Final = content_item.get("content")
|
||||
|
||||
if isinstance(tool_result_content, str):
|
||||
return ExtractedInput(
|
||||
scanned=(ScannedText(tool_result_content, ToolResultStringTarget(msg_idx, content_idx)),),
|
||||
images=(),
|
||||
)
|
||||
if not isinstance(tool_result_content, list):
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
|
||||
blocks: Final = tuple(
|
||||
(block_idx, block) for block_idx, block in enumerate(tool_result_content) if isinstance(block, dict)
|
||||
)
|
||||
return ExtractedInput(
|
||||
scanned=tuple(
|
||||
ScannedText(block["text"], ToolResultBlockTextTarget(msg_idx, content_idx, block_idx))
|
||||
for block_idx, block in blocks
|
||||
if isinstance(block.get("text"), str)
|
||||
),
|
||||
images=tuple(
|
||||
image for _, block in blocks if block.get("type") == "image" for image in cls._image_sources(block)
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _image_sources(block: Mapping[str, Any]) -> tuple[str, ...]:
|
||||
source: Final = block.get("source")
|
||||
if not isinstance(source, Mapping):
|
||||
return ()
|
||||
# Could be base64 or url
|
||||
data: Final = source.get("data")
|
||||
return (data,) if data else ()
|
||||
|
||||
async def _apply_guardrail_responses_to_input(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
responses: list[str],
|
||||
task_mappings: list[tuple[int, int | None]],
|
||||
scanned: tuple[ScannedText, ...],
|
||||
) -> None:
|
||||
"""
|
||||
Apply guardrail responses back to input messages.
|
||||
|
||||
Override this method to customize how responses are applied.
|
||||
"""
|
||||
for task_idx, guardrail_response in enumerate(responses):
|
||||
mapping = task_mappings[task_idx]
|
||||
msg_idx = cast(int, mapping[0])
|
||||
content_idx_optional = cast(int | None, mapping[1])
|
||||
|
||||
content = messages[msg_idx].get("content", None)
|
||||
for item, guardrail_response in zip(scanned, responses):
|
||||
target = item.target
|
||||
message = messages[target.msg_idx]
|
||||
content = message.get("content", None)
|
||||
if content is None:
|
||||
continue
|
||||
|
||||
if isinstance(content, str) and content_idx_optional is None:
|
||||
# Replace string content with guardrail response
|
||||
messages[msg_idx]["content"] = guardrail_response
|
||||
|
||||
elif isinstance(content, list) and content_idx_optional is not None:
|
||||
# Replace specific text item in list content
|
||||
messages[msg_idx]["content"][content_idx_optional]["text"] = guardrail_response
|
||||
match target:
|
||||
case MessageContentTarget():
|
||||
if isinstance(content, str):
|
||||
message["content"] = (
|
||||
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
case ContentBlockTextTarget(content_idx=content_idx):
|
||||
if isinstance(content, list):
|
||||
content[content_idx]["text"] = (
|
||||
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
case ToolResultStringTarget(content_idx=content_idx):
|
||||
if isinstance(content, list):
|
||||
content[content_idx]["content"] = (
|
||||
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx):
|
||||
if isinstance(content, list):
|
||||
content[content_idx]["content"][block_idx]["text"] = (
|
||||
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
case _:
|
||||
assert_never(target)
|
||||
|
||||
async def process_output_response(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -592,7 +592,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
# Anthropic requires additionalProperties=false for object schemas
|
||||
# See: https://docs.anthropic.com/en/docs/build-with-claude/structured-outputs
|
||||
if result.get("type") == "object" and "additionalProperties" not in result:
|
||||
if result.get("type") == "object":
|
||||
result["additionalProperties"] = False
|
||||
|
||||
return result
|
||||
|
|
|
|||
|
|
@ -1,5 +1,10 @@
|
|||
from collections.abc import AsyncIterator, Coroutine, Iterator
|
||||
from typing import Any, Final, cast
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Final,
|
||||
cast,
|
||||
)
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -21,6 +26,10 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
|
|||
from litellm.types.utils import ModelResponse
|
||||
from litellm.utils import get_model_info
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.router import Router
|
||||
|
||||
# Anthropic-only keys already mapped by the translator; strip on extra_kwargs re-merge.
|
||||
ANTHROPIC_ONLY_REQUEST_KEYS: Final[frozenset[str]] = frozenset({"output_config"})
|
||||
|
||||
|
|
@ -37,6 +46,14 @@ def _messages_have_compaction_block(messages: list[dict]) -> bool:
|
|||
return False
|
||||
|
||||
|
||||
def _proxy_router_fallback() -> "Router | None":
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router as _proxy_router
|
||||
except Exception:
|
||||
return None
|
||||
return _proxy_router
|
||||
|
||||
|
||||
def _extract_proxy_litellm_metadata(kwargs: dict[str, Any]) -> dict[str, Any] | None:
|
||||
"""Return ``kwargs["litellm_metadata"]`` when it's a dict; ``None`` otherwise.
|
||||
|
||||
|
|
@ -64,8 +81,8 @@ async def _prepare_context_managed_request(
|
|||
context_management_spec: Any,
|
||||
litellm_metadata: dict | None,
|
||||
additional_drop_params: list[str] | None,
|
||||
llm_router: Any,
|
||||
user_api_key_auth: Any = None,
|
||||
llm_router: "Router | None",
|
||||
user_api_key_auth: "UserAPIKeyAuth | None" = None,
|
||||
) -> PolyfillResult | None:
|
||||
"""Apply client compaction history, then optional context_management polyfill."""
|
||||
from litellm.llms.anthropic.experimental_pass_through.context_management.editors.compact import (
|
||||
|
|
@ -149,7 +166,7 @@ def _polyfill_will_run(
|
|||
COMPACT_EDIT_TYPE,
|
||||
)
|
||||
|
||||
return any(isinstance(edit, dict) and edit.get("type") == COMPACT_EDIT_TYPE for edit in edits)
|
||||
return any(edit.get("type") == COMPACT_EDIT_TYPE for edit in edits)
|
||||
|
||||
|
||||
def _spec_has_non_compact_edits(
|
||||
|
|
@ -175,10 +192,7 @@ def _spec_has_non_compact_edits(
|
|||
COMPACT_EDIT_TYPE,
|
||||
)
|
||||
|
||||
return any(
|
||||
isinstance(edit, dict) and isinstance(edit.get("type"), str) and edit.get("type") != COMPACT_EDIT_TYPE
|
||||
for edit in edits
|
||||
)
|
||||
return any(isinstance(edit.get("type"), str) and edit.get("type") != COMPACT_EDIT_TYPE for edit in edits)
|
||||
|
||||
|
||||
def _context_management_explicitly_dropped(additional_drop_params: list[str] | None) -> bool:
|
||||
|
|
@ -228,8 +242,8 @@ async def _run_polyfill_if_enabled(
|
|||
context_management_spec: Any,
|
||||
litellm_metadata: dict | None,
|
||||
additional_drop_params: list[str] | None,
|
||||
llm_router: Any,
|
||||
user_api_key_auth: Any = None,
|
||||
llm_router: "Router | None",
|
||||
user_api_key_auth: "UserAPIKeyAuth | None" = None,
|
||||
) -> PolyfillResult | None:
|
||||
"""Run the async context_management polyfill if a spec is present.
|
||||
|
||||
|
|
@ -339,7 +353,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
reasoning_effort: Final = completion_kwargs.get("reasoning_effort")
|
||||
summary: Final = thinking.get("summary")
|
||||
if isinstance(reasoning_effort, str) and reasoning_effort:
|
||||
reasoning_dict: Final[dict[str, Any]] = {"effort": reasoning_effort}
|
||||
reasoning_dict: Final[dict[str, object]] = {"effort": reasoning_effort}
|
||||
if summary:
|
||||
reasoning_dict["summary"] = summary
|
||||
elif auto_summary:
|
||||
|
|
@ -528,21 +542,17 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
top_p: float | None = None,
|
||||
output_format: dict | None = None,
|
||||
**kwargs,
|
||||
) -> AnthropicMessagesResponse | AsyncIterator[Any] | Iterator[bytes]:
|
||||
) -> AnthropicMessagesResponse | AsyncIterator[bytes] | Iterator[bytes]:
|
||||
"""Handle non-Anthropic models asynchronously using the adapter"""
|
||||
context_management: Final = kwargs.pop("context_management", None)
|
||||
additional_drop_params: Final[list[str] | None] = kwargs.get("additional_drop_params", None)
|
||||
litellm_router = kwargs.pop("litellm_router", None)
|
||||
if litellm_router is None:
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router as _proxy_router
|
||||
|
||||
litellm_router = _proxy_router
|
||||
except Exception:
|
||||
pass
|
||||
requested_router: Final[Router | None] = kwargs.pop("litellm_router", None)
|
||||
litellm_router: Final[Router | None] = (
|
||||
requested_router if requested_router is not None else _proxy_router_fallback()
|
||||
)
|
||||
|
||||
proxy_litellm_metadata: Final = _extract_proxy_litellm_metadata(kwargs)
|
||||
user_api_key_auth: Final = (
|
||||
user_api_key_auth: Final[UserAPIKeyAuth | None] = (
|
||||
proxy_litellm_metadata.get("user_api_key_auth") if proxy_litellm_metadata is not None else None
|
||||
)
|
||||
|
||||
|
|
@ -626,8 +636,8 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
) -> (
|
||||
AnthropicMessagesResponse
|
||||
| Iterator[bytes]
|
||||
| AsyncIterator[Any]
|
||||
| Coroutine[Any, Any, AnthropicMessagesResponse | AsyncIterator[Any] | Iterator[bytes]]
|
||||
| AsyncIterator[bytes]
|
||||
| Coroutine[None, None, AnthropicMessagesResponse | AsyncIterator[bytes] | Iterator[bytes]]
|
||||
):
|
||||
"""Handle non-Anthropic models using the adapter."""
|
||||
if _is_async is True:
|
||||
|
|
@ -667,7 +677,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
# ``llm_router`` is ``None``, which is safe to call from the bridged
|
||||
# loop. The async ``async_anthropic_messages_handler`` path is
|
||||
# unaffected because it ``await``s within the original event loop.
|
||||
litellm_router: Final = kwargs.pop("litellm_router", None)
|
||||
litellm_router: Final[Router | None] = kwargs.pop("litellm_router", None)
|
||||
|
||||
# Skip the async bridge entirely when there is nothing for either the
|
||||
# polyfill or the client-history slice-only fallback to do. The vast
|
||||
|
|
@ -679,7 +689,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
polyfill_result: PolyfillResult | None = None
|
||||
else:
|
||||
proxy_litellm_metadata: Final = _extract_proxy_litellm_metadata(kwargs)
|
||||
user_api_key_auth: Final = (
|
||||
user_api_key_auth: Final[UserAPIKeyAuth | None] = (
|
||||
proxy_litellm_metadata.get("user_api_key_auth") if proxy_litellm_metadata is not None else None
|
||||
)
|
||||
polyfill_result = run_async_function(
|
||||
|
|
|
|||
|
|
@ -4,8 +4,15 @@ import copy
|
|||
import json
|
||||
import traceback
|
||||
from collections import deque
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, get_args
|
||||
from collections.abc import AsyncIterator, Iterator, Sequence
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Final,
|
||||
Literal,
|
||||
Protocol,
|
||||
get_args,
|
||||
)
|
||||
|
||||
from typing_extensions import assert_never
|
||||
|
||||
|
|
@ -14,7 +21,9 @@ from litellm._uuid import uuid
|
|||
from litellm.types.llms.anthropic import (
|
||||
AppliedEdit,
|
||||
CompactionBlock,
|
||||
ContentBlockDelta,
|
||||
ContextManagementResponse,
|
||||
MessageBlockDelta,
|
||||
StreamingContentBlockDeltaType,
|
||||
UsageDelta,
|
||||
UsageIteration,
|
||||
|
|
@ -28,6 +37,25 @@ if TYPE_CHECKING:
|
|||
_STREAMING_DELTA_TYPES: Final = frozenset(get_args(StreamingContentBlockDeltaType))
|
||||
|
||||
|
||||
class _UsageDeltaWithIterations(UsageDelta, total=False):
|
||||
iterations: list[UsageIteration]
|
||||
|
||||
|
||||
class _ChunkStream(Protocol):
|
||||
def __iter__(self) -> "Iterator[ModelResponseStream]": ...
|
||||
|
||||
def __aiter__(self) -> "AsyncIterator[ModelResponseStream]": ...
|
||||
|
||||
|
||||
def _optional_attr(obj: object, name: str) -> object:
|
||||
return getattr(obj, name, None)
|
||||
|
||||
|
||||
def _optional_attr_sequence(obj: object, name: str) -> Sequence[object]:
|
||||
value: Final = getattr(obj, name, None)
|
||||
return value if value else ()
|
||||
|
||||
|
||||
def _delta_payload_field(delta_type: StreamingContentBlockDeltaType) -> str:
|
||||
match delta_type:
|
||||
case "text_delta":
|
||||
|
|
@ -62,29 +90,29 @@ class _CombinedChunkSplitter:
|
|||
would advance them out of sync.
|
||||
"""
|
||||
|
||||
def __init__(self, completion_stream: Any):
|
||||
self._stream = completion_stream
|
||||
self._sync_iter: Iterator[Any] | None = None
|
||||
self._async_iter: AsyncIterator[Any] | None = None
|
||||
self._buffer: deque = deque()
|
||||
def __init__(self, completion_stream: _ChunkStream):
|
||||
self._stream: _ChunkStream = completion_stream
|
||||
self._sync_iter: Iterator[ModelResponseStream] | None = None
|
||||
self._async_iter: AsyncIterator[ModelResponseStream] | None = None
|
||||
self._buffer: deque[ModelResponseStream] = deque()
|
||||
|
||||
@staticmethod
|
||||
def _is_combined(chunk: Any) -> bool:
|
||||
def _is_combined(chunk: "ModelResponseStream") -> bool:
|
||||
"""True if ``chunk`` carries response content AND a finish_reason."""
|
||||
choices: Final = getattr(chunk, "choices", None)
|
||||
choices: Final = _optional_attr_sequence(chunk, "choices")
|
||||
if not choices:
|
||||
return False
|
||||
choice: Final = choices[0]
|
||||
if getattr(choice, "finish_reason", None) is None:
|
||||
if _optional_attr(choice, "finish_reason") is None:
|
||||
return False
|
||||
delta: Final = getattr(choice, "delta", None)
|
||||
delta: Final = _optional_attr(choice, "delta")
|
||||
if delta is None:
|
||||
return False
|
||||
return bool(
|
||||
getattr(delta, "content", None)
|
||||
or getattr(delta, "tool_calls", None)
|
||||
or getattr(delta, "reasoning_content", None)
|
||||
or getattr(delta, "thinking_blocks", None)
|
||||
_optional_attr(delta, "content")
|
||||
or _optional_attr(delta, "tool_calls")
|
||||
or _optional_attr(delta, "reasoning_content")
|
||||
or _optional_attr(delta, "thinking_blocks")
|
||||
)
|
||||
|
||||
_PAYLOAD_FIELD_GROUPS: "tuple[tuple[str, ...], ...]" = (
|
||||
|
|
@ -119,21 +147,21 @@ class _CombinedChunkSplitter:
|
|||
normalized to ``reasoning_content`` so the synthesized block start
|
||||
stays empty and the thinking text is emitted exactly once.
|
||||
"""
|
||||
choices: Final = getattr(chunk, "choices", None)
|
||||
if not choices or len(choices) != 1:
|
||||
choices: Final = _optional_attr_sequence(chunk, "choices")
|
||||
if len(choices) != 1:
|
||||
return (chunk,)
|
||||
delta: Final = getattr(choices[0], "delta", None)
|
||||
delta: Final = _optional_attr(choices[0], "delta")
|
||||
if delta is None:
|
||||
return (chunk,)
|
||||
tool_calls: Final = getattr(delta, "tool_calls", None)
|
||||
tool_calls: Final = _optional_attr_sequence(delta, "tool_calls")
|
||||
if tool_calls and not any(
|
||||
getattr(getattr(tool_call, "function", None), "name", None) for tool_call in tool_calls
|
||||
_optional_attr(_optional_attr(tool_call, "function"), "name") for tool_call in tool_calls
|
||||
):
|
||||
return (chunk,)
|
||||
present_groups: Final = tuple(
|
||||
group
|
||||
for group in _CombinedChunkSplitter._PAYLOAD_FIELD_GROUPS
|
||||
if any(getattr(delta, field, None) for field in group)
|
||||
if any(_optional_attr(delta, field) for field in group)
|
||||
)
|
||||
if len(present_groups) <= 1:
|
||||
return (chunk,)
|
||||
|
|
@ -172,7 +200,7 @@ class _CombinedChunkSplitter:
|
|||
return {"reasoning_content": thinking_text}
|
||||
|
||||
@staticmethod
|
||||
def _split(chunk: Any) -> list[Any]:
|
||||
def _split(chunk: "ModelResponseStream") -> "list[ModelResponseStream]":
|
||||
"""Return ``[chunk]``, or ``[content_chunk, finish_chunk]`` if combined."""
|
||||
if not _CombinedChunkSplitter._is_combined(chunk):
|
||||
return [chunk]
|
||||
|
|
@ -194,10 +222,10 @@ class _CombinedChunkSplitter:
|
|||
finish_delta.thinking_blocks = None
|
||||
return [content_chunk, finish_chunk]
|
||||
|
||||
def __iter__(self) -> "Iterator[Any]":
|
||||
def __iter__(self) -> "Iterator[ModelResponseStream]":
|
||||
return self
|
||||
|
||||
def __next__(self) -> Any:
|
||||
def __next__(self) -> "ModelResponseStream":
|
||||
if self._buffer:
|
||||
return self._buffer.popleft()
|
||||
if self._sync_iter is None:
|
||||
|
|
@ -210,10 +238,10 @@ class _CombinedChunkSplitter:
|
|||
)
|
||||
return self._buffer.popleft()
|
||||
|
||||
def __aiter__(self) -> "AsyncIterator[Any]":
|
||||
def __aiter__(self) -> "AsyncIterator[ModelResponseStream]":
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> Any:
|
||||
async def __anext__(self) -> "ModelResponseStream":
|
||||
if self._buffer:
|
||||
return self._buffer.popleft()
|
||||
if self._async_iter is None:
|
||||
|
|
@ -246,14 +274,14 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
sent_content_block_finish: bool = False
|
||||
current_content_block_type: Literal["text", "tool_use", "thinking"] = "text"
|
||||
sent_last_message: bool = False
|
||||
holding_chunk: Any | None = None
|
||||
holding_stop_reason_chunk: Any | None = None
|
||||
holding_chunk: ContentBlockDelta | None = None
|
||||
holding_stop_reason_chunk: MessageBlockDelta | None = None
|
||||
queued_usage_chunk: bool = False
|
||||
current_content_block_index: int = 0
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
completion_stream: Any,
|
||||
completion_stream: _ChunkStream,
|
||||
model: str,
|
||||
tool_name_mapping: dict[str, str] | None = None,
|
||||
applied_edits: list[AppliedEdit] | None = None,
|
||||
|
|
@ -294,7 +322,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
text="",
|
||||
)
|
||||
|
||||
def _merge_usage_into_held_stop_reason_chunk(self, chunk: Any) -> dict[str, Any]:
|
||||
def _merge_usage_into_held_stop_reason_chunk(self, chunk: Any) -> MessageBlockDelta:
|
||||
"""Merge usage data from ``chunk`` into the held ``message_delta`` chunk.
|
||||
|
||||
Shared by both the sync ``__next__`` and async ``__anext__`` paths so
|
||||
|
|
@ -320,7 +348,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
merged_chunk["context_management"] = ContextManagementResponse(applied_edits=list(self.applied_edits))
|
||||
return self._augment_message_delta_usage(merged_chunk)
|
||||
|
||||
def _ensure_context_management_attached(self, message_delta_chunk: dict[str, Any]) -> dict[str, Any]:
|
||||
def _ensure_context_management_attached(self, message_delta_chunk: MessageBlockDelta) -> MessageBlockDelta:
|
||||
"""Attach ``context_management`` to a ``message_delta`` chunk if
|
||||
``self.applied_edits`` is non-empty and the chunk does not already
|
||||
carry it. Returns the (possibly new) chunk dict.
|
||||
|
|
@ -335,7 +363,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
augmented["context_management"] = ContextManagementResponse(applied_edits=list(self.applied_edits))
|
||||
return augmented
|
||||
|
||||
def _augment_message_delta_usage(self, message_delta_chunk: dict[str, Any]) -> dict[str, Any]:
|
||||
def _augment_message_delta_usage(self, message_delta_chunk: MessageBlockDelta) -> MessageBlockDelta:
|
||||
"""Attach polyfill compaction iteration usage to the final message_delta.
|
||||
|
||||
Also defensively re-attaches ``context_management`` so the direct
|
||||
|
|
@ -352,7 +380,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
input_tokens: Final = usage.get("input_tokens", 0) or 0
|
||||
output_tokens: Final = usage.get("output_tokens", 0) or 0
|
||||
augmented: Final = message_delta_chunk.copy()
|
||||
augmented_usage: Final = dict(usage)
|
||||
augmented_usage: Final[_UsageDeltaWithIterations] = {**usage}
|
||||
iterations: Final[list[UsageIteration]] = list(self.iterations_usage)
|
||||
# Only emit a ``message`` iteration when we have real token data.
|
||||
# Without a separate usage chunk (e.g. provider sent finish_reason
|
||||
|
|
|
|||
|
|
@ -427,88 +427,95 @@ class BaseAzureLLM(BaseOpenAILLM):
|
|||
f"|azure_password={hashlib.sha256(_azure_password.encode()).hexdigest() if isinstance(_azure_password, str) else None}"
|
||||
f"|azure_scope={_lp.get('azure_scope')}"
|
||||
)
|
||||
if client is None:
|
||||
cached_client: Final = self.get_cached_openai_client(
|
||||
client_initialization_params=client_initialization_params,
|
||||
client_type="azure",
|
||||
)
|
||||
if cached_client:
|
||||
if isinstance(cached_client, (AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI)):
|
||||
return cached_client
|
||||
|
||||
azure_client_params: Final = self.initialize_azure_sdk_client(
|
||||
litellm_params=litellm_params or {},
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
model_name=model,
|
||||
api_version=api_version,
|
||||
is_async=_is_async,
|
||||
)
|
||||
|
||||
# For Azure v1 API, use standard OpenAI client instead of AzureOpenAI
|
||||
# See: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#api-specs
|
||||
if self._is_azure_v1_api_version(api_version):
|
||||
# Extract only params that OpenAI client accepts
|
||||
# Always use /openai/v1/ regardless of whether user passed "v1", "latest", or "preview"
|
||||
# The OpenAI client accepts a callable for `api_key` and re-invokes it
|
||||
# on every request (via `_refresh_api_key`), so passing
|
||||
# `azure_ad_token_provider` directly preserves Azure AD token refresh
|
||||
# behavior that the regular AzureOpenAI client provides.
|
||||
v1_api_key: str | Callable[[], Any] | None = (
|
||||
azure_client_params.get("api_key")
|
||||
or azure_client_params.get("azure_ad_token_provider")
|
||||
or azure_client_params.get("azure_ad_token")
|
||||
)
|
||||
if _is_async is True and callable(v1_api_key):
|
||||
# AsyncOpenAI expects an async provider; wrap the sync provider
|
||||
# returned by azure-identity. Offload to a thread so a token
|
||||
# refresh (blocking HTTP call to AAD on cache miss) does not
|
||||
# stall the event loop.
|
||||
_sync_provider: Final = v1_api_key
|
||||
|
||||
async def _async_v1_api_key() -> str:
|
||||
return await asyncio.to_thread(_sync_provider)
|
||||
|
||||
v1_api_key = _async_v1_api_key
|
||||
|
||||
v1_params: Final[dict[str, Any]] = {
|
||||
"api_key": v1_api_key,
|
||||
"base_url": f"{api_base}/openai/v1/",
|
||||
}
|
||||
if "timeout" in azure_client_params:
|
||||
v1_params["timeout"] = azure_client_params["timeout"]
|
||||
if "max_retries" in azure_client_params:
|
||||
v1_params["max_retries"] = azure_client_params["max_retries"]
|
||||
if "http_client" in azure_client_params:
|
||||
v1_params["http_client"] = azure_client_params["http_client"]
|
||||
|
||||
verbose_logger.debug("Using Azure v1 API with base_url: %s", v1_params["base_url"])
|
||||
|
||||
if _is_async is True:
|
||||
openai_client = AsyncOpenAI(**v1_params)
|
||||
else:
|
||||
openai_client = OpenAI(**v1_params)
|
||||
else:
|
||||
# Traditional Azure API uses AzureOpenAI client
|
||||
if _is_async is True:
|
||||
openai_client = AsyncAzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
openai_client = AzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
openai_client = client
|
||||
if client is not None:
|
||||
if (
|
||||
api_version is not None
|
||||
and isinstance(openai_client, (AzureOpenAI, AsyncAzureOpenAI))
|
||||
and isinstance(openai_client._custom_query, dict)
|
||||
and isinstance(client, (AzureOpenAI, AsyncAzureOpenAI))
|
||||
and isinstance(client._custom_query, dict)
|
||||
):
|
||||
# set api_version to version passed by user
|
||||
openai_client._custom_query.setdefault("api-version", api_version)
|
||||
client._custom_query.setdefault("api-version", api_version)
|
||||
self.set_cached_openai_client(
|
||||
openai_client=client,
|
||||
client_initialization_params=client_initialization_params,
|
||||
client_type="azure",
|
||||
litellm_owned_client=False,
|
||||
)
|
||||
return client
|
||||
|
||||
cached_client: Final = self.get_cached_openai_client(
|
||||
client_initialization_params=client_initialization_params,
|
||||
client_type="azure",
|
||||
)
|
||||
if cached_client:
|
||||
if isinstance(cached_client, (AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI)):
|
||||
return cached_client
|
||||
|
||||
azure_client_params: Final = self.initialize_azure_sdk_client(
|
||||
litellm_params=litellm_params or {},
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
model_name=model,
|
||||
api_version=api_version,
|
||||
is_async=_is_async,
|
||||
)
|
||||
|
||||
# For Azure v1 API, use standard OpenAI client instead of AzureOpenAI
|
||||
# See: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#api-specs
|
||||
if self._is_azure_v1_api_version(api_version):
|
||||
# Extract only params that OpenAI client accepts
|
||||
# Always use /openai/v1/ regardless of whether user passed "v1", "latest", or "preview"
|
||||
# The OpenAI client accepts a callable for `api_key` and re-invokes it
|
||||
# on every request (via `_refresh_api_key`), so passing
|
||||
# `azure_ad_token_provider` directly preserves Azure AD token refresh
|
||||
# behavior that the regular AzureOpenAI client provides.
|
||||
v1_api_key: str | Callable[[], Any] | None = (
|
||||
azure_client_params.get("api_key")
|
||||
or azure_client_params.get("azure_ad_token_provider")
|
||||
or azure_client_params.get("azure_ad_token")
|
||||
)
|
||||
if _is_async is True and callable(v1_api_key):
|
||||
# AsyncOpenAI expects an async provider; wrap the sync provider
|
||||
# returned by azure-identity. Offload to a thread so a token
|
||||
# refresh (blocking HTTP call to AAD on cache miss) does not
|
||||
# stall the event loop.
|
||||
_sync_provider: Final = v1_api_key
|
||||
|
||||
async def _async_v1_api_key() -> str:
|
||||
return await asyncio.to_thread(_sync_provider)
|
||||
|
||||
v1_api_key = _async_v1_api_key
|
||||
|
||||
v1_params: Final[dict[str, Any]] = {
|
||||
"api_key": v1_api_key,
|
||||
"base_url": f"{api_base}/openai/v1/",
|
||||
}
|
||||
if "timeout" in azure_client_params:
|
||||
v1_params["timeout"] = azure_client_params["timeout"]
|
||||
if "max_retries" in azure_client_params:
|
||||
v1_params["max_retries"] = azure_client_params["max_retries"]
|
||||
if "http_client" in azure_client_params:
|
||||
v1_params["http_client"] = azure_client_params["http_client"]
|
||||
|
||||
verbose_logger.debug("Using Azure v1 API with base_url: %s", v1_params["base_url"])
|
||||
|
||||
if _is_async is True:
|
||||
openai_client = AsyncOpenAI(**v1_params)
|
||||
else:
|
||||
openai_client = OpenAI(**v1_params)
|
||||
else:
|
||||
# Traditional Azure API uses AzureOpenAI client
|
||||
if _is_async is True:
|
||||
openai_client = AsyncAzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
openai_client = AzureOpenAI(**azure_client_params)
|
||||
|
||||
# save client in-memory cache
|
||||
self.set_cached_openai_client(
|
||||
openai_client=openai_client,
|
||||
client_initialization_params=client_initialization_params,
|
||||
client_type="azure",
|
||||
litellm_owned_client=self.owns_wrapped_http_client(azure_client_params.get("http_client")),
|
||||
)
|
||||
return openai_client
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Final
|
||||
from collections.abc import Callable, Iterator, Sequence
|
||||
from typing import Any, Final, TypeVar
|
||||
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
|
@ -113,13 +114,131 @@ def effective_skip_tool_message_for_guardrail(guardrail_to_apply: Any) -> bool:
|
|||
return bool(getattr(litellm, "skip_tool_message_in_guardrail", False))
|
||||
|
||||
|
||||
def _message_role(message: AllMessageValues) -> str:
|
||||
return str((message or {}).get("role") or "").lower()
|
||||
|
||||
|
||||
def openai_messages_without_system(
|
||||
messages: list[AllMessageValues],
|
||||
) -> list[AllMessageValues]:
|
||||
return [m for m in messages if str((m or {}).get("role") or "").lower() != "system"]
|
||||
messages: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
return tuple(m for m in messages if _message_role(m) != "system")
|
||||
|
||||
|
||||
def openai_messages_without_tool(
|
||||
messages: list[AllMessageValues],
|
||||
messages: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
return tuple(m for m in messages if _message_role(m) != "tool")
|
||||
|
||||
|
||||
def effective_scan_only_tool_results_for_guardrail(guardrail_to_apply: object) -> bool:
|
||||
return getattr(guardrail_to_apply, "scan_only_tool_results", None) is True
|
||||
|
||||
|
||||
def role_out_of_guardrail_scope(
|
||||
role: str,
|
||||
*,
|
||||
skip_system_message: bool,
|
||||
skip_tool_message: bool,
|
||||
scan_only_tool_results: bool = False,
|
||||
) -> bool:
|
||||
if skip_system_message and role == "system":
|
||||
return True
|
||||
if skip_tool_message and role == "tool":
|
||||
return True
|
||||
return scan_only_tool_results and role not in ("tool", "function")
|
||||
|
||||
|
||||
def scoped_structured_message_indices(
|
||||
messages: Sequence[AllMessageValues],
|
||||
*,
|
||||
scan_only_tool_results: bool,
|
||||
skip_system: bool,
|
||||
skip_tool: bool,
|
||||
) -> tuple[int, ...]:
|
||||
return tuple(
|
||||
index
|
||||
for index, message in enumerate(messages)
|
||||
if not role_out_of_guardrail_scope(
|
||||
_message_role(message),
|
||||
skip_system_message=skip_system,
|
||||
skip_tool_message=skip_tool,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
ToolT = TypeVar("ToolT")
|
||||
|
||||
|
||||
def openai_tool_name(tool: object) -> str | None:
|
||||
if not isinstance(tool, dict):
|
||||
return None
|
||||
function: Final = tool.get("function")
|
||||
if isinstance(function, dict):
|
||||
function_name: Final = function.get("name")
|
||||
return function_name if isinstance(function_name, str) else None
|
||||
flat_name: Final = tool.get("name")
|
||||
return flat_name if isinstance(flat_name, str) else None
|
||||
|
||||
|
||||
def anthropic_tool_name(tool: object) -> str | None:
|
||||
name: Final = tool.get("name") if isinstance(tool, dict) else None
|
||||
return name if isinstance(name, str) else None
|
||||
|
||||
|
||||
def merge_returned_tools_into_request_tools(
|
||||
request_tools: Sequence[ToolT] | None,
|
||||
returned_tools: Sequence[ToolT],
|
||||
tool_name: Callable[[ToolT], str | None],
|
||||
) -> list[ToolT]:
|
||||
"""Union of the request's tools and guardrail-returned tools, keyed by name.
|
||||
|
||||
Under ``scan_only_tool_results`` the guardrail never saw the request's
|
||||
tools, so a returned list can neither replace them (it would drop every
|
||||
user-defined function) nor be discarded (it may carry a tool the guardrail
|
||||
synthesized and told the model to call, like Compresr's retrieve tool).
|
||||
Keep every request tool and append only returned tools whose names aren't
|
||||
already taken by a request tool or an earlier returned tool.
|
||||
"""
|
||||
originals: Final = tuple(request_tools or ())
|
||||
taken_names: Final = frozenset(name for tool in originals if (name := tool_name(tool)) is not None)
|
||||
additions: Final = tuple(
|
||||
tool
|
||||
for index, tool in enumerate(returned_tools)
|
||||
if (name := tool_name(tool)) not in taken_names
|
||||
and (name is None or all(tool_name(earlier) != name for earlier in returned_tools[:index]))
|
||||
)
|
||||
return [*originals, *additions]
|
||||
|
||||
|
||||
def merge_guardrailed_scoped_messages(
|
||||
full_messages: Sequence[AllMessageValues],
|
||||
scoped_indices: Sequence[int],
|
||||
guardrailed_scoped: Sequence[AllMessageValues],
|
||||
) -> list[AllMessageValues]:
|
||||
return [m for m in messages if str((m or {}).get("role") or "").lower() != "tool"]
|
||||
"""Substitute guardrail-returned messages back into the full conversation.
|
||||
|
||||
Guardrails only ever see the scoped subset of messages, so a replacement
|
||||
list they hand back describes that subset, not the whole request. Writing
|
||||
it over ``data["messages"]`` wholesale would silently drop every
|
||||
out-of-scope message (system prompt, prior turns). Instead, swap each
|
||||
returned message into the position its scoped original came from; extra
|
||||
returned messages land after the last scoped position, and scoped
|
||||
originals without a counterpart are treated as removed by the guardrail.
|
||||
When nothing was filtered out this degenerates to the returned list
|
||||
itself, preserving wholesale-replacement behavior for unscoped guardrails.
|
||||
"""
|
||||
replacements: Final = dict(zip(scoped_indices, guardrailed_scoped))
|
||||
removed: Final = frozenset(scoped_indices[len(guardrailed_scoped) :])
|
||||
appended: Final = tuple(guardrailed_scoped[len(scoped_indices) :])
|
||||
last_scoped_index: Final = scoped_indices[-1] if scoped_indices else None
|
||||
|
||||
def _merged() -> Iterator[AllMessageValues]:
|
||||
for index, message in enumerate(full_messages):
|
||||
if index in removed:
|
||||
continue
|
||||
yield replacements.get(index, message)
|
||||
if index == last_scoped_index:
|
||||
yield from appended
|
||||
|
||||
return list(_merged())
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@
|
|||
import base64
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, TypeVar, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, TypeVar, cast
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm.llms.base_llm.managed_resources.isolation import (
|
||||
|
|
@ -23,7 +23,7 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy.utils import PrismaClient as _PrismaClient
|
||||
from litellm.router import Router as _Router
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
Span = _Span | Any
|
||||
InternalUsageCache = _InternalUsageCache
|
||||
PrismaClient = _PrismaClient
|
||||
Router = _Router
|
||||
|
|
|
|||
|
|
@ -859,6 +859,7 @@ class BaseAWSLLM:
|
|||
"Action": [
|
||||
"bedrock:InvokeModel",
|
||||
"bedrock:InvokeModelWithResponseStream",
|
||||
"bedrock:CountTokens",
|
||||
"bedrock:ApplyGuardrail",
|
||||
"bedrock:GetGuardrail",
|
||||
"bedrock:ListGuardrails",
|
||||
|
|
|
|||
|
|
@ -12,7 +12,6 @@ from litellm.litellm_core_utils.cloud_storage_security import (
|
|||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.bedrock import (
|
||||
BedrockCreateBatchRequest,
|
||||
BedrockCreateBatchResponse,
|
||||
|
|
@ -29,7 +28,7 @@ from litellm.types.llms.openai import (
|
|||
from litellm.types.utils import LiteLLMBatch, LlmProviders
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
from ..common_utils import CommonBatchFilesUtils
|
||||
from ..common_utils import CommonBatchFilesUtils, resolve_s3_encryption_key_id
|
||||
|
||||
# Bedrock batch input files are uploaded as
|
||||
# s3://bucket/litellm-bedrock-files-{model, ":" -> "-"}-{uuid4}.jsonl (see
|
||||
|
|
@ -200,7 +199,10 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
|
|||
)
|
||||
|
||||
# Add optional KMS encryption key ID if provided
|
||||
s3_encryption_key_id = litellm_params.get("s3_encryption_key_id") or get_secret_str("AWS_S3_ENCRYPTION_KEY_ID")
|
||||
s3_encryption_key_id = resolve_s3_encryption_key_id(
|
||||
litellm_params=litellm_params,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
if s3_encryption_key_id:
|
||||
s3_output_config["s3EncryptionKeyId"] = s3_encryption_key_id
|
||||
|
||||
|
|
|
|||
|
|
@ -57,7 +57,7 @@ class AmazonCohereChatConfig:
|
|||
Reference - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-cohere-command-r-plus.html
|
||||
"""
|
||||
|
||||
documents: List[Document] | None = None
|
||||
documents: list[Document] | None = None
|
||||
search_queries_only: bool | None = None
|
||||
preamble: str | None = None
|
||||
max_tokens: int | None = None
|
||||
|
|
@ -69,12 +69,12 @@ class AmazonCohereChatConfig:
|
|||
presence_penalty: float | None = None
|
||||
seed: int | None = None
|
||||
return_prompt: bool | None = None
|
||||
stop_sequences: List[str] | None = None
|
||||
stop_sequences: list[str] | None = None
|
||||
raw_prompting: bool | None = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
documents: List[Document] | None = None,
|
||||
documents: list[Document] | None = None,
|
||||
search_queries_only: bool | None = None,
|
||||
preamble: str | None = None,
|
||||
max_tokens: int | None = None,
|
||||
|
|
@ -112,7 +112,7 @@ class AmazonCohereChatConfig:
|
|||
and v is not None
|
||||
}
|
||||
|
||||
def get_supported_openai_params(self) -> List[str]:
|
||||
def get_supported_openai_params(self) -> list[str]:
|
||||
return [
|
||||
"max_tokens",
|
||||
"max_completion_tokens",
|
||||
|
|
@ -325,7 +325,7 @@ class AWSEventStreamDecoder:
|
|||
|
||||
self.model = model
|
||||
self.parser = EventStreamJSONParser()
|
||||
self.content_blocks: List[ContentBlockDeltaEvent] = []
|
||||
self.content_blocks: list[ContentBlockDeltaEvent] = []
|
||||
self.tool_calls_index: int | None = None
|
||||
self.response_id: str | None = None
|
||||
self.json_mode = json_mode
|
||||
|
|
@ -362,13 +362,13 @@ class AWSEventStreamDecoder:
|
|||
|
||||
def translate_thinking_blocks(
|
||||
self, thinking_block: BedrockConverseReasoningContentBlockDelta
|
||||
) -> List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None:
|
||||
) -> list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None:
|
||||
"""
|
||||
Translate the thinking blocks to a string
|
||||
"""
|
||||
|
||||
thinking_blocks_list: Final[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]] = []
|
||||
_thinking_block: Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock] | None = None
|
||||
thinking_blocks_list: Final[list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock]] = []
|
||||
_thinking_block: ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock | None = None
|
||||
|
||||
if "text" in thinking_block:
|
||||
_thinking_block = ChatCompletionThinkingBlock(type="thinking")
|
||||
|
|
@ -402,12 +402,12 @@ class AWSEventStreamDecoder:
|
|||
) -> tuple[
|
||||
ChatCompletionToolCallChunk | None,
|
||||
dict,
|
||||
List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None,
|
||||
list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None,
|
||||
]:
|
||||
"""Handle 'start' event in converse chunk parsing."""
|
||||
tool_use: ChatCompletionToolCallChunk | None = None
|
||||
provider_specific_fields: dict = {}
|
||||
thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
|
||||
|
||||
self.content_blocks = [] # reset
|
||||
if start_obj is not None:
|
||||
|
|
@ -450,14 +450,14 @@ class AWSEventStreamDecoder:
|
|||
ChatCompletionToolCallChunk | None,
|
||||
dict,
|
||||
str | None,
|
||||
List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None,
|
||||
list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None,
|
||||
]:
|
||||
"""Handle 'delta' event in converse chunk parsing."""
|
||||
text = ""
|
||||
tool_use: ChatCompletionToolCallChunk | None = None
|
||||
provider_specific_fields: dict = {}
|
||||
reasoning_content: str | None = None
|
||||
thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
|
||||
|
||||
self.content_blocks.append(delta_obj)
|
||||
if "text" in delta_obj:
|
||||
|
|
@ -535,7 +535,7 @@ class AWSEventStreamDecoder:
|
|||
usage: Usage | None = None
|
||||
provider_specific_fields: dict = {}
|
||||
reasoning_content: str | None = None
|
||||
thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
|
||||
|
||||
content_block_index: Final = int(chunk_data.get("contentBlockIndex", 0))
|
||||
if "start" in chunk_data:
|
||||
|
|
@ -590,7 +590,7 @@ class AWSEventStreamDecoder:
|
|||
except Exception as e:
|
||||
raise Exception(f"Received streaming error - {e}")
|
||||
|
||||
def _chunk_parser(self, chunk_data: dict) -> Union[GChunk, ModelResponseStream, dict]:
|
||||
def _chunk_parser(self, chunk_data: dict) -> GChunk | ModelResponseStream | dict:
|
||||
text = ""
|
||||
is_finished = False
|
||||
finish_reason = ""
|
||||
|
|
@ -645,7 +645,7 @@ class AWSEventStreamDecoder:
|
|||
tool_use=None,
|
||||
)
|
||||
|
||||
def iter_bytes(self, iterator: Iterator[bytes]) -> Iterator[Union[GChunk, ModelResponseStream, dict]]:
|
||||
def iter_bytes(self, iterator: Iterator[bytes]) -> Iterator[GChunk | ModelResponseStream | dict]:
|
||||
"""Given an iterator that yields lines, iterate over it & yield every event encountered"""
|
||||
from botocore.eventstream import EventStreamBuffer
|
||||
|
||||
|
|
@ -659,9 +659,7 @@ class AWSEventStreamDecoder:
|
|||
_data = json.loads(message)
|
||||
yield self._chunk_parser(chunk_data=_data)
|
||||
|
||||
async def aiter_bytes(
|
||||
self, iterator: AsyncIterator[bytes]
|
||||
) -> AsyncIterator[Union[GChunk, ModelResponseStream, dict]]:
|
||||
async def aiter_bytes(self, iterator: AsyncIterator[bytes]) -> AsyncIterator[GChunk | ModelResponseStream | dict]:
|
||||
"""Given an async iterator that yields lines, iterate over it & yield every event encountered"""
|
||||
from botocore.eventstream import EventStreamBuffer
|
||||
|
||||
|
|
@ -741,7 +739,7 @@ class AmazonDeepSeekR1StreamDecoder(AWSEventStreamDecoder):
|
|||
sync_stream=sync_stream,
|
||||
)
|
||||
|
||||
def _chunk_parser(self, chunk_data: dict) -> Union[GChunk, ModelResponseStream, dict]:
|
||||
def _chunk_parser(self, chunk_data: dict) -> GChunk | ModelResponseStream | dict:
|
||||
return self.deepseek_model_response_iterator.chunk_parser(chunk=chunk_data)
|
||||
|
||||
|
||||
|
|
@ -756,7 +754,7 @@ class MockResponseIterator: # for returning ai21 streaming responses
|
|||
return self
|
||||
|
||||
def _handle_json_mode_chunk(
|
||||
self, text: str, tool_calls: List[ChatCompletionToolCallChunk] | None
|
||||
self, text: str, tool_calls: list[ChatCompletionToolCallChunk] | None
|
||||
) -> tuple[str, ChatCompletionToolCallChunk | None]:
|
||||
"""
|
||||
If JSON mode is enabled, convert the tool call to a message.
|
||||
|
|
@ -789,7 +787,7 @@ class MockResponseIterator: # for returning ai21 streaming responses
|
|||
text = chunk_data.choices[0].message.content or ""
|
||||
tool_use = None
|
||||
_model_response_tool_call: Final = cast(
|
||||
List[ChatCompletionMessageToolCall] | None,
|
||||
list[ChatCompletionMessageToolCall] | None,
|
||||
cast(Choices, chunk_data.choices[0]).message.tool_calls,
|
||||
)
|
||||
if self.json_mode is True:
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ from litellm.llms.base_llm.anthropic_messages.transformation import (
|
|||
)
|
||||
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.secret_managers.main import get_secret
|
||||
from litellm.secret_managers.main import get_secret, get_secret_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
|
@ -1304,6 +1304,23 @@ def get_anthropic_beta_from_headers(headers: dict) -> list[str]:
|
|||
return []
|
||||
|
||||
|
||||
def resolve_s3_encryption_key_id(
|
||||
litellm_params: Mapping[str, Any],
|
||||
optional_params: Mapping[str, Any] | None = None,
|
||||
) -> str | None:
|
||||
"""
|
||||
Resolve the SSE-KMS key configured for Bedrock batch/file S3 objects.
|
||||
|
||||
Precedence: `s3_encryption_key_id` in litellm_params, then optional_params
|
||||
(client-side / request params), then the AWS_S3_ENCRYPTION_KEY_ID env var.
|
||||
"""
|
||||
candidates: Final = tuple(
|
||||
source.get("s3_encryption_key_id") for source in (litellm_params, optional_params) if source is not None
|
||||
)
|
||||
explicit: Final = next((value for value in candidates if isinstance(value, str) and value), None)
|
||||
return explicit or get_secret_str("AWS_S3_ENCRYPTION_KEY_ID")
|
||||
|
||||
|
||||
class CommonBatchFilesUtils:
|
||||
"""
|
||||
Common utilities for Bedrock batch and file operations.
|
||||
|
|
|
|||
|
|
@ -34,7 +34,7 @@ class BedrockCohereEmbeddingConfig:
|
|||
new_transformed_request: Final = CohereEmbeddingRequest(
|
||||
input_type=transformed_request["input_type"],
|
||||
)
|
||||
for k in CohereEmbeddingRequest.__annotations__.keys():
|
||||
for k in CohereEmbeddingRequest.__annotations__:
|
||||
if k in transformed_request:
|
||||
new_transformed_request[k] = transformed_request[k]
|
||||
|
||||
|
|
|
|||
|
|
@ -46,7 +46,7 @@ from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums
|
|||
from litellm.utils import get_llm_provider
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
from ..common_utils import BedrockError
|
||||
from ..common_utils import BedrockError, resolve_s3_encryption_key_id
|
||||
|
||||
# litellm_params key used to hand the SigV4-signed GET headers from
|
||||
# `transform_file_content_request` to `validate_environment` (the only hook
|
||||
|
|
@ -733,6 +733,10 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
content=file_content,
|
||||
api_base=api_base,
|
||||
optional_params=optional_params,
|
||||
s3_encryption_key_id=resolve_s3_encryption_key_id(
|
||||
litellm_params=litellm_params,
|
||||
optional_params=optional_params,
|
||||
),
|
||||
)
|
||||
|
||||
litellm_params["upload_url"] = api_base
|
||||
|
|
@ -750,6 +754,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
content: str,
|
||||
api_base: str,
|
||||
optional_params: dict,
|
||||
s3_encryption_key_id: str | None = None,
|
||||
) -> tuple[dict, str]:
|
||||
"""
|
||||
Sign S3 PUT request using the same proven logic as S3Logger.
|
||||
|
|
@ -759,7 +764,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
import hashlib
|
||||
|
||||
import requests
|
||||
from botocore.auth import SigV4Auth
|
||||
from botocore.auth import S3SigV4Auth
|
||||
from botocore.awsrequest import AWSRequest
|
||||
except ImportError:
|
||||
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
|
||||
|
|
@ -782,12 +787,25 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
content_hash: Final = hashlib.sha256(content.encode("utf-8")).hexdigest()
|
||||
|
||||
# Prepare headers with required S3 headers (same as s3_v2.py)
|
||||
request_headers: Final = {
|
||||
"Content-Type": "application/json", # JSONL files are JSON content
|
||||
"x-amz-content-sha256": content_hash, # REQUIRED by S3
|
||||
"Content-Language": "en",
|
||||
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
|
||||
}
|
||||
sse_headers: Final = (
|
||||
MappingProxyType(
|
||||
{
|
||||
"x-amz-server-side-encryption": "aws:kms",
|
||||
"x-amz-server-side-encryption-aws-kms-key-id": s3_encryption_key_id,
|
||||
}
|
||||
)
|
||||
if s3_encryption_key_id
|
||||
else MappingProxyType({})
|
||||
)
|
||||
request_headers: Final = MappingProxyType(
|
||||
{
|
||||
"Content-Type": "application/json", # JSONL files are JSON content
|
||||
"x-amz-content-sha256": content_hash, # REQUIRED by S3
|
||||
"Content-Language": "en",
|
||||
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
|
||||
**sse_headers,
|
||||
}
|
||||
)
|
||||
|
||||
# Use requests.Request to prepare the request (same pattern as s3_v2.py)
|
||||
req: Final = requests.Request("PUT", api_base, data=content, headers=request_headers)
|
||||
|
|
@ -804,7 +822,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
# Get region name for non-LLM API calls (same as s3_v2.py)
|
||||
signing_region: Final = self.get_aws_region_name_for_non_llm_api_calls(aws_region_name=aws_region_name)
|
||||
|
||||
SigV4Auth(credentials, "s3", signing_region).add_auth(aws_request)
|
||||
S3SigV4Auth(credentials, "s3", signing_region).add_auth(aws_request)
|
||||
|
||||
# Return signed headers and body
|
||||
signed_body = aws_request.body
|
||||
|
|
@ -1015,7 +1033,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
try:
|
||||
import hashlib
|
||||
|
||||
from botocore.auth import SigV4Auth
|
||||
from botocore.auth import S3SigV4Auth
|
||||
from botocore.awsrequest import AWSRequest
|
||||
except ImportError:
|
||||
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
|
||||
|
|
@ -1038,7 +1056,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
url=api_base,
|
||||
headers={"x-amz-content-sha256": empty_body_hash},
|
||||
)
|
||||
auth: Final = SigV4Auth(credentials, "s3", aws_region_name) # any-ok: botocore untyped
|
||||
auth: Final = S3SigV4Auth(credentials, "s3", aws_region_name) # any-ok: botocore untyped
|
||||
auth.add_auth(aws_request) # any-ok: botocore request mutation is untyped
|
||||
return dict(aws_request.headers) # any-ok: botocore headers are untyped
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -49,12 +49,12 @@ class BedrockImagePreparedRequest(BaseModel):
|
|||
data: dict
|
||||
|
||||
|
||||
BedrockImageConfigClass = Union[
|
||||
type[AmazonTitanImageGenerationConfig],
|
||||
type[AmazonNovaCanvasConfig],
|
||||
type[AmazonStability3Config],
|
||||
type[AmazonStabilityConfig],
|
||||
]
|
||||
BedrockImageConfigClass = (
|
||||
type[AmazonTitanImageGenerationConfig]
|
||||
| type[AmazonNovaCanvasConfig]
|
||||
| type[AmazonStability3Config]
|
||||
| type[AmazonStabilityConfig]
|
||||
)
|
||||
|
||||
|
||||
class BedrockImageGeneration(BaseAWSLLM):
|
||||
|
|
|
|||
|
|
@ -160,10 +160,10 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
|
|||
aws_filters: dict | None = None
|
||||
|
||||
if isinstance(value, dict):
|
||||
if "operator" in value.keys():
|
||||
if "operator" in value:
|
||||
# Single operator - map directly (no wrapping needed)
|
||||
aws_filters = self._map_operator_filter(value)
|
||||
elif "and" in value.keys() or "or" in value.keys():
|
||||
elif "and" in value or "or" in value:
|
||||
aws_filters = self._map_and_or_filters(value)
|
||||
else:
|
||||
# Assume it's already in AWS KB format
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue