mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Merge remote-tracking branch 'origin/main' into litellm_stream_served_service_tier
Some checks failed
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled
Some checks failed
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled
This commit is contained in:
commit
cc446c60ce
564 changed files with 10547 additions and 3874 deletions
|
|
@ -5,8 +5,10 @@ flag="${1:?usage: unit_selection.sh <codecov flag>}"
|
|||
|
||||
legacy_flags=(
|
||||
caching-local
|
||||
core-utils
|
||||
enterprise-package
|
||||
enterprise-routing
|
||||
integrations
|
||||
llm-other-providers
|
||||
llm-vertex-ai
|
||||
mcp-integration
|
||||
|
|
@ -31,6 +33,7 @@ legacy_flags=(
|
|||
legacy_paths() {
|
||||
case "$1" in
|
||||
caching-local) echo tests/unit/caching ;;
|
||||
core-utils) echo tests/unit/litellm_core_utils ;;
|
||||
enterprise-package)
|
||||
echo tests/unit/enterprise/integrations
|
||||
echo tests/unit/enterprise/proxy/auth
|
||||
|
|
@ -41,6 +44,8 @@ legacy_paths() {
|
|||
echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;;
|
||||
enterprise-routing)
|
||||
echo tests/unit/google_genai
|
||||
echo tests/unit/router_strategy
|
||||
echo tests/unit/router_utils
|
||||
echo tests/unit/enterprise/enterprise_callbacks/send_emails
|
||||
echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py
|
||||
|
|
@ -52,6 +57,7 @@ legacy_paths() {
|
|||
echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py
|
||||
echo tests/unit/enterprise/proxy/test_managed_files_access_check.py
|
||||
echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;;
|
||||
integrations) echo tests/unit/integrations ;;
|
||||
llm-other-providers) find tests/unit/llms -name 'test_*.py' -not -path 'tests/unit/llms/vertex_ai/*' ;;
|
||||
llm-vertex-ai) echo tests/unit/llms/vertex_ai ;;
|
||||
mcp-integration)
|
||||
|
|
@ -75,6 +81,8 @@ legacy_paths() {
|
|||
echo tests/unit/messages
|
||||
echo tests/unit/rag
|
||||
echo tests/unit/rerank_api
|
||||
echo tests/unit/rust_bridge
|
||||
echo tests/unit/secret_managers
|
||||
echo tests/unit/vector_stores
|
||||
echo tests/unit/videos ;;
|
||||
proxy-db-auth-checks)
|
||||
|
|
@ -139,7 +147,9 @@ legacy_paths() {
|
|||
proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;;
|
||||
proxy-extras) echo tests/unit/litellm_proxy_extras ;;
|
||||
proxy-infra) echo tests/unit/gateway ;;
|
||||
responses-caching-types) echo tests/unit/types ;;
|
||||
responses-caching-types)
|
||||
find tests/unit/responses -name 'test_*.py' -not -path 'tests/unit/responses/mcp/*'
|
||||
echo tests/unit/types ;;
|
||||
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
|
||||
esac
|
||||
}
|
||||
|
|
|
|||
|
|
@ -369,6 +369,20 @@ workflows:
|
|||
reruns: 2
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-core-utils
|
||||
flag: core-utils
|
||||
shards: 2
|
||||
reruns: 1
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-integrations
|
||||
flag: integrations
|
||||
shards: 2
|
||||
reruns: 3
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-misc
|
||||
flag: misc
|
||||
|
|
|
|||
8
.github/merge-smoke-tests.json
vendored
8
.github/merge-smoke-tests.json
vendored
|
|
@ -7,9 +7,9 @@
|
|||
"MODEL-DENY": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_denials_return_forbidden[key-key_model_access_denied]",
|
||||
"COST-EXPLICIT": "tests/unit/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
|
||||
"COST-ZERO": "tests/unit/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
|
||||
"LOG-CONTENT-ON": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
|
||||
"LOG-CONTENT-OFF": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
|
||||
"CALLBACK-SUCCESS": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
|
||||
"CALLBACK-FAILURE": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
|
||||
"LOG-CONTENT-ON": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
|
||||
"LOG-CONTENT-OFF": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
|
||||
"CALLBACK-SUCCESS": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
|
||||
"CALLBACK-FAILURE": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
|
||||
}
|
||||
}
|
||||
|
|
|
|||
12
.github/workflows/test-redis-compat.yml
vendored
12
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -12,9 +12,9 @@ on:
|
|||
- "litellm/caching/evicted_client_closer.py"
|
||||
- "tests/unit/test_redis.py"
|
||||
- "tests/local_testing/test_caching.py"
|
||||
- "tests/test_litellm/caching/test_redis_connection_pool.py"
|
||||
- "tests/test_litellm/caching/test_redis_cluster_cache.py"
|
||||
- "tests/test_litellm/caching/test_evicted_client_closer.py"
|
||||
- "tests/unit/caching/test_redis_connection_pool.py"
|
||||
- "tests/unit/caching/test_redis_cluster_cache.py"
|
||||
- "tests/unit/caching/test_evicted_client_closer.py"
|
||||
- ".github/workflows/test-redis-compat.yml"
|
||||
- "pyproject.toml"
|
||||
- "uv.lock"
|
||||
|
|
@ -85,9 +85,9 @@ jobs:
|
|||
redis-server --version
|
||||
uv run --no-sync pytest \
|
||||
tests/unit/test_redis.py \
|
||||
tests/test_litellm/caching/test_redis_connection_pool.py \
|
||||
tests/test_litellm/caching/test_redis_cluster_cache.py \
|
||||
tests/test_litellm/caching/test_evicted_client_closer.py \
|
||||
tests/unit/caching/test_redis_connection_pool.py \
|
||||
tests/unit/caching/test_redis_cluster_cache.py \
|
||||
tests/unit/caching/test_evicted_client_closer.py \
|
||||
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \
|
||||
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \
|
||||
--tb=short -vv \
|
||||
|
|
|
|||
6
.github/workflows/test-rust.yml
vendored
6
.github/workflows/test-rust.yml
vendored
|
|
@ -24,7 +24,7 @@ on:
|
|||
- ".github/actions/setup-uv-with-retries/**"
|
||||
- ".github/scripts/smoke_test_native_wheel.py"
|
||||
- ".github/scripts/verify_linux_native_wheel.py"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
- "tests/unit/rust_bridge/native_route_wheel_test.py"
|
||||
- ".github/workflows/test-rust.yml"
|
||||
pull_request:
|
||||
branches:
|
||||
|
|
@ -52,7 +52,7 @@ on:
|
|||
- ".github/actions/setup-uv-with-retries/**"
|
||||
- ".github/scripts/smoke_test_native_wheel.py"
|
||||
- ".github/scripts/verify_linux_native_wheel.py"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
- "tests/unit/rust_bridge/native_route_wheel_test.py"
|
||||
- ".github/workflows/test-rust.yml"
|
||||
|
||||
permissions:
|
||||
|
|
@ -171,7 +171,7 @@ jobs:
|
|||
env:
|
||||
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
|
||||
|
||||
- run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl
|
||||
- run: python tests/unit/rust_bridge/native_route_wheel_test.py dist/*.whl
|
||||
|
||||
- name: Run pytest tests/test_litellm_rust with the compiled extension
|
||||
run: make test-rust-extension
|
||||
|
|
|
|||
14
.github/workflows/test-unit.yml
vendored
14
.github/workflows/test-unit.yml
vendored
|
|
@ -62,6 +62,7 @@ jobs:
|
|||
- shard: core-utils
|
||||
artifact-name: core-utils
|
||||
test-path: "tests/test_litellm/litellm_core_utils"
|
||||
unit-flag: core-utils
|
||||
workers: 2
|
||||
reruns: 1
|
||||
timeout-minutes: 20
|
||||
|
|
@ -69,9 +70,7 @@ jobs:
|
|||
|
||||
- shard: enterprise-routing
|
||||
artifact-name: enterprise-routing
|
||||
test-path: >-
|
||||
tests/test_litellm/router_utils
|
||||
tests/test_litellm/router_strategy
|
||||
test-path: ""
|
||||
unit-flag: enterprise-routing
|
||||
workers: 2
|
||||
reruns: 2
|
||||
|
|
@ -80,7 +79,8 @@ jobs:
|
|||
|
||||
- shard: integrations
|
||||
artifact-name: integrations
|
||||
test-path: "tests/test_litellm/integrations"
|
||||
test-path: ""
|
||||
unit-flag: integrations
|
||||
workers: 2
|
||||
reruns: 3
|
||||
timeout-minutes: 20
|
||||
|
|
@ -107,11 +107,9 @@ jobs:
|
|||
- shard: misc
|
||||
artifact-name: misc
|
||||
test-path: >-
|
||||
tests/test_litellm/secret_managers
|
||||
tests/test_litellm/interactions
|
||||
tests/test_litellm/ocr
|
||||
tests/test_litellm/passthrough
|
||||
tests/test_litellm/rust_bridge
|
||||
tests/test_litellm/test_*.py
|
||||
unit-flag: misc
|
||||
workers: 2
|
||||
|
|
@ -228,9 +226,7 @@ jobs:
|
|||
|
||||
- shard: responses-caching-types
|
||||
artifact-name: responses-caching-types
|
||||
test-path: >-
|
||||
tests/test_litellm/responses
|
||||
tests/test_litellm/caching
|
||||
test-path: ""
|
||||
unit-flag: responses-caching-types
|
||||
workers: 2
|
||||
reruns: 2
|
||||
|
|
|
|||
8
Makefile
8
Makefile
|
|
@ -301,7 +301,7 @@ test-rust-extension:
|
|||
UV_PROJECT_ENVIRONMENT="$$temporary/venv" $(UV) sync --python 3.12 --frozen --no-install-project --all-groups --all-extras && \
|
||||
$(UV) pip install --python "$$temporary/venv/bin/python" --no-deps "$$1" && \
|
||||
"$$temporary/venv/bin/python" -I -m mypy.stubtest \
|
||||
--mypy-config-file tests/test_litellm/rust_bridge/stubtest.ini \
|
||||
--mypy-config-file tests/unit/rust_bridge/stubtest.ini \
|
||||
litellm.rust_bridge._native && \
|
||||
LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \
|
||||
"$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust
|
||||
|
|
@ -326,13 +326,13 @@ test-unit-proxy-misc: install-test-deps
|
|||
$(UV_RUN) pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/shutdown tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-integrations: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/integrations --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/integrations --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-core-utils: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/litellm_core_utils --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-unit-other: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/caching tests/unit/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/unit/router_strategy tests/unit/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-root: install-test-deps
|
||||
$(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
|
||||
|
||||
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard
|
||||
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/unit/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard
|
||||
|
||||
The first eleven rows need only `callbacks: ["prometheus"]`. The last three rows and the `litellm_admission_*` panels are emitted by other subsystems and stay empty until those are on: the service callback row needs `service_callback: ["prometheus_system"]` in `litellm_settings`, the circuit breaker row needs a Redis cache, the cleanup row needs spend log retention, and admission control needs its middleware enabled. Within the base rows, many panels only fill in once the matching feature is in use: budgets need keys, teams, users or orgs with `max_budget` set, cache panels need caching on, guardrail and MCP panels need those features configured, deployment health needs the router with more than one deployment or a failure to record, and `litellm_in_flight_requests` needs traffic at scrape time. An empty panel for a feature you do not use is expected
|
||||
|
||||
|
|
|
|||
|
|
@ -33,7 +33,7 @@ mod test_support {
|
|||
use crate::{LegacyLogging, LegacySurface, PublicCall};
|
||||
|
||||
/// The parameters of every `callbacks_legacy_python` function, as the real module declares them.
|
||||
/// `tests/test_litellm/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
|
||||
/// `tests/unit/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
|
||||
/// signatures, and [`namespace`] binds every fake call against it.
|
||||
pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json");
|
||||
|
||||
|
|
|
|||
|
|
@ -118,7 +118,7 @@ pub(super) struct ParityCase {
|
|||
pub(super) fn parity_cases() -> Vec<ParityCase> {
|
||||
serde_json::from_str(include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/../../../tests/test_litellm/secret_managers/hashicorp_vault_parity.json"
|
||||
"/../../../tests/unit/secret_managers/hashicorp_vault_parity.json"
|
||||
)))
|
||||
.unwrap()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -32,7 +32,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
|
||||
`_SecretManagerRuntime` is a private implementation detail, not a replacement SDK class. Its async methods return Futures; public `async def` methods retain lazy coroutine creation and `asyncio.create_task` support. Passing the same names and arguments is insufficient to claim parity until the remaining return-value, error, cache and configuration differences above are closed
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py)
|
||||
## [tests/unit/secret_managers/test_aws_secret_manager_replication.py](../../../tests/unit/secret_managers/test_aws_secret_manager_replication.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -48,7 +48,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_replicate_secret_http_error_raises` | [direct_replication_returns_response_or_service_error](../secrets-aws/tests/secret_manager/writes.rs) |
|
||||
| `test_replicate_secret_timeout_raises` | [write_and_replication_timeouts_remain_errors](../secrets-aws/tests/secret_manager/writes.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py)
|
||||
## [tests/unit/secret_managers/test_aws_secret_manager_rotation.py](../../../tests/unit/secret_managers/test_aws_secret_manager_rotation.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -59,7 +59,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_write_secret_to_name_inside_recovery_window_restores_and_stores_new_value` | [recovery_window_alias_is_restored_updated_and_tagged](../secrets-aws/tests/secret_manager/writes.rs) |
|
||||
| `test_write_secret_to_live_existing_name_still_fails_without_overwriting` | [create_failure_does_not_overwrite_an_alias_without_a_deletion_date](../secrets-aws/tests/secret_manager/writes.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py)
|
||||
## [tests/unit/secret_managers/test_aws_secret_manager_v2.py](../../../tests/unit/secret_managers/test_aws_secret_manager_v2.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -70,14 +70,14 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_prepare_request_explicit_bedrock_runtime_endpoint_param_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) |
|
||||
| `test_prepare_request_env_bedrock_runtime_endpoint_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_base_secret_manager.py](../../../tests/test_litellm/secret_managers/test_base_secret_manager.py)
|
||||
## [tests/unit/secret_managers/test_base_secret_manager.py](../../../tests/unit/secret_managers/test_base_secret_manager.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
| `test_raise_if_unsafe_secret_name_rejects_traversal_and_line_breaks` | [names_reject_path_traversal_and_control_characters](../secrets-types/tests/rotation.rs) |
|
||||
| `test_raise_if_unsafe_secret_name_allows_legitimate_aliases` | [names_allow_safe_values](../secrets-types/tests/rotation.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_custom_secret_manager.py](../../../tests/test_litellm/secret_managers/test_custom_secret_manager.py)
|
||||
## [tests/unit/secret_managers/test_custom_secret_manager.py](../../../tests/unit/secret_managers/test_custom_secret_manager.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -89,7 +89,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_custom_secret_manager_integration_with_litellm` | [manager_strings_are_coerced_like_literal_eval](../secrets/tests/resolution.rs) |
|
||||
| `test_minimal_custom_secret_manager` | Exercises the Python example subclass itself or Python default methods. Caller-authored Python implementations remain Python callbacks; resolver integration is covered by `manager_strings_are_coerced_like_literal_eval` |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_cyberark_secret_manager.py](../../../tests/test_litellm/secret_managers/test_cyberark_secret_manager.py)
|
||||
## [tests/unit/secret_managers/test_cyberark_secret_manager.py](../../../tests/unit/secret_managers/test_cyberark_secret_manager.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -97,7 +97,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_async_write_matches_parity_fixture` | [writes_match_python_parity_fixture](../secrets-cyberark/tests/secret_manager/writes.rs) |
|
||||
| `test_missing_credentials_raise_value_error` | [new_validates_credentials_before_license_and_configuration](../secrets-cyberark/tests/secret_manager/configuration.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py](../../../tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py)
|
||||
## [tests/unit/secret_managers/test_get_azure_ad_token_provider.py](../../../tests/unit/secret_managers/test_get_azure_ad_token_provider.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -115,7 +115,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_get_azure_ad_token_provider_prefers_workload_identity_over_managed_identity` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate |
|
||||
| `test_get_azure_ad_token_provider_defaults_to_default_azure_credential` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py](../../../tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py)
|
||||
## [tests/unit/secret_managers/test_hashicorp_secret_manager.py](../../../tests/unit/secret_managers/test_hashicorp_secret_manager.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -130,13 +130,13 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_tls_login_uses_login_namespace` | [tls_login_posts_the_role_and_uses_the_client_identity](../secrets-hashicorp/tests/secret_manager/configuration.rs) |
|
||||
| `test_configuration_matches_native_parity_fixture` | [configuration_matches_python_parity_fixture](../secrets-hashicorp/tests/secret_manager/configuration.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_secret_manager_handler.py](../../../tests/test_litellm/secret_managers/test_secret_manager_handler.py)
|
||||
## [tests/unit/secret_managers/test_secret_manager_handler.py](../../../tests/unit/secret_managers/test_secret_manager_handler.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
| `test_azure_key_vault_matches_rust_parity_fixture` | [parity_fixture_matches_python_backend_contract](../secrets-azure/tests/key_vault.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_secret_managers_main.py](../../../tests/test_litellm/secret_managers/test_secret_managers_main.py)
|
||||
## [tests/unit/secret_managers/test_secret_managers_main.py](../../../tests/unit/secret_managers/test_secret_managers_main.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ The HashiCorp Vault backend is enabled with the `hashicorp` feature and reads KV
|
|||
|
||||
Native backends consistently distinguish absence from failure instead of swallowing provider errors. Python-compatible resolution maps these results back to the Python handler contract before applying fallback
|
||||
|
||||
`hosted_keys` excludes a name for every backend. Python's handler recognizes Azure `SecretClient` and Google `KeyManagementServiceClient` instances before the `local` branch, allowing excluded names to reach those providers. Rust treats that as a routing bug. `test_rust_hosted_keys_exclude_azure_sdk_clients_too` in `tests/test_litellm/rust_bridge/ocr/test_secrets.py` pins this behavior
|
||||
`hosted_keys` excludes a name for every backend. Python's handler recognizes Azure `SecretClient` and Google `KeyManagementServiceClient` instances before the `local` branch, allowing excluded names to reach those providers. Rust treats that as a routing bug. `test_rust_hosted_keys_exclude_azure_sdk_clients_too` in `tests/unit/rust_bridge/ocr/test_secrets.py` pins this behavior
|
||||
|
||||
Google rejects malformed base64 and mismatched CRC32C values instead of accepting corrupted payloads. Python currently ignores the checksum and uses permissive base64 decoding. Rust follows [RFC 4648](https://www.rfc-editor.org/rfc/rfc4648#section-3.3) and [Google's integrity guidance](https://docs.cloud.google.com/secret-manager/docs/data-integrity); `failed_or_missing_reads_are_not_cached` covers rejection and recovery
|
||||
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ from litellm.types.integrations.datadog import DatadogInitParams
|
|||
from litellm.types.integrations.newrelic import NewRelicInitParams
|
||||
from litellm.litellm_core_utils.core_helpers import drop_params_env_flag
|
||||
from litellm.types.integrations.pointfive import PointFiveInitParams
|
||||
from litellm.types.integrations.zerobus import ZerobusInitParams
|
||||
from litellm._logging import (
|
||||
set_verbose,
|
||||
_turn_on_debug,
|
||||
|
|
@ -157,6 +158,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
|
|||
"deepeval",
|
||||
"s3_v2",
|
||||
"pointfive",
|
||||
"zerobus",
|
||||
"aws_sqs",
|
||||
"vector_store_pre_call_hook",
|
||||
"dotprompt",
|
||||
|
|
@ -442,6 +444,7 @@ datadog_llm_observability_params: Optional[Union[DatadogLLMObsInitParams, Dict]]
|
|||
datadog_params: Optional[Union[DatadogInitParams, Dict]] = None
|
||||
newrelic_params: Optional[Union[NewRelicInitParams, Dict]] = None
|
||||
pointfive_params: Optional[Union[PointFiveInitParams, Mapping[str, object]]] = None
|
||||
zerobus_params: Optional[Union[ZerobusInitParams, Mapping[str, object]]] = None
|
||||
aws_sqs_callback_params: Optional[Dict] = None
|
||||
generic_logger_headers: Optional[Dict] = None
|
||||
default_key_generate_params: Optional[Dict] = None
|
||||
|
|
|
|||
|
|
@ -2,7 +2,9 @@
|
|||
|
||||
from .exception_mapping_utils import (
|
||||
ANTHROPIC_ERROR_TYPE_MAP,
|
||||
AnthropicErrorSseFrame,
|
||||
AnthropicExceptionMapping,
|
||||
anthropic_error_sse_frame,
|
||||
)
|
||||
from .exceptions import (
|
||||
AnthropicErrorDetail,
|
||||
|
|
@ -14,6 +16,8 @@ __all__ = [
|
|||
"ANTHROPIC_ERROR_TYPE_MAP",
|
||||
"AnthropicErrorDetail",
|
||||
"AnthropicErrorResponse",
|
||||
"AnthropicErrorSseFrame",
|
||||
"AnthropicErrorType",
|
||||
"AnthropicExceptionMapping",
|
||||
"anthropic_error_sse_frame",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -4,11 +4,12 @@ Utilities for mapping exceptions to Anthropic error format.
|
|||
Similar to litellm/litellm_core_utils/exception_mapping_utils.py but for Anthropic response format.
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Final
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
|
||||
from .exceptions import AnthropicErrorResponse, AnthropicErrorType
|
||||
from .exceptions import AnthropicErrorDetail, AnthropicErrorResponse, AnthropicErrorType
|
||||
|
||||
# HTTP status code -> Anthropic error type
|
||||
# Source: https://docs.anthropic.com/en/api/errors
|
||||
|
|
@ -166,3 +167,36 @@ class AnthropicExceptionMapping:
|
|||
message=message,
|
||||
request_id=request_id,
|
||||
)
|
||||
|
||||
|
||||
class AnthropicErrorSseFrame(str):
|
||||
"""One `event: error` frame, for a stream that fails once the response headers are out.
|
||||
|
||||
Anthropic clients pick stream events by the `event:` name, so a frame carrying only a `data:`
|
||||
line is skipped and the failure never reaches the caller. The frame remembers the status and
|
||||
body it was built from, so a stream that fails before its first byte can still answer as a
|
||||
JSON error with that exact status instead of a 200 that only says `api_error`
|
||||
"""
|
||||
|
||||
status_code: int
|
||||
error_response: AnthropicErrorResponse
|
||||
|
||||
def __new__(cls, status_code: int, error_response: AnthropicErrorResponse) -> "AnthropicErrorSseFrame":
|
||||
frame: Final = super().__new__(cls, f"event: error\ndata: {json.dumps(error_response)}\n\n")
|
||||
frame.status_code = status_code
|
||||
frame.error_response = error_response
|
||||
return frame
|
||||
|
||||
def json_body(self, call_id: str | None) -> AnthropicErrorResponse:
|
||||
if call_id is None:
|
||||
return self.error_response
|
||||
detail: Final[AnthropicErrorDetail] = {**self.error_response["error"], "litellm_call_id": call_id}
|
||||
body: Final[AnthropicErrorResponse] = {**self.error_response, "error": detail}
|
||||
return body
|
||||
|
||||
|
||||
def anthropic_error_sse_frame(status_code: int, raw_message: str) -> AnthropicErrorSseFrame:
|
||||
return AnthropicErrorSseFrame(
|
||||
status_code,
|
||||
AnthropicExceptionMapping.transform_to_anthropic_error(status_code=status_code, raw_message=raw_message),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -706,6 +706,10 @@ def _get_batch_job_usage_from_response_body(
|
|||
if ResponseAPILoggingUtils._is_response_api_usage(_usage_dict):
|
||||
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(_usage_dict)
|
||||
usage: Final[Usage] = Usage(**_usage_dict)
|
||||
if custom_llm_provider == "xai":
|
||||
from litellm.llms.xai.chat.transformation import XAIChatConfig
|
||||
|
||||
XAIChatConfig.fold_reasoning_tokens_into_completion(usage)
|
||||
return usage
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -31,6 +31,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
|||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.openai.openai import OpenAIBatchesAPI
|
||||
from litellm.llms.vertex_ai.batches.handler import VertexAIBatchPrediction
|
||||
from litellm.llms.xai.batches.handler import XAIBatchesHandler
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import (
|
||||
CancelBatchRequest,
|
||||
|
|
@ -59,6 +60,7 @@ openai_batches_instance: Final = OpenAIBatchesAPI()
|
|||
azure_batches_instance: Final = AzureBatchesAPI()
|
||||
vertex_ai_batches_instance: Final = VertexAIBatchPrediction(gcs_bucket_name="")
|
||||
anthropic_batches_instance: Final = AnthropicBatchesHandler()
|
||||
xai_batches_instance: Final = XAIBatchesHandler()
|
||||
base_llm_http_handler = BaseLLMHTTPHandler()
|
||||
#################################################
|
||||
|
||||
|
|
@ -105,10 +107,22 @@ def _resolve_timeout(
|
|||
@client
|
||||
async def acreate_batch(
|
||||
completion_window: Literal["24h"],
|
||||
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"],
|
||||
endpoint: Literal[
|
||||
"/v1/chat/completions",
|
||||
"/v1/embeddings",
|
||||
"/v1/completions",
|
||||
"/v1/responses",
|
||||
"/v1/ocr",
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits",
|
||||
"/v1/videos/generations",
|
||||
"/v1/videos",
|
||||
"/v1/videos/edits",
|
||||
"/v1/videos/extensions",
|
||||
],
|
||||
input_file_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral"
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral", "xai"
|
||||
] = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -157,10 +171,22 @@ async def acreate_batch(
|
|||
@client
|
||||
def create_batch(
|
||||
completion_window: Literal["24h"],
|
||||
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"],
|
||||
endpoint: Literal[
|
||||
"/v1/chat/completions",
|
||||
"/v1/embeddings",
|
||||
"/v1/completions",
|
||||
"/v1/responses",
|
||||
"/v1/ocr",
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits",
|
||||
"/v1/videos/generations",
|
||||
"/v1/videos",
|
||||
"/v1/videos/edits",
|
||||
"/v1/videos/extensions",
|
||||
],
|
||||
input_file_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral"
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral", "xai"
|
||||
] = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -243,6 +269,14 @@ def create_batch(
|
|||
model=model,
|
||||
)
|
||||
return response
|
||||
if custom_llm_provider == LlmProviders.XAI.value:
|
||||
return xai_batches_instance.create_batch(
|
||||
_is_async=_is_async,
|
||||
create_batch_data=_create_batch_request,
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
timeout=timeout,
|
||||
)
|
||||
api_base: str | None = None
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
# for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there
|
||||
|
|
@ -345,7 +379,7 @@ def create_batch(
|
|||
async def aretrieve_batch(
|
||||
batch_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral"
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral", "xai"
|
||||
] = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -393,10 +427,18 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
_retrieve_batch_request: RetrieveBatchRequest,
|
||||
_is_async: bool,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral"
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral", "xai"
|
||||
] = "openai",
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
):
|
||||
if custom_llm_provider == LlmProviders.XAI.value:
|
||||
return xai_batches_instance.retrieve_batch(
|
||||
_is_async=_is_async,
|
||||
batch_id=batch_id,
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
timeout=timeout,
|
||||
)
|
||||
api_base: str | None = None
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
# for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there
|
||||
|
|
@ -518,7 +560,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
def retrieve_batch(
|
||||
batch_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral"
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral", "xai"
|
||||
] = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -741,6 +783,15 @@ def list_batches(
|
|||
timeout = 600.0
|
||||
|
||||
_is_async: Final = kwargs.pop("alist_batches", False) is True
|
||||
if custom_llm_provider == LlmProviders.XAI.value:
|
||||
return xai_batches_instance.list_batches(
|
||||
_is_async=_is_async,
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
timeout=timeout,
|
||||
after=after,
|
||||
limit=limit,
|
||||
)
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
# for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there
|
||||
api_base = (
|
||||
|
|
@ -837,7 +888,7 @@ def list_batches(
|
|||
async def acancel_batch(
|
||||
batch_id: str,
|
||||
model: str | None = None,
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy"] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy", "xai"] = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
extra_body: dict[str, str] | None = None,
|
||||
|
|
@ -883,7 +934,7 @@ async def acancel_batch(
|
|||
def cancel_batch(
|
||||
batch_id: str,
|
||||
model: str | None = None,
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy"] | str = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy", "xai"] | str = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
extra_body: dict[str, str] | None = None,
|
||||
|
|
@ -933,6 +984,14 @@ def cancel_batch(
|
|||
)
|
||||
|
||||
_is_async: Final = kwargs.pop("acancel_batch", False) is True
|
||||
if custom_llm_provider == LlmProviders.XAI.value:
|
||||
return xai_batches_instance.cancel_batch(
|
||||
_is_async=_is_async,
|
||||
batch_id=batch_id,
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
timeout=timeout,
|
||||
)
|
||||
api_base: str | None = None
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
api_base = (
|
||||
|
|
|
|||
|
|
@ -1952,6 +1952,15 @@ SENTRY_DENYLIST: Final = [
|
|||
"auth_token",
|
||||
"jwt_token",
|
||||
"private_key",
|
||||
"authorization",
|
||||
"api-key",
|
||||
"x-api-key",
|
||||
"x-goog-api-key",
|
||||
"ocp-apim-subscription-key",
|
||||
"x-litellm-api-key",
|
||||
"x-mcp-auth",
|
||||
"cookie",
|
||||
"set-cookie",
|
||||
"SLACK_WEBHOOK_URL",
|
||||
"ALERTING_WEBHOOK_URL",
|
||||
"webhook_url",
|
||||
|
|
@ -1974,6 +1983,12 @@ SENTRY_DENYLIST: Final = [
|
|||
]
|
||||
SENTRY_PII_DENYLIST: Final = [
|
||||
"user_id",
|
||||
"user_email",
|
||||
"end_user_id",
|
||||
"user_api_key_hash",
|
||||
"user_api_key_user_id",
|
||||
"user_api_key_user_email",
|
||||
"user_api_key_end_user_id",
|
||||
"email",
|
||||
"phone",
|
||||
"address",
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import os
|
|||
from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Sequence
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
from functools import partial
|
||||
from importlib.metadata import version
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeAlias, TypeVar, cast
|
||||
|
||||
|
|
@ -34,8 +35,16 @@ _TransportContext: TypeAlias = AbstractAsyncContextManager[_TransportStreams]
|
|||
from mcp.types import (
|
||||
METHOD_NOT_FOUND,
|
||||
REQUEST_TIMEOUT,
|
||||
ClientCapabilities,
|
||||
ElicitationCapability,
|
||||
FormElicitationCapability,
|
||||
GetPromptRequestParams,
|
||||
GetPromptResult,
|
||||
Implementation,
|
||||
InitializedNotification,
|
||||
InitializeRequest,
|
||||
InitializeRequestParams,
|
||||
InitializeResult,
|
||||
InputRequiredResult,
|
||||
ListPromptsResult,
|
||||
ListResourcesResult,
|
||||
|
|
@ -44,12 +53,14 @@ from mcp.types import (
|
|||
PaginatedResult,
|
||||
Prompt,
|
||||
ResourceTemplate,
|
||||
SamplingCapability,
|
||||
ServerNotification,
|
||||
UrlElicitationCapability,
|
||||
)
|
||||
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
|
||||
from mcp.types import CallToolResult as MCPCallToolResult
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl
|
||||
from pydantic import AnyUrl, TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import (
|
||||
|
|
@ -64,11 +75,13 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_er
|
|||
from litellm.proxy._experimental.mcp_server.result_conversion import error_text_result
|
||||
from litellm.types.llms.custom_http import VerifyTypes
|
||||
from litellm.types.mcp import (
|
||||
MCP_LEGACY_VERSIONS,
|
||||
MCPAuth,
|
||||
MCPAuthType,
|
||||
MCPStdioConfig,
|
||||
MCPTransport,
|
||||
MCPTransportType,
|
||||
MCPUpstreamProtocol,
|
||||
credential_redirect_hook,
|
||||
has_header,
|
||||
without_header,
|
||||
|
|
@ -386,7 +399,9 @@ class MCPClient:
|
|||
sampling_callback: Callable | None = None,
|
||||
elicitation_callback: Callable | None = None,
|
||||
logging_callback: Callable | None = None,
|
||||
protocol_version: MCPUpstreamProtocol = "auto",
|
||||
):
|
||||
self.protocol_version: MCPUpstreamProtocol = TypeAdapter(MCPUpstreamProtocol).validate_python(protocol_version)
|
||||
self.server_url: str = server_url
|
||||
self.transport_type: MCPTransport = transport_type
|
||||
self.auth_type: MCPAuthType = auth_type
|
||||
|
|
@ -525,6 +540,35 @@ class MCPClient:
|
|||
|
||||
return safe_env
|
||||
|
||||
async def _initialize_session(self, session: ClientSession) -> InitializeResult:
|
||||
if self.protocol_version == "auto":
|
||||
automatic: Final = await session.initialize()
|
||||
if automatic.protocol_version not in MCP_LEGACY_VERSIONS:
|
||||
raise MCPError(code=-32022, message="Upstream selected an unsupported MCP protocol version")
|
||||
return automatic
|
||||
result: Final = await session.send_request(
|
||||
InitializeRequest(
|
||||
params=InitializeRequestParams(
|
||||
protocol_version=self.protocol_version,
|
||||
client_info=Implementation(name="litellm", version=version("litellm")),
|
||||
capabilities=ClientCapabilities(
|
||||
sampling=SamplingCapability() if self._sampling_callback is not None else None,
|
||||
elicitation=ElicitationCapability(
|
||||
form=FormElicitationCapability(), url=UrlElicitationCapability()
|
||||
)
|
||||
if self._elicitation_callback is not None
|
||||
else None,
|
||||
),
|
||||
)
|
||||
),
|
||||
InitializeResult,
|
||||
)
|
||||
if result.protocol_version != self.protocol_version:
|
||||
raise MCPError(code=-32022, message="Upstream did not accept the configured MCP protocol version")
|
||||
session.adopt(result)
|
||||
await session.send_notification(InitializedNotification())
|
||||
return result
|
||||
|
||||
async def _execute_session_operation(
|
||||
self,
|
||||
transport_ctx: _TransportContext,
|
||||
|
|
@ -579,7 +623,7 @@ class MCPClient:
|
|||
)
|
||||
session: Final = await session_ctx.__aenter__()
|
||||
try:
|
||||
init_result: Final = await session.initialize()
|
||||
init_result: Final = await self._initialize_session(session)
|
||||
instructions: Final = getattr(init_result, "instructions", None)
|
||||
self._last_initialize_instructions = (
|
||||
instructions.strip() or None if isinstance(instructions, str) else None
|
||||
|
|
|
|||
|
|
@ -28,12 +28,15 @@ FileCreateProvider = Literal[
|
|||
"manus",
|
||||
"anthropic",
|
||||
"mistral",
|
||||
"xai",
|
||||
]
|
||||
FileRetrieveProvider = Literal[
|
||||
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic", "mistral"
|
||||
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic", "mistral", "xai"
|
||||
]
|
||||
FileDeleteProvider = Literal["openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic", "mistral"]
|
||||
FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic", "mistral"]
|
||||
FileDeleteProvider = Literal[
|
||||
"openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic", "mistral", "xai"
|
||||
]
|
||||
FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic", "mistral", "xai"]
|
||||
import litellm
|
||||
from litellm import get_secret_str
|
||||
from litellm.files.streaming import FileContentStreamingResponse
|
||||
|
|
@ -49,6 +52,8 @@ from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
|||
from litellm.llms.openai.common_utils import get_openai_credentials
|
||||
from litellm.llms.openai.openai import FileDeleted, FileObject, OpenAIFilesAPI
|
||||
from litellm.llms.vertex_ai.files.handler import VertexAIFilesHandler
|
||||
from litellm.llms.xai.batches.handler import XAIBatchesHandler
|
||||
from litellm.llms.xai.batches.transformation import is_xai_batch_results_id
|
||||
from litellm.types.llms.openai import (
|
||||
CreateFileRequest,
|
||||
FileContentRequest,
|
||||
|
|
@ -103,6 +108,7 @@ openai_files_instance: Final = OpenAIFilesAPI()
|
|||
azure_files_instance: Final = AzureOpenAIFilesAPI()
|
||||
vertex_ai_files_instance: Final = VertexAIFilesHandler()
|
||||
bedrock_files_instance: Final = BedrockFilesHandler()
|
||||
xai_batch_results_instance: Final = XAIBatchesHandler()
|
||||
#################################################
|
||||
|
||||
|
||||
|
|
@ -920,6 +926,15 @@ def file_content(
|
|||
client=client,
|
||||
)
|
||||
|
||||
if custom_llm_provider == LlmProviders.XAI.value and is_xai_batch_results_id(file_id):
|
||||
return xai_batch_results_instance.batch_results_content(
|
||||
_is_async=_is_async,
|
||||
batch_id=file_id,
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
# Check if provider has a custom files config (e.g., Anthropic, Manus)
|
||||
provider_config: Final = ProviderConfigManager.get_provider_files_config(
|
||||
model="",
|
||||
|
|
|
|||
|
|
@ -406,6 +406,45 @@
|
|||
},
|
||||
"description": "PointFive Logging Integration"
|
||||
},
|
||||
{
|
||||
"id": "zerobus",
|
||||
"displayName": "Databricks Zerobus",
|
||||
"logo": "databricks.svg",
|
||||
"supports_key_team_logging": false,
|
||||
"dynamic_params": {
|
||||
"ZEROBUS_WORKSPACE_URL": {
|
||||
"type": "text",
|
||||
"ui_name": "Workspace URL",
|
||||
"description": "Databricks workspace URL, e.g. https://dbc-a1b2c3d4-e5f6.cloud.databricks.com",
|
||||
"required": true
|
||||
},
|
||||
"ZEROBUS_SERVER_ENDPOINT": {
|
||||
"type": "text",
|
||||
"ui_name": "Zerobus Endpoint",
|
||||
"description": "Zerobus ingest endpoint, e.g. https://<workspace-id>.zerobus.<region>.cloud.databricks.com",
|
||||
"required": true
|
||||
},
|
||||
"ZEROBUS_CLIENT_ID": {
|
||||
"type": "text",
|
||||
"ui_name": "Service Principal Client ID",
|
||||
"description": "OAuth client id of a service principal with USE CATALOG, USE SCHEMA, SELECT and MODIFY on the table",
|
||||
"required": true
|
||||
},
|
||||
"ZEROBUS_CLIENT_SECRET": {
|
||||
"type": "password",
|
||||
"ui_name": "Service Principal Client Secret",
|
||||
"description": "OAuth client secret of the service principal",
|
||||
"required": true
|
||||
},
|
||||
"ZEROBUS_TABLE_NAME": {
|
||||
"type": "text",
|
||||
"ui_name": "Table",
|
||||
"description": "Fully qualified Unity Catalog table, catalog.schema.table, created with the LiteLLM trace schema",
|
||||
"required": true
|
||||
}
|
||||
},
|
||||
"description": "Databricks Zerobus Ingest Logging Integration"
|
||||
},
|
||||
{
|
||||
"id": "s3",
|
||||
"displayName": "S3",
|
||||
|
|
|
|||
|
|
@ -92,7 +92,7 @@ litellm/integrations/levo/
|
|||
|
||||
## Testing
|
||||
|
||||
See the test files in `tests/test_litellm/integrations/levo/`:
|
||||
See the test files in `tests/unit/integrations/levo/`:
|
||||
- `test_levo.py`: Unit tests for configuration
|
||||
- `test_levo_integration.py`: Integration tests for callback registration
|
||||
|
||||
|
|
|
|||
5
litellm/integrations/zerobus/__init__.py
Normal file
5
litellm/integrations/zerobus/__init__.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
"""Databricks Zerobus logging integration for LiteLLM."""
|
||||
|
||||
from litellm.integrations.zerobus.logger import ZerobusLogger
|
||||
|
||||
__all__ = ("ZerobusLogger",)
|
||||
161
litellm/integrations/zerobus/client.py
Normal file
161
litellm/integrations/zerobus/client.py
Normal file
|
|
@ -0,0 +1,161 @@
|
|||
"""
|
||||
Writes rows to a Unity Catalog table through the Zerobus Ingest REST API.
|
||||
|
||||
Zerobus only accepts a Databricks OAuth token minted for its own resource and scoped to
|
||||
the target table's privileges, so the client mints that token itself with the service
|
||||
principal's client credentials and reuses it until shortly before it expires.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import time
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.types.integrations.zerobus import (
|
||||
RETRYABLE_INGEST_STATUS_CODES,
|
||||
TOKEN_REFRESH_LEEWAY_SECONDS,
|
||||
ZerobusAccessToken,
|
||||
ZerobusConnection,
|
||||
ZerobusIngestFailure,
|
||||
)
|
||||
|
||||
TOKEN_PATH: Final = "/oidc/v1/token"
|
||||
OAUTH_SCOPE: Final = "all-apis"
|
||||
|
||||
|
||||
class _TokenResponse(BaseModel):
|
||||
access_token: str
|
||||
expires_in: float = 3600
|
||||
|
||||
|
||||
class ZerobusIngestError(Exception):
|
||||
"""A batch could not be written and the failure is worth retrying."""
|
||||
|
||||
|
||||
def zerobus_resource(workspace_id: str) -> str:
|
||||
return f"api://databricks/workspaces/{workspace_id}/zerobusDirectWriteApi"
|
||||
|
||||
|
||||
def authorization_details(table_name: str) -> str:
|
||||
"""The Unity Catalog privileges Zerobus requires the token to carry, as the token endpoint expects them."""
|
||||
catalog, schema, _table = table_name.split(".", 2)
|
||||
return json.dumps(
|
||||
(
|
||||
{
|
||||
"type": "unity_catalog_privileges",
|
||||
"privileges": ("USE CATALOG",),
|
||||
"object_type": "CATALOG",
|
||||
"object_full_path": catalog,
|
||||
},
|
||||
{
|
||||
"type": "unity_catalog_privileges",
|
||||
"privileges": ("USE SCHEMA",),
|
||||
"object_type": "SCHEMA",
|
||||
"object_full_path": f"{catalog}.{schema}",
|
||||
},
|
||||
{
|
||||
"type": "unity_catalog_privileges",
|
||||
"privileges": ("SELECT", "MODIFY"),
|
||||
"object_type": "TABLE",
|
||||
"object_full_path": table_name,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def insert_url(connection: ZerobusConnection) -> str:
|
||||
return f"{connection.server_endpoint.rstrip('/')}/zerobus/v1/tables/{connection.table_name}/insert"
|
||||
|
||||
|
||||
def token_url(connection: ZerobusConnection) -> str:
|
||||
return f"{connection.workspace_url.rstrip('/')}{TOKEN_PATH}"
|
||||
|
||||
|
||||
def _basic_auth(client_id: str, client_secret: str) -> str:
|
||||
return "Basic " + base64.b64encode(f"{client_id}:{client_secret}".encode()).decode()
|
||||
|
||||
|
||||
def _status_failure(what: str, error: httpx.HTTPStatusError) -> ZerobusIngestFailure:
|
||||
status: Final = error.response.status_code
|
||||
return ZerobusIngestFailure(
|
||||
detail=f"{what} returned {status}: {error.response.text}"[:500],
|
||||
retryable=status in RETRYABLE_INGEST_STATUS_CODES,
|
||||
)
|
||||
|
||||
|
||||
class ZerobusIngestClient:
|
||||
def __init__(
|
||||
self,
|
||||
connection: ZerobusConnection,
|
||||
http_client: AsyncHTTPHandler,
|
||||
clock: Callable[[], float] = time.time,
|
||||
) -> None:
|
||||
self.connection: Final = connection
|
||||
self.http_client: Final = http_client
|
||||
self.clock: Final = clock
|
||||
self._token: ZerobusAccessToken | None = None
|
||||
self._token_lock: Final = asyncio.Lock()
|
||||
|
||||
async def insert(self, rows: Sequence[Mapping[str, object]]) -> ZerobusIngestFailure | None:
|
||||
"""Write ``rows`` as one request. ``None`` means Zerobus accepted every row."""
|
||||
token: Final = await self.access_token()
|
||||
if isinstance(token, ZerobusIngestFailure):
|
||||
return token
|
||||
try:
|
||||
await self.http_client.post(
|
||||
insert_url(self.connection),
|
||||
content=json.dumps([dict(row) for row in rows]).encode(),
|
||||
headers={"Content-Type": "application/json", "Authorization": f"Bearer {token.value}"},
|
||||
)
|
||||
except httpx.HTTPStatusError as error:
|
||||
if error.response.status_code == 401:
|
||||
self._token = None
|
||||
return ZerobusIngestFailure(detail="insert returned 401, token discarded", retryable=True)
|
||||
return _status_failure("insert", error)
|
||||
except (httpx.HTTPError, litellm.Timeout) as error:
|
||||
return ZerobusIngestFailure(detail=f"insert failed: {error}", retryable=True)
|
||||
return None
|
||||
|
||||
async def access_token(self) -> ZerobusAccessToken | ZerobusIngestFailure:
|
||||
"""The cached token while it has more than the leeway left, otherwise a fresh one."""
|
||||
async with self._token_lock:
|
||||
cached: Final = self._token
|
||||
if cached is not None and cached.expires_at - self.clock() > TOKEN_REFRESH_LEEWAY_SECONDS:
|
||||
return cached
|
||||
minted: Final = await self._mint_token()
|
||||
if isinstance(minted, ZerobusAccessToken):
|
||||
self._token = minted
|
||||
return minted
|
||||
|
||||
async def _mint_token(self) -> ZerobusAccessToken | ZerobusIngestFailure:
|
||||
connection: Final = self.connection
|
||||
try:
|
||||
response: Final = await self.http_client.post(
|
||||
token_url(connection),
|
||||
data={
|
||||
"grant_type": "client_credentials",
|
||||
"scope": OAUTH_SCOPE,
|
||||
"resource": zerobus_resource(connection.workspace_id),
|
||||
"authorization_details": authorization_details(connection.table_name),
|
||||
},
|
||||
headers={
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Authorization": _basic_auth(connection.client_id, connection.client_secret),
|
||||
},
|
||||
)
|
||||
except httpx.HTTPStatusError as error:
|
||||
return _status_failure("token request", error)
|
||||
except (httpx.HTTPError, litellm.Timeout) as error:
|
||||
return ZerobusIngestFailure(detail=f"token request failed: {error}", retryable=True)
|
||||
try:
|
||||
parsed: Final = _TokenResponse.model_validate_json(response.text)
|
||||
except ValidationError as error:
|
||||
return ZerobusIngestFailure(detail=f"token response was not understood: {error}", retryable=False)
|
||||
return ZerobusAccessToken(value=parsed.access_token, expires_at=self.clock() + parsed.expires_in)
|
||||
230
litellm/integrations/zerobus/logger.py
Normal file
230
litellm/integrations/zerobus/logger.py
Normal file
|
|
@ -0,0 +1,230 @@
|
|||
"""Databricks Zerobus logging integration."""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
from litellm.integrations.zerobus.client import ZerobusIngestClient, ZerobusIngestError
|
||||
from litellm.integrations.zerobus.row import trace_row
|
||||
from litellm.litellm_core_utils.redact_messages import (
|
||||
redacted_standard_logging_payload,
|
||||
should_redact_message_logging,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client, httpxSpecialProvider
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.integrations.zerobus import ZerobusConnection, ZerobusInitParams
|
||||
|
||||
_ENV_REFERENCE_PREFIX: Final = "os.environ/"
|
||||
|
||||
|
||||
def _resolved_secret(value: str | None) -> str | None:
|
||||
"""Resolve a config value that may name a secret; an unset ``os.environ/NAME`` stays unresolved."""
|
||||
if value is None:
|
||||
return None
|
||||
resolved: Final = get_secret_str(value)
|
||||
if resolved:
|
||||
return resolved
|
||||
return None if value.startswith(_ENV_REFERENCE_PREFIX) else value
|
||||
|
||||
|
||||
def _configured_params() -> ZerobusInitParams:
|
||||
configured: Final = litellm.zerobus_params
|
||||
if isinstance(configured, ZerobusInitParams):
|
||||
return configured
|
||||
if isinstance(configured, Mapping):
|
||||
return ZerobusInitParams.model_validate(configured)
|
||||
return ZerobusInitParams()
|
||||
|
||||
|
||||
def _setting(configured: str | None, env_var: str) -> str:
|
||||
"""Prefer the configured value, falling back to the environment the proxy UI writes."""
|
||||
value: Final = _resolved_secret(configured) or get_secret_str(env_var)
|
||||
if not value:
|
||||
raise ValueError(
|
||||
f"zerobus logging requires {env_var}. Set it in the environment, or "
|
||||
f"litellm_settings.zerobus_params.{env_var.removeprefix('ZEROBUS_').lower()} in config.yaml"
|
||||
)
|
||||
return value
|
||||
|
||||
|
||||
def _workspace_id(server_endpoint: str) -> str:
|
||||
"""The Zerobus endpoint is ``https://<workspace_id>.zerobus.<region>.<cloud>``, so the id is its first label."""
|
||||
host: Final = urlsplit(server_endpoint).hostname or ""
|
||||
workspace_id: Final = host.split(".", 1)[0]
|
||||
if not workspace_id.isdigit():
|
||||
raise ValueError(
|
||||
f"ZEROBUS_SERVER_ENDPOINT {server_endpoint!r} does not look like "
|
||||
"https://<workspace_id>.zerobus.<region>.cloud.databricks.com"
|
||||
)
|
||||
return workspace_id
|
||||
|
||||
|
||||
def _table_name(configured: str | None) -> str:
|
||||
table_name: Final = _setting(configured, "ZEROBUS_TABLE_NAME")
|
||||
if table_name.count(".") != 2:
|
||||
raise ValueError(f"ZEROBUS_TABLE_NAME {table_name!r} must be fully qualified as catalog.schema.table")
|
||||
return table_name
|
||||
|
||||
|
||||
def connection_for(params: ZerobusInitParams) -> ZerobusConnection:
|
||||
"""The connection configured right now, so a UI edit takes effect without a restart."""
|
||||
server_endpoint: Final = _setting(params.server_endpoint, "ZEROBUS_SERVER_ENDPOINT")
|
||||
return ZerobusConnection(
|
||||
workspace_url=_setting(params.workspace_url, "ZEROBUS_WORKSPACE_URL"),
|
||||
workspace_id=_workspace_id(server_endpoint),
|
||||
server_endpoint=server_endpoint,
|
||||
client_id=_setting(params.client_id, "ZEROBUS_CLIENT_ID"),
|
||||
client_secret=_setting(params.client_secret, "ZEROBUS_CLIENT_SECRET"),
|
||||
table_name=_table_name(params.table_name),
|
||||
)
|
||||
|
||||
|
||||
class ZerobusLogger(CustomBatchLogger):
|
||||
preserve_events_added_during_flush = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
params: ZerobusInitParams | None = None,
|
||||
client: ZerobusIngestClient | None = None,
|
||||
start_periodic_flush: bool = True,
|
||||
) -> None:
|
||||
resolved: Final = params if params is not None else _configured_params()
|
||||
self.params: Final = resolved
|
||||
self.given_client: Final = client
|
||||
self._cached_client: ZerobusIngestClient | None = None
|
||||
if client is None:
|
||||
connection_for(resolved)
|
||||
super().__init__(
|
||||
flush_lock=asyncio.Lock(),
|
||||
batch_size=resolved.batch_size,
|
||||
flush_interval=resolved.flush_interval,
|
||||
turn_off_message_logging=bool(resolved.turn_off_message_logging),
|
||||
)
|
||||
self._flushing: bool = False
|
||||
self._batch_flush_task: asyncio.Task[None] | None = None
|
||||
self._periodic_flush_task: asyncio.Task[None] | None = (
|
||||
self._start_periodic_flush_task() if start_periodic_flush else None
|
||||
)
|
||||
|
||||
@property
|
||||
def client(self) -> ZerobusIngestClient:
|
||||
"""A client for the current connection, kept while the connection is unchanged so its token is reused."""
|
||||
if self.given_client is not None:
|
||||
return self.given_client
|
||||
connection: Final = connection_for(self.params)
|
||||
cached: Final = self._cached_client
|
||||
if cached is not None and cached.connection == connection:
|
||||
return cached
|
||||
fresh: Final = ZerobusIngestClient(
|
||||
connection=connection,
|
||||
http_client=get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback),
|
||||
)
|
||||
self._cached_client = fresh
|
||||
return fresh
|
||||
|
||||
def _start_periodic_flush_task(self) -> asyncio.Task[None] | None:
|
||||
try:
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return None
|
||||
return loop.create_task(self.periodic_flush())
|
||||
|
||||
def _start_batch_flush_task(self) -> None:
|
||||
if self._batch_flush_task is not None and not self._batch_flush_task.done():
|
||||
return
|
||||
try:
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return
|
||||
self._batch_flush_task = loop.create_task(self.flush_queue(skip_if_flushing=True))
|
||||
|
||||
def _flush_task_is_alive(self) -> bool:
|
||||
task: Final = self._periodic_flush_task
|
||||
return task is not None and not task.done() and not task.get_loop().is_closed()
|
||||
|
||||
async def async_log_success_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> None:
|
||||
await self._enqueue(kwargs)
|
||||
|
||||
async def async_log_failure_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> None:
|
||||
await self._enqueue(kwargs)
|
||||
|
||||
async def _enqueue(self, kwargs: Mapping[str, object]) -> None:
|
||||
try:
|
||||
if not self._flush_task_is_alive():
|
||||
self._periodic_flush_task = self._start_periodic_flush_task()
|
||||
|
||||
payload: Final = self._payload_for(kwargs)
|
||||
if payload is None:
|
||||
verbose_logger.debug("zerobus: event carried no standard_logging_object, skipping")
|
||||
return
|
||||
|
||||
if self._flushing and len(self.log_queue) >= self.max_queue_size:
|
||||
verbose_logger.warning("zerobus: queue at %s rows during a flush, dropped a row", self.max_queue_size)
|
||||
return
|
||||
|
||||
self.log_queue.append(trace_row(payload))
|
||||
self._drop_overflow()
|
||||
if len(self.log_queue) >= self.batch_size:
|
||||
self._start_batch_flush_task()
|
||||
except Exception: # noqa: BLE001 # logging must never break the request path
|
||||
verbose_logger.exception("zerobus: failed to queue an event")
|
||||
|
||||
def _payload_for(self, kwargs: Mapping[str, object]) -> Mapping[str, object] | None:
|
||||
"""The payload to buffer, redacted the way the framework redacts the success path."""
|
||||
details: Final = self.redact_standard_logging_payload_from_model_call_details(
|
||||
dict(kwargs) # mutable-ok: both framework helpers take the call details as a dict
|
||||
)
|
||||
payload: Final = details.get("standard_logging_object")
|
||||
if not isinstance(payload, dict):
|
||||
return None
|
||||
if should_redact_message_logging(details):
|
||||
return redacted_standard_logging_payload(payload)
|
||||
return payload
|
||||
|
||||
def _drop_overflow(self) -> None:
|
||||
"""Trim the oldest rows, except mid flush when the in-flight batch is the head of the queue."""
|
||||
if self._flushing:
|
||||
return
|
||||
overflow: Final = len(self.log_queue) - self.max_queue_size
|
||||
if overflow <= 0:
|
||||
return
|
||||
del self.log_queue[:overflow]
|
||||
verbose_logger.warning("zerobus: queue over %s rows, dropped %s oldest", self.max_queue_size, overflow)
|
||||
|
||||
async def flush_queue(self, skip_if_flushing: bool = False) -> None:
|
||||
if skip_if_flushing and self._flushing:
|
||||
return
|
||||
self._flushing = True
|
||||
try:
|
||||
await super().flush_queue()
|
||||
finally:
|
||||
self._flushing = False
|
||||
|
||||
async def async_send_batch(self) -> None:
|
||||
"""A retryable failure propagates so the rows are kept; a permanent one drops them so the queue moves on."""
|
||||
rows: Final = tuple(self.log_queue)
|
||||
if not rows:
|
||||
return
|
||||
failure: Final = await self.client.insert(rows)
|
||||
if failure is None:
|
||||
return
|
||||
if failure.retryable:
|
||||
raise ZerobusIngestError(failure.detail)
|
||||
verbose_logger.error("zerobus: dropping %s rows, %s", len(rows), failure.detail)
|
||||
156
litellm/integrations/zerobus/row.py
Normal file
156
litellm/integrations/zerobus/row.py
Normal file
|
|
@ -0,0 +1,156 @@
|
|||
"""
|
||||
Shape of one Delta table row per LiteLLM request.
|
||||
|
||||
Zerobus validates every record against the target table and rejects unknown columns, so
|
||||
the row is a fixed set of scalar columns for filtering plus JSON-encoded ``VARIANT``
|
||||
columns for anything nested. ``create_table_sql`` renders the matching DDL.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
TRACE_TABLE_COLUMNS: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
"id": "STRING",
|
||||
"trace_id": "STRING",
|
||||
"session_id": "STRING",
|
||||
"litellm_call_id": "STRING",
|
||||
"call_type": "STRING",
|
||||
"status": "STRING",
|
||||
"model": "STRING",
|
||||
"model_group": "STRING",
|
||||
"model_id": "STRING",
|
||||
"custom_llm_provider": "STRING",
|
||||
"api_base": "STRING",
|
||||
"stream": "BOOLEAN",
|
||||
"cache_hit": "BOOLEAN",
|
||||
"start_time": "TIMESTAMP",
|
||||
"end_time": "TIMESTAMP",
|
||||
"completion_start_time": "TIMESTAMP",
|
||||
"response_time": "DOUBLE",
|
||||
"prompt_tokens": "LONG",
|
||||
"completion_tokens": "LONG",
|
||||
"total_tokens": "LONG",
|
||||
"response_cost": "DOUBLE",
|
||||
"saved_cache_cost": "DOUBLE",
|
||||
"api_key_hash": "STRING",
|
||||
"api_key_alias": "STRING",
|
||||
"team_id": "STRING",
|
||||
"team_alias": "STRING",
|
||||
"user_id": "STRING",
|
||||
"org_id": "STRING",
|
||||
"end_user": "STRING",
|
||||
"requester_ip_address": "STRING",
|
||||
"user_agent": "STRING",
|
||||
"request_tags": "VARIANT",
|
||||
"messages": "VARIANT",
|
||||
"response": "VARIANT",
|
||||
"error_str": "STRING",
|
||||
"error_information": "VARIANT",
|
||||
"metadata": "VARIANT",
|
||||
"model_parameters": "VARIANT",
|
||||
"hidden_params": "VARIANT",
|
||||
"guardrail_information": "VARIANT",
|
||||
"cost_breakdown": "VARIANT",
|
||||
}
|
||||
)
|
||||
|
||||
_MICROSECONDS: Final = 1_000_000
|
||||
|
||||
|
||||
def create_table_sql(table_name: str) -> str:
|
||||
columns: Final = ",\n".join(f" {name} {delta_type}" for name, delta_type in TRACE_TABLE_COLUMNS.items())
|
||||
return f"CREATE TABLE {table_name} (\n{columns}\n);"
|
||||
|
||||
|
||||
def _text(payload: Mapping[str, object], key: str) -> str | None:
|
||||
value: Final = payload.get(key)
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
def _flag(payload: Mapping[str, object], key: str) -> bool | None:
|
||||
value: Final = payload.get(key)
|
||||
return value if isinstance(value, bool) else None
|
||||
|
||||
|
||||
def _number(payload: Mapping[str, object], key: str) -> float | None:
|
||||
value: Final = payload.get(key)
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
return None
|
||||
return float(value)
|
||||
|
||||
|
||||
def _count(payload: Mapping[str, object], key: str) -> int | None:
|
||||
value: Final = _number(payload, key)
|
||||
return None if value is None else int(value)
|
||||
|
||||
|
||||
def _timestamp_micros(payload: Mapping[str, object], key: str) -> int | None:
|
||||
"""Delta ``TIMESTAMP`` over Zerobus is epoch microseconds; LiteLLM keeps epoch seconds."""
|
||||
seconds: Final = _number(payload, key)
|
||||
if seconds is None or seconds <= 0:
|
||||
return None
|
||||
return int(seconds * _MICROSECONDS)
|
||||
|
||||
|
||||
def _json(payload: Mapping[str, object], key: str) -> str | None:
|
||||
value: Final = payload.get(key)
|
||||
return None if value is None else safe_dumps(value)
|
||||
|
||||
|
||||
def _metadata(payload: Mapping[str, object]) -> Mapping[str, object]:
|
||||
value: Final = payload.get("metadata")
|
||||
return value if isinstance(value, Mapping) else MappingProxyType({})
|
||||
|
||||
|
||||
def trace_row(payload: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""One ``TRACE_TABLE_COLUMNS`` row for a ``StandardLoggingPayload``."""
|
||||
metadata: Final = _metadata(payload)
|
||||
return MappingProxyType(
|
||||
{
|
||||
"id": _text(payload, "id"),
|
||||
"trace_id": _text(payload, "trace_id"),
|
||||
"session_id": _text(payload, "session_id"),
|
||||
"litellm_call_id": _text(payload, "litellm_call_id"),
|
||||
"call_type": _text(payload, "call_type"),
|
||||
"status": _text(payload, "status"),
|
||||
"model": _text(payload, "model"),
|
||||
"model_group": _text(payload, "model_group"),
|
||||
"model_id": _text(payload, "model_id"),
|
||||
"custom_llm_provider": _text(payload, "custom_llm_provider"),
|
||||
"api_base": _text(payload, "api_base"),
|
||||
"stream": _flag(payload, "stream"),
|
||||
"cache_hit": _flag(payload, "cache_hit"),
|
||||
"start_time": _timestamp_micros(payload, "startTime"),
|
||||
"end_time": _timestamp_micros(payload, "endTime"),
|
||||
"completion_start_time": _timestamp_micros(payload, "completionStartTime"),
|
||||
"response_time": _number(payload, "response_time"),
|
||||
"prompt_tokens": _count(payload, "prompt_tokens"),
|
||||
"completion_tokens": _count(payload, "completion_tokens"),
|
||||
"total_tokens": _count(payload, "total_tokens"),
|
||||
"response_cost": _number(payload, "response_cost"),
|
||||
"saved_cache_cost": _number(payload, "saved_cache_cost"),
|
||||
"api_key_hash": _text(metadata, "user_api_key_hash"),
|
||||
"api_key_alias": _text(metadata, "user_api_key_alias"),
|
||||
"team_id": _text(metadata, "user_api_key_team_id"),
|
||||
"team_alias": _text(metadata, "user_api_key_team_alias"),
|
||||
"user_id": _text(metadata, "user_api_key_user_id"),
|
||||
"org_id": _text(metadata, "user_api_key_org_id"),
|
||||
"end_user": _text(payload, "end_user"),
|
||||
"requester_ip_address": _text(payload, "requester_ip_address"),
|
||||
"user_agent": _text(payload, "user_agent"),
|
||||
"request_tags": _json(payload, "request_tags"),
|
||||
"messages": _json(payload, "messages"),
|
||||
"response": _json(payload, "response"),
|
||||
"error_str": _text(payload, "error_str"),
|
||||
"error_information": _json(payload, "error_information"),
|
||||
"metadata": _json(payload, "metadata"),
|
||||
"model_parameters": _json(payload, "model_parameters"),
|
||||
"hidden_params": _json(payload, "hidden_params"),
|
||||
"guardrail_information": _json(payload, "guardrail_information"),
|
||||
"cost_breakdown": _json(payload, "cost_breakdown"),
|
||||
}
|
||||
)
|
||||
|
|
@ -52,6 +52,7 @@ from litellm.integrations.vantage.vantage_logger import VantageLogger
|
|||
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
|
||||
VectorStorePreCallHook,
|
||||
)
|
||||
from litellm.integrations.zerobus import ZerobusLogger
|
||||
from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler
|
||||
from litellm.proxy.hooks.dynamic_rate_limiter_v3 import _PROXY_DynamicRateLimitHandlerV3
|
||||
|
||||
|
|
@ -97,6 +98,7 @@ class CustomLoggerRegistry:
|
|||
"deepeval": DeepEvalLogger,
|
||||
"s3_v2": S3Logger,
|
||||
"pointfive": PointFiveLogger,
|
||||
"zerobus": ZerobusLogger,
|
||||
"aws_sqs": SQSLogger,
|
||||
"dynamic_rate_limiter": _PROXY_DynamicRateLimitHandler,
|
||||
"dynamic_rate_limiter_v3": _PROXY_DynamicRateLimitHandlerV3,
|
||||
|
|
|
|||
|
|
@ -114,10 +114,9 @@ class HealthCheckHelpers:
|
|||
"""
|
||||
Health check for batch mode.
|
||||
|
||||
Calls list_batches for providers that support it (openai, hosted_vllm, azure,
|
||||
vertex_ai). For all other providers (e.g. bedrock) the batch API surface doesn't
|
||||
include list_batches, so we fall back to acompletion to verify connectivity and
|
||||
credential validity instead.
|
||||
Calls list_batches for providers that support it. For all other providers (e.g. bedrock)
|
||||
the batch API surface doesn't include list_batches, so we fall back to acompletion to
|
||||
verify connectivity and credential validity instead.
|
||||
"""
|
||||
import litellm
|
||||
|
||||
|
|
@ -132,10 +131,9 @@ class HealthCheckHelpers:
|
|||
litellm_params={"api_base": api_base} if api_base else None,
|
||||
)
|
||||
|
||||
if custom_llm_provider in LIST_BATCHES_SUPPORTED_PROVIDERS:
|
||||
return await litellm.alist_batches(**filtered_model_params)
|
||||
else:
|
||||
if custom_llm_provider not in LIST_BATCHES_SUPPORTED_PROVIDERS:
|
||||
return await litellm.acompletion(**model_params)
|
||||
return await litellm.alist_batches(**{**filtered_model_params, "custom_llm_provider": custom_llm_provider})
|
||||
|
||||
@staticmethod
|
||||
async def _image_edit_health_check(edit_request: Callable[[], Awaitable["ImageResponse"]]) -> "ImageResponse":
|
||||
|
|
|
|||
|
|
@ -20,11 +20,7 @@ from httpx import Response
|
|||
from pydantic import BaseModel, JsonValue
|
||||
|
||||
import litellm
|
||||
from litellm import (
|
||||
_custom_logger_compatible_callbacks_literal,
|
||||
json_logs,
|
||||
turn_off_message_logging,
|
||||
)
|
||||
from litellm import _custom_logger_compatible_callbacks_literal
|
||||
from litellm._logging import (
|
||||
_is_debugging_on,
|
||||
_redact_string,
|
||||
|
|
@ -43,8 +39,7 @@ from litellm.constants import (
|
|||
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT,
|
||||
EMPTY_MAPPING,
|
||||
PROVIDER_REQUEST_ID_HEADERS,
|
||||
SENTRY_DENYLIST,
|
||||
SENTRY_PII_DENYLIST,
|
||||
REDACTED_BY_LITELLM,
|
||||
)
|
||||
from litellm.cost_calculator import (
|
||||
RealtimeAPITokenUsageProcessor,
|
||||
|
|
@ -213,6 +208,7 @@ from ..integrations.s3 import S3Logger
|
|||
from ..integrations.s3_v2 import S3Logger as S3V2Logger
|
||||
from ..integrations.supabase import Supabase
|
||||
from ..integrations.traceloop import TraceloopLogger
|
||||
from ..integrations.zerobus import ZerobusLogger
|
||||
from .exception_mapping_utils import _get_response_headers
|
||||
from .initialize_dynamic_callback_params import (
|
||||
get_trusted_callback_params,
|
||||
|
|
@ -380,9 +376,12 @@ _DEPLOYMENT_PRICING_KEYS: Final = (
|
|||
"output_cost_per_token",
|
||||
"input_cost_per_token_batches",
|
||||
"output_cost_per_token_batches",
|
||||
"input_cost_per_token_above_200k_tokens_batches",
|
||||
"input_cost_per_token_above_272k_tokens_batches",
|
||||
"output_cost_per_token_above_200k_tokens_batches",
|
||||
"output_cost_per_token_above_272k_tokens_batches",
|
||||
"cache_read_input_token_cost_batches",
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches",
|
||||
"cache_read_input_token_cost_above_272k_tokens_batches",
|
||||
"cache_creation_input_token_cost_batches",
|
||||
"cache_creation_input_token_cost_above_272k_tokens_batches",
|
||||
|
|
@ -1356,10 +1355,19 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
_litellm_params: Final = self.model_call_details.get("litellm_params", {})
|
||||
_metadata: Final = _litellm_params.get("metadata", {}) or {}
|
||||
try:
|
||||
# [Non-blocking Extra Debug Information in metadata]
|
||||
if turn_off_message_logging is True:
|
||||
_metadata["raw_request"] = "redacted by litellm. \
|
||||
'litellm.turn_off_message_logging=True'"
|
||||
self.model_call_details["raw_request_typed_dict"] = RawRequestTypedDict(
|
||||
raw_request_api_base=self._get_masked_api_base(str(additional_args.get("api_base") or "")),
|
||||
raw_request_body=self._get_raw_request_body(additional_args.get("complete_input_dict", {})),
|
||||
# NOTE: setting ignore_sensitive_headers to True will cause
|
||||
# the Authorization header to be leaked when calls to the health
|
||||
# endpoint are made and fail.
|
||||
raw_request_headers=self._get_masked_headers(
|
||||
additional_args.get("headers", {}) or {},
|
||||
),
|
||||
error=None,
|
||||
)
|
||||
if should_redact_message_logging(self.model_call_details):
|
||||
_metadata["raw_request"] = REDACTED_BY_LITELLM
|
||||
else:
|
||||
curl_command: Final = self._get_request_curl_command(
|
||||
api_base=additional_args.get("api_base", ""),
|
||||
|
|
@ -1367,20 +1375,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
additional_args=additional_args,
|
||||
data=additional_args.get("complete_input_dict", {}),
|
||||
)
|
||||
|
||||
_metadata["raw_request"] = _redact_string(str(curl_command))
|
||||
# split up, so it's easier to parse in the UI
|
||||
self.model_call_details["raw_request_typed_dict"] = RawRequestTypedDict(
|
||||
raw_request_api_base=self._get_masked_api_base(str(additional_args.get("api_base") or "")),
|
||||
raw_request_body=self._get_raw_request_body(additional_args.get("complete_input_dict", {})),
|
||||
# NOTE: setting ignore_sensitive_headers to True will cause
|
||||
# the Authorization header to be leaked when calls to the health
|
||||
# endpoint are made and fail.
|
||||
raw_request_headers=self._get_masked_headers(
|
||||
additional_args.get("headers", {}) or {},
|
||||
),
|
||||
error=None,
|
||||
)
|
||||
except Exception as e:
|
||||
self.model_call_details["raw_request_typed_dict"] = RawRequestTypedDict(
|
||||
error=str(e),
|
||||
|
|
@ -1474,7 +1469,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
def _print_llm_call_debugging_log(
|
||||
self,
|
||||
api_base: str,
|
||||
headers: dict,
|
||||
headers: dict | None,
|
||||
additional_args: dict,
|
||||
):
|
||||
"""
|
||||
|
|
@ -1483,8 +1478,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
Prints the RAW curl command sent from LiteLLM
|
||||
"""
|
||||
if _is_debugging_on() or self.litellm_request_debug:
|
||||
if json_logs:
|
||||
masked_headers: Final = self._get_masked_headers(headers)
|
||||
if litellm.json_logs:
|
||||
masked_headers: Final = self._get_masked_headers(headers or {})
|
||||
masked_api_base: Final = self._get_masked_api_base(str(api_base or ""))
|
||||
if self.litellm_request_debug:
|
||||
verbose_logger.warning( # .warning ensures this shows up in all environments
|
||||
|
|
@ -1561,20 +1556,12 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
else:
|
||||
attr = "debug"
|
||||
|
||||
if json_logs:
|
||||
callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug
|
||||
callattr(
|
||||
"RAW RESPONSE:\n{}\n\n".format(
|
||||
self.model_call_details.get("original_response", self.model_call_details)
|
||||
),
|
||||
)
|
||||
else:
|
||||
callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug
|
||||
callattr(
|
||||
"RAW RESPONSE:\n{}\n\n".format(
|
||||
self.model_call_details.get("original_response", self.model_call_details)
|
||||
)
|
||||
callattr: Final = verbose_logger.warning if attr == "warning" else verbose_logger.debug
|
||||
callattr(
|
||||
"RAW RESPONSE:\n{}\n\n".format(
|
||||
self.model_call_details.get("original_response", self.model_call_details)
|
||||
)
|
||||
)
|
||||
if getattr(self, "logger_fn", None) and callable(self.logger_fn):
|
||||
try:
|
||||
self.logger_fn(
|
||||
|
|
@ -4215,6 +4202,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
json_mode=False,
|
||||
litellm_params={},
|
||||
)
|
||||
elif result is None:
|
||||
verbose_logger.warning("LiteLLM: the anthropic_messages stream assembled no response, logging an empty one")
|
||||
return litellm.ModelResponse(model=self.model)
|
||||
else:
|
||||
from litellm.types.llms.anthropic import AnthropicResponse
|
||||
|
||||
|
|
@ -4423,21 +4413,10 @@ def set_callbacks(callback_list, function_id=None):
|
|||
print_verbose("Package 'sentry_sdk' is missing. Installing it...")
|
||||
subprocess.check_call([sys.executable, "-m", "pip", "install", "sentry_sdk"])
|
||||
import sentry_sdk
|
||||
from sentry_sdk.scrubber import EventScrubber
|
||||
from litellm.litellm_core_utils.sentry_scrubbing import build_sentry_init_options
|
||||
|
||||
sentry_sdk_instance = sentry_sdk
|
||||
sentry_trace_rate = os.environ.get("SENTRY_API_TRACE_RATE", "1.0")
|
||||
sentry_sample_rate = (
|
||||
os.environ.get("SENTRY_API_SAMPLE_RATE") if "SENTRY_API_SAMPLE_RATE" in os.environ else "1.0"
|
||||
)
|
||||
sentry_sdk_instance.init(
|
||||
dsn=os.environ.get("SENTRY_DSN"),
|
||||
traces_sample_rate=float(sentry_trace_rate),
|
||||
sample_rate=float(sentry_sample_rate if sentry_sample_rate else 1.0),
|
||||
send_default_pii=False, # Prevent sending Personal Identifiable Information
|
||||
event_scrubber=EventScrubber(denylist=SENTRY_DENYLIST, pii_denylist=SENTRY_PII_DENYLIST),
|
||||
environment=os.environ.get("SENTRY_ENVIRONMENT", "production"),
|
||||
)
|
||||
sentry_sdk_instance.init(**build_sentry_init_options(os.environ))
|
||||
capture_exception = sentry_sdk_instance.capture_exception
|
||||
add_breadcrumb = sentry_sdk_instance.add_breadcrumb
|
||||
elif callback == "slack":
|
||||
|
|
@ -4660,6 +4639,14 @@ def _init_custom_logger_compatible_class(
|
|||
_pointfive_logger: Final = PointFiveLogger()
|
||||
_in_memory_loggers.append(_pointfive_logger)
|
||||
return _pointfive_logger
|
||||
elif logging_integration == "zerobus":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, ZerobusLogger):
|
||||
return callback
|
||||
|
||||
_zerobus_logger: Final = ZerobusLogger()
|
||||
_in_memory_loggers.append(_zerobus_logger)
|
||||
return _zerobus_logger
|
||||
elif logging_integration == "aws_sqs":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, SQSLogger):
|
||||
|
|
@ -5352,6 +5339,10 @@ def get_custom_logger_compatible_class(
|
|||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, PointFiveLogger):
|
||||
return callback
|
||||
elif logging_integration == "zerobus":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, ZerobusLogger):
|
||||
return callback
|
||||
elif logging_integration == "aws_sqs":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, SQSLogger):
|
||||
|
|
|
|||
152
litellm/litellm_core_utils/sentry_scrubbing.py
Normal file
152
litellm/litellm_core_utils/sentry_scrubbing.py
Normal file
|
|
@ -0,0 +1,152 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from functools import reduce
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias, cast
|
||||
|
||||
from pydantic import JsonValue
|
||||
from sentry_sdk.scrubber import DEFAULT_DENYLIST, DEFAULT_PII_DENYLIST, EventScrubber
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.constants import (
|
||||
LENGTH_OF_LITELLM_GENERATED_KEY,
|
||||
MINIMUM_CUSTOM_KEY_LENGTH,
|
||||
SENTRY_DENYLIST,
|
||||
SENTRY_PII_DENYLIST,
|
||||
)
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sentry_sdk.types import Event, Hint
|
||||
|
||||
EventScrubFn: TypeAlias = "Callable[[Event, Hint], Event]"
|
||||
JsonPath: TypeAlias = tuple[str, ...]
|
||||
|
||||
FILTERED: Final = "[Filtered]"
|
||||
SEND_DEFAULT_PII_ENV: Final = "SENTRY_SEND_DEFAULT_PII"
|
||||
SECRET_FIELD_NAMES: Final = tuple(DEFAULT_DENYLIST) + tuple(SENTRY_DENYLIST)
|
||||
PII_FIELD_NAMES: Final = tuple(DEFAULT_PII_DENYLIST) + tuple(SENTRY_PII_DENYLIST)
|
||||
|
||||
KEY_PREFIX: Final = "sk-"
|
||||
|
||||
|
||||
def build_key_pattern(custom_key_minimum: int, generated_key_bytes: int) -> re.Pattern[str]:
|
||||
generated_suffix_length: Final = (generated_key_bytes * 4 + 2) // 3
|
||||
floor: Final = min(custom_key_minimum - len(KEY_PREFIX), generated_suffix_length)
|
||||
return re.compile(rf"{KEY_PREFIX}[A-Za-z0-9_-]{{{floor},}}")
|
||||
|
||||
|
||||
LITELLM_KEY_PATTERN: Final = build_key_pattern(MINIMUM_CUSTOM_KEY_LENGTH, LENGTH_OF_LITELLM_GENERATED_KEY)
|
||||
SOURCE_CONTEXT_KEYS: Final = frozenset({"pre_context", "context_line", "post_context"})
|
||||
STACK_FRAME_PATHS: Final = frozenset(
|
||||
{
|
||||
("exception", "values", "*", "stacktrace", "frames", "*"),
|
||||
("threads", "values", "*", "stacktrace", "frames", "*"),
|
||||
("stacktrace", "frames", "*"),
|
||||
}
|
||||
)
|
||||
MAX_SCRUB_DEPTH: Final = 64
|
||||
EMAIL_PATTERN: Final = re.compile(r"[A-Za-z0-9._%+-]+@[A-Za-z0-9-]+(?:\.[A-Za-z0-9-]+)*\.[A-Za-z]{2,}")
|
||||
SHA256_HEX_PATTERN: Final = re.compile(r"(?<![0-9A-Za-z])[0-9a-f]{64}(?![0-9A-Za-z])")
|
||||
QUOTED_VALUE: Final = r"'(?:[^'\\]|\\.)*'|\"(?:[^\"\\]|\\.)*\""
|
||||
BRACKET_ATOM: Final = rf"(?:{QUOTED_VALUE})|[^\[\]{{}}()'\"]"
|
||||
NESTED_BRACKET_LEVELS: Final = 3
|
||||
BRACKETED_VALUE: Final = reduce(
|
||||
lambda inner, _: rf"[\[{{(](?:{BRACKET_ATOM}|{inner})*[\]}})]",
|
||||
range(NESTED_BRACKET_LEVELS),
|
||||
rf"[\[{{(](?:{BRACKET_ATOM})*[\]}})]",
|
||||
)
|
||||
BARE_VALUE: Final = r"(?!None(?![0-9A-Za-z_]))[^,)\]}\s]+"
|
||||
|
||||
|
||||
class SentryInitOptions(TypedDict):
|
||||
dsn: ReadOnly[str | None]
|
||||
traces_sample_rate: ReadOnly[float]
|
||||
sample_rate: ReadOnly[float]
|
||||
send_default_pii: ReadOnly[bool]
|
||||
event_scrubber: ReadOnly[EventScrubber]
|
||||
before_send: ReadOnly[EventScrubFn]
|
||||
before_send_transaction: ReadOnly[EventScrubFn]
|
||||
environment: ReadOnly[str]
|
||||
|
||||
|
||||
def build_repr_field_pattern(field_names: Sequence[str]) -> re.Pattern[str]:
|
||||
names: Final = "|".join(re.escape(name) for name in field_names)
|
||||
return re.compile(
|
||||
rf"(?P<field>(?<![0-9A-Za-z_])(?:{names})=|['\"](?:{names})['\"]:\s*)(?P<value>{QUOTED_VALUE}|{BRACKETED_VALUE}|{BARE_VALUE})",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def build_string_scrubber(send_default_pii: bool) -> Callable[[str], str]:
|
||||
field_names: Final = SECRET_FIELD_NAMES if send_default_pii else SECRET_FIELD_NAMES + PII_FIELD_NAMES
|
||||
field_pattern: Final = build_repr_field_pattern(field_names)
|
||||
value_patterns: Final = (
|
||||
(LITELLM_KEY_PATTERN,) if send_default_pii else (LITELLM_KEY_PATTERN, EMAIL_PATTERN, SHA256_HEX_PATTERN)
|
||||
)
|
||||
|
||||
def scrub(text: str) -> str:
|
||||
fields_scrubbed: Final = field_pattern.sub(_filtered_field, text)
|
||||
return _substitute_all(value_patterns, fields_scrubbed)
|
||||
|
||||
return scrub
|
||||
|
||||
|
||||
def _filtered_field(match: re.Match[str]) -> str:
|
||||
quote: Final = '"' if match.group("value").startswith('"') else "'"
|
||||
return f"{match.group('field')}{quote}{FILTERED}{quote}"
|
||||
|
||||
|
||||
def _substitute_all(patterns: Sequence[re.Pattern[str]], text: str) -> str:
|
||||
return reduce(lambda scrubbed, pattern: pattern.sub(FILTERED, scrubbed), patterns, text)
|
||||
|
||||
|
||||
def scrub_json_strings(value: JsonValue, scrub: Callable[[str], str], path: JsonPath = ()) -> JsonValue:
|
||||
if len(path) > MAX_SCRUB_DEPTH:
|
||||
return FILTERED
|
||||
if isinstance(value, str):
|
||||
return scrub(value)
|
||||
if isinstance(value, dict):
|
||||
unscrubbed_keys: Final = SOURCE_CONTEXT_KEYS if path in STACK_FRAME_PATHS else frozenset[str]()
|
||||
return { # mutable-ok: JSON object
|
||||
key: item if key in unscrubbed_keys else scrub_json_strings(item, scrub, (*path, key))
|
||||
for key, item in value.items()
|
||||
}
|
||||
if isinstance(value, list):
|
||||
return [scrub_json_strings(item, scrub, (*path, "*")) for item in value] # mutable-ok: JSON array
|
||||
return value
|
||||
|
||||
|
||||
def build_event_scrubber(send_default_pii: bool) -> EventScrubFn:
|
||||
scrub: Final = build_string_scrubber(send_default_pii)
|
||||
|
||||
def scrub_event(event: Event, _hint: Hint) -> Event:
|
||||
json_event: Final = cast("JsonValue", event) # cast-ok: [LIT006] the SDK serialized the event to JSON already
|
||||
return cast("Event", scrub_json_strings(json_event, scrub)) # cast-ok: [LIT006] same JSON shape going back
|
||||
|
||||
return scrub_event
|
||||
|
||||
|
||||
def send_default_pii_from_env(env: Mapping[str, str]) -> bool:
|
||||
return str_to_bool(env.get(SEND_DEFAULT_PII_ENV)) is True
|
||||
|
||||
|
||||
def build_sentry_init_options(env: Mapping[str, str]) -> SentryInitOptions:
|
||||
send_default_pii: Final = send_default_pii_from_env(env)
|
||||
scrub_event: Final = build_event_scrubber(send_default_pii)
|
||||
return SentryInitOptions(
|
||||
dsn=env.get("SENTRY_DSN"),
|
||||
traces_sample_rate=float(env.get("SENTRY_API_TRACE_RATE") or "1.0"),
|
||||
sample_rate=float(env.get("SENTRY_API_SAMPLE_RATE") or "1.0"),
|
||||
send_default_pii=send_default_pii,
|
||||
event_scrubber=EventScrubber(
|
||||
denylist=list(SECRET_FIELD_NAMES), # mutable-ok: EventScrubber appends pii_denylist onto denylist in place
|
||||
pii_denylist=list(PII_FIELD_NAMES), # mutable-ok: EventScrubber takes List[str]
|
||||
recursive=True,
|
||||
send_default_pii=send_default_pii,
|
||||
),
|
||||
before_send=scrub_event,
|
||||
before_send_transaction=scrub_event,
|
||||
environment=env.get("SENTRY_ENVIRONMENT", "production"),
|
||||
)
|
||||
|
|
@ -77,13 +77,13 @@ class _CallerHeadersView(TypedDict):
|
|||
headers: ReadOnly[dict[str, str]]
|
||||
|
||||
|
||||
# Globally-routable IPs that are cloud-internal. Everything else
|
||||
# non-public is caught by ``not ip.is_global`` (RFC 6890, as implemented by
|
||||
# Python's ``ipaddress`` module). This list only holds IPs that are
|
||||
# publicly routable *and* point to cloud-fabric services reachable from
|
||||
# inside a VM via special in-fabric routing.
|
||||
# Cloud-internal IPs that ``ip.is_global`` can report as public. Everything
|
||||
# else non-public is caught by ``not ip.is_global`` (RFC 6890, as implemented
|
||||
# by Python's ``ipaddress`` module). Older Python patch releases (3.12.2, for
|
||||
# one) treat most of 192.0.0.0/24 as global, so it is listed to block it everywhere.
|
||||
_CLOUD_METADATA_EXCEPTIONS: Final = [
|
||||
ip_network("168.63.129.16/32"), # Azure Wire Server
|
||||
ip_network("192.0.0.0/24"),
|
||||
]
|
||||
|
||||
_ALLOWED_SCHEMES: Final = ("http", "https")
|
||||
|
|
|
|||
|
|
@ -70,15 +70,11 @@ def _error_status_and_message(exc: Exception) -> tuple[int, str]:
|
|||
|
||||
def _mid_stream_error_sse_event(exc: Exception) -> bytes:
|
||||
from litellm.anthropic_interface.exceptions.exception_mapping_utils import (
|
||||
AnthropicExceptionMapping,
|
||||
anthropic_error_sse_frame,
|
||||
)
|
||||
|
||||
status_code, message = _error_status_and_message(exc)
|
||||
error_response = AnthropicExceptionMapping.transform_to_anthropic_error(
|
||||
status_code=status_code,
|
||||
raw_message=message,
|
||||
)
|
||||
return f"event: error\ndata: {json.dumps(error_response)}\n\n".encode()
|
||||
return anthropic_error_sse_frame(status_code=status_code, raw_message=message).encode()
|
||||
|
||||
|
||||
def _delta_payload_field(delta_type: StreamingContentBlockDeltaType) -> str:
|
||||
|
|
|
|||
|
|
@ -69,7 +69,9 @@ def make_sync_call(
|
|||
completion_stream: Any = MockResponseIterator(model_response=model_response, json_mode=json_mode)
|
||||
else:
|
||||
decoder: Final = AWSEventStreamDecoder(model=model, json_mode=json_mode)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
|
||||
completion_stream = decoder.iter_bytes(
|
||||
response.iter_bytes(chunk_size=stream_chunk_size), response_headers=response.headers
|
||||
)
|
||||
|
||||
# LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import types
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Final, cast
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
|
@ -51,7 +51,11 @@ from ..common_utils import (
|
|||
bedrock_tool_name_mappings: Final[InMemoryCache] = InMemoryCache(max_size_in_memory=50, default_ttl=600)
|
||||
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from botocore.eventstream import EventStreamMessage
|
||||
|
||||
converse_config: Final = AmazonConverseConfig()
|
||||
_STREAM_HEAD_BYTES: Final = 200
|
||||
NOVA_INVOKE_STREAM_EVENT_TYPES: Final = (
|
||||
"messageStart",
|
||||
"contentBlockStart",
|
||||
|
|
@ -162,6 +166,22 @@ class AmazonCohereChatConfig:
|
|||
return optional_params
|
||||
|
||||
|
||||
def _stream_decoder(
|
||||
bedrock_invoke_provider: litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL | None,
|
||||
*,
|
||||
model: str,
|
||||
json_mode: bool | None,
|
||||
sync_stream: bool,
|
||||
) -> "AWSEventStreamDecoder":
|
||||
if bedrock_invoke_provider == "anthropic":
|
||||
return AmazonAnthropicClaudeStreamDecoder(model=model, sync_stream=sync_stream, json_mode=json_mode)
|
||||
if bedrock_invoke_provider == "deepseek_r1":
|
||||
return AmazonDeepSeekR1StreamDecoder(model=model, sync_stream=sync_stream)
|
||||
if bedrock_invoke_provider == "moonshot":
|
||||
return AmazonOpenAICompatibleStreamDecoder(model=model, sync_stream=sync_stream)
|
||||
return AWSEventStreamDecoder(model=model, json_mode=json_mode)
|
||||
|
||||
|
||||
async def make_call(
|
||||
client: AsyncHTTPHandler | None,
|
||||
api_base: str,
|
||||
|
|
@ -218,28 +238,13 @@ async def make_call(
|
|||
completion_stream: MockResponseIterator | AsyncIterator[GChunk | ModelResponseStream | dict] = (
|
||||
MockResponseIterator(model_response=model_response, json_mode=json_mode)
|
||||
)
|
||||
elif bedrock_invoke_provider == "anthropic":
|
||||
decoder: AWSEventStreamDecoder = AmazonAnthropicClaudeStreamDecoder(
|
||||
model=model,
|
||||
sync_stream=False,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size))
|
||||
elif bedrock_invoke_provider == "deepseek_r1":
|
||||
decoder = AmazonDeepSeekR1StreamDecoder(
|
||||
model=model,
|
||||
sync_stream=False,
|
||||
)
|
||||
completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size))
|
||||
elif bedrock_invoke_provider == "moonshot":
|
||||
decoder = AmazonOpenAICompatibleStreamDecoder(
|
||||
model=model,
|
||||
sync_stream=False,
|
||||
)
|
||||
completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size))
|
||||
else:
|
||||
decoder = AWSEventStreamDecoder(model=model, json_mode=json_mode)
|
||||
completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size))
|
||||
decoder: Final = _stream_decoder(
|
||||
bedrock_invoke_provider, model=model, json_mode=json_mode, sync_stream=False
|
||||
)
|
||||
completion_stream = decoder.aiter_bytes(
|
||||
response.aiter_bytes(chunk_size=stream_chunk_size), response_headers=response.headers
|
||||
)
|
||||
|
||||
# LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -322,28 +327,13 @@ def make_sync_call(
|
|||
completion_stream: MockResponseIterator | Iterator[GChunk | ModelResponseStream | dict] = (
|
||||
MockResponseIterator(model_response=model_response, json_mode=json_mode)
|
||||
)
|
||||
elif bedrock_invoke_provider == "anthropic":
|
||||
decoder: AWSEventStreamDecoder = AmazonAnthropicClaudeStreamDecoder(
|
||||
model=model,
|
||||
sync_stream=True,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
|
||||
elif bedrock_invoke_provider == "deepseek_r1":
|
||||
decoder = AmazonDeepSeekR1StreamDecoder(
|
||||
model=model,
|
||||
sync_stream=True,
|
||||
)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
|
||||
elif bedrock_invoke_provider == "moonshot":
|
||||
decoder = AmazonOpenAICompatibleStreamDecoder(
|
||||
model=model,
|
||||
sync_stream=True,
|
||||
)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
|
||||
else:
|
||||
decoder = AWSEventStreamDecoder(model=model, json_mode=json_mode)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
|
||||
decoder: Final = _stream_decoder(
|
||||
bedrock_invoke_provider, model=model, json_mode=json_mode, sync_stream=True
|
||||
)
|
||||
completion_stream = decoder.iter_bytes(
|
||||
response.iter_bytes(chunk_size=stream_chunk_size), response_headers=response.headers
|
||||
)
|
||||
|
||||
# LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -370,6 +360,49 @@ def make_sync_call(
|
|||
raise BedrockError(status_code=500, message=str(e))
|
||||
|
||||
|
||||
def _response_header(response_headers: Mapping[str, str] | None, name: str) -> str | None:
|
||||
return None if response_headers is None else response_headers.get(name)
|
||||
|
||||
|
||||
class _EventStreamTally:
|
||||
def __init__(self) -> None:
|
||||
self.bytes_received = 0
|
||||
self.bytes_decoded = 0
|
||||
self.events = 0
|
||||
self.head = b""
|
||||
|
||||
def add_chunk(self, chunk: bytes) -> None:
|
||||
self.bytes_received += len(chunk)
|
||||
if len(self.head) < _STREAM_HEAD_BYTES:
|
||||
self.head = (self.head + chunk)[:_STREAM_HEAD_BYTES]
|
||||
|
||||
def add_event(self, event: "EventStreamMessage") -> None:
|
||||
self.events += 1
|
||||
self.bytes_decoded += event.prelude.total_length
|
||||
|
||||
def undecoded_stream_error(self, response_headers: Mapping[str, str] | None) -> BedrockError | None:
|
||||
undecoded: Final = self.bytes_received - self.bytes_decoded
|
||||
if self.events and not undecoded:
|
||||
return None
|
||||
detail: Final = (
|
||||
f"content-type={_response_header(response_headers, 'content-type')!r}, "
|
||||
f"x-amzn-requestid={_response_header(response_headers, 'x-amzn-requestid')!r}, "
|
||||
f"{self.bytes_received} bytes received"
|
||||
)
|
||||
if not self.events:
|
||||
return BedrockError(
|
||||
status_code=502,
|
||||
message=(
|
||||
"Bedrock answered the stream with HTTP 200 but its body decoded to no events "
|
||||
f"({detail}, first bytes={self.head!r})"
|
||||
),
|
||||
)
|
||||
return BedrockError(
|
||||
status_code=502,
|
||||
message=f"Bedrock stream ended with {undecoded} undecoded bytes after {self.events} events ({detail})",
|
||||
)
|
||||
|
||||
|
||||
class AWSEventStreamDecoder:
|
||||
def __init__(self, model: str, json_mode: bool | None = False) -> None:
|
||||
from botocore.parsers import EventStreamJSONParser
|
||||
|
|
@ -709,32 +742,48 @@ class AWSEventStreamDecoder:
|
|||
tool_use=None,
|
||||
)
|
||||
|
||||
def iter_bytes(self, iterator: Iterator[bytes]) -> Iterator[GChunk | ModelResponseStream | dict]:
|
||||
def iter_bytes(
|
||||
self, iterator: Iterator[bytes], *, response_headers: Mapping[str, str] | None = None
|
||||
) -> Iterator[GChunk | ModelResponseStream | dict]:
|
||||
"""Given an iterator that yields lines, iterate over it & yield every event encountered"""
|
||||
from botocore.eventstream import EventStreamBuffer
|
||||
|
||||
event_stream_buffer: Final = EventStreamBuffer()
|
||||
tally: Final = _EventStreamTally()
|
||||
for chunk in iterator:
|
||||
event_stream_buffer.add_data(chunk)
|
||||
tally.add_chunk(chunk)
|
||||
for event in event_stream_buffer:
|
||||
tally.add_event(event)
|
||||
message = self._parse_message_from_event(event)
|
||||
if message:
|
||||
# sse_event = ServerSentEvent(data=message, event="completion")
|
||||
_data = json.loads(message)
|
||||
yield self._chunk_parser(chunk_data=_data)
|
||||
undecoded_stream_error: Final = tally.undecoded_stream_error(response_headers)
|
||||
if undecoded_stream_error is not None:
|
||||
raise undecoded_stream_error
|
||||
|
||||
async def aiter_bytes(self, iterator: AsyncIterator[bytes]) -> AsyncIterator[GChunk | ModelResponseStream | dict]:
|
||||
async def aiter_bytes(
|
||||
self, iterator: AsyncIterator[bytes], *, response_headers: Mapping[str, str] | None = None
|
||||
) -> AsyncIterator[GChunk | ModelResponseStream | dict]:
|
||||
"""Given an async iterator that yields lines, iterate over it & yield every event encountered"""
|
||||
from botocore.eventstream import EventStreamBuffer
|
||||
|
||||
event_stream_buffer: Final = EventStreamBuffer()
|
||||
tally: Final = _EventStreamTally()
|
||||
async for chunk in iterator:
|
||||
event_stream_buffer.add_data(chunk)
|
||||
tally.add_chunk(chunk)
|
||||
for event in event_stream_buffer:
|
||||
tally.add_event(event)
|
||||
message = self._parse_message_from_event(event)
|
||||
if message:
|
||||
_data = json.loads(message)
|
||||
yield self._chunk_parser(chunk_data=_data)
|
||||
undecoded_stream_error: Final = tally.undecoded_stream_error(response_headers)
|
||||
if undecoded_stream_error is not None:
|
||||
raise undecoded_stream_error
|
||||
|
||||
def _parse_message_from_event(self, event) -> str | None:
|
||||
response_stream_shape: Final = get_bedrock_response_stream_shape()
|
||||
|
|
|
|||
|
|
@ -770,7 +770,9 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
aws_decoder: Final = AmazonAnthropicClaudeMessagesStreamDecoder(
|
||||
model=model,
|
||||
)
|
||||
completion_stream: Final = aws_decoder.aiter_bytes(httpx_response.aiter_bytes())
|
||||
completion_stream: Final = aws_decoder.aiter_bytes(
|
||||
httpx_response.aiter_bytes(), response_headers=httpx_response.headers
|
||||
)
|
||||
# Convert decoded Bedrock events to Server-Sent Events expected by Anthropic clients.
|
||||
return self.bedrock_sse_wrapper(
|
||||
completion_stream=completion_stream,
|
||||
|
|
|
|||
|
|
@ -53,14 +53,18 @@ class VertexAIFilesHandler(GCSBucketBase):
|
|||
|
||||
Sources them from the deployment's ``litellm_params`` (``gcs_bucket_name`` /
|
||||
``bucket_name`` and ``vertex_credentials``), mirroring the write path in
|
||||
``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the global
|
||||
``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` env vars. This lets Vertex batch
|
||||
run entirely at the model-group level, so output written to a per-model bucket is
|
||||
readable without setting the global env vars.
|
||||
``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the
|
||||
``GCS_BATCH_BUCKET_NAME`` then ``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT``
|
||||
env vars. This lets Vertex batch run entirely at the model-group level, so output
|
||||
written to a per-model bucket is readable without setting the global env vars.
|
||||
"""
|
||||
params: Final[Mapping[str, object]] = litellm_params or {}
|
||||
bucket_candidate: Final = params.get("gcs_bucket_name") or params.get("bucket_name")
|
||||
configured_bucket_name = bucket_candidate if isinstance(bucket_candidate, str) else os.getenv("GCS_BUCKET_NAME")
|
||||
configured_bucket_name = (
|
||||
bucket_candidate
|
||||
if isinstance(bucket_candidate, str)
|
||||
else os.getenv("GCS_BATCH_BUCKET_NAME") or os.getenv("GCS_BUCKET_NAME")
|
||||
)
|
||||
|
||||
credentials: Final = params.get("vertex_credentials") or vertex_credentials
|
||||
if isinstance(credentials, dict):
|
||||
|
|
|
|||
|
|
@ -961,7 +961,10 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
|
||||
def _get_configured_bucket_name(self, litellm_params: dict) -> str:
|
||||
bucket_name: Final = (
|
||||
litellm_params.get("gcs_bucket_name") or litellm_params.get("bucket_name") or os.getenv("GCS_BUCKET_NAME")
|
||||
litellm_params.get("gcs_bucket_name")
|
||||
or litellm_params.get("bucket_name")
|
||||
or os.getenv("GCS_BATCH_BUCKET_NAME")
|
||||
or os.getenv("GCS_BUCKET_NAME")
|
||||
)
|
||||
if not bucket_name:
|
||||
raise ValueError("GCS bucket_name is required")
|
||||
|
|
|
|||
|
|
@ -122,41 +122,26 @@ class VertexAIRAGIngestion(BaseRAGIngestion):
|
|||
"""
|
||||
import litellm
|
||||
|
||||
# Set GCS_BUCKET_NAME env var for litellm.files.create_file
|
||||
# The handler uses this to determine where to upload
|
||||
original_bucket: Final = os.environ.get("GCS_BUCKET_NAME")
|
||||
if self.gcs_bucket:
|
||||
os.environ["GCS_BUCKET_NAME"] = self.gcs_bucket
|
||||
file_tuple: Final = (filename, file_content, content_type)
|
||||
|
||||
try:
|
||||
# Create file tuple for litellm.files.acreate_file
|
||||
file_tuple: Final = (filename, file_content, content_type)
|
||||
verbose_logger.debug(
|
||||
"Uploading file to GCS via litellm.files.acreate_file: %s (bucket: %s)", filename, self.gcs_bucket
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
"Uploading file to GCS via litellm.files.acreate_file: %s (bucket: %s)", filename, self.gcs_bucket
|
||||
)
|
||||
response: Final = await litellm.acreate_file(
|
||||
file=file_tuple,
|
||||
purpose="assistants",
|
||||
custom_llm_provider="vertex_ai",
|
||||
gcs_bucket_name=self.gcs_bucket,
|
||||
vertex_project=self.vertex_project,
|
||||
vertex_location=self.vertex_location,
|
||||
vertex_credentials=self.vertex_credentials,
|
||||
)
|
||||
|
||||
# Upload to GCS using LiteLLM's file upload
|
||||
response: Final = await litellm.acreate_file(
|
||||
file=file_tuple,
|
||||
purpose="assistants", # Purpose for file storage
|
||||
custom_llm_provider="vertex_ai",
|
||||
vertex_project=self.vertex_project,
|
||||
vertex_location=self.vertex_location,
|
||||
vertex_credentials=self.vertex_credentials,
|
||||
)
|
||||
gcs_uri: Final = response.id
|
||||
verbose_logger.info("Uploaded file to GCS: %s", gcs_uri)
|
||||
|
||||
# The response.id should be the GCS URI
|
||||
gcs_uri: Final = response.id
|
||||
verbose_logger.info("Uploaded file to GCS: %s", gcs_uri)
|
||||
|
||||
return gcs_uri
|
||||
finally:
|
||||
# Restore original env var
|
||||
if original_bucket is not None:
|
||||
os.environ["GCS_BUCKET_NAME"] = original_bucket
|
||||
elif "GCS_BUCKET_NAME" in os.environ:
|
||||
del os.environ["GCS_BUCKET_NAME"]
|
||||
return gcs_uri
|
||||
|
||||
async def _import_file_to_corpus_via_sdk(
|
||||
self,
|
||||
|
|
@ -259,6 +244,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion):
|
|||
content_type: str | None,
|
||||
chunks: list[str],
|
||||
embeddings: list[list[float]] | None,
|
||||
existing_file_id: str | None = None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""
|
||||
Store content in Vertex AI RAG corpus.
|
||||
|
|
@ -274,6 +260,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion):
|
|||
content_type: MIME type
|
||||
chunks: Ignored - Vertex AI handles chunking
|
||||
embeddings: Ignored - Vertex AI handles embedding
|
||||
existing_file_id: Existing provider file ID, unsupported for Vertex AI RAG Engine
|
||||
|
||||
Returns:
|
||||
Tuple of (corpus_id, gcs_uri)
|
||||
|
|
|
|||
195
litellm/llms/xai/batches/handler.py
Normal file
195
litellm/llms/xai/batches/handler.py
Normal file
|
|
@ -0,0 +1,195 @@
|
|||
from collections.abc import Coroutine
|
||||
from itertools import chain
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.types.llms.openai import CreateBatchRequest, HttpxBinaryResponseContent
|
||||
from litellm.types.utils import LiteLLMBatch, LlmProviders
|
||||
|
||||
from .transformation import (
|
||||
XAI_RESULTS_PAGE_SIZE,
|
||||
OpenAIBatchListResponse,
|
||||
XAIBatch,
|
||||
XAIBatchList,
|
||||
XAIBatchResult,
|
||||
XAIBatchResultsPage,
|
||||
get_xai_auth_headers,
|
||||
raise_for_xai_status,
|
||||
results_to_openai_jsonl,
|
||||
to_create_batch_body,
|
||||
to_litellm_batch,
|
||||
to_openai_batch_list,
|
||||
xai_batches_url,
|
||||
)
|
||||
|
||||
_JSONL_CONTENT_TYPE: Final = ("content-type", "application/jsonl")
|
||||
|
||||
|
||||
class _PageParams(TypedDict):
|
||||
limit: ReadOnly[int]
|
||||
pagination_token: NotRequired[ReadOnly[str]]
|
||||
|
||||
|
||||
def _results_params(after: str | None, limit: int | None) -> dict[str, object]: # mutable-ok: httpx params
|
||||
if after is None:
|
||||
return dict(_PageParams(limit=limit or XAI_RESULTS_PAGE_SIZE)) # mutable-ok: httpx params
|
||||
return dict(_PageParams(limit=limit or XAI_RESULTS_PAGE_SIZE, pagination_token=after)) # mutable-ok: httpx params
|
||||
|
||||
|
||||
def _flatten(pages: list[XAIBatchResultsPage]) -> tuple[XAIBatchResult, ...]:
|
||||
return tuple(chain.from_iterable(page.results for page in pages))
|
||||
|
||||
|
||||
def _jsonl_response(url: str, results: tuple[XAIBatchResult, ...]) -> HttpxBinaryResponseContent:
|
||||
return HttpxBinaryResponseContent(
|
||||
response=httpx.Response(
|
||||
status_code=200,
|
||||
content=results_to_openai_jsonl(results),
|
||||
headers=(_JSONL_CONTENT_TYPE,),
|
||||
request=httpx.Request(method="GET", url=url),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class XAIBatchesHandler:
|
||||
def __init__(self, sync_client: HTTPHandler | None = None, async_client: AsyncHTTPHandler | None = None) -> None:
|
||||
self._sync_client = sync_client
|
||||
self._async_client = async_client
|
||||
|
||||
def _sync(self, timeout: float | httpx.Timeout) -> HTTPHandler:
|
||||
return self._sync_client or HTTPHandler(timeout=timeout)
|
||||
|
||||
def _async(self, timeout: float | httpx.Timeout) -> AsyncHTTPHandler:
|
||||
return self._async_client or get_async_httpx_client(
|
||||
llm_provider=LlmProviders.XAI,
|
||||
params={"timeout": timeout}, # mutable-ok: get_async_httpx_client takes a dict
|
||||
)
|
||||
|
||||
def create_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
create_batch_data: CreateBatchRequest,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]:
|
||||
url: Final = xai_batches_url(api_base)
|
||||
headers: Final = get_xai_auth_headers(api_key=api_key)
|
||||
body: Final = dict(to_create_batch_body(create_batch_data)) # mutable-ok: httpx json body
|
||||
endpoint: Final = create_batch_data.get("endpoint") or "/v1/chat/completions"
|
||||
if _is_async:
|
||||
|
||||
async def _acreate() -> LiteLLMBatch:
|
||||
response: Final = await self._async(timeout).post(url, json=body, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()), endpoint)
|
||||
|
||||
return _acreate()
|
||||
response: Final = self._sync(timeout).post(url, json=body, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()), endpoint)
|
||||
|
||||
def retrieve_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
batch_id: str,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]:
|
||||
url: Final = xai_batches_url(api_base, batch_id)
|
||||
headers: Final = get_xai_auth_headers(api_key=api_key)
|
||||
if _is_async:
|
||||
|
||||
async def _aretrieve() -> LiteLLMBatch:
|
||||
response: Final = await self._async(timeout).get(url, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
return _aretrieve()
|
||||
response: Final = self._sync(timeout).get(url, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
def cancel_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
batch_id: str,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]:
|
||||
url: Final = xai_batches_url(api_base, batch_id, suffix=":cancel")
|
||||
headers: Final = get_xai_auth_headers(api_key=api_key)
|
||||
if _is_async:
|
||||
|
||||
async def _acancel() -> LiteLLMBatch:
|
||||
response: Final = await self._async(timeout).post(url, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
return _acancel()
|
||||
response: Final = self._sync(timeout).post(url, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
def list_batches(
|
||||
self,
|
||||
_is_async: bool,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
) -> OpenAIBatchListResponse | Coroutine[None, None, OpenAIBatchListResponse]:
|
||||
url: Final = xai_batches_url(api_base)
|
||||
headers: Final = get_xai_auth_headers(api_key=api_key)
|
||||
params: Final = _results_params(after, limit)
|
||||
if _is_async:
|
||||
|
||||
async def _alist() -> OpenAIBatchListResponse:
|
||||
response: Final = await self._async(timeout).get(url, params=params, headers=headers, timeout=timeout)
|
||||
return to_openai_batch_list(XAIBatchList.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
return _alist()
|
||||
response: Final = self._sync(timeout).get(url, params=params, headers=headers, timeout=timeout)
|
||||
return to_openai_batch_list(XAIBatchList.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
def batch_results_content(
|
||||
self,
|
||||
_is_async: bool,
|
||||
batch_id: str,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
) -> HttpxBinaryResponseContent | Coroutine[None, None, HttpxBinaryResponseContent]:
|
||||
url: Final = xai_batches_url(api_base, batch_id, suffix="/results")
|
||||
headers: Final = get_xai_auth_headers(api_key=api_key)
|
||||
if _is_async:
|
||||
|
||||
async def _aresults() -> HttpxBinaryResponseContent:
|
||||
client: Final = self._async(timeout)
|
||||
|
||||
async def _page(after: str | None) -> XAIBatchResultsPage:
|
||||
response: Final = await client.get(
|
||||
url, params=_results_params(after, None), headers=headers, timeout=timeout
|
||||
)
|
||||
return XAIBatchResultsPage.model_validate(raise_for_xai_status(response).json())
|
||||
|
||||
pages = [await _page(None)] # mutable-ok: page walk terminates on the cursor, not on a fixed count
|
||||
while pages[-1].pagination_token and pages[-1].results:
|
||||
pages.append(await _page(pages[-1].pagination_token))
|
||||
return _jsonl_response(url, _flatten(pages))
|
||||
|
||||
return _aresults()
|
||||
client: Final = self._sync(timeout)
|
||||
|
||||
def _page(after: str | None) -> XAIBatchResultsPage:
|
||||
response: Final = client.get(url, params=_results_params(after, None), headers=headers, timeout=timeout)
|
||||
return XAIBatchResultsPage.model_validate(raise_for_xai_status(response).json())
|
||||
|
||||
pages = [_page(None)] # mutable-ok: page walk terminates on the cursor, not on a fixed count
|
||||
while pages[-1].pagination_token and pages[-1].results:
|
||||
pages.append(_page(pages[-1].pagination_token))
|
||||
return _jsonl_response(url, _flatten(pages))
|
||||
278
litellm/llms/xai/batches/transformation.py
Normal file
278
litellm/llms/xai/batches/transformation.py
Normal file
|
|
@ -0,0 +1,278 @@
|
|||
"""
|
||||
xAI Batch API reference: https://docs.x.ai/developers/advanced-api-usage/batch-api
|
||||
|
||||
xAI batches carry request counters, not a status, and no output file: results are paged from
|
||||
``GET /v1/batches/{id}/results``, so LiteLLM hands back the batch id as ``output_file_id``.
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import httpx
|
||||
from openai.types.batch import BatchRequestCounts
|
||||
from openai.types.batch import Errors as BatchErrors
|
||||
from openai.types.batch_error import BatchError
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm.constants import XAI_API_BASE
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.xai.common_utils import XAIModelInfo
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import CreateBatchRequest
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
OpenAIBatchStatus: TypeAlias = Literal[
|
||||
"validating", "failed", "in_progress", "finalizing", "completed", "expired", "cancelling", "cancelled"
|
||||
]
|
||||
|
||||
XAI_BATCH_ID_PREFIX: Final = "batch_"
|
||||
XAI_RESULTS_PAGE_SIZE: Final = 1000
|
||||
DEFAULT_BATCH_NAME: Final = "litellm-batch"
|
||||
DEFAULT_BATCH_ENDPOINT: Final = "/v1/chat/completions"
|
||||
_EMPTY_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
|
||||
|
||||
|
||||
class XAIBatchesError(BaseLLMException):
|
||||
pass
|
||||
|
||||
|
||||
def xai_batches_error(
|
||||
error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers
|
||||
) -> XAIBatchesError:
|
||||
return XAIBatchesError(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers if isinstance(headers, httpx.Headers) else httpx.Headers(tuple(headers.items())),
|
||||
)
|
||||
|
||||
|
||||
def raise_for_xai_status(response: httpx.Response) -> httpx.Response:
|
||||
if response.status_code >= 400:
|
||||
raise xai_batches_error(response.text, response.status_code, response.headers)
|
||||
return response
|
||||
|
||||
|
||||
def get_xai_api_base(api_base: str | None) -> str:
|
||||
resolved: Final = (api_base or get_secret_str("XAI_API_BASE") or XAI_API_BASE).rstrip("/")
|
||||
return resolved.removesuffix("/v1")
|
||||
|
||||
|
||||
def get_xai_auth_headers(
|
||||
headers: Mapping[str, str] = _EMPTY_HEADERS, api_key: str | None = None
|
||||
) -> dict[str, str]: # mutable-ok: BaseConfig.validate_environment contract returns dict
|
||||
resolved_key: Final = XAIModelInfo.get_api_key(api_key)
|
||||
if resolved_key is None:
|
||||
raise xai_batches_error(
|
||||
"Missing xAI API Key. Pass api_key, set litellm.xai_key or XAI_API_KEY", 401, _EMPTY_HEADERS
|
||||
)
|
||||
return dict(headers, Authorization=f"Bearer {resolved_key}") # mutable-ok: BaseConfig contract returns dict
|
||||
|
||||
|
||||
def xai_batches_url(api_base: str | None, batch_id: str | None = None, suffix: str = "") -> str:
|
||||
base: Final = f"{get_xai_api_base(api_base)}/v1/batches"
|
||||
if batch_id is None:
|
||||
return base
|
||||
return f"{base}/{encode_url_path_segment(batch_id, field_name='batch_id')}{suffix}"
|
||||
|
||||
|
||||
def is_xai_batch_results_id(file_id: str) -> bool:
|
||||
return file_id.startswith(XAI_BATCH_ID_PREFIX)
|
||||
|
||||
|
||||
class XAICreateBatchRequest(TypedDict):
|
||||
name: ReadOnly[str]
|
||||
input_file_id: NotRequired[ReadOnly[str]]
|
||||
|
||||
|
||||
class XAIBatchState(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
num_requests: int = 0
|
||||
num_pending: int = 0
|
||||
num_success: int = 0
|
||||
num_error: int = 0
|
||||
num_cancelled: int = 0
|
||||
|
||||
|
||||
class XAIBatch(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
batch_id: str
|
||||
name: str = ""
|
||||
create_time: str | None = None
|
||||
expire_time: str | None = None
|
||||
cancel_time: str | None = None
|
||||
cancel_by_xai_message: str | None = None
|
||||
state: XAIBatchState = XAIBatchState()
|
||||
input_file_id: str | None = None
|
||||
|
||||
|
||||
class XAIBatchList(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
batches: tuple[XAIBatch, ...] = ()
|
||||
pagination_token: str | None = None
|
||||
|
||||
|
||||
class XAIBatchResultError(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
code: int | str | None = None
|
||||
message: str = ""
|
||||
|
||||
|
||||
class XAIBatchResultData(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
response: Mapping[str, Mapping[str, object]] | None = None
|
||||
error: XAIBatchResultError | None = None
|
||||
|
||||
|
||||
class XAIBatchResult(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
batch_request_id: str
|
||||
batch_result: XAIBatchResultData = XAIBatchResultData()
|
||||
|
||||
|
||||
class XAIBatchResultsPage(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
results: tuple[XAIBatchResult, ...] = ()
|
||||
pagination_token: str | None = None
|
||||
|
||||
|
||||
def _to_unix_timestamp(value: str | None) -> int | None:
|
||||
"""xAI returns RFC 3339 timestamps over gRPC but a bare ``YYYY-MM-DD`` over REST."""
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
parsed: Final = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
except ValueError:
|
||||
return None
|
||||
return int((parsed if parsed.tzinfo is not None else parsed.replace(tzinfo=timezone.utc)).timestamp())
|
||||
|
||||
|
||||
def xai_batch_status(batch: XAIBatch) -> OpenAIBatchStatus:
|
||||
"""xAI exposes counters, not a status. A batch xAI itself cancelled (input validation failed) is a failure,
|
||||
a caller-cancelled batch is cancelled, an empty batch is still validating its input file, and a batch
|
||||
with nothing pending has completed."""
|
||||
if batch.cancel_time is not None:
|
||||
return "failed" if batch.cancel_by_xai_message else "cancelled"
|
||||
if batch.state.num_requests == 0:
|
||||
return "validating"
|
||||
if batch.state.num_pending > 0:
|
||||
return "in_progress"
|
||||
return "completed"
|
||||
|
||||
|
||||
def to_litellm_batch(batch: XAIBatch, endpoint: str = DEFAULT_BATCH_ENDPOINT) -> LiteLLMBatch:
|
||||
status: Final = xai_batch_status(batch)
|
||||
created_at: Final = _to_unix_timestamp(batch.create_time)
|
||||
cancelled_at: Final = _to_unix_timestamp(batch.cancel_time)
|
||||
errors: Final = (
|
||||
BatchErrors(object="list", data=[BatchError(message=batch.cancel_by_xai_message)]) # mutable-ok: openai type
|
||||
if batch.cancel_by_xai_message
|
||||
else None
|
||||
)
|
||||
return LiteLLMBatch(
|
||||
id=batch.batch_id,
|
||||
object="batch",
|
||||
endpoint=endpoint,
|
||||
input_file_id=batch.input_file_id or "",
|
||||
completion_window="24h",
|
||||
status=status,
|
||||
created_at=created_at if created_at is not None else 0,
|
||||
expires_at=_to_unix_timestamp(batch.expire_time),
|
||||
failed_at=cancelled_at if status == "failed" else None,
|
||||
cancelled_at=cancelled_at if status == "cancelled" else None,
|
||||
output_file_id=batch.batch_id if status == "completed" else None,
|
||||
errors=errors,
|
||||
request_counts=BatchRequestCounts(
|
||||
total=batch.state.num_requests,
|
||||
completed=batch.state.num_success,
|
||||
failed=batch.state.num_error + batch.state.num_cancelled,
|
||||
),
|
||||
metadata={"name": batch.name} if batch.name else None, # mutable-ok: LiteLLMBatch.metadata is a dict
|
||||
)
|
||||
|
||||
|
||||
class OpenAIBatchListResponse(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
object: Literal["list"] = "list"
|
||||
data: tuple[LiteLLMBatch, ...]
|
||||
first_id: str | None
|
||||
last_id: str | None
|
||||
has_more: bool
|
||||
next_page_token: str | None = None
|
||||
|
||||
|
||||
def to_openai_batch_list(page: XAIBatchList) -> OpenAIBatchListResponse:
|
||||
data: Final = tuple(to_litellm_batch(b) for b in page.batches)
|
||||
return OpenAIBatchListResponse(
|
||||
data=data,
|
||||
first_id=data[0].id if data else None,
|
||||
last_id=data[-1].id if data else None,
|
||||
has_more=bool(page.pagination_token),
|
||||
next_page_token=page.pagination_token or None,
|
||||
)
|
||||
|
||||
|
||||
def to_create_batch_body(create_batch_data: CreateBatchRequest) -> XAICreateBatchRequest:
|
||||
input_file_id: Final = create_batch_data.get("input_file_id")
|
||||
if not input_file_id:
|
||||
raise xai_batches_error("input_file_id is required to create an xAI batch", 400, _EMPTY_HEADERS)
|
||||
metadata: Final = create_batch_data.get("metadata")
|
||||
name: Final = metadata.get("name") if metadata else None
|
||||
return XAICreateBatchRequest(name=name or DEFAULT_BATCH_NAME, input_file_id=input_file_id)
|
||||
|
||||
|
||||
class OpenAIBatchOutputError(TypedDict):
|
||||
code: ReadOnly[str]
|
||||
message: ReadOnly[str]
|
||||
|
||||
|
||||
class OpenAIBatchOutputResponse(TypedDict):
|
||||
status_code: ReadOnly[int]
|
||||
request_id: ReadOnly[object]
|
||||
body: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class OpenAIBatchOutputLine(TypedDict):
|
||||
id: ReadOnly[str]
|
||||
custom_id: ReadOnly[str]
|
||||
response: ReadOnly[OpenAIBatchOutputResponse | None]
|
||||
error: ReadOnly[OpenAIBatchOutputError | None]
|
||||
|
||||
|
||||
def _result_to_openai_line(result: XAIBatchResult) -> OpenAIBatchOutputLine:
|
||||
"""One output JSONL line. xAI wraps the body in a one-key map named after the endpoint
|
||||
(``chat_get_completion``, ``responses``, ``image_generation``, ...); the value is the OpenAI body."""
|
||||
error: Final = result.batch_result.error
|
||||
response: Final = result.batch_result.response
|
||||
body: Final = next(iter(response.values()), None) if response else None
|
||||
if body is None:
|
||||
message: Final = error.message if error is not None else "xAI returned no response for this request"
|
||||
code: Final = str(error.code) if error is not None and error.code is not None else "request_failed"
|
||||
return OpenAIBatchOutputLine(
|
||||
id=f"batch_req_{result.batch_request_id}",
|
||||
custom_id=result.batch_request_id,
|
||||
response=None,
|
||||
error=OpenAIBatchOutputError(code=code, message=message),
|
||||
)
|
||||
return OpenAIBatchOutputLine(
|
||||
id=f"batch_req_{result.batch_request_id}",
|
||||
custom_id=result.batch_request_id,
|
||||
response=OpenAIBatchOutputResponse(status_code=200, request_id=body.get("id"), body=body),
|
||||
error=None,
|
||||
)
|
||||
|
||||
|
||||
def results_to_openai_jsonl(results: Sequence[XAIBatchResult]) -> bytes:
|
||||
return "".join(f"{json.dumps(_result_to_openai_line(r), ensure_ascii=False)}\n" for r in results).encode()
|
||||
|
|
@ -296,7 +296,7 @@ class XAIChatConfig(OpenAIGPTConfig):
|
|||
except Exception as e:
|
||||
verbose_logger.debug("Error extracting X.AI web search usage: %s", e)
|
||||
|
||||
self._fold_reasoning_tokens_into_completion(response)
|
||||
self.fold_reasoning_tokens_into_completion(response)
|
||||
self._normalize_openai_compatible_usage_totals(getattr(response, "usage", None))
|
||||
restated_usage: Final = _usage_restated_from_xai_ticks(getattr(response, "usage", None))
|
||||
if restated_usage is not None:
|
||||
|
|
@ -304,7 +304,7 @@ class XAIChatConfig(OpenAIGPTConfig):
|
|||
return response
|
||||
|
||||
@staticmethod
|
||||
def _fold_reasoning_tokens_into_completion(
|
||||
def fold_reasoning_tokens_into_completion(
|
||||
target: ModelResponse | Usage | dict[str, Any] | None,
|
||||
) -> None:
|
||||
"""Reconcile xAI Usage to the OpenAI invariant.
|
||||
|
|
@ -426,7 +426,7 @@ class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler):
|
|||
chunk["choices"] = [{"index": 0, "delta": {}, "finish_reason": None}]
|
||||
|
||||
if "usage" in chunk and chunk["usage"] is not None:
|
||||
XAIChatConfig._fold_reasoning_tokens_into_completion(chunk["usage"])
|
||||
XAIChatConfig.fold_reasoning_tokens_into_completion(chunk["usage"])
|
||||
XAIChatConfig._normalize_openai_compatible_usage_totals(chunk["usage"])
|
||||
|
||||
parsed_chunk: Final = super().chunk_parser(chunk)
|
||||
|
|
|
|||
247
litellm/llms/xai/files/transformation.py
Normal file
247
litellm/llms/xai/files/transformation.py
Normal file
|
|
@ -0,0 +1,247 @@
|
|||
"""
|
||||
xAI Files API reference: https://docs.x.ai/developers/rest-api-reference/inference/files
|
||||
|
||||
xAI stores ``purpose`` as an empty string; LiteLLM reports uploads as ``batch``, the only purpose xAI files serve.
|
||||
"""
|
||||
|
||||
import time
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.files.transformation import BaseFilesConfig, LiteLLMLoggingObj
|
||||
from litellm.types.llms.openai import (
|
||||
CreateFileRequest,
|
||||
FileContentRequest,
|
||||
HttpxBinaryResponseContent,
|
||||
OpenAICreateFileRequestOptionalParams,
|
||||
OpenAIFileObject,
|
||||
OpenAIFilesPurpose,
|
||||
)
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
from ..batches.transformation import (
|
||||
get_xai_api_base,
|
||||
get_xai_auth_headers,
|
||||
raise_for_xai_status,
|
||||
xai_batches_error,
|
||||
)
|
||||
|
||||
_NO_QUERY_PARAMS: Final[dict[str, str]] = {} # mutable-ok: BaseFilesConfig request transforms return tuple[str, dict]
|
||||
_DEFAULT_PURPOSE: Final[OpenAIFilesPurpose] = "batch"
|
||||
|
||||
|
||||
class XAIMultipartUpload(TypedDict):
|
||||
file: ReadOnly[tuple[str, object, str]]
|
||||
purpose: ReadOnly[tuple[None, str]]
|
||||
|
||||
|
||||
class XAIFile(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
id: str
|
||||
bytes: int = 0
|
||||
created_at: int | None = None
|
||||
filename: str = ""
|
||||
purpose: str = ""
|
||||
expires_at: int | None = None
|
||||
|
||||
|
||||
class XAIFileList(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
data: tuple[XAIFile, ...] = ()
|
||||
pagination_token: str | None = None
|
||||
|
||||
|
||||
class XAIFileDeleted(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
id: str
|
||||
deleted: bool = True
|
||||
|
||||
|
||||
def _to_openai_file_object(file: XAIFile) -> OpenAIFileObject:
|
||||
return OpenAIFileObject(
|
||||
id=file.id,
|
||||
bytes=file.bytes,
|
||||
created_at=file.created_at if file.created_at is not None else int(time.time()),
|
||||
filename=file.filename,
|
||||
object="file",
|
||||
purpose=_DEFAULT_PURPOSE,
|
||||
status="uploaded",
|
||||
expires_at=file.expires_at,
|
||||
)
|
||||
|
||||
|
||||
def _api_base_from(litellm_params: Mapping[str, object]) -> str:
|
||||
api_base: Final = litellm_params.get("api_base")
|
||||
return get_xai_api_base(api_base if isinstance(api_base, str) else None)
|
||||
|
||||
|
||||
class XAIFilesConfig(BaseFilesConfig):
|
||||
@property
|
||||
def custom_llm_provider(self) -> LlmProviders:
|
||||
return LlmProviders.XAI
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
model: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
stream: bool | None = None,
|
||||
) -> str:
|
||||
return f"{get_xai_api_base(api_base)}/v1/files"
|
||||
|
||||
def _file_url(self, file_id: str, litellm_params: Mapping[str, object], suffix: str = "") -> str:
|
||||
encoded_file_id: Final = encode_url_path_segment(file_id, field_name="file_id")
|
||||
return f"{_api_base_from(litellm_params)}/v1/files/{encoded_file_id}{suffix}"
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers
|
||||
) -> BaseLLMException:
|
||||
return xai_batches_error(error_message, status_code, headers)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Mapping[str, str],
|
||||
model: str,
|
||||
messages: Sequence[object],
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict[str, str]: # mutable-ok: BaseFilesConfig signature
|
||||
return get_xai_auth_headers(headers, api_key)
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> list[OpenAICreateFileRequestOptionalParams]: # mutable-ok: BaseFilesConfig signature
|
||||
return ["purpose"] # mutable-ok: BaseFilesConfig signature
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: Mapping[str, object],
|
||||
optional_params: dict[str, object], # mutable-ok: BaseConfig signature, returned as-is
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict[str, object]: # mutable-ok: BaseConfig signature
|
||||
return optional_params
|
||||
|
||||
def transform_create_file_request(
|
||||
self,
|
||||
model: str,
|
||||
create_file_data: CreateFileRequest,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> dict[str, object]: # mutable-ok: BaseFilesConfig signature
|
||||
if "file" not in create_file_data:
|
||||
raise ValueError("File data is required")
|
||||
extracted: Final = extract_file_data(create_file_data["file"])
|
||||
filename: Final = extracted["filename"] or f"file_{int(time.time())}.jsonl"
|
||||
content_type: Final = extracted.get("content_type") or "application/octet-stream"
|
||||
upload: Final = XAIMultipartUpload(
|
||||
file=(filename, extracted["content"], content_type),
|
||||
purpose=(None, create_file_data.get("purpose") or _DEFAULT_PURPOSE),
|
||||
)
|
||||
return dict(upload) # mutable-ok: BaseFilesConfig signature
|
||||
|
||||
def transform_create_file_response(
|
||||
self,
|
||||
model: str | None,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> OpenAIFileObject:
|
||||
return _to_openai_file_object(XAIFile.model_validate(raise_for_xai_status(raw_response).json()))
|
||||
|
||||
def transform_retrieve_file_request(
|
||||
self,
|
||||
file_id: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature
|
||||
return self._file_url(file_id, litellm_params), _NO_QUERY_PARAMS
|
||||
|
||||
def transform_retrieve_file_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> OpenAIFileObject:
|
||||
return _to_openai_file_object(XAIFile.model_validate(raise_for_xai_status(raw_response).json()))
|
||||
|
||||
def transform_delete_file_request(
|
||||
self,
|
||||
file_id: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature
|
||||
return self._file_url(file_id, litellm_params), _NO_QUERY_PARAMS
|
||||
|
||||
def transform_delete_file_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> FileDeleted:
|
||||
deleted: Final = XAIFileDeleted.model_validate(raise_for_xai_status(raw_response).json())
|
||||
return FileDeleted(id=deleted.id, deleted=deleted.deleted, object="file")
|
||||
|
||||
def transform_list_files_request(
|
||||
self,
|
||||
purpose: str | None,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature
|
||||
return f"{_api_base_from(litellm_params)}/v1/files", _NO_QUERY_PARAMS
|
||||
|
||||
def transform_list_files_next_request(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, str]] | None: # mutable-ok: BaseFilesConfig signature
|
||||
page: Final = XAIFileList.model_validate(raw_response.json())
|
||||
if not page.pagination_token or not page.data:
|
||||
return None
|
||||
return f"{_api_base_from(litellm_params)}/v1/files", {"pagination_token": page.pagination_token}
|
||||
|
||||
def transform_list_files_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> list[OpenAIFileObject]: # mutable-ok: BaseFilesConfig signature
|
||||
return [ # mutable-ok: BaseFilesConfig signature
|
||||
_to_openai_file_object(f)
|
||||
for f in XAIFileList.model_validate(raise_for_xai_status(raw_response).json()).data
|
||||
]
|
||||
|
||||
def transform_file_content_request(
|
||||
self,
|
||||
file_content_request: FileContentRequest,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature
|
||||
file_id: Final = file_content_request.get("file_id")
|
||||
if file_id is None:
|
||||
raise ValueError("file_id is required to download file content")
|
||||
return self._file_url(file_id, litellm_params, suffix="/content"), _NO_QUERY_PARAMS
|
||||
|
||||
def transform_file_content_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> HttpxBinaryResponseContent:
|
||||
return HttpxBinaryResponseContent(response=raw_response)
|
||||
|
|
@ -51436,13 +51436,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -51450,8 +51453,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
@ -51480,9 +51486,13 @@
|
|||
"xai/grok-4.3": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
|
|
@ -51490,6 +51500,8 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -51502,9 +51514,13 @@
|
|||
"xai/grok-4.3-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
|
|
@ -51512,6 +51528,8 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59483,13 +59501,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309-non-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59497,20 +59518,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent-0309": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": false,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59519,8 +59546,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -62787,13 +62817,16 @@
|
|||
},
|
||||
"xai/grok-4.20": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62801,21 +62834,27 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62823,21 +62862,27 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-reasoning-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62845,8 +62890,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
@ -63067,13 +63115,16 @@
|
|||
},
|
||||
"xai/grok-4.20-non-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -63081,20 +63132,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-non-reasoning-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -63102,20 +63159,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/responses"
|
||||
|
|
@ -63127,20 +63190,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/responses"
|
||||
|
|
@ -63152,8 +63221,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
|
|
@ -75459,13 +75531,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -75473,8 +75548,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ from typing import Final
|
|||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.litellm_core_utils.request_timeout_resolver import get_configured_request_timeout
|
||||
|
||||
DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS: Final = 600.0
|
||||
|
||||
_SECONDS: Final = TypeAdapter(float)
|
||||
|
|
@ -48,8 +50,8 @@ def resolve_llm_passthrough_timeout(
|
|||
Anthropic /v1/messages).
|
||||
|
||||
Non-streaming precedence: kwargs timeout/request_timeout -> litellm_params
|
||||
timeout/request_timeout -> router_timeout -> general_settings.pass_through_request_timeout
|
||||
-> 600s default.
|
||||
timeout/request_timeout -> router_timeout -> litellm.request_timeout (litellm_settings.request_timeout,
|
||||
when explicitly set) -> general_settings.pass_through_request_timeout -> 600s default.
|
||||
|
||||
Streaming (``kwargs["stream"]`` truthy) resolves ``stream_timeout`` at every level before
|
||||
any generic timeout, matching ``Router._get_stream_timeout`` on the completion route:
|
||||
|
|
@ -73,6 +75,7 @@ def resolve_llm_passthrough_timeout(
|
|||
deployment.get("timeout"),
|
||||
deployment.get("request_timeout"),
|
||||
router_timeout,
|
||||
get_configured_request_timeout(),
|
||||
)
|
||||
winner: Final = next((val for val in candidates if val is not None), None)
|
||||
return resolve_pass_through_request_timeout() if winner is None else _SECONDS.validate_python(winner)
|
||||
|
|
|
|||
145
litellm/proxy/_experimental/mcp_server/capabilities.py
Normal file
145
litellm/proxy/_experimental/mcp_server/capabilities.py
Normal file
|
|
@ -0,0 +1,145 @@
|
|||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from itertools import product
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
|
||||
from mcp.server.context import CallNext, HandlerResult, ServerRequestContext
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.types import DiscoverResult, InitializeRequestParams, InitializeResult, ServerCapabilities
|
||||
from mcp_types.methods import CLIENT_REQUESTS
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.types.mcp import MCP_LEGACY_VERSIONS, MCPAdvertisedVersions, MCPLegacyVersion, MCPSpecVersion, MCPTransport
|
||||
|
||||
GATEWAY_OPERATIONS: Final = frozenset(
|
||||
{
|
||||
"tools/list",
|
||||
"tools/call",
|
||||
"prompts/list",
|
||||
"prompts/get",
|
||||
"resources/list",
|
||||
"resources/read",
|
||||
"resources/templates/list",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RevisionSupport:
|
||||
transports: frozenset[MCPTransport]
|
||||
operations: frozenset[str]
|
||||
results: frozenset[Literal["complete", "input_required"]]
|
||||
extensions: frozenset[str]
|
||||
completed: bool
|
||||
|
||||
|
||||
REVISION_SUPPORT: Final[Mapping[str, RevisionSupport]] = MappingProxyType(
|
||||
{
|
||||
version.value: RevisionSupport(
|
||||
transports=frozenset(MCPTransport)
|
||||
if version.value in HANDSHAKE_PROTOCOL_VERSIONS
|
||||
else frozenset({MCPTransport.http, MCPTransport.stdio}),
|
||||
operations=frozenset(method for method in GATEWAY_OPERATIONS if (method, version.value) in CLIENT_REQUESTS),
|
||||
results=frozenset({"complete"})
|
||||
if version.value in HANDSHAKE_PROTOCOL_VERSIONS
|
||||
else frozenset({"complete", "input_required"}),
|
||||
extensions=frozenset(),
|
||||
completed=version.value in HANDSHAKE_PROTOCOL_VERSIONS,
|
||||
)
|
||||
for version in MCPSpecVersion
|
||||
}
|
||||
)
|
||||
_COMPLETED_REVISIONS: Final = tuple(version for version, support in REVISION_SUPPORT.items() if support.completed)
|
||||
TRANSLATION_PAIRS: Final = frozenset(product(_COMPLETED_REVISIONS, repeat=2))
|
||||
_ADVERTISED_VERSIONS: Final[TypeAdapter[tuple[MCPLegacyVersion, ...]]] = TypeAdapter(MCPAdvertisedVersions)
|
||||
|
||||
|
||||
def configured_versions() -> tuple[str, ...]:
|
||||
from litellm.proxy.proxy_server import general_settings_view
|
||||
|
||||
configured: Final = general_settings_view().get("mcp_advertised_versions")
|
||||
return _ADVERTISED_VERSIONS.validate_python(MCP_LEGACY_VERSIONS if configured is None else configured)
|
||||
|
||||
|
||||
def build_discovery(
|
||||
*,
|
||||
configured: tuple[str, ...],
|
||||
revision: str,
|
||||
transport: MCPTransport,
|
||||
authorized_operations: frozenset[str],
|
||||
upstream_versions: frozenset[str],
|
||||
capabilities: ServerCapabilities,
|
||||
client_extensions: frozenset[str] = frozenset(),
|
||||
upstream_extensions: frozenset[str] = frozenset(),
|
||||
instructions: str | None = None,
|
||||
) -> DiscoverResult:
|
||||
supported: Final = tuple(
|
||||
version
|
||||
for version, support in REVISION_SUPPORT.items()
|
||||
if version in configured and support.completed and transport in support.transports
|
||||
)
|
||||
revision_support: Final = REVISION_SUPPORT.get(revision)
|
||||
operations: Final[frozenset[str]] = (
|
||||
authorized_operations & revision_support.operations
|
||||
if revision in supported
|
||||
and revision_support is not None
|
||||
and any((revision, upstream) in TRANSLATION_PAIRS for upstream in upstream_versions)
|
||||
else frozenset()
|
||||
)
|
||||
extensions: Final[frozenset[str]] = (
|
||||
revision_support.extensions & client_extensions & upstream_extensions
|
||||
if operations and revision_support is not None
|
||||
else frozenset()
|
||||
)
|
||||
caller_capabilities: Final = capabilities.model_copy(deep=True)
|
||||
return DiscoverResult(
|
||||
supported_versions=list(supported),
|
||||
capabilities=ServerCapabilities(
|
||||
tools=caller_capabilities.tools if {"tools/list", "tools/call"} <= operations else None,
|
||||
prompts=caller_capabilities.prompts if {"prompts/list", "prompts/get"} <= operations else None,
|
||||
resources=caller_capabilities.resources if {"resources/list", "resources/read"} <= operations else None,
|
||||
extensions={
|
||||
key: value for key, value in (caller_capabilities.extensions or {}).items() if key in extensions
|
||||
}
|
||||
or None,
|
||||
),
|
||||
instructions=instructions,
|
||||
cache_scope="private",
|
||||
ttl_ms=0,
|
||||
)
|
||||
|
||||
|
||||
class GatewayVersionPolicy:
|
||||
def __init__(self, versions: Callable[[], tuple[str, ...]] = configured_versions) -> None:
|
||||
self._versions = versions
|
||||
|
||||
async def __call__(self, ctx: ServerRequestContext[object, object], call_next: CallNext) -> HandlerResult:
|
||||
versions: Final = self._versions()
|
||||
requested: Final = (
|
||||
InitializeRequestParams.model_validate(ctx.params or {}).protocol_version
|
||||
if ctx.method == "initialize"
|
||||
else ctx.protocol_version
|
||||
)
|
||||
negotiated: Final = (
|
||||
(requested if requested in HANDSHAKE_PROTOCOL_VERSIONS else LATEST_HANDSHAKE_VERSION)
|
||||
if ctx.method == "initialize"
|
||||
else requested
|
||||
)
|
||||
if negotiated not in versions:
|
||||
raise MCPError(code=-32022, message="Unsupported MCP protocol version", data={"supported": list(versions)})
|
||||
result: Final = await call_next(ctx)
|
||||
if ctx.method != "initialize":
|
||||
return result
|
||||
initialized: Final = InitializeResult.model_validate(result)
|
||||
discovery: Final = build_discovery(
|
||||
configured=versions,
|
||||
revision=initialized.protocol_version,
|
||||
transport=MCPTransport.http,
|
||||
authorized_operations=GATEWAY_OPERATIONS,
|
||||
upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS),
|
||||
capabilities=initialized.capabilities,
|
||||
instructions=initialized.instructions,
|
||||
)
|
||||
return initialized.model_copy(update={"capabilities": discovery.capabilities})
|
||||
|
|
@ -28,6 +28,7 @@ class OperationContext:
|
|||
client_ip: str | None = None
|
||||
mcp_proxy_mode: bool = False
|
||||
wire_compat: WireCompat = WireCompat.LEGACY
|
||||
protocol_version: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "_caller", copy_caller(self._caller))
|
||||
|
|
|
|||
|
|
@ -193,6 +193,7 @@ from litellm.types.mcp import (
|
|||
MCPAuth,
|
||||
MCPStdioConfig,
|
||||
MCPTokenEndpointAuthMethod,
|
||||
MCPUpstreamProtocol,
|
||||
has_header,
|
||||
without_header,
|
||||
)
|
||||
|
|
@ -340,6 +341,7 @@ class MCPServerConfig(TypedDict, total=False):
|
|||
whatever the admin wrote, and each read applies its own default."""
|
||||
|
||||
server_id: ReadOnly[str]
|
||||
protocol_version: ReadOnly[MCPUpstreamProtocol]
|
||||
alias: str
|
||||
description: str
|
||||
mcp_info: MCPInfo
|
||||
|
|
@ -2549,6 +2551,9 @@ class MCPServerManager:
|
|||
new_server = MCPServer(
|
||||
server_id=server_id,
|
||||
name=name_for_prefix,
|
||||
protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python(
|
||||
server_config.get("protocol_version", mcp_info.get("protocol_version", "auto"))
|
||||
),
|
||||
alias=alias,
|
||||
server_name=server_name,
|
||||
spec_path=server_config.get("spec_path", None),
|
||||
|
|
@ -3109,6 +3114,9 @@ class MCPServerManager:
|
|||
new_server: Final = MCPServer(
|
||||
server_id=mcp_server.server_id,
|
||||
name=name_for_prefix,
|
||||
protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python(
|
||||
_mcp_info.get("protocol_version", "auto")
|
||||
),
|
||||
alias=getattr(mcp_server, "alias", None),
|
||||
server_name=getattr(mcp_server, "server_name", None),
|
||||
url=mcp_server.url,
|
||||
|
|
@ -4145,6 +4153,7 @@ class MCPServerManager:
|
|||
cred_provider: UpstreamCredentialProvider | None = None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
protocol_version_override: MCPUpstreamProtocol | None = None,
|
||||
) -> MCPClient:
|
||||
"""
|
||||
Create an MCPClient instance for the given server.
|
||||
|
|
@ -4168,6 +4177,9 @@ class MCPServerManager:
|
|||
"""
|
||||
record_auth_resolution(server.server_id, AuthResolution.unresolved)
|
||||
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
|
||||
protocol_version: Final = (
|
||||
protocol_version_override if protocol_version_override is not None else resolved_server.protocol_version
|
||||
)
|
||||
transport: Final = resolved_server.transport or MCPTransport.sse
|
||||
spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server)
|
||||
provider: Final = cred_provider or self._cred_provider
|
||||
|
|
@ -4249,6 +4261,7 @@ class MCPServerManager:
|
|||
return MCPClient(
|
||||
server_url="", # Not used for stdio
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=resolved_server.auth_type,
|
||||
auth_value=auth_value,
|
||||
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
|
||||
|
|
@ -4281,6 +4294,7 @@ class MCPServerManager:
|
|||
MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=resolved_server.auth_type,
|
||||
timeout=(
|
||||
resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
|
|
@ -4324,6 +4338,7 @@ class MCPServerManager:
|
|||
MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=resolved_server.auth_type,
|
||||
auth_value=auth_value,
|
||||
auth_header_name=auth_header_name,
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ from mcp.types import (
|
|||
CallToolRequest,
|
||||
CallToolRequestParams,
|
||||
CallToolResult,
|
||||
DiscoverRequest,
|
||||
DiscoverResult,
|
||||
GetPromptRequest,
|
||||
GetPromptRequestParams,
|
||||
GetPromptResult,
|
||||
|
|
@ -28,10 +30,14 @@ from mcp.types import (
|
|||
ListToolsResult,
|
||||
PaginatedRequestParams,
|
||||
Prompt,
|
||||
PromptsCapability,
|
||||
ReadResourceRequest,
|
||||
ReadResourceRequestParams,
|
||||
ResourcesCapability,
|
||||
ResourceTemplate,
|
||||
ServerCapabilities,
|
||||
TextContent,
|
||||
ToolsCapability,
|
||||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl, ConfigDict, Field, TypeAdapter
|
||||
|
|
@ -51,6 +57,11 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
|
|||
cache_byok_credential,
|
||||
get_cached_byok_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.capabilities import (
|
||||
GATEWAY_OPERATIONS,
|
||||
build_discovery,
|
||||
configured_versions,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.contracts import (
|
||||
AuthorizedToolCall,
|
||||
OperationContext,
|
||||
|
|
@ -122,7 +133,9 @@ from litellm.proxy.litellm_pre_call_utils import (
|
|||
)
|
||||
from litellm.types.mcp import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
MCP_LEGACY_VERSIONS,
|
||||
MCPAuth,
|
||||
MCPTransport,
|
||||
without_header,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
|
||||
|
|
@ -2657,7 +2670,11 @@ class _McpDeniedDetail(TypedDict):
|
|||
|
||||
|
||||
async def _execute_handle_list_tools(
|
||||
context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None
|
||||
context: OperationContext,
|
||||
params: PaginatedRequestParams,
|
||||
host_progress_callback: ProgressCallback | None = None,
|
||||
*,
|
||||
log_list_tools_to_spendlogs: bool = True,
|
||||
) -> ListToolsResult:
|
||||
try:
|
||||
(
|
||||
|
|
@ -2700,7 +2717,7 @@ async def _execute_handle_list_tools(
|
|||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
log_list_tools_to_spendlogs=True,
|
||||
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
|
||||
list_tools_log_source="mcp_protocol",
|
||||
client_ip=_client_ip,
|
||||
)
|
||||
|
|
@ -3065,6 +3082,7 @@ def prepare_context(
|
|||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
protocol_version: str | None = None,
|
||||
) -> OperationContext:
|
||||
return OperationContext(
|
||||
_caller=user_api_key_auth,
|
||||
|
|
@ -3076,11 +3094,13 @@ def prepare_context(
|
|||
client_ip=client_ip,
|
||||
mcp_proxy_mode=mcp_proxy_mode,
|
||||
wire_compat=wire_compat,
|
||||
protocol_version=protocol_version,
|
||||
)
|
||||
|
||||
|
||||
GatewayOperation: TypeAlias = (
|
||||
AuthorizedToolCall
|
||||
| DiscoverRequest
|
||||
| ListToolsRequest
|
||||
| CallToolRequest
|
||||
| ListPromptsRequest
|
||||
|
|
@ -3090,7 +3110,8 @@ GatewayOperation: TypeAlias = (
|
|||
| ReadResourceRequest
|
||||
)
|
||||
GatewayResult: TypeAlias = (
|
||||
ListToolsResult
|
||||
DiscoverResult
|
||||
| ListToolsResult
|
||||
| CallToolResult
|
||||
| InputRequiredResult
|
||||
| ListPromptsResult
|
||||
|
|
@ -3105,6 +3126,9 @@ class GatewayOperations:
|
|||
def __init__(self, host_progress_callback: ProgressCallback | None = None) -> None:
|
||||
self._host_progress_callback = host_progress_callback
|
||||
|
||||
@overload
|
||||
async def execute(self, operation: DiscoverRequest, context: OperationContext) -> DiscoverResult: ...
|
||||
|
||||
@overload
|
||||
async def execute(
|
||||
self, operation: AuthorizedToolCall, context: OperationContext
|
||||
|
|
@ -3137,6 +3161,51 @@ class GatewayOperations:
|
|||
|
||||
async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
|
||||
match operation:
|
||||
case DiscoverRequest():
|
||||
listings: Final = (
|
||||
()
|
||||
if context.mcp_proxy_mode
|
||||
else (ListPromptsRequest(), ListResourcesRequest(), ListResourceTemplatesRequest())
|
||||
)
|
||||
tasks: Final = (
|
||||
asyncio.create_task(
|
||||
_execute_handle_list_tools(
|
||||
context,
|
||||
PaginatedRequestParams(),
|
||||
self._host_progress_callback,
|
||||
log_list_tools_to_spendlogs=False,
|
||||
)
|
||||
),
|
||||
*(asyncio.create_task(self.execute(listing, context)) for listing in listings),
|
||||
)
|
||||
try:
|
||||
results: Final = await asyncio.gather(*tasks)
|
||||
finally:
|
||||
for task in tasks:
|
||||
task.cancel()
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
return build_discovery(
|
||||
configured=configured_versions(),
|
||||
revision=context.protocol_version or "2025-11-25",
|
||||
transport=MCPTransport.http,
|
||||
authorized_operations=GATEWAY_OPERATIONS,
|
||||
upstream_versions=frozenset(MCP_LEGACY_VERSIONS),
|
||||
capabilities=ServerCapabilities(
|
||||
tools=ToolsCapability()
|
||||
if any(isinstance(result, ListToolsResult) and result.tools for result in results)
|
||||
else None,
|
||||
prompts=PromptsCapability()
|
||||
if any(isinstance(result, ListPromptsResult) and result.prompts for result in results)
|
||||
else None,
|
||||
resources=ResourcesCapability()
|
||||
if any(
|
||||
(isinstance(result, ListResourcesResult) and result.resources)
|
||||
or (isinstance(result, ListResourceTemplatesResult) and result.resource_templates)
|
||||
for result in results
|
||||
)
|
||||
else None,
|
||||
),
|
||||
)
|
||||
case AuthorizedToolCall():
|
||||
auth, token, _servers, server_headers, oauth_headers, headers, _client_ip = context.legacy_auth()
|
||||
return await _execute_mcp_tool(
|
||||
|
|
|
|||
|
|
@ -1375,7 +1375,16 @@ if MCP_AVAILABLE:
|
|||
and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY)
|
||||
else None
|
||||
)
|
||||
return _StagedServerTest(request=request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers)
|
||||
preview_request: Final = (
|
||||
request.model_copy(
|
||||
update={"mcp_info": {**(request.mcp_info or {}), "protocol_version": saved_server.protocol_version}}
|
||||
)
|
||||
if saved_server is not None and "protocol_version" not in (request.mcp_info or {})
|
||||
else request
|
||||
)
|
||||
return _StagedServerTest(
|
||||
request=preview_request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers
|
||||
)
|
||||
|
||||
async def _list_tools_within(client: MCPClient, deadline: float) -> list[MCPTool] | None:
|
||||
with anyio.move_on_after(deadline):
|
||||
|
|
@ -1512,6 +1521,7 @@ if MCP_AVAILABLE:
|
|||
extra_headers=merged_headers,
|
||||
stdio_env=stdio_env,
|
||||
cred_provider=preview_cred_provider,
|
||||
protocol_version_override=server_model.protocol_version,
|
||||
)
|
||||
|
||||
return await operation(client)
|
||||
|
|
|
|||
|
|
@ -125,12 +125,14 @@ def unsupported_protocol_version(scope: Scope) -> str | None:
|
|||
``HANDSHAKE_PROTOCOL_VERSIONS`` to the modern single-exchange path, which
|
||||
bypasses litellm's session/auth model, so the ASGI entry rejects it.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.capabilities import configured_versions
|
||||
|
||||
headers: Final[Iterable[tuple[bytes, bytes]]] = scope.get("headers") or ()
|
||||
values: Final = tuple(
|
||||
raw.decode("latin-1").strip() for key, raw in headers if key.lower() == _MCP_PROTOCOL_VERSION_HEADER
|
||||
)
|
||||
for value in values:
|
||||
if value and value not in HANDSHAKE_PROTOCOL_VERSIONS:
|
||||
if value and value not in configured_versions():
|
||||
return value
|
||||
return None
|
||||
|
||||
|
|
@ -149,7 +151,10 @@ try:
|
|||
from mcp.server.session import ServerSession as _McpServerSession
|
||||
from mcp.types import (
|
||||
BlobResourceContents,
|
||||
DiscoverRequest,
|
||||
DiscoverResult,
|
||||
GetPromptResult,
|
||||
RequestParams,
|
||||
ResourceTemplate,
|
||||
TextResourceContents,
|
||||
)
|
||||
|
|
@ -526,11 +531,11 @@ if MCP_AVAILABLE:
|
|||
PaginatedRequestParams,
|
||||
ReadResourceRequestParams,
|
||||
)
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import (
|
||||
MCPAuthenticatedUser,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.capabilities import GatewayVersionPolicy, configured_versions
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
global_mcp_server_manager,
|
||||
|
|
@ -585,6 +590,7 @@ if MCP_AVAILABLE:
|
|||
name=LITELLM_MCP_SERVER_NAME,
|
||||
version=LITELLM_MCP_SERVER_VERSION,
|
||||
)
|
||||
server.middleware.append(GatewayVersionPolicy())
|
||||
server.create_initialization_options = types.MethodType(_gateway_create_initialization_options, server)
|
||||
sse: Final[SseServerTransport] = SseServerTransport("/sse/messages")
|
||||
|
||||
|
|
@ -830,6 +836,7 @@ if MCP_AVAILABLE:
|
|||
client_ip,
|
||||
_mcp_proxy_mode.get(),
|
||||
wire_compat_for(ctx.protocol_version),
|
||||
ctx.protocol_version,
|
||||
)
|
||||
|
||||
async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult:
|
||||
|
|
@ -948,6 +955,11 @@ if MCP_AVAILABLE:
|
|||
ReadResourceRequest(params=params), context
|
||||
)
|
||||
|
||||
async def discover(ctx: ServerRequestContext, params: RequestParams) -> DiscoverResult:
|
||||
async with _legacy_operation_context(ctx, trace=False) as context:
|
||||
return await operations.GatewayOperations().execute(DiscoverRequest(params=params), context)
|
||||
|
||||
server.add_request_handler("server/discover", RequestParams, discover)
|
||||
server.add_request_handler("tools/list", PaginatedRequestParams, handle_list_tools)
|
||||
server.add_request_handler("tools/call", CallToolRequestParams, mcp_server_tool_call)
|
||||
server.add_request_handler("prompts/list", PaginatedRequestParams, list_prompts)
|
||||
|
|
@ -1954,7 +1966,7 @@ if MCP_AVAILABLE:
|
|||
reject_disallowed_mcp_origin(StarletteRequest(scope))
|
||||
bad_version: Final = unsupported_protocol_version(scope)
|
||||
if bad_version is not None:
|
||||
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
|
||||
supported: Final = ", ".join(configured_versions())
|
||||
await JSONResponse(
|
||||
status_code=400,
|
||||
content={ # mutable-ok: JSON-RPC error payload
|
||||
|
|
@ -2299,7 +2311,7 @@ if MCP_AVAILABLE:
|
|||
reject_disallowed_mcp_origin(StarletteRequest(scope))
|
||||
bad_version: Final = unsupported_protocol_version(scope)
|
||||
if bad_version is not None:
|
||||
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
|
||||
supported: Final = ", ".join(configured_versions())
|
||||
await JSONResponse(
|
||||
status_code=400,
|
||||
content={ # mutable-ok: JSON-RPC error payload
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ from litellm.types.llms.openai import (
|
|||
ResponsesAPIResponse,
|
||||
)
|
||||
from litellm.types.mcp import (
|
||||
MCPAdvertisedVersions,
|
||||
MCPAllowedClient,
|
||||
MCPAuth,
|
||||
MCPAuthType,
|
||||
|
|
@ -2998,6 +2999,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).",
|
||||
)
|
||||
mcp_advertised_versions: MCPAdvertisedVersions | None = Field(
|
||||
None,
|
||||
description="MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. "
|
||||
"Modern protocol serving and Apps/Tasks remain disabled.",
|
||||
)
|
||||
mcp_allowed_clients: list[MCPAllowedClient] | None = Field(
|
||||
None,
|
||||
description="MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted.",
|
||||
|
|
@ -4013,6 +4019,18 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
|
|||
],
|
||||
)
|
||||
|
||||
zerobus: CallbackOnUI = CallbackOnUI(
|
||||
litellm_callback_name="zerobus",
|
||||
ui_callback_name="Databricks Zerobus",
|
||||
litellm_callback_params=[ # mutable-ok: the registry field is typed list
|
||||
"ZEROBUS_WORKSPACE_URL",
|
||||
"ZEROBUS_SERVER_ENDPOINT",
|
||||
"ZEROBUS_CLIENT_ID",
|
||||
"ZEROBUS_CLIENT_SECRET",
|
||||
"ZEROBUS_TABLE_NAME",
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
class HTTPExceptionErrorDetail(TypedDict):
|
||||
"""The `{"error": <message>}` shape most proxy endpoints raise as `HTTPException.detail`."""
|
||||
|
|
@ -5349,11 +5367,12 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
default=False,
|
||||
description=(
|
||||
"When True, users whose JWT contains no team claims are authenticated "
|
||||
"using their database team memberships instead of receiving HTTP 403. "
|
||||
"Usage is attributed to the user's first resolvable DB team, or to the "
|
||||
"team specified via the x-litellm-team-id request header (validated "
|
||||
"against DB membership). Requires user_id_upsert=True so that user "
|
||||
"records exist before the fallback runs."
|
||||
"using their database team memberships instead of receiving HTTP 403, "
|
||||
"with usage attributed to the user's first resolvable DB team. Whether or "
|
||||
"not the JWT carries team claims, the x-litellm-team-id request header may "
|
||||
"select any team the user is a member of in the database (validated against "
|
||||
"DB membership); without the header the JWT team stays the default. Requires "
|
||||
"user_id_upsert=True so that user records exist before the fallback runs."
|
||||
),
|
||||
)
|
||||
issuers: list[JWTIssuerConfig] | None = Field(
|
||||
|
|
|
|||
|
|
@ -1930,12 +1930,12 @@ class JWTAuthManager:
|
|||
) -> HeaderTeam | None:
|
||||
"""
|
||||
The team named by x-litellm-team-id, which may carry a team id or a team
|
||||
alias. A value that is already an allowed team id (or, under the DB
|
||||
fallback, an existing team id) never costs an alias lookup; an alias is
|
||||
accepted only when the team it names would have been accepted by id.
|
||||
Under the DB fallback only a team row that is provably absent falls
|
||||
through to the alias lookup; a read that failed for any other reason
|
||||
keeps the membership denial the id path already gives.
|
||||
alias. A value that is already an allowed team id never costs a lookup;
|
||||
under the DB fallback any other value is accepted provisionally, by id
|
||||
or alias, for the membership check auth_builder runs later. Under the
|
||||
DB fallback only a team row that is provably absent falls through to
|
||||
the alias lookup; a read that failed for any other reason keeps the
|
||||
membership denial the id path already gives.
|
||||
|
||||
Raises:
|
||||
HTTPException: 403 when neither the value nor the team it aliases is
|
||||
|
|
@ -1948,7 +1948,11 @@ class JWTAuthManager:
|
|||
if not header_value:
|
||||
return None
|
||||
|
||||
if fallback_to_db_teams and not allowed_team_ids:
|
||||
if header_value in allowed_team_ids:
|
||||
verbose_proxy_logger.debug("Using team_id from x-litellm-team-id header: %s", header_value)
|
||||
return HeaderTeam(header_value=header_value, team_id=header_value)
|
||||
|
||||
if fallback_to_db_teams:
|
||||
try:
|
||||
await get_team_object(
|
||||
team_id=header_value,
|
||||
|
|
@ -1969,10 +1973,6 @@ class JWTAuthManager:
|
|||
JWTAuthManager._raise_header_team_membership_denial(header_value)
|
||||
return HeaderTeam(header_value=header_value, team_id=header_value)
|
||||
|
||||
if header_value in allowed_team_ids:
|
||||
verbose_proxy_logger.debug("Using team_id from x-litellm-team-id header: %s", header_value)
|
||||
return HeaderTeam(header_value=header_value, team_id=header_value)
|
||||
|
||||
team_id_by_alias: Final = await JWTAuthManager._team_id_by_alias(
|
||||
header_value, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj
|
||||
)
|
||||
|
|
@ -2353,9 +2353,9 @@ class JWTAuthManager:
|
|||
header_value: str,
|
||||
) -> None:
|
||||
"""
|
||||
A provisional team_id from the x-litellm-team-id header (accepted without
|
||||
JWT-team validation when the JWT carries no team claims) must exist in the
|
||||
user's DB team memberships before it becomes request context. The denial
|
||||
A provisional team_id from the x-litellm-team-id header (accepted under
|
||||
fallback_to_db_teams because it is outside the JWT's teams) must exist in
|
||||
the user's DB team memberships before it becomes request context. The denial
|
||||
names `header_value`, the id or alias the caller sent, not `team_id`.
|
||||
"""
|
||||
user_team_ids: Final = user_object.teams if user_object else []
|
||||
|
|
@ -2587,22 +2587,30 @@ class JWTAuthManager:
|
|||
if specific_team_id and not db_team_fallback:
|
||||
all_team_ids.add(specific_team_id)
|
||||
|
||||
header_db_fallback: Final = handler.litellm_jwtauth.fallback_to_db_teams and team_id is None
|
||||
|
||||
header_team: Final = await JWTAuthManager.resolve_team_from_header(
|
||||
request_headers=request_headers,
|
||||
allowed_team_ids=all_team_ids,
|
||||
fallback_to_db_teams=db_team_fallback,
|
||||
fallback_to_db_teams=header_db_fallback,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
provisional_header_team: Final = (
|
||||
header_team
|
||||
if header_team is not None and header_db_fallback and header_team.team_id not in all_team_ids
|
||||
else None
|
||||
)
|
||||
if header_team:
|
||||
team_id = header_team.team_id
|
||||
# A provisional header team (accepted only because the JWT carries no
|
||||
# team claims) is validated against DB membership further down; never
|
||||
# upsert it here or an attacker-supplied x-litellm-team-id would create
|
||||
# an orphaned team row before that check runs. A genuine membership team
|
||||
# already exists, so suppressing the upsert in that case costs nothing.
|
||||
# A provisional header team (accepted because it is outside the
|
||||
# JWT's teams under fallback_to_db_teams) is validated against DB
|
||||
# membership further down; never upsert it here or an
|
||||
# attacker-supplied x-litellm-team-id would create an orphaned team
|
||||
# row before that check runs. A genuine membership team already
|
||||
# exists, so suppressing the upsert in that case costs nothing.
|
||||
try:
|
||||
team_object = await get_team_object(
|
||||
team_id=team_id,
|
||||
|
|
@ -2610,10 +2618,10 @@ class JWTAuthManager:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_id_upsert=(team_id_upsert and not db_team_fallback),
|
||||
team_id_upsert=(team_id_upsert and provisional_header_team is None),
|
||||
)
|
||||
except HTTPException:
|
||||
if not db_team_fallback:
|
||||
if provisional_header_team is None:
|
||||
raise
|
||||
JWTAuthManager._raise_header_team_membership_denial(header_team.header_value)
|
||||
elif not team_id and not db_team_fallback:
|
||||
|
|
@ -2756,11 +2764,11 @@ class JWTAuthManager:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_id_upsert=team_id_upsert,
|
||||
)
|
||||
elif db_team_fallback and header_team is not None and team_id == header_team.team_id:
|
||||
elif provisional_header_team is not None and team_id == provisional_header_team.team_id:
|
||||
JWTAuthManager._validate_header_team_in_db_membership(
|
||||
team_id=team_id,
|
||||
user_object=user_object,
|
||||
header_value=header_team.header_value,
|
||||
header_value=provisional_header_team.header_value,
|
||||
)
|
||||
if not JWTAuthManager._is_team_route_allowed(
|
||||
route=route,
|
||||
|
|
@ -2770,7 +2778,7 @@ class JWTAuthManager:
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=(
|
||||
f"Team '{header_team.header_value}' (from x-litellm-team-id header) "
|
||||
f"Team '{provisional_header_team.header_value}' (from x-litellm-team-id header) "
|
||||
f"is not allowed to access route '{route}'."
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from starlette.types import Receive, Scope, Send
|
|||
import litellm
|
||||
from litellm._logging import redact_internal_details_from_client_message, verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.anthropic_interface.exceptions import AnthropicErrorSseFrame, anthropic_error_sse_frame
|
||||
from litellm.constants import (
|
||||
DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE,
|
||||
DEFAULT_MAX_RECURSE_DEPTH,
|
||||
|
|
@ -102,8 +103,11 @@ from litellm.proxy.common_utils.openai_error_payload import (
|
|||
)
|
||||
from litellm.proxy.common_utils.sse_keepalive import (
|
||||
SSE_COMMENT_PING_BYTES,
|
||||
SSE_STREAM_START_TAIL,
|
||||
advance_sse_tail,
|
||||
coerce_keepalive_interval,
|
||||
resolve_ttft_keepalive_interval,
|
||||
seal_open_sse_frame,
|
||||
wrap_sse_stream_with_keepalive_pings,
|
||||
)
|
||||
from litellm.proxy.dd_span_tagger import DDSpanTagger
|
||||
|
|
@ -999,6 +1003,17 @@ async def create_response(
|
|||
first_chunk_value = await _buffer_first_chunk_honoring_disconnect(generator, request)
|
||||
resolved_headers: Final = await _resolve_stream_headers(headers, refresh_headers)
|
||||
|
||||
if isinstance(first_chunk_value, AnthropicErrorSseFrame):
|
||||
with contextlib.suppress(Exception):
|
||||
await generator.aclose()
|
||||
return JSONResponse(
|
||||
status_code=first_chunk_value.status_code,
|
||||
content=first_chunk_value.json_body(
|
||||
error_body_call_id(general_settings, resolved_headers.get(LITELLM_CALL_ID_HEADER))
|
||||
),
|
||||
headers=resolved_headers,
|
||||
)
|
||||
|
||||
if first_chunk_value is not None:
|
||||
try:
|
||||
error_code_from_chunk: Final = await _parse_event_data_for_error(first_chunk_value)
|
||||
|
|
@ -2781,15 +2796,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
request=request,
|
||||
)
|
||||
if route_type == "aresponses":
|
||||
# Streaming /v1/responses returns here without
|
||||
# reaching the non-streaming ownership tail below.
|
||||
# Wrap the SSE generator so container ownership is
|
||||
# written once the upstream iterator finishes
|
||||
# assembling ``completed_response`` — otherwise
|
||||
# code-interpreter containers created during the
|
||||
# stream stay unregistered and follow-up file API
|
||||
# calls 403. Covers the background-polling path
|
||||
# too, which loops ``body_iterator`` end-to-end.
|
||||
selected_data_generator = (
|
||||
ProxyBaseLLMRequestProcessing._wrap_responses_stream_for_container_ownership(
|
||||
original_stream_response=response,
|
||||
|
|
@ -3011,50 +3017,50 @@ class ProxyBaseLLMRequestProcessing:
|
|||
wrapped_generator: Any,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
):
|
||||
"""Forward SSE chunks, then record container ownership at stream end.
|
||||
"""Forward SSE chunks and record container ownership before the terminal chunk goes out.
|
||||
|
||||
Streaming ``/v1/responses`` short-circuits out of
|
||||
``base_process_llm_request`` before the non-streaming ownership
|
||||
tail runs, so without this wrap the
|
||||
``LiteLLM_ManagedObjectTable`` row for any container created
|
||||
during the stream is never written and follow-up file API calls
|
||||
return 403.
|
||||
tail runs. The OpenAI SDK closes the connection at ``data: [DONE]``
|
||||
and starlette cancels the body task on disconnect, so a write that
|
||||
waits for the generator to finish never lands. The iterator sets
|
||||
``completed_response`` before it hands over its terminal chunk, so
|
||||
the ``LiteLLM_ManagedObjectTable`` row is written the moment it
|
||||
appears, ahead of the chunk carrying ``response.completed``.
|
||||
"""
|
||||
try:
|
||||
async for chunk in wrapped_generator:
|
||||
async for chunk in wrapped_generator:
|
||||
completed_obj = ProxyBaseLLMRequestProcessing._extract_completed_responses_response(
|
||||
original_stream_response
|
||||
)
|
||||
if completed_obj is None:
|
||||
yield chunk
|
||||
finally:
|
||||
try:
|
||||
completed_obj: Final = ProxyBaseLLMRequestProcessing._extract_completed_responses_response(
|
||||
original_stream_response
|
||||
)
|
||||
if completed_obj is not None:
|
||||
await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed(
|
||||
response=completed_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
else:
|
||||
# Silent skip caused #30210: the proxy's Router wrapper
|
||||
# of the responses streaming iterator wasn't propagating
|
||||
# ``completed_response``, so this hook recorded nothing
|
||||
# and follow-up /v1/containers/<id>/files calls 403'd
|
||||
# for non-admin keys with no proxy-side hint. Log a
|
||||
# warning so future regressions of the same shape
|
||||
# surface in operator logs.
|
||||
verbose_proxy_logger.warning(
|
||||
"Container ownership recording skipped on streaming "
|
||||
"/v1/responses: no completed_response on stream "
|
||||
"iterator %s. If this stream created any tool "
|
||||
"container (e.g. code_interpreter), follow-up "
|
||||
"/v1/containers/<id>/files calls will 403 for "
|
||||
"non-admin keys.",
|
||||
type(original_stream_response).__name__,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"Container ownership recording failed after streaming responses call: %s",
|
||||
e,
|
||||
)
|
||||
continue
|
||||
await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed(
|
||||
response=completed_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
yield chunk
|
||||
async for remaining_chunk in wrapped_generator:
|
||||
yield remaining_chunk
|
||||
return
|
||||
late_completed_obj: Final = ProxyBaseLLMRequestProcessing._extract_completed_responses_response(
|
||||
original_stream_response
|
||||
)
|
||||
if late_completed_obj is not None:
|
||||
await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed(
|
||||
response=late_completed_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
return
|
||||
verbose_proxy_logger.warning(
|
||||
"Container ownership recording skipped on streaming "
|
||||
"/v1/responses: no completed_response on stream "
|
||||
"iterator %s. If this stream created any tool "
|
||||
"container (e.g. code_interpreter), follow-up "
|
||||
"/v1/containers/<id>/files calls will 403 for "
|
||||
"non-admin keys.",
|
||||
type(original_stream_response).__name__,
|
||||
)
|
||||
|
||||
async def base_passthrough_process_llm_request(
|
||||
self,
|
||||
|
|
@ -3861,6 +3867,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
serialize_error: StreamErrorSerializer,
|
||||
request: Request | None = None,
|
||||
flush_tail: Callable[[], bytes] | None = None,
|
||||
seal_open_frame: Callable[[bytes], str] | None = None,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
Shared streaming data generator: runs proxy iterator hook, per-chunk hook,
|
||||
|
|
@ -3870,6 +3877,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
``flush_tail`` runs once after the upstream iterator completes cleanly and
|
||||
its non-empty result is yielded, so a serializer that buffers bytes across
|
||||
chunks can emit anything still held at end of stream.
|
||||
|
||||
``seal_open_frame`` is given the tail of what has been yielded when the
|
||||
error frame goes out, and what it returns is written first. A passthrough
|
||||
relays raw upstream bytes, so an upstream that hangs up mid-frame leaves the
|
||||
client inside an open frame, where an error frame would be swallowed or
|
||||
misparsed instead of raised.
|
||||
"""
|
||||
verbose_proxy_logger.debug("inside generator")
|
||||
# Resolve per-stream (not per-chunk) whether the heavy per-chunk path
|
||||
|
|
@ -3886,6 +3899,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
stream_completed = False
|
||||
client_disconnected = False
|
||||
delivered_chunk = False
|
||||
recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the yielded bytes
|
||||
try:
|
||||
str_so_far = ""
|
||||
async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook(
|
||||
|
|
@ -3931,7 +3945,9 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# False and refunds. A keepalive ping carries no provider output,
|
||||
# so it must not suppress that refund.
|
||||
delivered_chunk = delivered_chunk or chunk != STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
yield serialize_chunk(chunk)
|
||||
serialized = serialize_chunk(chunk)
|
||||
recent_tail = advance_sse_tail(recent_tail, serialized)
|
||||
yield serialized
|
||||
held_tail: Final = flush_tail() if flush_tail is not None else b""
|
||||
if held_tail:
|
||||
yield serialize_chunk(held_tail)
|
||||
|
|
@ -3979,7 +3995,9 @@ class ProxyBaseLLMRequestProcessing:
|
|||
code=stream_error_status,
|
||||
)
|
||||
stream_completed = True
|
||||
yield serialize_error(proxy_exception)
|
||||
error_frame: Final = serialize_error(proxy_exception)
|
||||
seal: Final = "" if seal_open_frame is None else seal_open_frame(recent_tail)
|
||||
yield seal + error_frame if seal else error_frame
|
||||
finally:
|
||||
await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup(
|
||||
request=request,
|
||||
|
|
@ -4001,7 +4019,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
restamp_model: str | None = None,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
Anthropic /messages and Google /generateContent streaming data generator require SSE events.
|
||||
Anthropic /messages streaming data generator, which requires SSE events.
|
||||
|
||||
Returns the underlying ``async_streaming_data_generator`` configured with
|
||||
SSE serializers directly (rather than re-wrapping it in another
|
||||
|
|
@ -4019,11 +4037,13 @@ class ProxyBaseLLMRequestProcessing:
|
|||
request_data=request_data,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
serialize_chunk=ProxyBaseLLMRequestProcessing._sse_chunk_serializer(restamper),
|
||||
serialize_error=lambda proxy_exc: (
|
||||
f"{STREAM_SSE_DATA_PREFIX}{json.dumps({'error': proxy_exc.to_dict()})}\n\n"
|
||||
serialize_error=lambda proxy_exc: anthropic_error_sse_frame(
|
||||
status_code=error_status_code(proxy_exc, status.HTTP_500_INTERNAL_SERVER_ERROR),
|
||||
raw_message=proxy_exc.message,
|
||||
),
|
||||
request=request,
|
||||
flush_tail=None if restamper is None else restamper.flush,
|
||||
seal_open_frame=seal_open_sse_frame,
|
||||
)
|
||||
|
||||
@overload
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ SSE_COMMENT_PING_BYTES: Final = SSE_COMMENT_PING.encode()
|
|||
# terminates a line with CRLF, LF or CR, so a blank line is any of these three.
|
||||
_SSE_FRAME_DELIMITERS: Final = (b"\r\n\r\n", b"\n\n", b"\r\r")
|
||||
_SSE_DELIMITER_LOOKBACK: Final = max(len(delimiter) for delimiter in _SSE_FRAME_DELIMITERS)
|
||||
_STREAM_START_TAIL: Final = b"\n\n"
|
||||
SSE_STREAM_START_TAIL: Final = b"\n\n"
|
||||
_SSE_MEDIA_TYPE: Final = "text/event-stream"
|
||||
|
||||
|
||||
|
|
@ -128,7 +128,7 @@ async def _keepalive_ping_byte_stream(
|
|||
# Seeded as a delimiter because a stream starts at a frame boundary, and kept
|
||||
# across chunks because a delimiter can be split between two transport reads,
|
||||
# which testing only the latest chunk would miss for the rest of the stream.
|
||||
recent_tail = _STREAM_START_TAIL # rebind-ok: rolling window over the relayed bytes
|
||||
recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the relayed bytes
|
||||
try:
|
||||
while True:
|
||||
await asyncio.wait((pending,), timeout=ping_interval_seconds)
|
||||
|
|
@ -155,6 +155,28 @@ async def _keepalive_ping_byte_stream(
|
|||
await stream.aclose()
|
||||
|
||||
|
||||
def advance_sse_tail(recent_tail: bytes, chunk: object) -> bytes:
|
||||
written: Final = _sse_tail_bytes(chunk)
|
||||
if not written:
|
||||
return recent_tail
|
||||
return (recent_tail + written)[-_SSE_DELIMITER_LOOKBACK:]
|
||||
|
||||
|
||||
def _sse_tail_bytes(chunk: object) -> bytes:
|
||||
if isinstance(chunk, bytes):
|
||||
return chunk[-_SSE_DELIMITER_LOOKBACK:]
|
||||
if isinstance(chunk, str):
|
||||
return chunk[-_SSE_DELIMITER_LOOKBACK:].encode()
|
||||
return b""
|
||||
|
||||
|
||||
def seal_open_sse_frame(recent_tail: bytes) -> str:
|
||||
if recent_tail.endswith(_SSE_FRAME_DELIMITERS):
|
||||
return ""
|
||||
line_break: Final = "" if recent_tail.endswith((b"\n", b"\r")) else "\n"
|
||||
return f"{line_break}{ANTHROPIC_PING_SSE_CHUNK}"
|
||||
|
||||
|
||||
def resolve_ttft_keepalive_interval(
|
||||
deployments: Iterable[Mapping[str, object]],
|
||||
global_interval: float | str | None,
|
||||
|
|
|
|||
13
litellm/proxy/common_utils/validation_error_body.py
Normal file
13
litellm/proxy/common_utils/validation_error_body.py
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
from collections.abc import Sequence
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
|
||||
class ValidationErrorDetail(TypedDict):
|
||||
type: ReadOnly[str]
|
||||
loc: ReadOnly[tuple[int | str, ...]]
|
||||
msg: ReadOnly[str]
|
||||
|
||||
|
||||
def public_validation_errors(errors: Sequence[ValidationErrorDetail]) -> tuple[ValidationErrorDetail, ...]:
|
||||
return tuple(ValidationErrorDetail(type=error["type"], loc=error["loc"], msg=error["msg"]) for error in errors)
|
||||
|
|
@ -8,8 +8,8 @@ from fastapi import Request
|
|||
from fastapi.dependencies.utils import get_flat_params
|
||||
from fastapi.params import ParamTypes
|
||||
from fastapi.responses import JSONResponse
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.proxy.common_utils.validation_error_body import ValidationErrorDetail
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import (
|
||||
ListLinks,
|
||||
PageLinks,
|
||||
|
|
@ -58,14 +58,6 @@ def escape_like(value: str) -> str:
|
|||
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
|
||||
|
||||
|
||||
class ValidationErrorDetail(TypedDict):
|
||||
"""The keys of a pydantic/FastAPI validation error a problem document needs."""
|
||||
|
||||
type: ReadOnly[str]
|
||||
loc: ReadOnly[tuple[int | str, ...]]
|
||||
msg: ReadOnly[str]
|
||||
|
||||
|
||||
def _is_length_error_of_rejected_items(error: ValidationErrorDetail, errors: Sequence[ValidationErrorDetail]) -> bool:
|
||||
"""pydantic counts only items that validated, so a bad item also trips the parent's min_length."""
|
||||
return error["type"] == "too_short" and any(
|
||||
|
|
|
|||
|
|
@ -476,6 +476,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
|
|||
project_spend_counter_key,
|
||||
tag_cache_key,
|
||||
)
|
||||
from litellm.proxy.common_utils.validation_error_body import public_validation_errors
|
||||
from litellm.proxy.config_resolvers import (
|
||||
FieldSource,
|
||||
SettingsStore,
|
||||
|
|
@ -551,7 +552,6 @@ from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_sp
|
|||
from litellm.proxy.image_endpoints.endpoints import router as image_router
|
||||
from litellm.proxy.list_api.common import (
|
||||
ManagementProblem,
|
||||
ValidationErrorDetail,
|
||||
problem_response,
|
||||
request_validation_problem,
|
||||
)
|
||||
|
|
@ -1983,16 +1983,14 @@ class _ExceptionRow(TypedDict, total=False):
|
|||
|
||||
@app.exception_handler(RequestValidationError)
|
||||
async def otel_request_validation_exception_handler(request: Request, exc: RequestValidationError):
|
||||
public_errors: Final = public_validation_errors(exc.errors())
|
||||
public_exc: Final = RequestValidationError(public_errors).with_traceback(exc.__traceback__)
|
||||
if request.url.path.startswith(MANAGEMENT_V1_PREFIX):
|
||||
validation_errors: Final[Sequence[ValidationErrorDetail]] = exc.errors()
|
||||
problem: Final = request_validation_problem(validation_errors)
|
||||
_close_dangling_otel_server_span(request, problem.status, exc=exc)
|
||||
problem: Final = request_validation_problem(public_errors)
|
||||
_close_dangling_otel_server_span(request, problem.status, exc=public_exc)
|
||||
return problem_response(problem)
|
||||
_close_dangling_otel_server_span(request, 422, exc=exc)
|
||||
return JSONResponse(
|
||||
status_code=422,
|
||||
content={"detail": jsonable_encoder(exc.errors())},
|
||||
)
|
||||
_close_dangling_otel_server_span(request, 422, exc=public_exc)
|
||||
return JSONResponse(status_code=422, content={"detail": public_errors})
|
||||
|
||||
|
||||
@app.exception_handler(Exception)
|
||||
|
|
@ -6342,6 +6340,11 @@ class ProxyConfig:
|
|||
if general_settings is None:
|
||||
general_settings = {}
|
||||
|
||||
if general_settings.get("mcp_advertised_versions") is not None:
|
||||
from litellm.types.mcp import MCPAdvertisedVersions
|
||||
|
||||
TypeAdapter(MCPAdvertisedVersions).validate_python(general_settings["mcp_advertised_versions"])
|
||||
|
||||
if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None:
|
||||
warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1"))
|
||||
if declared_proxy_ranges(general_settings) is None:
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys
|
|||
from litellm.proxy._experimental.mcp_server.tool_search import MCP_TOOL_SEARCH_SETTINGS_KEY
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.validation_error_body import public_validation_errors
|
||||
from litellm.proxy.config_resolvers import FieldSource, SettingsStore, source_for
|
||||
from litellm.proxy.config_resolvers.settings_store import ConfigOwnedKeyError
|
||||
from litellm.proxy.config_resolvers.sso import (
|
||||
|
|
@ -1875,7 +1876,7 @@ async def update_ui_settings(
|
|||
try:
|
||||
settings: Final = effective_cls.model_validate(settings_body)
|
||||
except ValidationError as e:
|
||||
raise HTTPException(status_code=422, detail=e.errors())
|
||||
raise HTTPException(status_code=422, detail=public_validation_errors(e.errors()))
|
||||
|
||||
unsupported_team_fields: Final = sorted(
|
||||
frozenset(settings.team_admin_editable_team_fields) - SUPPORTED_TEAM_ADMIN_PERMISSIONS
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
|
|||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, overload, runtime_checkable
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Protocol, overload, runtime_checkable
|
||||
|
||||
import httpx
|
||||
from openai._streaming import SSEDecoder
|
||||
|
|
@ -265,6 +265,9 @@ def _mid_stream_fallback_eligible(mapped_exception: Exception) -> bool:
|
|||
return not isinstance(status_code, int) or status_code >= 500 or status_code == 429
|
||||
|
||||
|
||||
_PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: Final = frozenset({"response.created", "response.in_progress", "response.queued"})
|
||||
|
||||
|
||||
class BaseResponsesAPIStreamingIterator:
|
||||
"""
|
||||
Base class for streaming iterators that process responses from the Responses API.
|
||||
|
|
@ -292,6 +295,7 @@ class BaseResponsesAPIStreamingIterator:
|
|||
self.start_time = getattr(logging_obj, "start_time", datetime.now())
|
||||
self._failure_handled = False # Track if failure handler has been called
|
||||
self._yielded_first_chunk = False
|
||||
self._output_started = False
|
||||
self._generated_content = ""
|
||||
self._generated_tool_arguments = ""
|
||||
self._completed_response_cached = False
|
||||
|
|
@ -879,6 +883,46 @@ class BaseResponsesAPIStreamingIterator:
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
def _note_yielded_event(self, event: ResponsesAPIStreamingResponse) -> None:
|
||||
self._yielded_first_chunk = True
|
||||
if event.type not in _PRE_OUTPUT_LIFECYCLE_EVENT_TYPES:
|
||||
self._output_started = True
|
||||
|
||||
def _fallback_error(self, original: Exception) -> MidStreamFallbackError:
|
||||
return MidStreamFallbackError(
|
||||
message=str(original),
|
||||
model=self.model or "",
|
||||
llm_provider=self.custom_llm_provider or "",
|
||||
original_exception=original,
|
||||
generated_content="",
|
||||
is_pre_first_chunk=not self._yielded_first_chunk,
|
||||
)
|
||||
|
||||
def _stream_ended_early_error(self) -> litellm.APIConnectionError:
|
||||
return litellm.APIConnectionError(
|
||||
message=(
|
||||
f"{self.custom_llm_provider or 'provider'} closed the responses stream before any terminal event "
|
||||
"(response.completed, response.incomplete or response.failed)"
|
||||
),
|
||||
llm_provider=self.custom_llm_provider or "",
|
||||
model=self.model or "",
|
||||
)
|
||||
|
||||
def _raise_if_ended_without_terminal_event(self) -> None:
|
||||
if self.completed_response is not None:
|
||||
return
|
||||
error: Final = self._stream_ended_early_error()
|
||||
self._handle_failure(error)
|
||||
if self._output_started:
|
||||
raise error
|
||||
raise self._fallback_error(error) from error
|
||||
|
||||
def _raise_for_transport_error(self, error: httpx.ReadError | httpx.RemoteProtocolError) -> NoReturn:
|
||||
self._handle_failure(error)
|
||||
if self._output_started:
|
||||
raise error
|
||||
raise self._fallback_error(error) from error
|
||||
|
||||
|
||||
async def call_post_streaming_hooks_for_testing(
|
||||
iterator: object, chunk: ResponsesAPIStreamingResponse
|
||||
|
|
@ -934,12 +978,14 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
sse = await self.stream_iterator.__anext__()
|
||||
except StopAsyncIteration:
|
||||
self.finished = True
|
||||
self._raise_if_ended_without_terminal_event()
|
||||
raise StopAsyncIteration
|
||||
|
||||
self._check_max_streaming_duration()
|
||||
result = self._process_chunk(sse.data)
|
||||
|
||||
if self.finished:
|
||||
self._raise_if_ended_without_terminal_event()
|
||||
raise StopAsyncIteration
|
||||
elif result is not None:
|
||||
self._maybe_raise_for_error_event(result)
|
||||
|
|
@ -948,7 +994,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
result = await self._call_post_streaming_deployment_hook(
|
||||
chunk=result,
|
||||
)
|
||||
self._yielded_first_chunk = True
|
||||
self._note_yielded_event(result)
|
||||
return result
|
||||
# If result is None, continue the loop to get the next chunk
|
||||
|
||||
|
|
@ -957,10 +1003,9 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
raise
|
||||
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
|
||||
self.finished = True
|
||||
if self.completed_response is None:
|
||||
self._handle_failure(e)
|
||||
raise
|
||||
raise StopAsyncIteration from e
|
||||
if self.completed_response is not None:
|
||||
raise StopAsyncIteration from e
|
||||
self._raise_for_transport_error(e)
|
||||
except httpx.HTTPError as e:
|
||||
# Handle HTTP errors
|
||||
self.finished = True
|
||||
|
|
@ -1016,12 +1061,14 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
sse = next(self.stream_iterator)
|
||||
except StopIteration:
|
||||
self.finished = True
|
||||
self._raise_if_ended_without_terminal_event()
|
||||
raise StopIteration
|
||||
|
||||
self._check_max_streaming_duration()
|
||||
result = self._process_chunk(sse.data)
|
||||
|
||||
if self.finished:
|
||||
self._raise_if_ended_without_terminal_event()
|
||||
raise StopIteration
|
||||
elif result is not None:
|
||||
self._maybe_raise_for_error_event(result)
|
||||
|
|
@ -1030,7 +1077,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
async_function=self._call_post_streaming_deployment_hook,
|
||||
chunk=result,
|
||||
)
|
||||
self._yielded_first_chunk = True
|
||||
self._note_yielded_event(result)
|
||||
return result
|
||||
# If result is None, continue the loop to get the next chunk
|
||||
|
||||
|
|
@ -1039,10 +1086,9 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
raise
|
||||
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
|
||||
self.finished = True
|
||||
if self.completed_response is None:
|
||||
self._handle_failure(e)
|
||||
raise
|
||||
raise StopIteration from e
|
||||
if self.completed_response is not None:
|
||||
raise StopIteration from e
|
||||
self._raise_for_transport_error(e)
|
||||
except httpx.HTTPError as e:
|
||||
# Handle HTTP errors
|
||||
self.finished = True
|
||||
|
|
|
|||
|
|
@ -502,12 +502,12 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
|
||||
response_cost: Final[float] = standard_logging_payload.get("response_cost", 0)
|
||||
model_id: Final[str] = str(standard_logging_payload.get("model_id", ""))
|
||||
custom_llm_provider: Final[str] = kwargs.get("litellm_params", {}).get("custom_llm_provider", None)
|
||||
if custom_llm_provider is None:
|
||||
raise ValueError("custom_llm_provider is required")
|
||||
custom_llm_provider: Final[str | None] = standard_logging_payload.get("custom_llm_provider")
|
||||
|
||||
budget_config: Final = self._get_budget_config_for_provider(custom_llm_provider)
|
||||
if budget_config:
|
||||
budget_config: Final = (
|
||||
self._get_budget_config_for_provider(custom_llm_provider) if custom_llm_provider is not None else None
|
||||
)
|
||||
if custom_llm_provider is not None and budget_config is not None:
|
||||
# increment spend for provider
|
||||
spend_key: Final = f"provider_spend:{custom_llm_provider}:{budget_config.budget_duration}"
|
||||
start_time_key: Final = f"provider_budget_start_time:{custom_llm_provider}"
|
||||
|
|
|
|||
|
|
@ -202,15 +202,26 @@ def capability_classifier_system_prompt(mode: Literal["json_schema", "json_objec
|
|||
)
|
||||
|
||||
|
||||
def unwrap_classifier_json(content: str) -> str:
|
||||
"""Remove the optional Markdown fence without repairing or weakening verdict JSON."""
|
||||
text: Final = content.strip()
|
||||
if not text.startswith("```"):
|
||||
return text
|
||||
unfenced: Final = text.removeprefix("```").removeprefix("json").lstrip("\n\r")
|
||||
return unfenced.removesuffix("```").strip()
|
||||
_JSON_DECODER: Final = json.JSONDecoder()
|
||||
|
||||
|
||||
def _complete_json_object_at(content: str, start: int) -> str | None:
|
||||
try:
|
||||
_, end = _JSON_DECODER.raw_decode(content, start)
|
||||
except (ValueError, RecursionError):
|
||||
return None
|
||||
return content[start:end]
|
||||
|
||||
|
||||
def extract_classifier_json(content: str) -> str:
|
||||
"""Return the first complete JSON object in the reply, whatever prose or fence surrounds it.
|
||||
|
||||
A reply with no complete object comes back stripped so the caller's validation names the defect."""
|
||||
object_starts: Final = (index for index, char in enumerate(content) if char == "{")
|
||||
candidates: Final = (_complete_json_object_at(content, start) for start in object_starts)
|
||||
return next((candidate for candidate in candidates if candidate is not None), content.strip())
|
||||
|
||||
|
||||
def parse_capability_classifier_verdict(content: str) -> CapabilityClassifierVerdict:
|
||||
"""Parse raw JSON or the fenced JSON shape tolerated by Switchyard."""
|
||||
return CapabilityClassifierVerdict.model_validate_json(unwrap_classifier_json(content))
|
||||
"""Parse the verdict object out of a bare, fenced, or prose-wrapped reply."""
|
||||
return CapabilityClassifierVerdict.model_validate_json(extract_classifier_json(content))
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ from types import MappingProxyType
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast
|
||||
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError, create_model
|
||||
from pydantic_core import ErrorDetails
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.affinity_cache import claim_affinity_pin
|
||||
|
|
@ -85,8 +86,8 @@ from .capability_classifier import (
|
|||
CapabilityClassifierForecast,
|
||||
capability_classifier_response_format,
|
||||
capability_classifier_system_prompt,
|
||||
extract_classifier_json,
|
||||
parse_capability_classifier_verdict,
|
||||
unwrap_classifier_json,
|
||||
)
|
||||
from .classification_rubrics import BUSINESS_TIER_CRITERIA, calibration_examples_section
|
||||
from .config import (
|
||||
|
|
@ -427,6 +428,41 @@ def _effective_turn_off_message_logging(request_kwargs: Mapping[str, object] | N
|
|||
)
|
||||
|
||||
|
||||
def _classifier_reply_is_private(request_kwargs: Mapping[str, object] | None) -> bool:
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
|
||||
initialize_standard_callback_dynamic_params,
|
||||
)
|
||||
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
|
||||
|
||||
kwargs: Final = dict(request_kwargs) if request_kwargs else {}
|
||||
try:
|
||||
return should_redact_message_logging(
|
||||
{
|
||||
"litellm_params": kwargs,
|
||||
"standard_callback_dynamic_params": initialize_standard_callback_dynamic_params(kwargs),
|
||||
}
|
||||
)
|
||||
except AttributeError:
|
||||
return True
|
||||
|
||||
|
||||
def _validation_problem(detail: ErrorDetails) -> str:
|
||||
location: Final = ".".join(str(part) for part in detail["loc"])
|
||||
return f"{location}: {detail['msg']}" if location else detail["msg"]
|
||||
|
||||
|
||||
def _log_rejected_classifier_verdict(
|
||||
error: ValidationError, content: str, request_kwargs: Mapping[str, object] | None
|
||||
) -> None:
|
||||
problems: Final = "; ".join(_validation_problem(detail) for detail in error.errors())
|
||||
reply: Final = (
|
||||
"raw reply withheld (message logging is off)"
|
||||
if _classifier_reply_is_private(request_kwargs)
|
||||
else f"raw reply: {content!r}"
|
||||
)
|
||||
verbose_router_logger.warning("ComplexityRouter: classifier verdict rejected (%s); %s", problems, reply)
|
||||
|
||||
|
||||
_REMINDER_OPEN: Final = "<system-reminder>"
|
||||
_REMINDER_CLOSE: Final = "</system-reminder>"
|
||||
_DEFAULT_REMINDER_MARKERS: Final = ((_REMINDER_OPEN, _REMINDER_CLOSE),)
|
||||
|
|
@ -2040,7 +2076,7 @@ class ComplexityRouter(CustomLogger):
|
|||
except Exception as e: # noqa: BLE001 -- every unavailable or invalid judge verdict must fail closed
|
||||
if breaker is not None and permit is not None:
|
||||
breaker.record_failure(permit, is_timeout=_is_classifier_timeout(e))
|
||||
return self._capability_classifier_failure_outcome(f"capability classifier failed ({e})")
|
||||
return self._capability_classifier_failure_outcome(f"capability classifier failed ({type(e).__name__})")
|
||||
|
||||
def _capability_classifier_failure_outcome(self, reason: str, signal: str | None = None) -> ClassificationOutcome:
|
||||
"""Fail closed to the configured capable tier without consulting another taxonomy."""
|
||||
|
|
@ -2449,7 +2485,11 @@ class ComplexityRouter(CustomLogger):
|
|||
content, classifier_cost = await self._call_classifier_model(
|
||||
messages_for_call, request_kwargs, encrypted_task=encrypted_task
|
||||
)
|
||||
raw_tier: Final = _LabeledTierClassification.model_validate_json(content).tier
|
||||
try:
|
||||
raw_tier: Final = _LabeledTierClassification.model_validate_json(extract_classifier_json(content)).tier
|
||||
except ValidationError as error:
|
||||
_log_rejected_classifier_verdict(error, content, request_kwargs)
|
||||
raise
|
||||
tier: Final = self.config.resolve_classified_tier(raw_tier)
|
||||
if tier is None:
|
||||
raise ValueError(f"LLM classifier returned an unrecognized tier: {raw_tier!r}")
|
||||
|
|
@ -2508,7 +2548,11 @@ class ComplexityRouter(CustomLogger):
|
|||
max_output_tokens=capability.max_output_tokens,
|
||||
encrypted_task=encrypted_task,
|
||||
)
|
||||
verdict: Final = parse_capability_classifier_verdict(content)
|
||||
try:
|
||||
verdict: Final = parse_capability_classifier_verdict(content)
|
||||
except ValidationError as error:
|
||||
_log_rejected_classifier_verdict(error, content, request_kwargs)
|
||||
raise
|
||||
threshold: Final = verdict.routing_threshold(capability.base_threshold, capability.threshold_step)
|
||||
calibration: Final = capability.calibration
|
||||
forecast: Final = CapabilityClassifierForecast(
|
||||
|
|
@ -2563,8 +2607,9 @@ class ComplexityRouter(CustomLogger):
|
|||
messages_for_call, request_kwargs, encrypted_task=encrypted, max_output_tokens=v2.max_output_tokens
|
||||
)
|
||||
try:
|
||||
verdict: Final = LLMV2Verdict.model_validate_json(unwrap_classifier_json(content))
|
||||
except ValidationError:
|
||||
verdict: Final = LLMV2Verdict.model_validate_json(extract_classifier_json(content))
|
||||
except ValidationError as error:
|
||||
_log_rejected_classifier_verdict(error, content, request_kwargs)
|
||||
return self._classifier_failure_outcome("Invalid LLM V2 forecast", prompt, system_prompt)._replace(
|
||||
classifier_cost=classifier_cost
|
||||
)
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from litellm.llms.base_llm.base_utils import (
|
|||
from litellm.router_strategy.complexity_router.fuse_presets import ProfileText, resolve_fuse_profile
|
||||
|
||||
ShortText: TypeAlias = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1, max_length=512)]
|
||||
VerdictText: TypeAlias = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)]
|
||||
|
||||
|
||||
class _SolverProfile(TypedDict):
|
||||
|
|
@ -90,7 +91,7 @@ class LLMV2Demands(BaseModel):
|
|||
class LLMV2SolverForecast(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid", frozen=True)
|
||||
|
||||
likely_failure: ShortText
|
||||
likely_failure: VerdictText
|
||||
p_solve: StrictFloat = Field(ge=0.0, le=1.0)
|
||||
|
||||
|
||||
|
|
@ -104,7 +105,7 @@ class LLMV2SolverForecasts(BaseModel):
|
|||
class LLMV2Verdict(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid", frozen=True)
|
||||
|
||||
crux: ShortText
|
||||
crux: VerdictText
|
||||
demands: LLMV2Demands
|
||||
verification: Literal["relevant", "partial", "unavailable", "unknown"]
|
||||
forecasts: LLMV2SolverForecasts
|
||||
|
|
|
|||
53
litellm/types/integrations/zerobus.py
Normal file
53
litellm/types/integrations/zerobus.py
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
from dataclasses import dataclass, field
|
||||
from typing import Final
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams
|
||||
|
||||
RETRYABLE_INGEST_STATUS_CODES: Final = frozenset({408, 429, 500, 502, 503, 504})
|
||||
|
||||
TOKEN_REFRESH_LEEWAY_SECONDS: Final = 60
|
||||
|
||||
|
||||
class ZerobusInitParams(StandardCustomLoggerInitParams):
|
||||
"""
|
||||
Params for initializing a Databricks Zerobus logger on litellm.
|
||||
|
||||
Every connection field falls back to its ``ZEROBUS_*`` environment variable, which is
|
||||
what the proxy UI writes. ``table_name`` is the fully qualified ``catalog.schema.table``.
|
||||
"""
|
||||
|
||||
workspace_url: str | None = None
|
||||
server_endpoint: str | None = None
|
||||
client_id: str | None = None
|
||||
client_secret: str | None = None
|
||||
table_name: str | None = None
|
||||
batch_size: int = Field(default=100, gt=0)
|
||||
flush_interval: int = Field(default=10, gt=0)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ZerobusConnection:
|
||||
"""Everything needed to mint a token for one table and post rows to it."""
|
||||
|
||||
workspace_url: str
|
||||
workspace_id: str
|
||||
server_endpoint: str
|
||||
client_id: str
|
||||
client_secret: str = field(repr=False)
|
||||
table_name: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ZerobusAccessToken:
|
||||
value: str = field(repr=False)
|
||||
expires_at: float
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ZerobusIngestFailure:
|
||||
"""Why a batch could not be written, and whether a later attempt could still succeed."""
|
||||
|
||||
detail: str
|
||||
retryable: bool
|
||||
|
|
@ -522,7 +522,19 @@ class CreateBatchRequest(TypedDict, total=False):
|
|||
"""
|
||||
|
||||
completion_window: Literal["24h"]
|
||||
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"]
|
||||
endpoint: Literal[
|
||||
"/v1/chat/completions",
|
||||
"/v1/embeddings",
|
||||
"/v1/completions",
|
||||
"/v1/responses",
|
||||
"/v1/ocr",
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits",
|
||||
"/v1/videos/generations",
|
||||
"/v1/videos",
|
||||
"/v1/videos/edits",
|
||||
"/v1/videos/extensions",
|
||||
]
|
||||
input_file_id: str
|
||||
metadata: dict[str, str] | None
|
||||
output_expires_after: FileExpiresAfter
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import enum
|
|||
import re
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
|
|
@ -34,6 +34,8 @@ class MCPSpecVersion(str, enum.Enum):
|
|||
nov_2024 = "2024-11-05"
|
||||
mar_2025 = "2025-03-26"
|
||||
jun_2025 = "2025-06-18"
|
||||
nov_2025 = "2025-11-25"
|
||||
jul_2026 = "2026-07-28"
|
||||
|
||||
|
||||
class MCPAuth(str, enum.Enum):
|
||||
|
|
@ -59,7 +61,17 @@ DEFAULT_SUBJECT_TOKEN_TYPE: Final = "urn:ietf:params:oauth:token-type:access_tok
|
|||
|
||||
# MCP Literals
|
||||
MCPTransportType = Literal[MCPTransport.sse, MCPTransport.http, MCPTransport.stdio]
|
||||
MCPSpecVersionType = Literal[MCPSpecVersion.nov_2024, MCPSpecVersion.mar_2025, MCPSpecVersion.jun_2025]
|
||||
MCPLegacyVersion = Literal["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"]
|
||||
MCP_LEGACY_VERSIONS: Final[tuple[MCPLegacyVersion, ...]] = ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25")
|
||||
MCPUpstreamProtocol = MCPLegacyVersion | Literal["auto"]
|
||||
MCPAdvertisedVersions = Annotated[tuple[MCPLegacyVersion, ...], Field(min_length=1)]
|
||||
MCPSpecVersionType = Literal[
|
||||
MCPSpecVersion.nov_2024,
|
||||
MCPSpecVersion.mar_2025,
|
||||
MCPSpecVersion.jun_2025,
|
||||
MCPSpecVersion.nov_2025,
|
||||
MCPSpecVersion.jul_2026,
|
||||
]
|
||||
MCPAuthType = (
|
||||
Literal[
|
||||
MCPAuth.none,
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from datetime import datetime
|
||||
from typing import Any, Final, Literal
|
||||
from typing import Annotated, Any, Final, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
from pydantic import AfterValidator, BaseModel, ConfigDict, Field, TypeAdapter, field_validator, model_validator
|
||||
from typing_extensions import Self
|
||||
|
||||
from litellm.types.mcp import (
|
||||
|
|
@ -10,11 +10,19 @@ from litellm.types.mcp import (
|
|||
MCPAuthType,
|
||||
MCPTokenEndpointAuthMethod,
|
||||
MCPTransportType,
|
||||
MCPUpstreamProtocol,
|
||||
normalize_upstream_header_name,
|
||||
)
|
||||
|
||||
|
||||
# MCPInfo now allows arbitrary additional fields for custom metadata
|
||||
MCPInfo = dict[str, Any]
|
||||
def _validate_mcp_protocol_metadata(value: dict[str, object]) -> dict[str, object]:
|
||||
if "protocol_version" in value:
|
||||
TypeAdapter(MCPUpstreamProtocol).validate_python(value["protocol_version"])
|
||||
return value
|
||||
|
||||
|
||||
MCPInfo = Annotated[dict[str, Any], AfterValidator(_validate_mcp_protocol_metadata)]
|
||||
|
||||
|
||||
class MCPOAuthMetadata(BaseModel):
|
||||
|
|
@ -66,6 +74,7 @@ class MCPServer(BaseModel):
|
|||
server_name: str | None = None
|
||||
url: str | None = None
|
||||
transport: MCPTransportType
|
||||
protocol_version: MCPUpstreamProtocol = "auto"
|
||||
spec_path: str | None = None
|
||||
auth_type: MCPAuthType | None = None
|
||||
authentication_token: str | None = None
|
||||
|
|
@ -246,6 +255,14 @@ class MCPServer(BaseModel):
|
|||
"""
|
||||
return self.oauth2_flow == "client_credentials"
|
||||
|
||||
@model_validator(mode="after")
|
||||
def resolve_protocol_version(self) -> Self:
|
||||
if "protocol_version" not in self.model_fields_set and self.mcp_info is not None:
|
||||
self.protocol_version = TypeAdapter(MCPUpstreamProtocol).validate_python(
|
||||
self.mcp_info.get("protocol_version", "auto")
|
||||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_identity_binding_mode(self) -> Self:
|
||||
binding: Final = self.oauth_identity_binding
|
||||
|
|
|
|||
|
|
@ -345,6 +345,7 @@ class CredentialLiteLLMParams(BaseModel):
|
|||
|
||||
## OBJECT STORAGE (files / batches) ##
|
||||
gcs_bucket_name: str | None = None
|
||||
bucket_name: str | None = None
|
||||
|
||||
## AWS BEDROCK / SAGEMAKER ##
|
||||
aws_access_key_id: str | None = None
|
||||
|
|
|
|||
|
|
@ -299,6 +299,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
cache_read_input_token_cost_above_272k_tokens_flex: float | None
|
||||
cache_read_input_token_cost_above_512k_tokens: float | None
|
||||
cache_read_input_token_cost_batches: ReadOnly[float | None]
|
||||
cache_read_input_token_cost_above_200k_tokens_batches: ReadOnly[float | None]
|
||||
cache_read_input_token_cost_above_272k_tokens_batches: ReadOnly[float | None]
|
||||
cache_creation_input_token_cost_batches: ReadOnly[float | None]
|
||||
cache_creation_input_token_cost_above_272k_tokens_batches: ReadOnly[float | None]
|
||||
|
|
@ -327,8 +328,10 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
input_cost_per_second: float | None # for OpenAI Speech models
|
||||
input_cost_per_token_batches: float | None
|
||||
input_cost_per_video_token_batches: ReadOnly[float | None]
|
||||
input_cost_per_token_above_200k_tokens_batches: ReadOnly[float | None]
|
||||
input_cost_per_token_above_272k_tokens_batches: ReadOnly[float | None]
|
||||
output_cost_per_token_batches: float | None
|
||||
output_cost_per_token_above_200k_tokens_batches: ReadOnly[float | None]
|
||||
output_cost_per_token_above_272k_tokens_batches: ReadOnly[float | None]
|
||||
output_cost_per_token: Required[float | None]
|
||||
output_cost_per_token_flex: float | None # OpenAI flex service tier pricing
|
||||
|
|
@ -3729,6 +3732,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
cache_read_input_token_cost_above_272k_tokens_priority: float | None = None
|
||||
cache_read_input_token_cost_above_272k_tokens_flex: float | None = None
|
||||
cache_read_input_token_cost_batches: float | None = None
|
||||
cache_read_input_token_cost_above_200k_tokens_batches: float | None = None
|
||||
cache_read_input_token_cost_above_272k_tokens_batches: float | None = None
|
||||
cache_creation_input_token_cost_batches: float | None = None
|
||||
cache_creation_input_token_cost_above_272k_tokens_batches: float | None = None
|
||||
|
|
@ -3742,6 +3746,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
input_cost_per_token_above_200k_tokens_priority: float | None = None
|
||||
input_cost_per_token_above_272k_tokens_priority: float | None = None
|
||||
input_cost_per_token_above_272k_tokens_flex: float | None = None
|
||||
input_cost_per_token_above_200k_tokens_batches: float | None = None
|
||||
input_cost_per_token_above_272k_tokens_batches: float | None = None
|
||||
input_cost_per_query: float | None = None
|
||||
input_cost_per_image: float | None = None
|
||||
|
|
@ -3766,6 +3771,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
output_cost_per_token_above_200k_tokens_priority: float | None = None
|
||||
output_cost_per_token_above_272k_tokens_priority: float | None = None
|
||||
output_cost_per_token_above_272k_tokens_flex: float | None = None
|
||||
output_cost_per_token_above_200k_tokens_batches: float | None = None
|
||||
output_cost_per_token_above_272k_tokens_batches: float | None = None
|
||||
output_cost_per_character_above_128k_tokens: float | None = None
|
||||
output_cost_per_image: float | None = None
|
||||
|
|
@ -4140,7 +4146,7 @@ FILE_CONTENT_STREAMING_PROVIDERS: Final[frozenset[str]] = frozenset(
|
|||
|
||||
LITELLM_EXECUTED_BATCH_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.HOSTED_VLLM.value})
|
||||
|
||||
ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai"]
|
||||
ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai", "xai"]
|
||||
|
||||
LIST_BATCHES_SUPPORTED_PROVIDERS: Final[frozenset[str]] = frozenset(get_args(ListBatchesSupportedProvider))
|
||||
|
||||
|
|
|
|||
|
|
@ -6160,6 +6160,9 @@ def _get_model_info_helper(
|
|||
cache_read_input_token_cost_priority=_model_info.get("cache_read_input_token_cost_priority", None),
|
||||
cache_read_input_token_cost_ultrafast=_model_info.get("cache_read_input_token_cost_ultrafast", None),
|
||||
cache_read_input_token_cost_batches=_model_info.get("cache_read_input_token_cost_batches"),
|
||||
cache_read_input_token_cost_above_200k_tokens_batches=_model_info.get(
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches"
|
||||
),
|
||||
cache_read_input_token_cost_above_272k_tokens_batches=_model_info.get(
|
||||
"cache_read_input_token_cost_above_272k_tokens_batches"
|
||||
),
|
||||
|
|
@ -6197,10 +6200,16 @@ def _get_model_info_helper(
|
|||
input_cost_per_video_per_second=_model_info.get("input_cost_per_video_per_second", None),
|
||||
input_cost_per_token_batches=_model_info.get("input_cost_per_token_batches"),
|
||||
input_cost_per_video_token_batches=_model_info.get("input_cost_per_video_token_batches", None),
|
||||
input_cost_per_token_above_200k_tokens_batches=_model_info.get(
|
||||
"input_cost_per_token_above_200k_tokens_batches"
|
||||
),
|
||||
input_cost_per_token_above_272k_tokens_batches=_model_info.get(
|
||||
"input_cost_per_token_above_272k_tokens_batches"
|
||||
),
|
||||
output_cost_per_token_batches=_model_info.get("output_cost_per_token_batches"),
|
||||
output_cost_per_token_above_200k_tokens_batches=_model_info.get(
|
||||
"output_cost_per_token_above_200k_tokens_batches"
|
||||
),
|
||||
output_cost_per_token_above_272k_tokens_batches=_model_info.get(
|
||||
"output_cost_per_token_above_272k_tokens_batches"
|
||||
),
|
||||
|
|
@ -9357,6 +9366,10 @@ class ProviderConfigManager:
|
|||
from litellm.llms.mistral.files.transformation import MistralFilesConfig
|
||||
|
||||
return MistralFilesConfig()
|
||||
elif LlmProviders.XAI == provider:
|
||||
from litellm.llms.xai.files.transformation import XAIFilesConfig
|
||||
|
||||
return XAIFilesConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -51436,13 +51436,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -51450,8 +51453,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
@ -51480,9 +51486,13 @@
|
|||
"xai/grok-4.3": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
|
|
@ -51490,6 +51500,8 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -51502,9 +51514,13 @@
|
|||
"xai/grok-4.3-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
|
|
@ -51512,6 +51528,8 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59483,13 +59501,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309-non-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59497,20 +59518,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent-0309": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": false,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59519,8 +59546,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -62787,13 +62817,16 @@
|
|||
},
|
||||
"xai/grok-4.20": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62801,21 +62834,27 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62823,21 +62862,27 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-reasoning-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62845,8 +62890,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
@ -63067,13 +63115,16 @@
|
|||
},
|
||||
"xai/grok-4.20-non-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -63081,20 +63132,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-non-reasoning-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -63102,20 +63159,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/responses"
|
||||
|
|
@ -63127,20 +63190,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/responses"
|
||||
|
|
@ -63152,8 +63221,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
|
|
@ -75459,13 +75531,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -75473,8 +75548,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
|
|||
|
|
@ -170,6 +170,11 @@
|
|||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"cache_read_input_token_cost_above_200k_tokens_priority": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
|
|
@ -351,6 +356,11 @@
|
|||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"input_cost_per_token_above_200k_tokens_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"input_cost_per_token_above_200k_tokens_priority": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
|
|
@ -708,6 +718,11 @@
|
|||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"output_cost_per_token_above_200k_tokens_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"output_cost_per_token_above_200k_tokens_priority": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
|
|
|
|||
|
|
@ -249,6 +249,7 @@ proxy-dev = [
|
|||
"prisma==0.11.0",
|
||||
"hypercorn==0.17.3",
|
||||
"prometheus-client==0.20.0",
|
||||
"sentry-sdk==2.21.0",
|
||||
"opentelemetry-api==1.33.1",
|
||||
"opentelemetry-sdk==1.33.1",
|
||||
"opentelemetry-exporter-otlp==1.33.1",
|
||||
|
|
|
|||
|
|
@ -72,6 +72,7 @@ IGNORE_FUNCTIONS = [
|
|||
"_string_leaves", # bounded by the nesting depth of a safe_json_structure output (a finite JSON tree, no cycles possible).
|
||||
"_replace_string_leaves", # bounded by the nesting depth of a safe_json_structure output (a finite JSON tree, no cycles possible).
|
||||
"_sort_processed_sets", # bounded by the nesting depth of the log-record extra it walks (a finite JSON tree, no cycles possible).
|
||||
"scrub_json_strings", # max depth set (MAX_SCRUB_DEPTH); fails closed by returning "[Filtered]" for anything nested past the cap.
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -86,6 +86,7 @@
|
|||
- {id: llm.responses.azure_openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Responses w/ Azure OpenAI (smoke)"}
|
||||
- {id: llm.responses.azure_openai.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Responses tool calls w/ Azure OpenAI"}
|
||||
- {id: llm.responses.azure_openai.code_interpreter.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: code_interpreter, streaming: nonstream, assertions: [works], source: "llm_translation/test_containers_e2e.py", rationale: "An implicit code_interpreter container on an Azure deployment that carries its own api_base must serve GET /v1/containers/{id}/files/{fid}/content by its native cntr_ id to a team service-account key, the customer's shape (#27921, #28990)", fail_before_fix: proven}
|
||||
- {id: llm.responses.azure_openai.code_interpreter.stream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: code_interpreter, streaming: stream, assertions: [works], source: "llm_translation/test_containers_e2e.py", rationale: "A container created by a streamed /v1/responses code_interpreter call must serve /v1/containers/{id}/files to the same service-account key right after the OpenAI SDK closes at [DONE]; the ownership row used to be written after the stream and the disconnect cancelled it (LIT-8612)", fail_before_fix: proven}
|
||||
- {id: llm.chat_completions.together_ai.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning surfaces as reasoning_content (LIT-5960)"}
|
||||
- {id: llm.chat_completions.together_ai.thinking.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: stream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning deltas stream as reasoning_content"}
|
||||
- {id: llm.chat_completions.together_ai.thinking.nonstream.template_kwargs_forwarded, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [template_kwargs_forwarded], source: "llm_translation/test_together_ai_e2e.py", rationale: "chat_template_kwargs reaches Together and turns thinking off"}
|
||||
|
|
@ -103,6 +104,8 @@
|
|||
- {id: llm.chat_completions.anthropic.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic over /chat/completions: cost header and spend row agree"}
|
||||
- {id: llm.chat_completions.anthropic.multi_turn.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic tool result round trip over /chat/completions"}
|
||||
- {id: llm.messages.anthropic.multi_turn.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: anthropic, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic tool result round trip over /v1/messages"}
|
||||
- {id: llm.messages.anthropic.upstream_stream_failure.stream.error_event, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: upstream_stream_failure, streaming: stream, assertions: [error_event], source: "customer report", rationale: "An upstream that hangs up mid-stream must reach Anthropic clients as an event: error frame, not an OpenAI-shaped data-only error they silently drop"}
|
||||
- {id: llm.messages.anthropic.upstream_stream_failure.stream.error_status, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: upstream_stream_failure, streaming: stream, assertions: [error_status], source: "customer report", rationale: "An upstream that hangs up before its first byte must answer as a JSON error carrying its status, so Anthropic clients raise the status-specific error and retry on it instead of reading a 200 stream that only carries an error event"}
|
||||
- {id: llm.messages.openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI models served on the Anthropic Messages contract"}
|
||||
- {id: llm.messages.openai.basic.stream.works, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI over /v1/messages streams the Anthropic event grammar"}
|
||||
- {id: llm.messages.openai.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI over /v1/messages: cost header and spend row agree"}
|
||||
|
|
|
|||
|
|
@ -86,6 +86,7 @@ LlmCapability = Literal[
|
|||
"tool_search",
|
||||
"tool_search_history",
|
||||
"tool_use",
|
||||
"upstream_stream_failure",
|
||||
"vision",
|
||||
"web_search",
|
||||
"web_search_server_tool",
|
||||
|
|
|
|||
|
|
@ -48,7 +48,7 @@ most likely to silently break and the one a mock can't prove works.
|
|||
|----------|---------------|-----------|------------|-------------|--------|
|
||||
| Chat | live (spend suite) | live (spend suite) | gap | live | partial |
|
||||
| Embeddings | live (spend suite) | n/a | n/a | live | covered |
|
||||
| Responses (Azure code_interpreter container files) | live | gap | live | gap | partial |
|
||||
| Responses (Azure code_interpreter container files) | live | live | live | gap | partial |
|
||||
| Image / audio / rerank / realtime | - | - | - | - | gap |
|
||||
|
||||
## This suite's files
|
||||
|
|
@ -63,6 +63,7 @@ most likely to silently break and the one a mock can't prove works.
|
|||
| `test_anthropic_passthrough_tool_call_logs_cost` | anthropic native, tool call, cost |
|
||||
| `test_vertex_passthrough_via_managed_model_logs_cost` | vertex_ai native, non-stream, cost |
|
||||
| `test_service_account_key_reads_container_file_by_native_id` | azure responses code_interpreter, non-stream, native container id, service-account key |
|
||||
| `test_service_account_key_reads_container_file_created_by_a_streamed_response` | azure responses code_interpreter, stream, native container id, service-account key, upload right after `[DONE]` |
|
||||
|
||||
Vertex keeps the credential on the proxy like gemini/anthropic, but the deployment is
|
||||
added at runtime instead of declared in the gateway config: the test POSTs `/model/new`
|
||||
|
|
|
|||
|
|
@ -31,10 +31,11 @@ A proxy whose env carries ``AZURE_API_BASE`` for the same resource masks the
|
|||
second regression, since the global-credential fallback then reaches the
|
||||
container anyway.
|
||||
|
||||
The streaming variant is not here: a streamed ``/v1/responses`` writes the
|
||||
container ownership row only after the ``[DONE]`` frame, and the OpenAI SDK
|
||||
closes the connection at ``[DONE]``, so the write is cancelled and every
|
||||
follow-up container call 403s (LIT-8612). That cell comes with its fix.
|
||||
The streaming cell repeats the flow with ``stream=True`` and uploads right
|
||||
after the last event. The OpenAI SDK closes the connection at ``[DONE]``, so an
|
||||
ownership row written after the stream is cancelled with the body task and every
|
||||
follow-up container call 403s (LIT-8612); the row has to land before the
|
||||
``response.completed`` frame goes out.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -52,7 +53,7 @@ from lifecycle import ResourceManager
|
|||
from management.management_client import ManagementClient, build_client
|
||||
from models import KeyGenerateBody, KeyGenerateResponse, LiteLLMParamsBody, TeamNewBody, UserNewBody
|
||||
from openai import OpenAI
|
||||
from openai.types.responses import Response, ResponseCodeInterpreterToolCall
|
||||
from openai.types.responses import Response, ResponseCodeInterpreterToolCall, ResponseCompletedEvent
|
||||
from openai.types.responses.tool_param import CodeInterpreter
|
||||
from proxy_client import ProxyClient
|
||||
from sdk_clients import NO_PROXY_CACHE, SdkClients
|
||||
|
|
@ -120,6 +121,24 @@ def _response_with_code_interpreter(client: OpenAI, model: str) -> Response:
|
|||
)
|
||||
|
||||
|
||||
def _streamed_response_with_code_interpreter(client: OpenAI, model: str) -> Response:
|
||||
events: Final = tuple(
|
||||
client.with_options(timeout=CODE_INTERPRETER_TIMEOUT).responses.create(
|
||||
model=model,
|
||||
input=PROMPT,
|
||||
tools=[CODE_INTERPRETER],
|
||||
tool_choice="required",
|
||||
stream=True,
|
||||
extra_body=NO_PROXY_CACHE,
|
||||
)
|
||||
)
|
||||
assert events, "responses stream returned no events"
|
||||
assert isinstance(events[-1], ResponseCompletedEvent), (
|
||||
f"responses stream did not terminate with response.completed: {events[-1].type}"
|
||||
)
|
||||
return events[-1].response
|
||||
|
||||
|
||||
def _container_id(response: Response) -> str:
|
||||
calls: Final = tuple(item for item in response.output if isinstance(item, ResponseCodeInterpreterToolCall))
|
||||
assert calls, f"no code_interpreter_call in the responses output: {response.output!r}"
|
||||
|
|
@ -165,3 +184,17 @@ class TestAzureContainerFiles:
|
|||
f"container id is not the provider's own id: {native_id}"
|
||||
)
|
||||
_assert_file_round_trip(client, native_id, marker)
|
||||
|
||||
@pytest.mark.covers("llm.responses.azure_openai.code_interpreter.stream.works")
|
||||
def test_service_account_key_reads_container_file_created_by_a_streamed_response(
|
||||
self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients
|
||||
) -> None:
|
||||
marker: Final = unique_marker()
|
||||
model: Final = _register_two_azure_deployments(proxy, resources, marker)
|
||||
key: Final = _service_account_key(proxy, resources, build_client(proxy), marker, model)
|
||||
client: Final = sdk.openai(key)
|
||||
native_id: Final = _native_container_id(
|
||||
_container_id(_streamed_response_with_code_interpreter(client, model))
|
||||
)
|
||||
resources.defer(lambda: client.containers.delete(native_id, extra_query=AZURE_PROVIDER_QUERY))
|
||||
_assert_file_round_trip(client, native_id, marker)
|
||||
|
|
|
|||
|
|
@ -11,8 +11,11 @@ litellm-regression-tests/tests/test_inference_endpoints.py.
|
|||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import anthropic
|
||||
import pytest
|
||||
from anthropic import Anthropic
|
||||
from anthropic.types import (
|
||||
|
|
@ -30,12 +33,21 @@ from anthropic.types import (
|
|||
ToolParam,
|
||||
ToolUseBlock,
|
||||
)
|
||||
from e2e_config import STREAM_MIN_LEAD_SECONDS, provider_edge_base, provider_paces_stream, unique_marker
|
||||
from e2e_config import (
|
||||
PROVIDER_EDGE_ADVERTISE_HOST,
|
||||
PROVIDER_EDGE_BIND_HOST,
|
||||
STREAM_MIN_LEAD_SECONDS,
|
||||
provider_edge_base,
|
||||
provider_paces_stream,
|
||||
unique_marker,
|
||||
)
|
||||
from e2e_http import assert_client_error
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatMessage, LiteLLMParamsBody, SpendLogRow
|
||||
from models import AnthropicErrorEvent, AnthropicMessagesBody, ChatMessage, LiteLLMParamsBody, SpendLogRow
|
||||
from provider_edge import EDGE_MOUNTS, LiveEdge, RunningEdge, StreamCut, start_provider_edge
|
||||
from provider_edge_bedrock import bedrock_signer
|
||||
from proxy_client import ProxyClient
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError
|
||||
from sdk_clients import NO_PROXY_CACHE, SdkClients, response_header
|
||||
|
||||
pytestmark = [pytest.mark.e2e, pytest.mark.replayable]
|
||||
|
|
@ -385,3 +397,245 @@ class TestOpenAIMessagesToolContinuation:
|
|||
)
|
||||
assert _text(continuation).strip() == receipt, "continuation did not consume the correlated tool result"
|
||||
assert all(not isinstance(block, ToolUseBlock) for block in continuation.content)
|
||||
|
||||
|
||||
BEDROCK_BACKEND: Final = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
|
||||
BEDROCK_EDGE_REGION: Final = "us-east-1"
|
||||
_STREAM_FAILURE_PROMPT: Final = "Count from 1 to 100, one number per line."
|
||||
_FRAME_PAYLOAD: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
_AT_FRAME_BOUNDARY: Final = StreamCut(after_content=True)
|
||||
_MID_FRAME: Final = StreamCut(after_content=True, mid_chunk=True)
|
||||
_BEFORE_FIRST_BYTE: Final = StreamCut(after_content=False)
|
||||
|
||||
type _CutRegistration = Callable[[ProxyClient, ResourceManager, StreamCut], tuple[str, str]]
|
||||
|
||||
|
||||
def _cut_edge(backend: LiveEdge, mount: str) -> RunningEdge:
|
||||
return start_provider_edge(
|
||||
backend,
|
||||
mounts=MappingProxyType({mount: EDGE_MOUNTS[mount]}),
|
||||
bind_host=PROVIDER_EDGE_BIND_HOST,
|
||||
advertise_host=PROVIDER_EDGE_ADVERTISE_HOST,
|
||||
)
|
||||
|
||||
|
||||
def _register_cut_bedrock(proxy: ProxyClient, resources: ResourceManager, cut: StreamCut) -> tuple[str, str]:
|
||||
mount: Final = f"bedrock/{BEDROCK_EDGE_REGION}"
|
||||
edge: Final = _cut_edge(LiveEdge(cut=cut, sign=bedrock_signer(BEDROCK_EDGE_REGION)), mount)
|
||||
resources.defer(edge.shutdown)
|
||||
return _register(
|
||||
proxy,
|
||||
resources,
|
||||
LiteLLMParamsBody(
|
||||
model=BEDROCK_BACKEND,
|
||||
api_base=edge.edge.api_base(mount),
|
||||
aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID",
|
||||
aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY",
|
||||
aws_region_name=BEDROCK_EDGE_REGION,
|
||||
),
|
||||
prefix="e2e-messages-cut",
|
||||
)
|
||||
|
||||
|
||||
def _register_cut_anthropic(proxy: ProxyClient, resources: ResourceManager, cut: StreamCut) -> tuple[str, str]:
|
||||
edge: Final = _cut_edge(LiveEdge(cut=cut), "anthropic")
|
||||
resources.defer(edge.shutdown)
|
||||
return _register(
|
||||
proxy,
|
||||
resources,
|
||||
LiteLLMParamsBody(
|
||||
model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY", api_base=edge.edge.api_base("anthropic")
|
||||
),
|
||||
prefix="e2e-messages-cut",
|
||||
)
|
||||
|
||||
|
||||
_DROPPED_UPSTREAMS: Final[tuple[tuple[str, _CutRegistration, StreamCut], ...]] = (
|
||||
("bedrock_at_a_frame_boundary", _register_cut_bedrock, _AT_FRAME_BOUNDARY),
|
||||
("anthropic_at_a_frame_boundary", _register_cut_anthropic, _AT_FRAME_BOUNDARY),
|
||||
("anthropic_mid_frame", _register_cut_anthropic, _MID_FRAME),
|
||||
)
|
||||
_DROPPED_BEFORE_FIRST_BYTE: Final[tuple[tuple[str, _CutRegistration, StreamCut], ...]] = (
|
||||
("bedrock_before_the_first_byte", _register_cut_bedrock, _BEFORE_FIRST_BYTE),
|
||||
("anthropic_before_the_first_byte", _register_cut_anthropic, _BEFORE_FIRST_BYTE),
|
||||
)
|
||||
|
||||
|
||||
def _payload(frame: str) -> JsonValue | None:
|
||||
try:
|
||||
return _FRAME_PAYLOAD.validate_json(frame)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _bare_error_frame(frame: str) -> bool:
|
||||
payload: Final = _payload(frame)
|
||||
return isinstance(payload, dict) and "error" in payload and payload.get("type") != "error"
|
||||
|
||||
|
||||
@pytest.mark.provider_edge_host
|
||||
@pytest.mark.provider_live
|
||||
class TestMessagesUpstreamStreamFailure:
|
||||
@pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_event")
|
||||
@pytest.mark.parametrize(
|
||||
("register", "cut"), [case[1:] for case in _DROPPED_UPSTREAMS], ids=[case[0] for case in _DROPPED_UPSTREAMS]
|
||||
)
|
||||
def test_interrupted_upstream_stream_raises_in_the_anthropic_sdk(
|
||||
self,
|
||||
proxy: ProxyClient,
|
||||
resources: ResourceManager,
|
||||
sdk: SdkClients,
|
||||
register: _CutRegistration,
|
||||
cut: StreamCut,
|
||||
) -> None:
|
||||
model, key = register(proxy, resources, cut)
|
||||
client: Final = sdk.anthropic(key)
|
||||
|
||||
stream: Final = client.messages.create(
|
||||
model=model,
|
||||
max_tokens=300,
|
||||
stream=True,
|
||||
messages=[_user_turn(_STREAM_FAILURE_PROMPT)],
|
||||
extra_body=NO_PROXY_CACHE,
|
||||
)
|
||||
first: Final = next(stream)
|
||||
assert first.type == "message_start", (
|
||||
f"the stream produced a first event that is not message_start, so this run proves a "
|
||||
f"startup failure, not an interrupted stream: {first!r}"
|
||||
)
|
||||
with pytest.raises(anthropic.APIStatusError) as raised:
|
||||
for _ in stream:
|
||||
pass
|
||||
try:
|
||||
AnthropicErrorEvent.model_validate(raised.value.body)
|
||||
except ValidationError:
|
||||
pytest.fail(
|
||||
f"the SDK raised on the interrupted stream but without the Anthropic error envelope a "
|
||||
f"client reads the failure from: body={raised.value.body!r} message={raised.value}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_event")
|
||||
@pytest.mark.parametrize(
|
||||
("register", "cut"), [case[1:] for case in _DROPPED_UPSTREAMS], ids=[case[0] for case in _DROPPED_UPSTREAMS]
|
||||
)
|
||||
def test_interrupted_upstream_stream_is_an_anthropic_error_event(
|
||||
self, proxy: ProxyClient, resources: ResourceManager, register: _CutRegistration, cut: StreamCut
|
||||
) -> None:
|
||||
model, key = register(proxy, resources, cut)
|
||||
|
||||
outcome: Final = proxy.messages_stream(
|
||||
key,
|
||||
AnthropicMessagesBody(
|
||||
model=model,
|
||||
max_tokens=300,
|
||||
stream=True,
|
||||
messages=[ChatMessage(role="user", content=_STREAM_FAILURE_PROMPT)],
|
||||
),
|
||||
)
|
||||
frames: Final = outcome.stream_events
|
||||
assert outcome.is_streaming, (
|
||||
f"/v1/messages did not answer with an SSE stream: status={outcome.status_code} body={outcome.body}"
|
||||
)
|
||||
assert frames, (
|
||||
f"the proxy sent no SSE data frames although the upstream hung up; stream_error={outcome.stream_error!r}"
|
||||
)
|
||||
assert outcome.stream_error == "event: error", (
|
||||
f"the interrupted stream was not announced by an 'event: error' line Anthropic clients read; "
|
||||
f"stream_error={outcome.stream_error!r} frames={frames}"
|
||||
)
|
||||
try:
|
||||
AnthropicErrorEvent.model_validate_json(frames[-1])
|
||||
except ValidationError:
|
||||
pytest.fail(
|
||||
f'the last SSE frame was not an Anthropic {{"type": "error", "error": ...}} envelope; frames={frames}'
|
||||
)
|
||||
torn: Final = tuple(index for index, frame in enumerate(frames) if _payload(frame) is None)
|
||||
expected_torn: Final = 1 if cut.mid_chunk else 0
|
||||
assert len(torn) == expected_torn, (
|
||||
f"expected {expected_torn} data line(s) that are not JSON, since the edge tears one only when it "
|
||||
f"cuts mid-frame, but the proxy relayed {[frames[index] for index in torn]}; all frames={frames}"
|
||||
)
|
||||
for index in torn:
|
||||
assert _payload(frames[index + 1]) == {"type": "ping"}, (
|
||||
f"the frame the upstream tore was not closed as a ping event before the error, so an "
|
||||
f"Anthropic client parses the error inside it: after {frames[index]!r} came "
|
||||
f"{frames[index + 1]!r}; all frames={frames}"
|
||||
)
|
||||
bare: Final = tuple(frame for frame in frames if _bare_error_frame(frame))
|
||||
assert not bare, (
|
||||
f"the proxy emitted error frames without the Anthropic envelope, which Anthropic clients drop: "
|
||||
f"{bare}; all frames={frames}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_status")
|
||||
@pytest.mark.parametrize(
|
||||
("register", "cut"),
|
||||
[case[1:] for case in _DROPPED_BEFORE_FIRST_BYTE],
|
||||
ids=[case[0] for case in _DROPPED_BEFORE_FIRST_BYTE],
|
||||
)
|
||||
def test_upstream_that_hangs_up_before_the_first_byte_raises_with_its_status_in_the_anthropic_sdk(
|
||||
self,
|
||||
proxy: ProxyClient,
|
||||
resources: ResourceManager,
|
||||
sdk: SdkClients,
|
||||
register: _CutRegistration,
|
||||
cut: StreamCut,
|
||||
) -> None:
|
||||
model, key = register(proxy, resources, cut)
|
||||
client: Final = sdk.anthropic(key)
|
||||
|
||||
with pytest.raises(anthropic.APIStatusError) as raised:
|
||||
client.messages.create(
|
||||
model=model,
|
||||
max_tokens=300,
|
||||
stream=True,
|
||||
messages=[_user_turn(_STREAM_FAILURE_PROMPT)],
|
||||
extra_body=NO_PROXY_CACHE,
|
||||
)
|
||||
assert 500 <= raised.value.status_code < 600, (
|
||||
f"an upstream that hung up before sending anything must answer with a server error status the SDK "
|
||||
f"retries on, not {raised.value.status_code}: {raised.value}"
|
||||
)
|
||||
try:
|
||||
AnthropicErrorEvent.model_validate(raised.value.body)
|
||||
except ValidationError:
|
||||
pytest.fail(
|
||||
f"the SDK raised with the right status but without the Anthropic error envelope a client reads "
|
||||
f"the failure from: body={raised.value.body!r} message={raised.value}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_status")
|
||||
@pytest.mark.parametrize(
|
||||
("register", "cut"),
|
||||
[case[1:] for case in _DROPPED_BEFORE_FIRST_BYTE],
|
||||
ids=[case[0] for case in _DROPPED_BEFORE_FIRST_BYTE],
|
||||
)
|
||||
def test_upstream_that_hangs_up_before_the_first_byte_is_a_json_error_with_its_status(
|
||||
self, proxy: ProxyClient, resources: ResourceManager, register: _CutRegistration, cut: StreamCut
|
||||
) -> None:
|
||||
model, key = register(proxy, resources, cut)
|
||||
|
||||
outcome: Final = proxy.messages_stream(
|
||||
key,
|
||||
AnthropicMessagesBody(
|
||||
model=model,
|
||||
max_tokens=300,
|
||||
stream=True,
|
||||
messages=[ChatMessage(role="user", content=_STREAM_FAILURE_PROMPT)],
|
||||
),
|
||||
)
|
||||
assert not outcome.is_streaming, (
|
||||
f"nothing had been streamed when the upstream hung up, yet /v1/messages opened a 200 SSE stream "
|
||||
f"instead of answering with the failure's status: stream_error={outcome.stream_error!r} "
|
||||
f"frames={outcome.stream_events}"
|
||||
)
|
||||
assert 500 <= outcome.status_code < 600, (
|
||||
f"/v1/messages answered {outcome.status_code} for an upstream that hung up before its first byte; "
|
||||
f"body={outcome.body}"
|
||||
)
|
||||
try:
|
||||
AnthropicErrorEvent.model_validate_json(outcome.body)
|
||||
except ValidationError:
|
||||
pytest.fail(
|
||||
f'the error body is not an Anthropic {{"type": "error", "error": ...}} envelope; body={outcome.body}'
|
||||
)
|
||||
|
|
|
|||
|
|
@ -609,6 +609,16 @@ class CountTokensResponse(BaseModel):
|
|||
input_tokens: int
|
||||
|
||||
|
||||
class AnthropicErrorBody(BaseModel):
|
||||
type: str
|
||||
message: str
|
||||
|
||||
|
||||
class AnthropicErrorEvent(BaseModel):
|
||||
type: Literal["error"]
|
||||
error: AnthropicErrorBody
|
||||
|
||||
|
||||
# ---------- mcp servers ----------
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ import hashlib
|
|||
import os
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
from collections import deque
|
||||
from collections.abc import Callable, Generator, Mapping, Sequence
|
||||
from contextlib import closing, contextmanager
|
||||
|
|
@ -56,6 +57,7 @@ from types import MappingProxyType
|
|||
from typing import Final, Literal, assert_never
|
||||
from urllib.parse import parse_qsl, urlsplit
|
||||
|
||||
from botocore.eventstream import EventStreamBuffer
|
||||
from e2e_http import (
|
||||
NetworkError,
|
||||
StreamChunk,
|
||||
|
|
@ -96,16 +98,18 @@ from fixture_mode import (
|
|||
)
|
||||
from fixture_profile import IneligibleRequest, MatchProfile, match_profile, strict_identity
|
||||
from provider_cache import (
|
||||
JSON_VALUE,
|
||||
SIGNATURE_HEADERS,
|
||||
CacheEdge,
|
||||
MountPolicy,
|
||||
RequestSigner,
|
||||
invoke_chunk_value,
|
||||
is_bedrock,
|
||||
scoped_edge_base,
|
||||
split_test_segment,
|
||||
)
|
||||
from provider_cache_routing import LIVE_PROVIDER_REQUIRED
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
from pydantic import JsonValue, TypeAdapter, ValidationError
|
||||
|
||||
BEDROCK_REGIONS: Final[tuple[str, ...]] = ("us-east-1",)
|
||||
|
||||
|
|
@ -537,10 +541,33 @@ class ReplayEdge:
|
|||
source: ReplaySource
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StreamCut:
|
||||
"""Where a live edge hangs up on a streamed upstream body: before its first byte, or with
|
||||
``after_content`` set, right after the first transfer chunk carrying assistant output (a
|
||||
``content_block_delta``). That frame is what commits the proxy's mid-stream fallback
|
||||
wrapper to the client: it holds the lifecycle frames before it back and drops them when
|
||||
the transport fails first, so a cut after a fixed number of chunks landed on either side
|
||||
of that commit depending on how the provider batched its frames. With ``mid_chunk`` set
|
||||
the hang-up comes part way through the next ``data:`` line the provider sends after that,
|
||||
so the client is left inside an SSE frame the way a dropped transport leaves it.
|
||||
|
||||
Whatever was relayed sits on the wire for ``_CUT_SETTLE_SECONDS`` before the hang-up, so
|
||||
the client has read it by then instead of receiving the data and the close in one burst,
|
||||
where its reader can surface the close before what it buffered."""
|
||||
|
||||
after_content: bool
|
||||
mid_chunk: bool = False
|
||||
|
||||
|
||||
_CUT_SETTLE_SECONDS: Final = 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LiveEdge:
|
||||
observe_request: Callable[[str, Mapping[str, str], bytes | None], None] | None = None
|
||||
sign: RequestSigner | None = None
|
||||
cut: StreamCut | None = None
|
||||
|
||||
|
||||
type EdgeBackend = RecordEdge | ReplayEdge | LiveEdge | CacheEdge
|
||||
|
|
@ -786,11 +813,128 @@ def _handle_record(
|
|||
assert_never(head)
|
||||
|
||||
|
||||
def _data_line_start(data: bytes) -> int:
|
||||
if data.startswith(b"data:"):
|
||||
return 0
|
||||
at_line_start: Final = data.find(b"\ndata:")
|
||||
return -1 if at_line_start < 0 else at_line_start + 1
|
||||
|
||||
|
||||
def _torn_prefix(data: bytes) -> bytes:
|
||||
start: Final = _data_line_start(data)
|
||||
line_end: Final = data.find(b"\n", start)
|
||||
end: Final = len(data) if line_end < 0 else line_end
|
||||
return data[: start + (end - start) // 2]
|
||||
|
||||
|
||||
class _DataLineTearer:
|
||||
__slots__ = ("_unfinished_line",)
|
||||
|
||||
_unfinished_line: bytes
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._unfinished_line = b""
|
||||
|
||||
def observe(self, data: bytes) -> None:
|
||||
self._unfinished_line = (self._unfinished_line + data).rsplit(b"\n", 1)[-1]
|
||||
|
||||
def tear(self, data: bytes) -> bytes | None:
|
||||
buffered: Final = self._unfinished_line + data
|
||||
if _data_line_start(buffered) < 0:
|
||||
self.observe(data)
|
||||
return None
|
||||
return _torn_prefix(buffered)[len(self._unfinished_line):]
|
||||
|
||||
|
||||
def _is_content_delta(value: JsonValue | None) -> bool:
|
||||
return isinstance(value, dict) and value.get("type") == "content_block_delta"
|
||||
|
||||
|
||||
def _sse_data_carries_content(line: bytes) -> bool:
|
||||
if not line.startswith(b"data:"):
|
||||
return False
|
||||
try:
|
||||
return _is_content_delta(JSON_VALUE.validate_json(line[len(b"data:"):].strip()))
|
||||
except ValidationError:
|
||||
return False
|
||||
|
||||
|
||||
class _AnthropicContentDetector:
|
||||
__slots__ = ("_unfinished_line",)
|
||||
|
||||
_unfinished_line: bytes
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._unfinished_line = b""
|
||||
|
||||
def __call__(self, data: bytes) -> bool:
|
||||
lines: Final = (self._unfinished_line + data).split(b"\n")
|
||||
self._unfinished_line = lines[-1]
|
||||
return any(_sse_data_carries_content(line.rstrip(b"\r")) for line in lines[:-1])
|
||||
|
||||
|
||||
def _invoke_frame_carries_content(payload: bytes) -> bool:
|
||||
try:
|
||||
return _is_content_delta(invoke_chunk_value(JSON_VALUE.validate_json(payload)))
|
||||
except ValidationError:
|
||||
return False
|
||||
|
||||
|
||||
def _bedrock_content_detector() -> Callable[[bytes], bool]:
|
||||
"""Bedrock's invoke stream wraps each Anthropic event in an eventstream frame that a
|
||||
transfer chunk can split, so the frames are reassembled across chunks before being read."""
|
||||
frames: Final = EventStreamBuffer()
|
||||
|
||||
def carries_content(data: bytes) -> bool:
|
||||
frames.add_data(data)
|
||||
return any(_invoke_frame_carries_content(frame.payload) for frame in frames)
|
||||
|
||||
return carries_content
|
||||
|
||||
|
||||
def _content_detector(mount: str) -> Callable[[bytes], bool]:
|
||||
return _bedrock_content_detector() if is_bedrock(mount) else _AnthropicContentDetector()
|
||||
|
||||
|
||||
def _cut_steps(
|
||||
steps: Generator[StreamStep, None, None], cut: StreamCut, carries_content: Callable[[bytes], bool]
|
||||
) -> Generator[StreamStep, None, None]:
|
||||
with closing(steps) as source:
|
||||
tearer: Final = _DataLineTearer()
|
||||
if cut.after_content:
|
||||
for step in source:
|
||||
yield step
|
||||
if isinstance(step, StreamTruncation):
|
||||
return
|
||||
tearer.observe(step.data)
|
||||
if carries_content(step.data):
|
||||
break
|
||||
else:
|
||||
return
|
||||
if cut.mid_chunk:
|
||||
for step in source:
|
||||
if isinstance(step, StreamTruncation):
|
||||
yield step
|
||||
return
|
||||
if (torn := tearer.tear(step.data)) is None:
|
||||
yield step
|
||||
continue
|
||||
if torn:
|
||||
yield StreamChunk(data=torn)
|
||||
break
|
||||
else:
|
||||
return
|
||||
if cut.after_content or cut.mid_chunk:
|
||||
time.sleep(_CUT_SETTLE_SECONDS)
|
||||
yield StreamTruncation(reason=f"edge cut the upstream stream: {cut!r}")
|
||||
|
||||
|
||||
def _handle_live(
|
||||
method: str, url: str, headers: Mapping[str, str], body: bytes | None, timeout: float,
|
||||
cache: CacheEdge | None = None, mount: str = "", test_key: str | None = None,
|
||||
observe_request: Callable[[str, Mapping[str, str], bytes | None], None] | None = None,
|
||||
sign: RequestSigner | None = None,
|
||||
cut: StreamCut | None = None,
|
||||
) -> EdgeOutcome:
|
||||
forwarded: Final = {
|
||||
name: value for name, value in headers.items() if name.lower() not in _REQUEST_DROPPED_HEADERS
|
||||
|
|
@ -805,6 +949,8 @@ def _handle_live(
|
|||
match head:
|
||||
case NetworkError(message=message):
|
||||
return _recorded_outcome(_network_error_response(message))
|
||||
case StreamHead() if cut is not None:
|
||||
return EdgeStream(head.status_code, _filtered_response_headers(head.headers), _cut_steps(head.steps, cut, _content_detector(mount)))
|
||||
case StreamHead() if _is_streamed(head.headers):
|
||||
return EdgeStream(head.status_code, _filtered_response_headers(head.headers), head.steps)
|
||||
case StreamHead():
|
||||
|
|
@ -875,10 +1021,10 @@ def handle_edge_request(
|
|||
method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout,
|
||||
backend, mount, test_key,
|
||||
)
|
||||
case LiveEdge(observe_request=observe_request, sign=sign):
|
||||
case LiveEdge(observe_request=observe_request, sign=sign, cut=cut):
|
||||
return _handle_live(
|
||||
method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout,
|
||||
observe_request=observe_request, sign=sign,
|
||||
mount=mount, observe_request=observe_request, sign=sign, cut=cut,
|
||||
)
|
||||
case RecordEdge():
|
||||
return _handle_record(
|
||||
|
|
|
|||
|
|
@ -54,11 +54,13 @@ from provider_edge import (
|
|||
EdgeBackend,
|
||||
EdgeReply,
|
||||
EdgeStream,
|
||||
LiveEdge,
|
||||
ProviderEdge,
|
||||
ProviderRequestObservation,
|
||||
RecordEdge,
|
||||
ReplayEdge,
|
||||
ReplaySource,
|
||||
StreamCut,
|
||||
edge_request,
|
||||
handle_edge_request,
|
||||
observed_provider_edge,
|
||||
|
|
@ -1000,6 +1002,36 @@ def stream_chunks(response: RecordedStreamedResponse) -> list[bytes]:
|
|||
return [base64.b64decode(chunk) for chunk in response.chunks_b64]
|
||||
|
||||
|
||||
SECOND_DATA_LINE: Final = b'data: {"type":"content_block_delta","delta":{"text":" two"}}'
|
||||
SPLIT_MARKER_CHUNKS: tuple[bytes, ...] = (
|
||||
b'data: {"type":"content_block_delta","delta":{"text":"one"}}\n\nda',
|
||||
b"ta" + SECOND_DATA_LINE[4:] + b"\n\nda",
|
||||
b'ta: {"type":"message_delta","usage":{"output_tokens":7}}\n\nda',
|
||||
b"ta: [DONE]\n\n",
|
||||
)
|
||||
|
||||
|
||||
class TestStreamCut:
|
||||
def test_a_mid_frame_cut_tears_a_data_line_whose_marker_is_split_across_chunks(self) -> None:
|
||||
"""Every ``data:`` marker after the first content delta straddles a transfer
|
||||
chunk boundary, so a tearer that inspects each chunk on its own never finds
|
||||
one and lets the stream finish cleanly instead of cutting it."""
|
||||
backend: Final = LiveEdge(cut=StreamCut(after_content=True, mid_chunk=True))
|
||||
with chunked_provider(chunks=SPLIT_MARKER_CHUNKS) as provider:
|
||||
with running_edge(backend, {"openai": provider_url(provider)}) as edge:
|
||||
head, chunks, ending = raw_stream_post(edge.port, STREAM_PATH, STREAM_BODY)
|
||||
|
||||
assert head.startswith("HTTP/1.1 200 OK")
|
||||
assert ending == "truncated"
|
||||
relayed: Final = b"".join(chunks)
|
||||
whole: Final = b"".join(SPLIT_MARKER_CHUNKS)
|
||||
assert whole.startswith(relayed) and relayed != whole
|
||||
assert relayed.startswith(SPLIT_MARKER_CHUNKS[0])
|
||||
torn_line: Final = relayed.rsplit(b"\n", 1)[-1]
|
||||
assert torn_line and SECOND_DATA_LINE.startswith(torn_line) and torn_line != SECOND_DATA_LINE
|
||||
assert b"[DONE]" not in relayed
|
||||
|
||||
|
||||
class TestStreamingFidelity:
|
||||
"""LIT-5742: a streamed response records and replays as the chunk sequence the
|
||||
provider actually sent, not as one coalesced body. The unit of fidelity is the
|
||||
|
|
|
|||
|
|
@ -84,3 +84,42 @@ def test_jsonrpc_error_and_malformed_tool_result_remain_errors(gateway: Gateway)
|
|||
control: Final = call_tool(gateway, key, identity, names["add"], {"a": 3, "b": 5})
|
||||
assert control.status_code == 200 and control.json()["isError"] is False, control.text
|
||||
assert control.json()["content"][0]["text"] == "8"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("ingress", ("http", "sse"))
|
||||
def test_configured_revision_blocks_unadvertised_handshake_and_keeps_allowed_control(
|
||||
gateway: Gateway, tmp_path, ingress: str
|
||||
) -> None:
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
|
||||
import yaml
|
||||
from integration._support.mcp import mcp_peer
|
||||
from integration._support.process import owned_proxy
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from mcp import MCPError
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
||||
with mcp_peer() as upstream, gateway.scenario() as scenario:
|
||||
alias: Final = "restricted" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, upstream, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["general_settings"]["mcp_advertised_versions"] = ["2024-11-05"]
|
||||
config_path: Final = tmp_path / "restricted.yaml"
|
||||
config_path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {"DISABLE_SCHEMA_UPDATE": "true"}, config=config_path) as restricted:
|
||||
endpoint: Final = str(restricted.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp")
|
||||
headers: Final = {"Authorization": f"Bearer {key}", "x-mcp-servers": identity}
|
||||
denied: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2025-11-25", extra_headers=headers)
|
||||
allowed: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2024-11-05", extra_headers=headers)
|
||||
|
||||
async def exercise() -> None:
|
||||
with pytest.raises(MCPError, match="Unsupported MCP protocol version"):
|
||||
await denied.list_tools(raise_on_error=True)
|
||||
assert f"{alias}-add" in tuple(tool.name for tool in await allowed.list_tools(raise_on_error=True))
|
||||
result: Final = await allowed.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 2, "b": 5}))
|
||||
assert result.is_error is False and result.content[0].text == "7"
|
||||
|
||||
asyncio.run(exercise())
|
||||
|
|
|
|||
|
|
@ -155,3 +155,43 @@ def test_server_initiated_sampling_and_elicitation_surface_as_errors_not_success
|
|||
assert outcome.error is not None, outcome.raw
|
||||
assert outcome.text is None or not outcome.text.startswith(("sampled:", "elicited:")), outcome.raw
|
||||
assert len(tool_calls(peer.drain())) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("downstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
|
||||
@pytest.mark.parametrize("upstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
|
||||
@pytest.mark.parametrize("peer_kind", ("http", "sse", "stdio"))
|
||||
@pytest.mark.parametrize("ingress", ("http", "sse"))
|
||||
def test_pinned_revision_pairs_list_and_call_through_gateway(
|
||||
gateway: Gateway, downstream: str, upstream: str, peer_kind: PeerKind, ingress: str
|
||||
) -> None:
|
||||
import asyncio
|
||||
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
with peer_of(peer_kind) as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "versions" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, mcp_info={"protocol_version": upstream})
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
endpoint: Final = str(gateway.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp")
|
||||
client: Final = MCPClient(
|
||||
server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version=downstream,
|
||||
extra_headers={"Authorization": f"Bearer {key}", "x-mcp-servers": identity}, timeout=15,
|
||||
)
|
||||
|
||||
async def exercise() -> None:
|
||||
tools: Final = await client.list_tools(raise_on_error=True)
|
||||
assert f"{alias}-add" in tuple(tool.name for tool in tools)
|
||||
result: Final = await client.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 3, "b": 4}))
|
||||
assert result.is_error is False
|
||||
assert result.content[0].text == "7"
|
||||
|
||||
peer.drain()
|
||||
asyncio.run(exercise())
|
||||
observed: Final = peer.drain()
|
||||
negotiations: Final = tuple(item["body"] for item in observed if item["body"].get("method") == "initialize")
|
||||
assert negotiations, "The operation must reach the upstream negotiation"
|
||||
assert all(request["params"]["protocolVersion"] == upstream for request in negotiations), negotiations
|
||||
assert len(tool_calls(observed)) == 1
|
||||
|
|
|
|||
|
|
@ -742,7 +742,7 @@ class BaseResponsesAPITest(ABC):
|
|||
Passes tools=[{"type": "shell", "environment": {"type": "container_auto"}}];
|
||||
validates that the request is accepted and returns a valid response.
|
||||
Only runs for OpenAI; offline coverage for the Azure route lives in
|
||||
tests/test_litellm/responses/test_responses_api_request_body.py.
|
||||
tests/unit/responses/test_responses_api_request_body.py.
|
||||
"""
|
||||
base_completion_call_args = self.get_base_completion_call_args()
|
||||
model = (
|
||||
|
|
|
|||
|
|
@ -381,6 +381,36 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
)
|
||||
raise
|
||||
|
||||
@staticmethod
|
||||
def _config_completing_after_one_delta() -> Mock:
|
||||
mock_config = Mock(spec=BaseResponsesAPIConfig)
|
||||
completed_response = ResponsesAPIResponse(
|
||||
id="resp_123",
|
||||
created_at=0,
|
||||
status="completed",
|
||||
model="gpt-5.5",
|
||||
object="response",
|
||||
output=[],
|
||||
usage=ResponseAPIUsage(input_tokens=1, output_tokens=1, total_tokens=2),
|
||||
)
|
||||
|
||||
def _transform(model, parsed_chunk, logging_obj):
|
||||
if parsed_chunk.get("type") == "response.completed":
|
||||
return ResponseCompletedEvent(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response=completed_response,
|
||||
)
|
||||
return OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id="msg_123",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
delta=parsed_chunk["delta"],
|
||||
)
|
||||
|
||||
mock_config.transform_streaming_response.side_effect = _transform
|
||||
return mock_config
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_async_iteration_not_logged_as_failure(self):
|
||||
"""
|
||||
|
|
@ -399,6 +429,7 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
|
||||
async def mock_aiter_bytes():
|
||||
yield b'data: {"type": "response.output_text.delta", "delta": "test"}\n\n'
|
||||
yield b'data: {"type": "response.completed", "response": {"id": "resp_123"}}\n\n'
|
||||
|
||||
mock_response.aiter_bytes = mock_aiter_bytes
|
||||
|
||||
|
|
@ -408,11 +439,7 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
mock_logging_obj.async_failure_handler = Mock()
|
||||
mock_logging_obj.failure_handler = Mock()
|
||||
|
||||
mock_config = Mock(spec=BaseResponsesAPIConfig)
|
||||
mock_delta_event = Mock()
|
||||
mock_delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA
|
||||
mock_delta_event.delta = "test"
|
||||
mock_config.transform_streaming_response.return_value = mock_delta_event
|
||||
mock_config = self._config_completing_after_one_delta()
|
||||
|
||||
# Create the iterator instance
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
|
|
@ -432,8 +459,9 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
except StopAsyncIteration:
|
||||
pass # This is expected
|
||||
|
||||
# Verify we got the chunk
|
||||
assert len(chunks_received) == 1
|
||||
# Verify we got the delta and the terminal event
|
||||
assert len(chunks_received) == 2
|
||||
assert iterator.completed_response is not None
|
||||
|
||||
# CRITICAL: Verify that failure handlers were NOT called
|
||||
# StopAsyncIteration is a normal end of stream, not a failure
|
||||
|
|
@ -460,6 +488,7 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
|
||||
def mock_iter_bytes():
|
||||
yield b'data: {"type": "response.output_text.delta", "delta": "test"}\n\n'
|
||||
yield b'data: {"type": "response.completed", "response": {"id": "resp_123"}}\n\n'
|
||||
|
||||
mock_response.iter_bytes = mock_iter_bytes
|
||||
|
||||
|
|
@ -469,11 +498,7 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
mock_logging_obj.async_failure_handler = Mock()
|
||||
mock_logging_obj.failure_handler = Mock()
|
||||
|
||||
mock_config = Mock(spec=BaseResponsesAPIConfig)
|
||||
mock_delta_event = Mock()
|
||||
mock_delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA
|
||||
mock_delta_event.delta = "test"
|
||||
mock_config.transform_streaming_response.return_value = mock_delta_event
|
||||
mock_config = self._config_completing_after_one_delta()
|
||||
|
||||
# Create the iterator instance
|
||||
iterator = SyncResponsesAPIStreamingIterator(
|
||||
|
|
@ -493,8 +518,9 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
except StopIteration:
|
||||
pass # This is expected
|
||||
|
||||
# Verify we got the chunk
|
||||
assert len(chunks_received) == 1
|
||||
# Verify we got the delta and the terminal event
|
||||
assert len(chunks_received) == 2
|
||||
assert iterator.completed_response is not None
|
||||
|
||||
# CRITICAL: Verify that failure handlers were NOT called
|
||||
# StopIteration is a normal end of stream, not a failure
|
||||
|
|
|
|||
|
|
@ -1800,7 +1800,7 @@ def test_gemini_image_size_limit_exceeded(monkeypatch):
|
|||
that could cause memory issues and pod crashes.
|
||||
|
||||
The image fetch is mocked (mirroring the LargeImageClient pattern in
|
||||
tests/test_litellm/litellm_core_utils/test_image_handling.py) so the test
|
||||
tests/unit/litellm_core_utils/test_image_handling.py) so the test
|
||||
deterministically exercises the size-limit rejection path without any
|
||||
external network dependency.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,867 +0,0 @@
|
|||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES, LLMCachingHandler
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_async_embedding_cached_response():
|
||||
llm_caching_handler = LLMCachingHandler(
|
||||
original_function=MagicMock(),
|
||||
request_kwargs={},
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
args = {
|
||||
"cached_result": [
|
||||
{
|
||||
"embedding": [-0.025122925639152527, -0.019487135112285614],
|
||||
"index": 0,
|
||||
"object": "embedding",
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
|
||||
final_embedding_cached_response=None,
|
||||
cached_result=args["cached_result"],
|
||||
kwargs={"model": "text-embedding-ada-002", "input": "test"},
|
||||
logging_obj=mock_logging_obj,
|
||||
start_time=datetime.now(),
|
||||
model="text-embedding-ada-002",
|
||||
)
|
||||
|
||||
assert cache_hit
|
||||
|
||||
print(f"response: {response}")
|
||||
assert len(response.data) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_embedding_cache_preserves_prompt_tokens_details():
|
||||
"""Test that prompt_tokens_details (including image_count) survives a full cache hit."""
|
||||
llm_caching_handler = LLMCachingHandler(
|
||||
original_function=MagicMock(),
|
||||
request_kwargs={},
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
cached_result = [
|
||||
{
|
||||
"embedding": [-0.025, -0.019],
|
||||
"index": 0,
|
||||
"object": "embedding",
|
||||
"model": "amazon.titan-embed-image-v1",
|
||||
"prompt_tokens_details": {"image_count": 1},
|
||||
}
|
||||
]
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
|
||||
final_embedding_cached_response=None,
|
||||
cached_result=cached_result,
|
||||
kwargs={"model": "amazon.titan-embed-image-v1", "input": "base64imagedata"},
|
||||
logging_obj=mock_logging_obj,
|
||||
start_time=datetime.now(),
|
||||
model="amazon.titan-embed-image-v1",
|
||||
)
|
||||
|
||||
assert cache_hit
|
||||
assert response.usage is not None
|
||||
assert response.usage.prompt_tokens_details is not None
|
||||
assert response.usage.prompt_tokens_details.image_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_embedding_cache_backward_compat_no_prompt_tokens_details():
|
||||
"""Test that old cached items without prompt_tokens_details still work."""
|
||||
llm_caching_handler = LLMCachingHandler(
|
||||
original_function=MagicMock(),
|
||||
request_kwargs={},
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
# Old-format cached item — no prompt_tokens_details field
|
||||
cached_result = [
|
||||
{
|
||||
"embedding": [-0.025, -0.019],
|
||||
"index": 0,
|
||||
"object": "embedding",
|
||||
"model": "text-embedding-ada-002",
|
||||
}
|
||||
]
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
|
||||
final_embedding_cached_response=None,
|
||||
cached_result=cached_result,
|
||||
kwargs={"model": "text-embedding-ada-002", "input": "test"},
|
||||
logging_obj=mock_logging_obj,
|
||||
start_time=datetime.now(),
|
||||
model="text-embedding-ada-002",
|
||||
)
|
||||
|
||||
assert cache_hit
|
||||
assert response.usage is not None
|
||||
assert response.usage.prompt_tokens_details is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_embedding_cache_aggregates_multiple_image_counts():
|
||||
"""Test that image_count is summed correctly across multiple cached items."""
|
||||
llm_caching_handler = LLMCachingHandler(
|
||||
original_function=MagicMock(),
|
||||
request_kwargs={},
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
cached_result = [
|
||||
{
|
||||
"embedding": [-0.025, -0.019],
|
||||
"index": 0,
|
||||
"object": "embedding",
|
||||
"model": "amazon.titan-embed-image-v1",
|
||||
"prompt_tokens_details": {"image_count": 1},
|
||||
},
|
||||
{
|
||||
"embedding": [0.031, 0.042],
|
||||
"index": 1,
|
||||
"object": "embedding",
|
||||
"model": "amazon.titan-embed-image-v1",
|
||||
"prompt_tokens_details": {"image_count": 1},
|
||||
},
|
||||
]
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
|
||||
final_embedding_cached_response=None,
|
||||
cached_result=cached_result,
|
||||
kwargs={
|
||||
"model": "amazon.titan-embed-image-v1",
|
||||
"input": ["img1", "img2"],
|
||||
},
|
||||
logging_obj=mock_logging_obj,
|
||||
start_time=datetime.now(),
|
||||
model="amazon.titan-embed-image-v1",
|
||||
)
|
||||
|
||||
assert cache_hit
|
||||
assert response.usage.prompt_tokens_details is not None
|
||||
assert response.usage.prompt_tokens_details.image_count == 2
|
||||
|
||||
|
||||
def test_combine_usage_merges_prompt_tokens_details():
|
||||
"""Test that combine_usage merges prompt_tokens_details from both Usage objects."""
|
||||
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
|
||||
|
||||
llm_caching_handler = LLMCachingHandler(
|
||||
original_function=MagicMock(),
|
||||
request_kwargs={},
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
usage1 = Usage(
|
||||
prompt_tokens=10,
|
||||
completion_tokens=0,
|
||||
total_tokens=10,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(image_count=1),
|
||||
)
|
||||
usage2 = Usage(
|
||||
prompt_tokens=20,
|
||||
completion_tokens=0,
|
||||
total_tokens=20,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(image_count=2),
|
||||
)
|
||||
|
||||
combined = llm_caching_handler.combine_usage(usage1, usage2)
|
||||
|
||||
assert combined.prompt_tokens == 30
|
||||
assert combined.total_tokens == 30
|
||||
assert combined.prompt_tokens_details is not None
|
||||
assert combined.prompt_tokens_details.image_count == 3
|
||||
|
||||
|
||||
def test_combine_usage_handles_none_details():
|
||||
"""Test that combine_usage works when one or both sides have null prompt_tokens_details."""
|
||||
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
|
||||
|
||||
llm_caching_handler = LLMCachingHandler(
|
||||
original_function=MagicMock(),
|
||||
request_kwargs={},
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
# Both null
|
||||
usage_a = Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10)
|
||||
usage_b = Usage(prompt_tokens=20, completion_tokens=0, total_tokens=20)
|
||||
combined = llm_caching_handler.combine_usage(usage_a, usage_b)
|
||||
assert combined.prompt_tokens_details is None
|
||||
|
||||
# Only first has details
|
||||
usage_c = Usage(
|
||||
prompt_tokens=10,
|
||||
completion_tokens=0,
|
||||
total_tokens=10,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(image_count=1),
|
||||
)
|
||||
combined = llm_caching_handler.combine_usage(usage_c, usage_b)
|
||||
assert combined.prompt_tokens_details is not None
|
||||
assert combined.prompt_tokens_details.image_count == 1
|
||||
|
||||
# Only second has details
|
||||
combined = llm_caching_handler.combine_usage(usage_a, usage_c)
|
||||
assert combined.prompt_tokens_details is not None
|
||||
assert combined.prompt_tokens_details.image_count == 1
|
||||
|
||||
|
||||
def test_is_chat_completion_cached_dict():
|
||||
from litellm.caching.caching_handler import _is_chat_completion_cached_dict
|
||||
|
||||
assert _is_chat_completion_cached_dict(
|
||||
{"id": "chatcmpl-abc", "object": "chat.completion", "choices": []}
|
||||
)
|
||||
assert _is_chat_completion_cached_dict(
|
||||
{"id": "other", "object": "chat.completion.chunk", "choices": []}
|
||||
)
|
||||
assert _is_chat_completion_cached_dict(
|
||||
{"id": "no-object", "choices": [{"index": 0}]}
|
||||
)
|
||||
assert not _is_chat_completion_cached_dict(
|
||||
{"id": "resp_abc", "object": "response", "output": []}
|
||||
)
|
||||
|
||||
|
||||
def _build_logging_obj(call_type: str, stream: bool):
|
||||
import uuid as _uuid
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
|
||||
return LiteLLMLogging(
|
||||
litellm_call_id=str(datetime.now()),
|
||||
call_type=call_type,
|
||||
model="gpt-5.4",
|
||||
messages=[],
|
||||
function_id=str(_uuid.uuid4()),
|
||||
stream=stream,
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
|
||||
def test_convert_cached_aresponses_bridge_chat_completion_stream():
|
||||
"""openai/responses chat-completions bridge: streaming cache hit replays as chat stream."""
|
||||
from litellm import aresponses
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
caching_handler = LLMCachingHandler(
|
||||
original_function=aresponses, request_kwargs={}, start_time=datetime.now()
|
||||
)
|
||||
cached_result = {
|
||||
"id": "chatcmpl-bridge-cache-test",
|
||||
"object": "chat.completion",
|
||||
"created": int(time.time()),
|
||||
"model": "gpt-5.4",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "Hi!"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18},
|
||||
}
|
||||
|
||||
result = caching_handler._convert_cached_result_to_model_response(
|
||||
cached_result=cached_result,
|
||||
call_type=CallTypes.aresponses.value,
|
||||
kwargs={
|
||||
"model": "gpt-5.4",
|
||||
"stream": True,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
},
|
||||
logging_obj=_build_logging_obj(CallTypes.aresponses.value, stream=True),
|
||||
model="gpt-5.4",
|
||||
args=(),
|
||||
)
|
||||
|
||||
assert isinstance(result, CustomStreamWrapper)
|
||||
|
||||
|
||||
def test_convert_cached_responses_bridge_chat_completion_nonstream():
|
||||
"""openai/responses chat-completions bridge: non-streaming cache hit replays as ModelResponse."""
|
||||
from litellm import responses
|
||||
from litellm.types.utils import CallTypes, ModelResponse
|
||||
|
||||
caching_handler = LLMCachingHandler(
|
||||
original_function=responses, request_kwargs={}, start_time=datetime.now()
|
||||
)
|
||||
cached_result = {
|
||||
"id": "chatcmpl-bridge-nonstream",
|
||||
"object": "chat.completion",
|
||||
"created": int(time.time()),
|
||||
"model": "gpt-5.4",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "Hi!"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18},
|
||||
}
|
||||
|
||||
result = caching_handler._convert_cached_result_to_model_response(
|
||||
cached_result=cached_result,
|
||||
call_type=CallTypes.responses.value,
|
||||
kwargs={
|
||||
"model": "gpt-5.4",
|
||||
"stream": False,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
},
|
||||
logging_obj=_build_logging_obj(CallTypes.responses.value, stream=False),
|
||||
model="gpt-5.4",
|
||||
args=(),
|
||||
)
|
||||
|
||||
assert isinstance(result, ModelResponse)
|
||||
assert result.choices[0].message.content == "Hi!"
|
||||
|
||||
|
||||
def test_convert_cached_responses_legacy_nonstream_path():
|
||||
"""Genuine ResponsesAPIResponse dict (no chatcmpl/choices) falls through legacy path."""
|
||||
from litellm import responses
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
caching_handler = LLMCachingHandler(
|
||||
original_function=responses, request_kwargs={}, start_time=datetime.now()
|
||||
)
|
||||
cached_result = {
|
||||
"id": "resp_legacy_nonstream",
|
||||
"created_at": int(time.time()),
|
||||
"status": "completed",
|
||||
"model": "gpt-4o",
|
||||
"object": "response",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_legacy",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "output_text",
|
||||
"text": "legacy response",
|
||||
"annotations": [],
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
result = caching_handler._convert_cached_result_to_model_response(
|
||||
cached_result=cached_result,
|
||||
call_type=CallTypes.responses.value,
|
||||
kwargs={"model": "gpt-4o", "input": "hi", "stream": False},
|
||||
logging_obj=_build_logging_obj(CallTypes.responses.value, stream=False),
|
||||
model="gpt-4o",
|
||||
args=(),
|
||||
)
|
||||
|
||||
assert isinstance(result, ResponsesAPIResponse)
|
||||
assert result.id == "resp_legacy_nonstream"
|
||||
|
||||
|
||||
def test_convert_cached_responses_legacy_stream_path():
|
||||
"""Genuine ResponsesAPIResponse dict (no chatcmpl/choices) on stream falls through legacy path."""
|
||||
from litellm import responses
|
||||
from litellm.responses.streaming_iterator import (
|
||||
CachedResponsesAPIStreamingIterator,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
caching_handler = LLMCachingHandler(
|
||||
original_function=responses, request_kwargs={}, start_time=datetime.now()
|
||||
)
|
||||
cached_result = {
|
||||
"id": "resp_legacy_stream",
|
||||
"created_at": int(time.time()),
|
||||
"status": "completed",
|
||||
"model": "gpt-4o",
|
||||
"object": "response",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_legacy_stream",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "output_text",
|
||||
"text": "legacy stream",
|
||||
"annotations": [],
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
result = caching_handler._convert_cached_result_to_model_response(
|
||||
cached_result=cached_result,
|
||||
call_type=CallTypes.responses.value,
|
||||
kwargs={"model": "gpt-4o", "input": "hi", "stream": True},
|
||||
logging_obj=_build_logging_obj(CallTypes.responses.value, stream=True),
|
||||
model="gpt-4o",
|
||||
args=(),
|
||||
)
|
||||
|
||||
assert isinstance(result, CachedResponsesAPIStreamingIterator)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_embedding_cache_restores_stored_prompt_tokens_for_image_input():
|
||||
"""Image-embedding cache hit restores prompt_tokens=0 from the stored value
|
||||
instead of recomputing a bogus count by tokenizing the base64 input."""
|
||||
llm_caching_handler = LLMCachingHandler(
|
||||
original_function=MagicMock(),
|
||||
request_kwargs={},
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
# base64-like blob — token_counter over this would return a large nonzero count
|
||||
image_input = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk" * 50
|
||||
|
||||
cached_result = [
|
||||
{
|
||||
"embedding": [-0.025, -0.019],
|
||||
"index": 0,
|
||||
"object": "embedding",
|
||||
"model": "amazon.titan-embed-image-v1",
|
||||
"prompt_tokens": 0,
|
||||
"prompt_tokens_details": {"image_count": 1},
|
||||
}
|
||||
]
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
|
||||
final_embedding_cached_response=None,
|
||||
cached_result=cached_result,
|
||||
kwargs={"model": "amazon.titan-embed-image-v1", "input": image_input},
|
||||
logging_obj=mock_logging_obj,
|
||||
start_time=datetime.now(),
|
||||
model="amazon.titan-embed-image-v1",
|
||||
)
|
||||
|
||||
assert cache_hit
|
||||
assert response.usage is not None
|
||||
assert response.usage.prompt_tokens == 0
|
||||
assert response.usage.total_tokens == 0
|
||||
assert response.usage.prompt_tokens_details.image_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_embedding_cache_sums_stored_prompt_tokens_across_items():
|
||||
"""A multi-item cache hit sums the stored per-item prompt_tokens back to the total."""
|
||||
llm_caching_handler = LLMCachingHandler(
|
||||
original_function=MagicMock(),
|
||||
request_kwargs={},
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
cached_result = [
|
||||
{
|
||||
"embedding": [-0.01],
|
||||
"index": 0,
|
||||
"object": "embedding",
|
||||
"model": "text-embedding-3-small",
|
||||
"prompt_tokens": 5,
|
||||
},
|
||||
{
|
||||
"embedding": [-0.02],
|
||||
"index": 1,
|
||||
"object": "embedding",
|
||||
"model": "text-embedding-3-small",
|
||||
"prompt_tokens": 4,
|
||||
},
|
||||
]
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
|
||||
final_embedding_cached_response=None,
|
||||
cached_result=cached_result,
|
||||
kwargs={"model": "text-embedding-3-small", "input": ["hello world", "foo bar"]},
|
||||
logging_obj=mock_logging_obj,
|
||||
start_time=datetime.now(),
|
||||
model="text-embedding-3-small",
|
||||
)
|
||||
|
||||
assert cache_hit
|
||||
assert response.usage.prompt_tokens == 9
|
||||
assert response.usage.total_tokens == 9
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_embedding_cache_falls_back_to_token_counter_for_legacy_entries():
|
||||
"""Legacy cache entries with no stored prompt_tokens still recompute via token_counter
|
||||
for str inputs (backward compatibility)."""
|
||||
llm_caching_handler = LLMCachingHandler(
|
||||
original_function=MagicMock(),
|
||||
request_kwargs={},
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
# No prompt_tokens key — pre-fix entry
|
||||
cached_result = [
|
||||
{
|
||||
"embedding": [-0.025, -0.019],
|
||||
"index": 0,
|
||||
"object": "embedding",
|
||||
"model": "text-embedding-ada-002",
|
||||
},
|
||||
]
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
|
||||
final_embedding_cached_response=None,
|
||||
cached_result=cached_result,
|
||||
kwargs={"model": "text-embedding-ada-002", "input": "hello world"},
|
||||
logging_obj=mock_logging_obj,
|
||||
start_time=datetime.now(),
|
||||
model="text-embedding-ada-002",
|
||||
)
|
||||
|
||||
assert cache_hit
|
||||
# token_counter over "hello world" yields a nonzero count — fallback path still runs
|
||||
assert response.usage.prompt_tokens > 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_embedding_cache_hit_sets_custom_llm_provider_on_logging_obj():
|
||||
"""A full embedding cache hit must stamp the resolved provider onto the logging
|
||||
obj so spend logs record the provider instead of None/unknown."""
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
llm_caching_handler = LLMCachingHandler(
|
||||
original_function=MagicMock(),
|
||||
request_kwargs={},
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
cached_result = [
|
||||
{
|
||||
"embedding": [-0.025, -0.019],
|
||||
"index": 0,
|
||||
"object": "embedding",
|
||||
"model": "text-embedding-3-small",
|
||||
"prompt_tokens": 5,
|
||||
}
|
||||
]
|
||||
|
||||
logging_obj = _build_logging_obj(CallTypes.aembedding.value, stream=False)
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
|
||||
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
|
||||
final_embedding_cached_response=None,
|
||||
cached_result=cached_result,
|
||||
kwargs={"model": "text-embedding-3-small", "input": "hello world"},
|
||||
logging_obj=logging_obj,
|
||||
start_time=datetime.now(),
|
||||
model="text-embedding-3-small",
|
||||
)
|
||||
|
||||
assert cache_hit
|
||||
assert logging_obj.model_call_details["custom_llm_provider"] == "openai"
|
||||
|
||||
|
||||
def test_sync_stream_responses_cache_hit_sets_custom_llm_provider_on_logging_obj(monkeypatch):
|
||||
import litellm
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
|
||||
kwargs = {"model": "azure/gpt-5.4-mini", "input": "hello", "stream": True}
|
||||
cached_response = {
|
||||
"id": "resp_sync_stream",
|
||||
"created_at": int(time.time()),
|
||||
"status": "completed",
|
||||
"model": "gpt-5.4-mini",
|
||||
"object": "response",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_sync_stream",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
|
||||
}
|
||||
],
|
||||
}
|
||||
litellm.cache.add_cache(json.dumps(cached_response), **kwargs)
|
||||
handler = LLMCachingHandler(original_function=litellm.responses, request_kwargs=kwargs, start_time=datetime.now())
|
||||
logging_obj = _build_logging_obj(CallTypes.responses.value, stream=True)
|
||||
|
||||
hit = handler._sync_get_cache(
|
||||
model="azure/gpt-5.4-mini",
|
||||
original_function=litellm.responses,
|
||||
logging_obj=logging_obj,
|
||||
start_time=datetime.now(),
|
||||
call_type=CallTypes.responses.value,
|
||||
kwargs=kwargs,
|
||||
args=(),
|
||||
)
|
||||
|
||||
assert hit.cached_result is not None
|
||||
assert logging_obj.model_call_details["custom_llm_provider"] == "azure"
|
||||
assert logging_obj.model_call_details["litellm_params"]["custom_llm_provider"] == "azure"
|
||||
|
||||
|
||||
def test_request_kwargs_does_not_retain_logging_obj():
|
||||
"""
|
||||
The caching handler lives on logging_obj._llm_caching_handler, so keeping
|
||||
litellm_logging_obj inside request_kwargs closes a reference cycle
|
||||
(Logging -> LLMCachingHandler -> kwargs -> Logging). That cycle keeps the
|
||||
full request payload alive until a generational GC pass instead of being
|
||||
freed by refcount when the request finishes; under bursts of large-token
|
||||
requests this presents as stepwise RSS growth that never returns to
|
||||
baseline. Other kwargs (messages included) must be preserved.
|
||||
"""
|
||||
logging_obj = MagicMock()
|
||||
kwargs = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"litellm_logging_obj": logging_obj,
|
||||
}
|
||||
|
||||
handler = LLMCachingHandler(
|
||||
original_function=MagicMock(),
|
||||
request_kwargs=kwargs,
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
|
||||
assert "litellm_logging_obj" not in handler.request_kwargs
|
||||
assert handler.request_kwargs["messages"] == kwargs["messages"]
|
||||
assert handler.request_kwargs["model"] == "gpt-4o"
|
||||
|
||||
|
||||
def test_async_cache_write_completes_when_asyncio_run_closes_the_loop(monkeypatch):
|
||||
"""
|
||||
Regression test for the SDK losing async cache writes in short-lived scripts:
|
||||
async_set_cache dispatched the write as a bare fire-and-forget task, so
|
||||
asyncio.run cancelled it at loop close before the write landed (LIT-6184,
|
||||
deterministic with hiredis installed). The write must survive loop shutdown.
|
||||
"""
|
||||
import litellm
|
||||
|
||||
writes = []
|
||||
|
||||
class _SlowWriteCache:
|
||||
supported_call_types = ["acompletion"]
|
||||
cache = None
|
||||
|
||||
async def async_add_cache(self, result, dynamic_cache_object=None, **kwargs):
|
||||
await asyncio.sleep(0.2)
|
||||
writes.append(result)
|
||||
|
||||
async def acompletion(**kwargs):
|
||||
return None
|
||||
|
||||
handler = LLMCachingHandler(
|
||||
original_function=acompletion,
|
||||
request_kwargs={},
|
||||
start_time=datetime.now(),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "cache", _SlowWriteCache())
|
||||
|
||||
async def _short_lived_script():
|
||||
await handler.async_set_cache(
|
||||
result=litellm.ModelResponse(),
|
||||
original_function=acompletion,
|
||||
kwargs={},
|
||||
)
|
||||
|
||||
asyncio.run(_short_lived_script())
|
||||
|
||||
assert len(writes) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cache_hit_records_the_looked_up_key_as_the_preset_cache_key(monkeypatch):
|
||||
"""The spend log for a cache hit must reuse the key the lookup already computed instead of hashing again."""
|
||||
import litellm
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
async def acompletion(**kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
|
||||
kwargs = {"model": "gpt-5.4", "messages": [{"role": "user", "content": "hello"}], "caching": True}
|
||||
await litellm.cache.async_add_cache(
|
||||
litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "hi"}}]), **kwargs
|
||||
)
|
||||
handler = LLMCachingHandler(original_function=acompletion, request_kwargs=kwargs, start_time=datetime.now())
|
||||
logging_obj = _build_logging_obj(CallTypes.acompletion.value, stream=False)
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
|
||||
hit = await handler._async_get_cache(
|
||||
model="gpt-5.4",
|
||||
original_function=acompletion,
|
||||
logging_obj=logging_obj,
|
||||
start_time=datetime.now(),
|
||||
call_type=CallTypes.acompletion.value,
|
||||
kwargs=kwargs,
|
||||
args=(),
|
||||
)
|
||||
|
||||
assert hit is not None and hit.cached_result is not None
|
||||
assert handler.preset_cache_key is not None
|
||||
assert logging_obj.litellm_params["preset_cache_key"] == handler.preset_cache_key
|
||||
assert hit.cached_result._hidden_params["cache_key"] == handler.preset_cache_key
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_converted_stream_cache_hit_replayed_as_plain_object_logs_at_hit_time(monkeypatch):
|
||||
import litellm
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
async def aanthropic_messages(**kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
|
||||
kwargs = {
|
||||
"model": "claude-sonnet-5",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"max_tokens": 16,
|
||||
"caching": True,
|
||||
"stream": False,
|
||||
"_websearch_interception_converted_stream": True,
|
||||
}
|
||||
cached_message = {
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "hi"}],
|
||||
}
|
||||
await litellm.cache.async_add_cache(cached_message, **kwargs)
|
||||
handler = LLMCachingHandler(original_function=aanthropic_messages, request_kwargs=kwargs, start_time=datetime.now())
|
||||
logging_obj = _build_logging_obj(CallTypes.aanthropic_messages.value, stream=False)
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock()
|
||||
|
||||
hit = await handler._async_get_cache(
|
||||
model="claude-sonnet-5",
|
||||
original_function=aanthropic_messages,
|
||||
logging_obj=logging_obj,
|
||||
start_time=datetime.now(),
|
||||
call_type=CallTypes.aanthropic_messages.value,
|
||||
kwargs=kwargs,
|
||||
args=(),
|
||||
)
|
||||
|
||||
assert hit is not None and hit.cached_result == cached_message
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once()
|
||||
assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_args.kwargs["cache_hit"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agentic_loop_followup_cache_hit_with_converted_stream_marker_replays_as_plain_object(monkeypatch):
|
||||
import litellm
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
async def acompletion(**kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
|
||||
kwargs = {
|
||||
"model": "gpt-5.6",
|
||||
"messages": [{"role": "user", "content": "run the code"}],
|
||||
"caching": True,
|
||||
"stream": False,
|
||||
"_code_interpreter_interception_converted_stream": True,
|
||||
"_agentic_loop_depth": 1,
|
||||
}
|
||||
await litellm.cache.async_add_cache(
|
||||
litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "done"}}]), **kwargs
|
||||
)
|
||||
handler = LLMCachingHandler(original_function=acompletion, request_kwargs=kwargs, start_time=datetime.now())
|
||||
logging_obj = _build_logging_obj(CallTypes.acompletion.value, stream=False)
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock()
|
||||
|
||||
hit = await handler._async_get_cache(
|
||||
model="gpt-5.6",
|
||||
original_function=acompletion,
|
||||
logging_obj=logging_obj,
|
||||
start_time=datetime.now(),
|
||||
call_type=CallTypes.acompletion.value,
|
||||
kwargs=kwargs,
|
||||
args=(),
|
||||
)
|
||||
|
||||
assert hit is not None and isinstance(hit.cached_result, litellm.ModelResponse)
|
||||
assert hit.cached_result.choices[0].message.content == "done"
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once()
|
||||
assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_args.kwargs["cache_hit"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_partial_embedding_cache_hit_sends_only_misses_and_keeps_input_order(monkeypatch):
|
||||
import litellm
|
||||
from litellm import CustomLLM
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.utils import Embedding, EmbeddingResponse
|
||||
|
||||
class RecordingEmbedder(CustomLLM):
|
||||
provider_inputs: tuple[tuple[str, ...], ...] = ()
|
||||
|
||||
async def aembedding(self, model, input, model_response, **kwargs) -> EmbeddingResponse:
|
||||
self.provider_inputs = (*self.provider_inputs, tuple(input))
|
||||
return EmbeddingResponse(
|
||||
model=model,
|
||||
data=[
|
||||
Embedding(embedding=[float(len(text))], index=idx, object="embedding")
|
||||
for idx, text in enumerate(input)
|
||||
],
|
||||
)
|
||||
|
||||
embedder = RecordingEmbedder()
|
||||
monkeypatch.setattr(litellm, "custom_provider_map", [{"provider": "recording-embedder", "custom_handler": embedder}])
|
||||
monkeypatch.setattr(litellm, "provider_list", [*litellm.provider_list, "recording-embedder"])
|
||||
monkeypatch.setattr(litellm, "_custom_providers", [*litellm._custom_providers, "recording-embedder"])
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
|
||||
|
||||
await litellm.aembedding(model="recording-embedder/m", input=["aa", "bbbb"])
|
||||
await asyncio.gather(*_PENDING_CACHE_WRITES)
|
||||
mixed_input = ["c", "aa", "ddd", "bbbb", "eeeee"]
|
||||
response = await litellm.aembedding(model="recording-embedder/m", input=mixed_input)
|
||||
await asyncio.gather(*_PENDING_CACHE_WRITES)
|
||||
|
||||
assert embedder.provider_inputs == (("aa", "bbbb"), ("c", "ddd", "eeeee")), embedder.provider_inputs
|
||||
assert [item["index"] for item in response.data] == [0, 1, 2, 3, 4]
|
||||
assert [item["embedding"] for item in response.data] == [[float(len(text))] for text in mixed_input]
|
||||
assert response._hidden_params["cache_hit"] is True, "a partial hit must still be reported as a cache hit"
|
||||
|
||||
repeat = await litellm.aembedding(model="recording-embedder/m", input=mixed_input)
|
||||
|
||||
assert len(embedder.provider_inputs) == 2, embedder.provider_inputs
|
||||
assert [item["embedding"] for item in repeat.data] == [[float(len(text))] for text in mixed_input]
|
||||
|
|
@ -22,19 +22,6 @@ import litellm
|
|||
from litellm import router as litellm_router_module
|
||||
from litellm import utils as litellm_utils_module
|
||||
from litellm._logging import ALL_LOGGERS
|
||||
from litellm.litellm_core_utils.cli_keyring import (
|
||||
KeyringDiscardsWrites,
|
||||
KeyringUnreachable,
|
||||
KeyringUnusable,
|
||||
SecretErase,
|
||||
SecretErased,
|
||||
SecretFound,
|
||||
SecretMissing,
|
||||
SecretRead,
|
||||
SecretStored,
|
||||
SecretStranded,
|
||||
SecretWrite,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates import (
|
||||
image_handling as image_handling_module,
|
||||
)
|
||||
|
|
@ -42,6 +29,7 @@ from litellm.llms.custom_httpx.async_client_cleanup import (
|
|||
close_litellm_async_clients,
|
||||
)
|
||||
from litellm.proxy.db import tool_registry_writer as tool_registry_writer_module
|
||||
from tests.unit.litellm_core_utils.fake_secret_vault import FakeSecretVault
|
||||
|
||||
|
||||
def _reset_module_level_aws_auth_caches():
|
||||
|
|
@ -128,60 +116,6 @@ def isolate_host_os_keychain(monkeypatch):
|
|||
monkeypatch.setenv("LITELLM_CLI_DISABLE_KEYRING", "1")
|
||||
|
||||
|
||||
class FakeSecretVault:
|
||||
"""In-memory stand-in for the OS keychain, injected wherever CLI credential storage is exercised.
|
||||
|
||||
`available=False` models a keychain that is locked or has no backend, `writable=False` one that
|
||||
refuses to store, `erasable=False` one that will not release what it already holds, and `failure`
|
||||
picks which unusable state those report. `discards=True` is keyring's null backend, which answers
|
||||
reads and erases like any other yet keeps nothing it is given, so only writes report it.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
blob: str | None = None,
|
||||
*,
|
||||
available: bool = True,
|
||||
writable: bool = True,
|
||||
erasable: bool = True,
|
||||
discards: bool = False,
|
||||
failure: KeyringUnusable = KeyringUnreachable(),
|
||||
) -> None:
|
||||
self.blob: str | None = blob
|
||||
self.available: bool = available
|
||||
self.writable: bool = writable
|
||||
self.erasable: bool = erasable
|
||||
self.discards: bool = discards
|
||||
self.failure: KeyringUnusable = failure
|
||||
self.reads: int = 0
|
||||
self.writes: list[str] = []
|
||||
self.erases: int = 0
|
||||
|
||||
def read(self) -> SecretRead:
|
||||
self.reads += 1
|
||||
if not self.available:
|
||||
return self.failure
|
||||
return SecretMissing() if self.blob is None else SecretFound(self.blob)
|
||||
|
||||
def write(self, blob: str) -> SecretWrite:
|
||||
self.writes.append(blob)
|
||||
if not (self.available and self.writable):
|
||||
return self.failure
|
||||
if self.discards:
|
||||
return KeyringDiscardsWrites()
|
||||
self.blob = blob
|
||||
return SecretStored()
|
||||
|
||||
def erase(self) -> SecretErase:
|
||||
self.erases += 1
|
||||
if not self.available:
|
||||
return self.failure
|
||||
if not self.erasable:
|
||||
return SecretStranded() if self.blob is not None else SecretErased()
|
||||
self.blob = None
|
||||
return SecretErased()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def secret_vault_factory():
|
||||
"""Build FakeSecretVault instances; see its docstring for the failure modes it can model."""
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
# Levo integration tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# This file makes the tests/litellm/litellm_core_utils directory a Python package
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -1,403 +1,20 @@
|
|||
import copy
|
||||
import os
|
||||
import pickle
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
import tiktoken
|
||||
from tokenizers import Tokenizer as ReferenceTokenizer
|
||||
|
||||
import litellm
|
||||
from litellm.caching._embedding_router import truncate_embedding_input
|
||||
from litellm.litellm_core_utils.tokenizer import HuggingFaceTokenizer, OpenAIEncoding
|
||||
from litellm.utils import claude_json_str
|
||||
from tests.test_litellm.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"name", ("cl100k_base", "o200k_base", "p50k_base", "p50k_edit", "r50k_base", "gpt2", "o200k_harmony")
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"text", ("hello world", "café 漢字 🙂", "", "a\ud800b", "\ud83d\ude42", "🙂\ud83d\ude42\udfff", " " * 64)
|
||||
from tests.unit.litellm_core_utils.test_tokenizer import (
|
||||
UNICODE_TEXTS,
|
||||
assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface,
|
||||
assert_openai_encoding_matches_python,
|
||||
)
|
||||
|
||||
NETWORK_ENCODINGS = ("r50k_base", "gpt2")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name", NETWORK_ENCODINGS)
|
||||
@pytest.mark.parametrize("text", UNICODE_TEXTS)
|
||||
def test_openai_encoding_matches_python_unicode_and_batches(name: str, text: str) -> None:
|
||||
reference: Final = tiktoken.get_encoding(name)
|
||||
encoding: Final = OpenAIEncoding.from_tiktoken(name)
|
||||
expected: Final = reference.encode(text)
|
||||
|
||||
assert encoding.encode(text) == expected
|
||||
assert encoding.count(text) == len(expected)
|
||||
assert encoding.encode_batch([text], num_threads=2) == reference.encode_batch([text], num_threads=2)
|
||||
assert encoding.encode_ordinary_batch([text]) == reference.encode_ordinary_batch([text])
|
||||
assert encoding.decode_batch([expected]) == reference.decode_batch([expected])
|
||||
assert encoding.decode_bytes_batch([expected]) == reference.decode_bytes_batch([expected])
|
||||
assert_openai_encoding_matches_python(name, text)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("allowed", (frozenset(), frozenset({"<|endoftext|>"}), "all"))
|
||||
@pytest.mark.parametrize("disallowed", (frozenset(), frozenset({"<|fim_prefix|>"}), "all"))
|
||||
def test_openai_special_token_options_match_python(
|
||||
allowed: frozenset[str] | Literal["all"], disallowed: frozenset[str] | Literal["all"]
|
||||
) -> None:
|
||||
reference: Final = tiktoken.get_encoding("cl100k_base")
|
||||
encoding: Final = OpenAIEncoding.from_tiktoken(reference.name)
|
||||
text: Final = "hello<|endoftext|><|fim_prefix|>world"
|
||||
allowed_set: Final = reference.special_tokens_set if allowed == "all" else allowed
|
||||
disallowed_set: Final = reference.special_tokens_set - allowed_set if disallowed == "all" else disallowed
|
||||
if any(token in text for token in disallowed_set):
|
||||
with pytest.raises(ValueError, match="disallowed special token"):
|
||||
encoding.encode(text, allowed_special=allowed, disallowed_special=disallowed)
|
||||
return
|
||||
assert encoding.encode(text, allowed_special=allowed, disallowed_special=disallowed) == reference.encode(
|
||||
text, allowed_special=allowed, disallowed_special=disallowed
|
||||
)
|
||||
assert encoding.special_tokens_set == reference.special_tokens_set
|
||||
assert encoding.eot_token == reference.eot_token
|
||||
|
||||
|
||||
@pytest.mark.parametrize("errors", ("replace", "ignore", "backslashreplace", "strict"))
|
||||
def test_openai_partial_token_decoding_preserves_error_policy(errors: str) -> None:
|
||||
reference: Final = tiktoken.get_encoding("cl100k_base")
|
||||
encoding: Final = OpenAIEncoding.from_tiktoken(reference.name)
|
||||
tokens: Final = reference.encode("🙂")[:1]
|
||||
assert encoding.decode_bytes(tokens) == reference.decode_bytes(tokens)
|
||||
if errors == "strict":
|
||||
with pytest.raises(UnicodeDecodeError):
|
||||
encoding.decode(tokens, errors=errors)
|
||||
return
|
||||
assert encoding.decode(tokens, errors=errors) == reference.decode(tokens, errors=errors)
|
||||
assert encoding.decode_tokens_bytes(tokens) == reference.decode_tokens_bytes(tokens)
|
||||
|
||||
|
||||
def test_public_encoding_and_semantic_cache_preserve_truncated_unicode() -> None:
|
||||
reference: Final = tiktoken.get_encoding(litellm.encoding.name)
|
||||
text: Final = "🙂"
|
||||
tokens: Final = reference.encode(text)
|
||||
|
||||
assert litellm.encoding.encode(text, disallowed_special=()) == tokens
|
||||
assert litellm.encoding.encode_batch([text]) == [tokens]
|
||||
assert litellm.decode(tokens=tokens[:1]) == reference.decode(tokens[:1])
|
||||
assert truncate_embedding_input(text, "", 1) == reference.decode(tokens[:1])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("add_special_tokens", (True, False))
|
||||
def test_huggingface_encoding_preserves_result_fields_and_serialization(add_special_tokens: bool) -> None:
|
||||
reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON)
|
||||
tokenizer: Final = HuggingFaceTokenizer.from_str(TOKENIZER_JSON)
|
||||
expected: Final = reference.encode("Hello World", add_special_tokens=add_special_tokens)
|
||||
actual: Final = tokenizer.encode("Hello World", add_special_tokens=add_special_tokens)
|
||||
|
||||
assert (actual.ids, actual.tokens, actual.type_ids, actual.offsets, actual.word_ids, actual.sequence_ids) == (
|
||||
expected.ids,
|
||||
expected.tokens,
|
||||
expected.type_ids,
|
||||
expected.offsets,
|
||||
expected.word_ids,
|
||||
expected.sequence_ids,
|
||||
)
|
||||
assert (actual.attention_mask, actual.special_tokens_mask, actual.n_sequences, len(actual)) == (
|
||||
expected.attention_mask,
|
||||
expected.special_tokens_mask,
|
||||
expected.n_sequences,
|
||||
len(expected),
|
||||
)
|
||||
assert copy.deepcopy(actual).ids == expected.ids
|
||||
assert pickle.loads(pickle.dumps(actual)).offsets == expected.offsets
|
||||
assert tokenizer.decode(actual.ids, skip_special_tokens=False) == reference.decode(
|
||||
expected.ids, skip_special_tokens=False
|
||||
)
|
||||
|
||||
|
||||
def test_huggingface_character_offsets_and_pretokenized_pairs_match_python() -> None:
|
||||
reference: Final = ReferenceTokenizer.from_str(claude_json_str)
|
||||
tokenizer: Final = HuggingFaceTokenizer.from_str(claude_json_str)
|
||||
text: Final = "café 漢字 🙂"
|
||||
actual: Final = tokenizer.encode(text)
|
||||
expected: Final = reference.encode(text)
|
||||
|
||||
assert actual.offsets == expected.offsets
|
||||
assert actual.ids == expected.ids
|
||||
assert (
|
||||
tokenizer.encode(["hello", "world"], ["again"], is_pretokenized=True).ids
|
||||
== reference.encode(["hello", "world"], ["again"], is_pretokenized=True).ids
|
||||
)
|
||||
|
||||
|
||||
def test_huggingface_batches_apply_padding_across_inputs() -> None:
|
||||
reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON)
|
||||
reference.enable_padding(pad_id=0, pad_token="[UNK]")
|
||||
tokenizer: Final = HuggingFaceTokenizer.from_str(reference.to_str())
|
||||
inputs: Final = ["Hello", ("Hello World", "World")]
|
||||
expected: Final = reference.encode_batch(inputs)
|
||||
actual: Final = tokenizer.encode_batch(inputs)
|
||||
fast: Final = tokenizer.encode_batch_fast(inputs)
|
||||
|
||||
assert [(item.ids, item.attention_mask, item.offsets) for item in actual] == [
|
||||
(item.ids, item.attention_mask, item.offsets) for item in expected
|
||||
]
|
||||
assert [item.ids for item in fast] == [item.ids for item in expected]
|
||||
assert tokenizer.decode_batch([item.ids for item in actual]) == reference.decode_batch(
|
||||
[item.ids for item in expected]
|
||||
)
|
||||
|
||||
|
||||
def test_caller_supplied_huggingface_tokenizer_preserves_public_encode_and_count() -> None:
|
||||
tokenizer: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON)
|
||||
custom: Final = {"type": "huggingface_tokenizer", "tokenizer": tokenizer}
|
||||
expected: Final = tokenizer.encode("Hello World").ids
|
||||
|
||||
assert litellm.encode(text="Hello World", custom_tokenizer=custom) == expected
|
||||
assert litellm.token_counter(text="Hello World", custom_tokenizer=custom) == len(expected)
|
||||
assert litellm.decode(tokens=expected, custom_tokenizer=custom) == "Hello World"
|
||||
|
||||
|
||||
def test_caller_supplied_tiktoken_treats_special_spellings_as_text() -> None:
|
||||
tokenizer: Final = tiktoken.get_encoding("cl100k_base")
|
||||
custom: Final = {"type": "openai_tokenizer", "tokenizer": tokenizer}
|
||||
text: Final = "<|endoftext|>"
|
||||
|
||||
assert litellm.encode(text=text, custom_tokenizer=custom) == tokenizer.encode(text, disallowed_special=())
|
||||
|
||||
|
||||
def test_public_tokenizer_objects_survive_pickle_and_deepcopy(tmp_path: Path) -> None:
|
||||
custom: Final = litellm.create_tokenizer(TOKENIZER_JSON)
|
||||
tokenizer: Final = custom["tokenizer"]
|
||||
path: Final = tmp_path / "tokenizer.json"
|
||||
tokenizer.save(str(path))
|
||||
|
||||
assert copy.deepcopy(custom)["tokenizer"].encode("Hello World").ids == tokenizer.encode("Hello World").ids
|
||||
assert (
|
||||
pickle.loads(pickle.dumps(custom))["tokenizer"].encode("Hello World").ids == tokenizer.encode("Hello World").ids
|
||||
)
|
||||
assert HuggingFaceTokenizer.from_file(str(path)).encode("Hello World").ids == tokenizer.encode("Hello World").ids
|
||||
assert copy.deepcopy(litellm.encoding).encode("hello") == litellm.encoding.encode("hello")
|
||||
assert pickle.loads(pickle.dumps(litellm.encoding)).encode("hello") == litellm.encoding.encode("hello")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("offline", ("0", "1"))
|
||||
def test_hub_loader_preserves_environment_auth_cache_and_offline(tmp_path: Path, offline: str) -> None:
|
||||
script: Final = """
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
sys.path.insert(0, sys.argv[1])
|
||||
import httpx
|
||||
import huggingface_hub
|
||||
from huggingface_hub.errors import LocalEntryNotFoundError
|
||||
import litellm
|
||||
payload = sys.argv[2].encode()
|
||||
offline = sys.argv[3] == "1"
|
||||
observed = []
|
||||
def handle(request):
|
||||
assert not offline, "offline loading issued a request"
|
||||
if request.url.path.endswith("/tokenizer.json"):
|
||||
observed.append(request.headers.get("authorization"))
|
||||
if request.headers.get("authorization") != "Bearer audit-fixture-token":
|
||||
return httpx.Response(401)
|
||||
return httpx.Response(200, headers={"content-length": str(len(payload)), "etag": '"fixture"', "x-repo-commit": "a" * 40}, content=payload if request.method == "GET" else b"")
|
||||
if not offline:
|
||||
huggingface_hub.set_client_factory(lambda: httpx.Client(transport=httpx.MockTransport(handle)))
|
||||
try:
|
||||
tokenizer = litellm.create_pretrained_tokenizer("test-fixture/tokenizer")["tokenizer"]
|
||||
except LocalEntryNotFoundError:
|
||||
assert offline
|
||||
assert observed == []
|
||||
else:
|
||||
assert not offline
|
||||
assert "Bearer audit-fixture-token" in observed
|
||||
assert tokenizer.decode(tokenizer.encode("Hello World").ids) == "Hello World"
|
||||
assert tuple(Path(sys.argv[4]).rglob("tokenizer.json"))
|
||||
print("compatible")
|
||||
"""
|
||||
result: Final = subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
"-I",
|
||||
"-c",
|
||||
script,
|
||||
str(Path(litellm.__file__).parent.parent),
|
||||
TOKENIZER_JSON,
|
||||
offline,
|
||||
str(tmp_path / "cache"),
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=30,
|
||||
env={
|
||||
**os.environ,
|
||||
"HF_HOME": str(tmp_path / "home"),
|
||||
"HF_HUB_CACHE": str(tmp_path / "cache"),
|
||||
"HF_ENDPOINT": "http://127.0.0.1:9",
|
||||
"HF_TOKEN": "audit-fixture-token",
|
||||
"HF_HUB_OFFLINE": offline,
|
||||
"HF_HUB_DISABLE_IMPLICIT_TOKEN": "0",
|
||||
"LITELLM_LOCAL_MODEL_COST_MAP": "True",
|
||||
},
|
||||
)
|
||||
assert result.returncode == 0, result.stdout + result.stderr
|
||||
assert result.stdout.strip() == "compatible"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rust", (None, "0", "1"))
|
||||
def test_tokenization_without_native_extension_stays_offline(tmp_path: Path, rust: str | None) -> None:
|
||||
script: Final = """
|
||||
import importlib.abc
|
||||
import sys
|
||||
sys.path.insert(0, sys.argv[1])
|
||||
def reject_network(event, args):
|
||||
if event == "socket.connect":
|
||||
raise AssertionError("tokenizer attempted a network connection")
|
||||
sys.addaudithook(reject_network)
|
||||
class Block(importlib.abc.MetaPathFinder):
|
||||
def find_spec(self, fullname, path=None, target=None):
|
||||
if fullname == "litellm.rust_bridge._native":
|
||||
raise ImportError("native extension is unavailable")
|
||||
sys.meta_path.insert(0, Block())
|
||||
import litellm
|
||||
from litellm.rust_bridge.tokenizer import get_encoding
|
||||
import tiktoken
|
||||
from tokenizers import Tokenizer
|
||||
assert isinstance(litellm.encoding, tiktoken.Encoding)
|
||||
for name in ("cl100k_base", "o200k_base", "o200k_harmony", "p50k_base", "p50k_edit"):
|
||||
encoding = get_encoding(name)
|
||||
text = "offline café 漢字 🙂" + " " * 64
|
||||
assert encoding.decode(encoding.encode(text)) == text
|
||||
ids = litellm.encode(text="hello world")
|
||||
assert litellm.decode(tokens=ids) == "hello world"
|
||||
assert litellm.token_counter(model=None, text="hello world") == len(ids)
|
||||
custom = litellm.create_tokenizer(sys.argv[2])
|
||||
assert isinstance(custom["tokenizer"], Tokenizer)
|
||||
custom["tokenizer"].enable_padding(pad_id=0, pad_token="[UNK]")
|
||||
assert litellm.decode(tokens=litellm.encode(text="Hello World", custom_tokenizer=custom), custom_tokenizer=custom) == "Hello World"
|
||||
print("compatible")
|
||||
"""
|
||||
result: Final = subprocess.run(
|
||||
[sys.executable, "-I", "-c", script, str(Path(litellm.__file__).parent.parent), TOKENIZER_JSON],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=30,
|
||||
cwd=tmp_path,
|
||||
env={
|
||||
**{key: value for key, value in os.environ.items() if key != "LITELLM_RUST"},
|
||||
**({"LITELLM_RUST": rust} if rust is not None else {}),
|
||||
"LITELLM_LOCAL_MODEL_COST_MAP": "True",
|
||||
"TIKTOKEN_CACHE_DIR": str(tmp_path / "unused-tokenizer-cache"),
|
||||
},
|
||||
)
|
||||
assert result.returncode == 0, result.stdout + result.stderr
|
||||
assert result.stdout.strip() == "compatible"
|
||||
assert not (tmp_path / "unused-tokenizer-cache").exists()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_pretokenized", (False, True))
|
||||
def test_huggingface_batch_sequence_containers_match_python(is_pretokenized: bool) -> None:
|
||||
reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON)
|
||||
tokenizer: Final = HuggingFaceTokenizer.from_str(TOKENIZER_JSON)
|
||||
inputs: Final = [["Hello", "World"], ("Hello", "World")]
|
||||
actual: Final = tokenizer.encode_batch(inputs, is_pretokenized=is_pretokenized)
|
||||
expected: Final = reference.encode_batch(inputs, is_pretokenized=is_pretokenized)
|
||||
assert [(item.ids, item.type_ids, item.sequence_ids) for item in actual] == [
|
||||
(item.ids, item.type_ids, item.sequence_ids) for item in expected
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name", ("cl100k_base", "o200k_base", "p50k_edit", "gpt2"))
|
||||
@pytest.mark.parametrize("name", ("gpt2",))
|
||||
def test_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name: str) -> None:
|
||||
reference: Final = tiktoken.get_encoding(name)
|
||||
encoding: Final = OpenAIEncoding.from_tiktoken(name)
|
||||
text: Final = "hello fanta"
|
||||
|
||||
assert repr(encoding) == repr(reference) == f"<Encoding {name!r}>"
|
||||
assert (encoding.name, encoding.n_vocab, encoding.max_token_value) == (
|
||||
reference.name,
|
||||
reference.n_vocab,
|
||||
reference.max_token_value,
|
||||
)
|
||||
assert encoding.token_byte_values() == reference.token_byte_values()
|
||||
assert encoding.encode_single_token("hello") == reference.encode_single_token("hello")
|
||||
assert encoding.encode_single_token(b"<|endoftext|>") == reference.eot_token
|
||||
assert [encoding.is_special_token(token) for token in (0, reference.eot_token)] == [False, True]
|
||||
assert encoding.decode_with_offsets(reference.encode(text)) == reference.decode_with_offsets(reference.encode(text))
|
||||
assert encoding.encode_to_numpy(text).tolist() == reference.encode_to_numpy(text).tolist()
|
||||
stable, completions = encoding.encode_with_unstable(text)
|
||||
expected_stable, expected_completions = reference.encode_with_unstable(text)
|
||||
assert (stable, sorted(completions)) == (expected_stable, sorted(expected_completions))
|
||||
with pytest.raises(KeyError):
|
||||
encoding.encode_single_token("<|not-a-token|>")
|
||||
|
||||
|
||||
def test_huggingface_tokenizer_exposes_the_tokenizers_vocabulary_surface() -> None:
|
||||
reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON)
|
||||
reference.enable_padding(pad_id=0, pad_token="[UNK]", length=4)
|
||||
reference.enable_truncation(max_length=3, stride=1, strategy="only_first", direction="left")
|
||||
tokenizer: Final = HuggingFaceTokenizer.from_str(reference.to_str())
|
||||
|
||||
assert tokenizer.token_to_id("Hello") == reference.token_to_id("Hello") == 1
|
||||
assert tokenizer.id_to_token(3) == reference.id_to_token(3) == "[BOS]"
|
||||
assert tokenizer.id_to_token(99) is None
|
||||
assert tokenizer.get_vocab() == reference.get_vocab()
|
||||
assert tokenizer.get_vocab(with_added_tokens=False) == reference.get_vocab(with_added_tokens=False)
|
||||
assert tokenizer.get_vocab_size() == reference.get_vocab_size() == 4
|
||||
assert tokenizer.get_vocab_size(with_added_tokens=False) == reference.get_vocab_size(with_added_tokens=False)
|
||||
added: Final = tokenizer.get_added_tokens_decoder()
|
||||
expected_added: Final = reference.get_added_tokens_decoder()
|
||||
assert {token_id: str(token) for token_id, token in added.items()} == {
|
||||
token_id: str(token) for token_id, token in expected_added.items()
|
||||
}
|
||||
assert added[3].special == expected_added[3].special
|
||||
assert tokenizer.num_special_tokens_to_add(False) == reference.num_special_tokens_to_add(False) == 1
|
||||
assert tokenizer.num_special_tokens_to_add(True) == reference.num_special_tokens_to_add(True) == 0
|
||||
assert tokenizer.padding == reference.padding
|
||||
assert tokenizer.truncation == reference.truncation
|
||||
assert tokenizer.encode_special_tokens == reference.encode_special_tokens is False
|
||||
assert HuggingFaceTokenizer.from_buffer(TOKENIZER_JSON.encode()).encode("Hello").ids == [3, 1]
|
||||
assert HuggingFaceTokenizer.from_str(TOKENIZER_JSON).padding is None
|
||||
assert HuggingFaceTokenizer.from_str(TOKENIZER_JSON).truncation is None
|
||||
|
||||
|
||||
def test_huggingface_encoding_exposes_the_tokenizers_lookup_and_mutation_surface() -> None:
|
||||
reference: Final = ReferenceTokenizer.from_str(claude_json_str)
|
||||
tokenizer: Final = HuggingFaceTokenizer.from_str(claude_json_str)
|
||||
text: Final = "hello wide world"
|
||||
actual: Final = tokenizer.encode(text, "again")
|
||||
expected: Final = reference.encode(text, "again")
|
||||
|
||||
lookups: Final = (
|
||||
lambda encoding: [encoding.token_to_chars(index) for index in range(len(encoding))],
|
||||
lambda encoding: [encoding.token_to_word(index) for index in range(len(encoding))],
|
||||
lambda encoding: [encoding.token_to_sequence(index) for index in range(len(encoding))],
|
||||
lambda encoding: [encoding.char_to_token(position) for position in range(len(text))],
|
||||
lambda encoding: [encoding.char_to_word(position) for position in range(len(text))],
|
||||
lambda encoding: [encoding.char_to_token(position, 1) for position in range(5)],
|
||||
lambda encoding: [encoding.word_to_tokens(word) for word in range(3)],
|
||||
lambda encoding: [encoding.word_to_chars(word) for word in range(3)],
|
||||
lambda encoding: [encoding.word_to_tokens(0, 1), encoding.word_to_chars(0, 1)],
|
||||
)
|
||||
for lookup in lookups:
|
||||
assert lookup(actual) == lookup(expected)
|
||||
assert repr(actual) == repr(expected)
|
||||
|
||||
actual.truncate(4, stride=1, direction="left")
|
||||
expected.truncate(4, stride=1, direction="left")
|
||||
assert (actual.ids, [item.ids for item in actual.overflowing]) == (
|
||||
expected.ids,
|
||||
[item.ids for item in expected.overflowing],
|
||||
)
|
||||
actual.pad(6, direction="left", pad_id=7, pad_type_id=1, pad_token="<pad>")
|
||||
expected.pad(6, direction="left", pad_id=7, pad_type_id=1, pad_token="<pad>")
|
||||
assert (actual.ids, actual.attention_mask, actual.type_ids, actual.tokens) == (
|
||||
expected.ids,
|
||||
expected.attention_mask,
|
||||
expected.type_ids,
|
||||
expected.tokens,
|
||||
)
|
||||
actual.set_sequence_id(3)
|
||||
expected.set_sequence_id(3)
|
||||
assert actual.sequence_ids == expected.sequence_ids
|
||||
merged: Final = type(actual).merge([actual, tokenizer.encode("more")])
|
||||
assert merged.ids == type(expected).merge([expected, reference.encode("more")]).ids
|
||||
assert merged.offsets == type(expected).merge([expected, reference.encode("more")]).offsets
|
||||
with pytest.raises(ValueError, match="direction"):
|
||||
actual.pad(8, direction="sideways")
|
||||
assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,106 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from mcp import Client
|
||||
from mcp.server import Server
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.types import PromptsCapability, ResourcesCapability, ServerCapabilities, ToolsCapability
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.capabilities import (
|
||||
GATEWAY_OPERATIONS,
|
||||
REVISION_SUPPORT,
|
||||
TRANSLATION_PAIRS,
|
||||
GatewayVersionPolicy,
|
||||
build_discovery,
|
||||
)
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision", HANDSHAKE_PROTOCOL_VERSIONS)
|
||||
@pytest.mark.parametrize("transport", tuple(MCPTransport))
|
||||
def test_discovery_only_exposes_authorized_completed_support(revision, transport):
|
||||
result = build_discovery(
|
||||
configured=(revision, "2026-07-28", "unknown"),
|
||||
revision=revision,
|
||||
transport=transport,
|
||||
authorized_operations=frozenset({"tools/list", "tools/call"}),
|
||||
upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS),
|
||||
capabilities=ServerCapabilities(
|
||||
tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability(),
|
||||
extensions={"io.modelcontextprotocol/ui": {}},
|
||||
),
|
||||
client_extensions=frozenset({"io.modelcontextprotocol/ui"}),
|
||||
upstream_extensions=frozenset({"io.modelcontextprotocol/ui"}),
|
||||
)
|
||||
assert result.supported_versions == [revision]
|
||||
assert result.capabilities.tools is not None
|
||||
assert result.capabilities.prompts is None
|
||||
assert result.capabilities.resources is None
|
||||
assert result.capabilities.extensions is None
|
||||
assert result.capabilities.tasks is None
|
||||
assert result.cache_scope == "private"
|
||||
assert result.ttl_ms == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("upstream", [frozenset(), frozenset({"unknown"}), frozenset({"2026-07-28"})])
|
||||
def test_unproven_translation_never_advertises_operations(upstream):
|
||||
result = build_discovery(
|
||||
configured=HANDSHAKE_PROTOCOL_VERSIONS,
|
||||
revision="2025-11-25",
|
||||
transport=MCPTransport.http,
|
||||
authorized_operations=GATEWAY_OPERATIONS,
|
||||
upstream_versions=upstream,
|
||||
capabilities=ServerCapabilities(tools=ToolsCapability()),
|
||||
)
|
||||
assert result.capabilities.tools is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", "2024-11-05"])
|
||||
def test_unadvertised_revision_never_gains_capabilities(revision):
|
||||
result = build_discovery(
|
||||
configured=("2025-11-25",), revision=revision, transport=MCPTransport.http,
|
||||
authorized_operations=GATEWAY_OPERATIONS, upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS),
|
||||
capabilities=ServerCapabilities(tools=ToolsCapability()),
|
||||
)
|
||||
assert result.capabilities.tools is None
|
||||
|
||||
|
||||
def test_discovery_results_do_not_share_mutable_capabilities():
|
||||
capabilities = ServerCapabilities(tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability())
|
||||
args = dict(
|
||||
configured=HANDSHAKE_PROTOCOL_VERSIONS, revision="2025-11-25", transport=MCPTransport.http,
|
||||
upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS), capabilities=capabilities,
|
||||
)
|
||||
allowed = build_discovery(**args, authorized_operations=GATEWAY_OPERATIONS)
|
||||
denied = build_discovery(**args, authorized_operations=frozenset())
|
||||
assert allowed.capabilities.prompts is not None
|
||||
assert allowed.capabilities.resources is not None
|
||||
assert denied.capabilities.model_dump(exclude_none=True) == {}
|
||||
assert allowed.capabilities.tools is not None
|
||||
allowed.capabilities.tools.list_changed = True
|
||||
assert capabilities.tools.list_changed is not True
|
||||
|
||||
|
||||
def test_modern_candidates_do_not_enable_public_serving():
|
||||
modern = REVISION_SUPPORT["2026-07-28"]
|
||||
assert modern.completed is False
|
||||
assert "input_required" in modern.results
|
||||
assert MCPTransport.sse not in modern.transports
|
||||
assert not any("2026-07-28" in pair for pair in TRANSLATION_PAIRS)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("versions,accepted", [(("2025-11-25",), True), (("2025-06-18",), False)])
|
||||
async def test_version_policy_gates_the_actual_sdk_handshake(versions, accepted):
|
||||
server: Final = Server("test-gateway", version="1")
|
||||
server.middleware.append(GatewayVersionPolicy(lambda: versions))
|
||||
if accepted:
|
||||
async with Client(server, mode="legacy") as client:
|
||||
assert client.protocol_version == "2025-11-25"
|
||||
result = await client.session.send_ping()
|
||||
assert result is not None
|
||||
else:
|
||||
with pytest.RaisesGroup(pytest.RaisesExc(MCPError, match="Unsupported MCP protocol version"), flatten_subgroups=True):
|
||||
async with Client(server, mode="legacy"):
|
||||
pytest.fail("The excluded revision must not initialize")
|
||||
|
|
@ -10444,11 +10444,12 @@ async def test_active_request_ctx_var_feeds_auth_resolution_recording(_mcp_reque
|
|||
)
|
||||
@pytest.mark.parametrize("handler", ("handle_streamable_http_mcp", "handle_sse_mcp"))
|
||||
async def test_streamable_http_rejects_modern_protocol_version(
|
||||
header_value: str, expected_rejected: bool, handler: str
|
||||
header_value: str, expected_rejected: bool, handler: str, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||||
scope: Scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
|
|
@ -10659,3 +10660,32 @@ async def test_legacy_sse_mount_emits_message_endpoint(
|
|||
await incoming.put({"type": "http.disconnect"})
|
||||
await asyncio.wait_for(task, 2)
|
||||
assert await post(initialization) == 404
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision,rejected", [("2024-11-05", False), ("2025-11-25", True), ("2026-07-28", True)])
|
||||
def test_protocol_header_respects_configured_advertisement(revision, rejected):
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version
|
||||
|
||||
with patch.dict(proxy_server.general_settings, {"mcp_advertised_versions": ["2024-11-05"]}):
|
||||
result = unsupported_protocol_version({"headers": [(b"mcp-protocol-version", revision.encode())]})
|
||||
assert result == (revision if rejected else None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ctx):
|
||||
from mcp.types import DiscoverResult, RequestParams, ServerCapabilities
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
|
||||
expected = DiscoverResult(supported_versions=["2025-11-25"], capabilities=ServerCapabilities())
|
||||
dispatched = AsyncMock(return_value=expected)
|
||||
auth = UserAPIKeyAuth(user_id="discover-caller")
|
||||
with (
|
||||
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None))),
|
||||
patch.object(server.operations.GatewayOperations, "execute", dispatched),
|
||||
):
|
||||
result = await server.discover(_mcp_request_ctx(), RequestParams())
|
||||
assert result is expected
|
||||
context = dispatched.await_args.args[1]
|
||||
assert context.user_api_key_auth.user_id == "discover-caller"
|
||||
assert context.mcp_servers == ("allowed",)
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.mcp import MCPAuth, MCPAuthType
|
||||
from litellm.types.mcp import MCPAuth, MCPAuthType, MCPUpstreamProtocol
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
|
|
@ -14637,3 +14637,27 @@ class TestSharedIdentifierPrefixWarning:
|
|||
assert "srv-b" in shared_warnings[0]
|
||||
assert "srv-c" not in shared_warnings[0]
|
||||
assert "'shared'" in shared_warnings[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("revision", ["auto", "2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"])
|
||||
async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_manager_factory, revision):
|
||||
manager = config_only_mcp_manager_factory()
|
||||
await manager.load_servers_from_config({"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}})
|
||||
server = next(iter(manager.config_mcp_servers.values()))
|
||||
client = await manager._create_mcp_client(server)
|
||||
assert server.protocol_version == revision
|
||||
assert client.protocol_version == revision
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision", ("auto", "2024-11-05", "2025-06-18"))
|
||||
@pytest.mark.parametrize("explicit", (None, "auto", "2025-11-25"))
|
||||
def test_runtime_protocol_metadata_preserves_explicit_precedence(
|
||||
revision: MCPUpstreamProtocol, explicit: MCPUpstreamProtocol | None
|
||||
) -> None:
|
||||
server: Final = MCPServer.model_validate({
|
||||
"server_id": "preview", "name": "preview", "transport": "http",
|
||||
"mcp_info": {"protocol_version": revision},
|
||||
**({"protocol_version": explicit} if explicit is not None else {}),
|
||||
})
|
||||
assert server.protocol_version == (explicit if explicit is not None else revision)
|
||||
|
|
|
|||
|
|
@ -542,3 +542,126 @@ async def test_local_tool_json_array_is_converted_once_for_the_caller_revision(c
|
|||
|
||||
assert [block.text for block in result.content] == [body]
|
||||
assert result.structured_content == (["a", "b"] if compat == "modern" else None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_preserves_caller_scope_and_proxy_restrictions():
|
||||
from mcp.types import DiscoverRequest, ListToolsResult, Tool
|
||||
|
||||
listed = AsyncMock(return_value=ListToolsResult(tools=[Tool(name="allowed", input_schema={"type": "object"})]))
|
||||
context = prepare_context(UserAPIKeyAuth(user_id="scoped"), mcp_servers=["only-this"], mcp_proxy_mode=True, protocol_version="2025-06-18")
|
||||
with patch("litellm.proxy._experimental.mcp_server.operations._execute_handle_list_tools", listed):
|
||||
result = await GatewayOperations().execute(DiscoverRequest(), context)
|
||||
assert result.capabilities.tools is not None
|
||||
assert result.capabilities.resources is None
|
||||
assert result.capabilities.prompts is None
|
||||
assert listed.await_args.args[0] is context
|
||||
assert listed.await_args.args[0].user_api_key_auth.user_id == "scoped"
|
||||
assert listed.await_args.args[0].mcp_servers == ("only-this",)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_denial_cannot_advertise_tools():
|
||||
from mcp.types import DiscoverRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
denied = AsyncMock(side_effect=HTTPException(status_code=403, detail="Forbidden"))
|
||||
with patch("litellm.proxy._experimental.mcp_server.operations._execute_handle_list_tools", denied):
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="denied")))
|
||||
assert error.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("available", ["none", "resources", "templates", "prompts"])
|
||||
async def test_discovery_lists_each_capability_with_the_same_caller(available):
|
||||
from mcp.types import (
|
||||
DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult,
|
||||
ListResourceTemplatesResult, Prompt, Resource, ResourceTemplate,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
|
||||
context = prepare_context(UserAPIKeyAuth(user_id="scoped"), mcp_servers=["authorized"])
|
||||
tools = AsyncMock(return_value=ListToolsResult(tools=[]))
|
||||
prompts = AsyncMock(return_value=ListPromptsResult(prompts=[Prompt(name="allowed")] if available == "prompts" else []))
|
||||
resources = AsyncMock(return_value=ListResourcesResult(resources=[Resource(name="allowed", uri="test://allowed")] if available == "resources" else []))
|
||||
templates = AsyncMock(return_value=ListResourceTemplatesResult(resource_templates=[ResourceTemplate(name="allowed", uri_template="test://{id}")] if available == "templates" else []))
|
||||
with (
|
||||
patch.object(operations, "_execute_handle_list_tools", tools),
|
||||
patch.object(operations, "_execute_list_prompts", prompts),
|
||||
patch.object(operations, "_execute_list_resources", resources),
|
||||
patch.object(operations, "_execute_list_resource_templates", templates),
|
||||
):
|
||||
result = await GatewayOperations().execute(DiscoverRequest(), context)
|
||||
assert result.capabilities.tools is None
|
||||
assert (result.capabilities.prompts is not None) == (available == "prompts")
|
||||
assert (result.capabilities.resources is not None) == (available in {"resources", "templates"})
|
||||
for listing in (tools, prompts, resources, templates):
|
||||
assert listing.await_args.args[0] is context
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("outcome", ["success", "failure", "cancel"])
|
||||
async def test_discovery_concurrent_listings_drain_on_failure_and_cancellation(outcome):
|
||||
from mcp.types import DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
|
||||
ready = [asyncio.Event() for _ in range(4)]
|
||||
closed = [asyncio.Event() for _ in range(4)]
|
||||
release = asyncio.Event()
|
||||
responses = (ListToolsResult(tools=[]), ListPromptsResult(prompts=[]), ListResourcesResult(resources=[]), ListResourceTemplatesResult(resource_templates=[]))
|
||||
|
||||
def listing(index):
|
||||
async def run(*args, **kwargs):
|
||||
ready[index].set()
|
||||
try:
|
||||
await release.wait()
|
||||
if index == 0 and outcome == "failure":
|
||||
raise ValueError("discovery failed")
|
||||
if outcome != "success":
|
||||
await asyncio.Event().wait()
|
||||
return responses[index]
|
||||
finally:
|
||||
closed[index].set()
|
||||
return run
|
||||
|
||||
with (
|
||||
patch.object(operations, "_execute_handle_list_tools", side_effect=listing(0)) as tools,
|
||||
patch.object(operations, "_execute_list_prompts", side_effect=listing(1)),
|
||||
patch.object(operations, "_execute_list_resources", side_effect=listing(2)),
|
||||
patch.object(operations, "_execute_list_resource_templates", side_effect=listing(3)),
|
||||
):
|
||||
task = asyncio.create_task(GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="scoped"))))
|
||||
try:
|
||||
await asyncio.wait_for(asyncio.gather(*(event.wait() for event in ready)), 1)
|
||||
if outcome == "cancel":
|
||||
task.cancel()
|
||||
else:
|
||||
release.set()
|
||||
if outcome == "success":
|
||||
result = await asyncio.wait_for(task, 1)
|
||||
assert result.capabilities.model_dump(exclude_none=True) == {}
|
||||
else:
|
||||
with pytest.raises(asyncio.CancelledError if outcome == "cancel" else ValueError):
|
||||
await asyncio.wait_for(task, 1)
|
||||
assert all(event.is_set() for event in closed)
|
||||
assert tools.call_args.kwargs["log_list_tools_to_spendlogs"] is False
|
||||
finally:
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("log_enabled", [False, True])
|
||||
async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled):
|
||||
from mcp.types import PaginatedRequestParams
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
|
||||
listing = AsyncMock(return_value=operations.AggregateToolListing(tools=[], outcomes={}))
|
||||
with patch.object(operations, "_list_mcp_tools", listing):
|
||||
result = await operations._execute_handle_list_tools(
|
||||
prepare_context(UserAPIKeyAuth(user_id="caller")), PaginatedRequestParams(),
|
||||
log_list_tools_to_spendlogs=log_enabled,
|
||||
)
|
||||
assert result.tools == []
|
||||
assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport, MCPUpstreamProtocol
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
_OK_TOOL_RESULT: Final = CallToolResult(content=[TextContent(type="text", text='{"result": "ok"}')], is_error=False)
|
||||
|
|
@ -1476,7 +1476,7 @@ class TestListToolsRestAPI:
|
|||
monkeypatch,
|
||||
):
|
||||
"""The REST tools/list path should include tools beyond the upstream first page."""
|
||||
from mcp.types import ListToolsResult, PaginatedRequestParams
|
||||
from mcp.types import Implementation, InitializeResult, ListToolsResult, PaginatedRequestParams, ServerCapabilities
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
import litellm.experimental_mcp_client.client as mcp_client_module
|
||||
|
|
@ -1512,7 +1512,11 @@ class TestListToolsRestAPI:
|
|||
|
||||
mock_session_ctx = AsyncMock()
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock(return_value=None)
|
||||
mock_session_instance.initialize = AsyncMock(return_value=InitializeResult(
|
||||
protocol_version="2025-11-25",
|
||||
capabilities=ServerCapabilities(),
|
||||
server_info=Implementation(name="stub", version="1"),
|
||||
))
|
||||
mock_session_instance.list_tools.side_effect = [
|
||||
ListToolsResult(
|
||||
tools=[
|
||||
|
|
@ -4628,3 +4632,74 @@ class TestClientAllowlistOnRestRoutes:
|
|||
assert denied.value.detail["error"] == "Forbidden"
|
||||
assert "'claude-code'" in denied.value.detail["details"]
|
||||
acting.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("revision", ("auto", "2024-11-05", "2025-06-18"))
|
||||
async def test_preview_client_honors_protocol_metadata(revision: MCPUpstreamProtocol) -> None:
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
|
||||
payload: Final = NewMCPServerRequest(
|
||||
server_name="preview", url="http://127.0.0.1:9/mcp", transport="http",
|
||||
auth_type=MCPAuth.none, mcp_info={"protocol_version": revision},
|
||||
)
|
||||
|
||||
async def inspect_client(client: MCPClient) -> dict[str, str]:
|
||||
return {"protocol_version": client.protocol_version}
|
||||
|
||||
result: Final = await rest_endpoints._execute_with_mcp_client(payload, inspect_client)
|
||||
assert result == {"protocol_version": revision}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type", (MCPAuth.none, MCPAuth.bearer_token, MCPAuth.oauth2))
|
||||
@pytest.mark.parametrize(
|
||||
("metadata", "expected"),
|
||||
(
|
||||
(None, "2025-11-25"),
|
||||
({}, "2025-11-25"),
|
||||
({"description": "edited"}, "2025-11-25"),
|
||||
({"protocol_version": "auto"}, "auto"),
|
||||
({"protocol_version": "2024-11-05"}, "2024-11-05"),
|
||||
({"protocol_version": "2025-06-18"}, "2025-06-18"),
|
||||
),
|
||||
)
|
||||
async def test_saved_preview_protocol_omission_and_explicit_edits(
|
||||
monkeypatch: pytest.MonkeyPatch, auth_type: MCPAuth,
|
||||
metadata: dict[str, str] | None, expected: MCPUpstreamProtocol,
|
||||
) -> None:
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy.management_endpoints import mcp_management_endpoints
|
||||
|
||||
saved: Final = MCPServer(
|
||||
server_id="saved-preview", name="preview", url="https://example.com/mcp",
|
||||
transport="http", auth_type=auth_type, protocol_version="2025-11-25",
|
||||
authentication_token="stored-token",
|
||||
authorization_url="https://example.com/authorize", token_url="https://example.com/token",
|
||||
)
|
||||
manager: Final = MCPServerManager()
|
||||
manager.registry = {saved.server_id: saved}
|
||||
monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(mcp_management_endpoints, "global_mcp_server_manager", manager)
|
||||
payload: Final = NewMCPServerRequest(
|
||||
server_id=saved.server_id, server_name=saved.name, url=saved.url, transport="http",
|
||||
auth_type=auth_type, mcp_info=metadata,
|
||||
authorization_url=saved.authorization_url, token_url=saved.token_url,
|
||||
)
|
||||
staged: Final = rest_endpoints._stage_server_test(
|
||||
payload, Headers({"x-litellm-api-key": "sk-admin", "authorization": "Bearer preview-token"})
|
||||
)
|
||||
|
||||
async def inspect_client(client: MCPClient) -> dict[str, str]:
|
||||
return {"protocol_version": client.protocol_version}
|
||||
|
||||
result: Final = await rest_endpoints._execute_with_mcp_client(
|
||||
staged.request, inspect_client,
|
||||
mcp_auth_header=staged.mcp_auth_header, oauth2_headers=staged.oauth2_headers,
|
||||
)
|
||||
assert result == {"protocol_version": expected}
|
||||
assert saved.protocol_version == "2025-11-25"
|
||||
assert payload.mcp_info == metadata
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue