mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge branch 'main' into litellm_lit8140_azure_flux2_megapixel_billing
This commit is contained in:
commit
9a7f0bafc8
968 changed files with 14815 additions and 11439 deletions
|
|
@ -14,12 +14,12 @@ while IFS= read -r file || [ -n "$file" ]; do
|
|||
[ -n "$file" ] || continue
|
||||
case "$file" in
|
||||
*.md | *.mdx) : ;;
|
||||
pyproject.toml | */pyproject.toml | uv.lock | uv.toml | .python-version | rust-toolchain.toml | litellm-rust/* | litellm/__init__.py | litellm/proxy/proxy_server.py | litellm/*mcp* | tests/*mcp* | litellm/integrations/arize/* | tests/base_sdk_tests/* | scripts/check_mcp_sdk_install.py | .github/workflows/test-mcp-dependency-resolution.yml | .github/actions/detect-changes/* | .github/actions/setup-uv-with-retries/* | .github/actions/cache-cargo-build/* | .github/scripts/detect_changes.sh | .github/scripts/uv_sync_with_retries.sh | .circleci/scripts/classify_changes.sh | tests/test_litellm/test_circleci_path_filter.py | tests/test_litellm/test_detect_changes.py)
|
||||
pyproject.toml | */pyproject.toml | uv.lock | uv.toml | .python-version | rust-toolchain.toml | litellm-rust/* | litellm/__init__.py | litellm/proxy/proxy_server.py | litellm/*mcp* | tests/*mcp* | litellm/integrations/arize/* | tests/base_sdk_tests/* | scripts/check_mcp_sdk_install.py | .github/workflows/test-mcp-dependency-resolution.yml | .github/actions/detect-changes/* | .github/actions/setup-uv-with-retries/* | .github/actions/cache-cargo-build/* | .github/scripts/detect_changes.sh | .github/scripts/uv_sync_with_retries.sh | .circleci/scripts/classify_changes.sh | tests/unit/test_circleci_path_filter.py | tests/unit/test_detect_changes.py)
|
||||
has_mcp_dependencies=true ;;
|
||||
esac
|
||||
case "$file" in
|
||||
tests/e2e/*/*.py) : ;;
|
||||
tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/test_litellm/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock)
|
||||
tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/unit/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock)
|
||||
has_provider_harness=true ;;
|
||||
esac
|
||||
case "$file" in
|
||||
|
|
|
|||
|
|
@ -7,7 +7,11 @@ legacy_flags=(
|
|||
caching-local
|
||||
enterprise-package
|
||||
enterprise-routing
|
||||
integrations
|
||||
llm-other-providers
|
||||
llm-vertex-ai
|
||||
mcp-integration
|
||||
misc
|
||||
proxy-db-auth-checks
|
||||
proxy-db-budgets
|
||||
proxy-db-custom-logging
|
||||
|
|
@ -22,6 +26,7 @@ legacy_flags=(
|
|||
proxy-db-proxy-utils
|
||||
proxy-extras
|
||||
proxy-infra
|
||||
responses-caching-types
|
||||
)
|
||||
|
||||
legacy_paths() {
|
||||
|
|
@ -36,6 +41,7 @@ legacy_paths() {
|
|||
echo tests/unit/enterprise/proxy/test_audit_logging_endpoints.py
|
||||
echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;;
|
||||
enterprise-routing)
|
||||
echo tests/unit/google_genai
|
||||
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
|
||||
|
|
@ -47,10 +53,33 @@ 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)
|
||||
echo tests/unit/experimental_mcp_client
|
||||
echo tests/unit/proxy/_experimental/mcp_server
|
||||
echo tests/unit/responses/mcp
|
||||
echo tests/mcp_tests/test_proxy_mcp_e2e.py ;;
|
||||
misc)
|
||||
find tests/unit -maxdepth 1 -name 'test_*.py'
|
||||
echo tests/unit/test_router
|
||||
echo tests/unit/a2a_protocol
|
||||
echo tests/unit/batches
|
||||
echo tests/unit/chat_completions
|
||||
echo tests/unit/completion_extras
|
||||
echo tests/unit/containers
|
||||
echo tests/unit/embeddings
|
||||
echo tests/unit/endpoints
|
||||
echo tests/unit/files
|
||||
echo tests/unit/images
|
||||
echo tests/unit/interactions
|
||||
echo tests/unit/messages
|
||||
echo tests/unit/rag
|
||||
echo tests/unit/rerank_api
|
||||
echo tests/unit/secret_managers
|
||||
echo tests/unit/vector_stores
|
||||
echo tests/unit/videos ;;
|
||||
proxy-db-auth-checks)
|
||||
echo tests/unit/proxy/auth/test_auth_checks.py
|
||||
echo tests/unit/proxy/auth/test_user_api_key_auth.py
|
||||
|
|
@ -113,6 +142,7 @@ 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 ;;
|
||||
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
|
||||
esac
|
||||
}
|
||||
|
|
|
|||
|
|
@ -341,6 +341,7 @@ workflows:
|
|||
flag:
|
||||
- enterprise-package
|
||||
- proxy-infra
|
||||
- responses-caching-types
|
||||
- proxy-db-auth-checks
|
||||
- proxy-db-jwt-and-keys
|
||||
- proxy-db-proxy-server-core
|
||||
|
|
@ -353,6 +354,35 @@ workflows:
|
|||
- proxy-db-endpoints-and-responses
|
||||
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-llm-vertex-ai
|
||||
flag: llm-vertex-ai
|
||||
shards: 2
|
||||
workers: 1
|
||||
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-llm-other-providers
|
||||
flag: llm-other-providers
|
||||
shards: 3
|
||||
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-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
|
||||
shards: 2
|
||||
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-proxy-db-proxy-utils
|
||||
flag: proxy-db-proxy-utils
|
||||
|
|
|
|||
10
.github/merge-smoke-tests.json
vendored
10
.github/merge-smoke-tests.json
vendored
|
|
@ -1,12 +1,12 @@
|
|||
{
|
||||
"cases": {
|
||||
"CHAT-JSON": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
|
||||
"CHAT-TEXT-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
|
||||
"CHAT-TOOL-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
|
||||
"CHAT-JSON": "tests/unit/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
|
||||
"CHAT-TEXT-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
|
||||
"CHAT-TOOL-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
|
||||
"MODEL-ALLOW": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_allows_listed_model_for_key",
|
||||
"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/test_litellm/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
|
||||
"COST-ZERO": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
|
||||
"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",
|
||||
|
|
|
|||
18
.github/workflows/_test-unit-base.yml
vendored
18
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -13,12 +13,13 @@ on:
|
|||
have its path existence-checked like any other token.
|
||||
required: true
|
||||
type: string
|
||||
fork-flag:
|
||||
unit-flag:
|
||||
description: >-
|
||||
Codecov flag of the `.circleci/tests.yml` job that now owns part of
|
||||
this shard. CircleCI does not run on pull requests from forks, so on
|
||||
those events this shard also runs the files
|
||||
`.circleci/scripts/unit_selection.sh` lists for the flag.
|
||||
this shard. The shard also runs the files
|
||||
`.circleci/scripts/unit_selection.sh` lists for the flag, on every
|
||||
event, because the CircleCI pipeline is manual-only while the tests
|
||||
migrate.
|
||||
required: false
|
||||
type: string
|
||||
default: ""
|
||||
|
|
@ -175,8 +176,7 @@ jobs:
|
|||
timeout-minutes: ${{ inputs.timeout-minutes }}
|
||||
env:
|
||||
TEST_PATH: ${{ inputs.test-path }}
|
||||
FORK_FLAG: ${{ inputs.fork-flag }}
|
||||
IS_FORK: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository }}
|
||||
UNIT_FLAG: ${{ inputs.unit-flag }}
|
||||
MAX_FAILURES: ${{ inputs.max-failures }}
|
||||
WORKERS: ${{ inputs.workers }}
|
||||
RERUNS: ${{ inputs.reruns }}
|
||||
|
|
@ -186,11 +186,11 @@ jobs:
|
|||
run: |
|
||||
echo "has-coverage=false" >> "$GITHUB_OUTPUT"
|
||||
selection="${TEST_PATH}"
|
||||
if [ "${IS_FORK}" = "true" ] && [ -n "${FORK_FLAG}" ]; then
|
||||
selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${FORK_FLAG}" | tr '\n' ' ')"
|
||||
if [ -n "${UNIT_FLAG}" ]; then
|
||||
selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${UNIT_FLAG}" | tr '\n' ' ')"
|
||||
fi
|
||||
if [ -z "${selection// /}" ]; then
|
||||
echo "shard selection is empty on this event (CircleCI flag ${FORK_FLAG:-none} owns it); nothing to run"
|
||||
echo "shard selection is empty; nothing to run"
|
||||
exit 0
|
||||
fi
|
||||
pytest_args=()
|
||||
|
|
|
|||
4
.github/workflows/test-redis-compat.yml
vendored
4
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -10,7 +10,7 @@ on:
|
|||
- "litellm/_redis_credential_provider.py"
|
||||
- "litellm/caching/redis_cache.py"
|
||||
- "litellm/caching/evicted_client_closer.py"
|
||||
- "tests/test_litellm/test_redis.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"
|
||||
|
|
@ -84,7 +84,7 @@ jobs:
|
|||
run: |
|
||||
redis-server --version
|
||||
uv run --no-sync pytest \
|
||||
tests/test_litellm/test_redis.py \
|
||||
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 \
|
||||
|
|
|
|||
33
.github/workflows/test-unit-proxy-db.yml
vendored
33
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -22,9 +22,10 @@ concurrency:
|
|||
#
|
||||
# `.circleci/tests.yml` runs each group's files on same-repo events under the
|
||||
# `proxy-db-<group>` Codecov flag; `.circleci/scripts/unit_selection.sh` holds
|
||||
# the file lists. CircleCI does not build pull requests from forks, so `fork-flag`
|
||||
# makes the shard run that list there. `test-path` keeps the files that still
|
||||
# reach real providers and never left tests/proxy_unit_tests.
|
||||
# the file lists. That pipeline is manual-only while the tests migrate, so
|
||||
# `unit-flag` makes the shard run that list on every event. `test-path` keeps
|
||||
# the files that still reach real providers and never left
|
||||
# tests/proxy_unit_tests.
|
||||
#
|
||||
# Design targets:
|
||||
# * Every shard runs in <= 7 minutes of wall-clock on the default runner.
|
||||
|
|
@ -78,7 +79,7 @@ jobs:
|
|||
# Must run serially — event-loop conflict with the logging worker.
|
||||
- test-group: key-generation
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-key-generation
|
||||
unit-flag: proxy-db-key-generation
|
||||
workers: 0
|
||||
dist: loadscope
|
||||
timeout: 20
|
||||
|
|
@ -86,13 +87,13 @@ jobs:
|
|||
# ---- auth: split into 2 shards ----
|
||||
- test-group: auth-checks
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-auth-checks
|
||||
unit-flag: proxy-db-auth-checks
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: jwt-and-keys
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-jwt-and-keys
|
||||
unit-flag: proxy-db-jwt-and-keys
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
|
@ -100,7 +101,7 @@ jobs:
|
|||
# ---- test_proxy_utils.py, single shard, worksteal distribution ----
|
||||
- test-group: proxy-utils
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-proxy-utils
|
||||
unit-flag: proxy-db-proxy-utils
|
||||
workers: 4
|
||||
dist: worksteal
|
||||
timeout: 15
|
||||
|
|
@ -108,13 +109,13 @@ jobs:
|
|||
# ---- proxy server: split into 2 shards ----
|
||||
- test-group: proxy-server-core
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py"
|
||||
fork-flag: proxy-db-proxy-server-core
|
||||
unit-flag: proxy-db-proxy-server-core
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: proxy-runtime
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-proxy-runtime
|
||||
unit-flag: proxy-db-proxy-runtime
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
|
@ -122,20 +123,20 @@ jobs:
|
|||
# ---- logging: split into 2 shards ----
|
||||
- test-group: custom-logging
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_custom_logger.py"
|
||||
fork-flag: proxy-db-custom-logging
|
||||
unit-flag: proxy-db-custom-logging
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: logging-misc
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-logging-misc
|
||||
unit-flag: proxy-db-logging-misc
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
- test-group: db-and-spend
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-db-and-spend
|
||||
unit-flag: proxy-db-db-and-spend
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
|
@ -143,27 +144,27 @@ jobs:
|
|||
# ---- guardrails + budget + hooks: split into 2 ----
|
||||
- test-group: guardrails-hooks
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-guardrails-hooks
|
||||
unit-flag: proxy-db-guardrails-hooks
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: budgets
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-budgets
|
||||
unit-flag: proxy-db-budgets
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
- test-group: endpoints-and-responses
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_exception_mapping.py"
|
||||
fork-flag: proxy-db-endpoints-and-responses
|
||||
unit-flag: proxy-db-endpoints-and-responses
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: ${{ matrix.test-path }}
|
||||
fork-flag: ${{ matrix.fork-flag }}
|
||||
unit-flag: ${{ matrix.unit-flag }}
|
||||
workers: ${{ matrix.workers }}
|
||||
reruns: 2
|
||||
timeout-minutes: ${{ matrix.timeout }}
|
||||
|
|
|
|||
46
.github/workflows/test-unit.yml
vendored
46
.github/workflows/test-unit.yml
vendored
|
|
@ -36,9 +36,9 @@ concurrency:
|
|||
# Folding it in here is a follow-up, together with generalising that guard into
|
||||
# assert_ci_coverage.py.
|
||||
#
|
||||
# `fork-flag` names the `.circleci/tests.yml` job that now runs part of the
|
||||
# shard under the same Codecov flag. CircleCI does not build pull requests from
|
||||
# forks, so the shard still runs those files there and skips them elsewhere.
|
||||
# `unit-flag` names the `.circleci/tests.yml` job that now runs part of the
|
||||
# shard under the same Codecov flag. That pipeline is manual-only while the
|
||||
# tests migrate, so the shard also runs those files on every event.
|
||||
jobs:
|
||||
unit:
|
||||
name: ${{ matrix.shard }}
|
||||
|
|
@ -52,8 +52,8 @@ jobs:
|
|||
include:
|
||||
- shard: mcp-integration
|
||||
artifact-name: mcp-integration
|
||||
test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client"
|
||||
fork-flag: mcp-integration
|
||||
test-path: "tests/mcp_tests"
|
||||
unit-flag: mcp-integration
|
||||
workers: 2
|
||||
reruns: 0
|
||||
timeout-minutes: 20
|
||||
|
|
@ -70,10 +70,9 @@ jobs:
|
|||
- shard: enterprise-routing
|
||||
artifact-name: enterprise-routing
|
||||
test-path: >-
|
||||
tests/test_litellm/google_genai
|
||||
tests/test_litellm/router_utils
|
||||
tests/test_litellm/router_strategy
|
||||
fork-flag: enterprise-routing
|
||||
unit-flag: enterprise-routing
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -81,7 +80,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
|
||||
|
|
@ -90,6 +90,7 @@ jobs:
|
|||
- shard: Vertex AI
|
||||
artifact-name: llm-vertex-ai
|
||||
test-path: "tests/test_litellm/llms/vertex_ai"
|
||||
unit-flag: llm-vertex-ai
|
||||
workers: 1
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -98,6 +99,7 @@ jobs:
|
|||
- shard: All Other Providers
|
||||
artifact-name: llm-other-providers
|
||||
test-path: "tests/test_litellm/llms --ignore=tests/test_litellm/llms/vertex_ai"
|
||||
unit-flag: llm-other-providers
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -106,26 +108,12 @@ jobs:
|
|||
- shard: misc
|
||||
artifact-name: misc
|
||||
test-path: >-
|
||||
tests/test_litellm/batches
|
||||
tests/test_litellm/secret_managers
|
||||
tests/test_litellm/a2a_protocol
|
||||
tests/test_litellm/chat_completions
|
||||
tests/test_litellm/completion_extras
|
||||
tests/test_litellm/containers
|
||||
tests/test_litellm/endpoints
|
||||
tests/test_litellm/files
|
||||
tests/test_litellm/images
|
||||
tests/test_litellm/interactions
|
||||
tests/test_litellm/messages
|
||||
tests/test_litellm/embeddings
|
||||
tests/test_litellm/ocr
|
||||
tests/test_litellm/passthrough
|
||||
tests/test_litellm/rag
|
||||
tests/test_litellm/rerank_api
|
||||
tests/test_litellm/rust_bridge
|
||||
tests/test_litellm/vector_stores
|
||||
tests/test_litellm/videos
|
||||
tests/test_litellm/test_*.py
|
||||
unit-flag: misc
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -205,7 +193,7 @@ jobs:
|
|||
tests/test_litellm/proxy/types_utils
|
||||
tests/test_litellm/proxy/logging_endpoints
|
||||
tests/test_litellm/proxy/test_*.py
|
||||
fork-flag: proxy-infra
|
||||
unit-flag: proxy-infra
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -214,7 +202,7 @@ jobs:
|
|||
- shard: caching-local
|
||||
artifact-name: caching-local
|
||||
test-path: ""
|
||||
fork-flag: caching-local
|
||||
unit-flag: caching-local
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -223,7 +211,7 @@ jobs:
|
|||
- shard: proxy-extras
|
||||
artifact-name: proxy-extras
|
||||
test-path: ""
|
||||
fork-flag: proxy-extras
|
||||
unit-flag: proxy-extras
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -232,7 +220,7 @@ jobs:
|
|||
- shard: enterprise-package
|
||||
artifact-name: enterprise-package
|
||||
test-path: ""
|
||||
fork-flag: enterprise-package
|
||||
unit-flag: enterprise-package
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -243,7 +231,7 @@ jobs:
|
|||
test-path: >-
|
||||
tests/test_litellm/responses
|
||||
tests/test_litellm/caching
|
||||
tests/test_litellm/types
|
||||
unit-flag: responses-caching-types
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -251,7 +239,7 @@ jobs:
|
|||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: ${{ matrix.test-path }}
|
||||
fork-flag: ${{ matrix.fork-flag || '' }}
|
||||
unit-flag: ${{ matrix.unit-flag || '' }}
|
||||
workers: ${{ matrix.workers }}
|
||||
reruns: ${{ matrix.reruns }}
|
||||
timeout-minutes: ${{ matrix.timeout-minutes }}
|
||||
|
|
|
|||
8
Makefile
8
Makefile
|
|
@ -314,7 +314,7 @@ test-unit: install-test-deps
|
|||
|
||||
# Matrix test targets (matching CI workflow groups)
|
||||
test-unit-llms: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/llms --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/llms --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-proxy-guardrails: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/proxy/guardrails tests/test_litellm/proxy/management_endpoints tests/test_litellm/proxy/management_helpers --tb=short -vv -n 4 --durations=20
|
||||
|
|
@ -326,16 +326,16 @@ 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
|
||||
|
||||
test-unit-other: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/unit/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/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/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-root: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
# Proxy unit tests (tests/unit/proxy split alphabetically)
|
||||
test-proxy-unit-a: install-test-deps
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
168
litellm-rust/Cargo.lock
generated
168
litellm-rust/Cargo.lock
generated
|
|
@ -73,6 +73,15 @@ version = "1.0.104"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470"
|
||||
|
||||
[[package]]
|
||||
name = "arbitrary"
|
||||
version = "1.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1"
|
||||
dependencies = [
|
||||
"derive_arbitrary",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "arc-swap"
|
||||
version = "1.9.2"
|
||||
|
|
@ -1470,6 +1479,17 @@ dependencies = [
|
|||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "derive_arbitrary"
|
||||
version = "1.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1e567bd82dcff979e4b03460c307b3cdc9e96fde3d73bed1496d2bc75d9dd62a"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "derive_builder"
|
||||
version = "0.20.2"
|
||||
|
|
@ -1643,6 +1663,16 @@ version = "2.5.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223"
|
||||
|
||||
[[package]]
|
||||
name = "filetime"
|
||||
version = "0.2.29"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5c287a33c7f0a620c38e641e7f60827713987b3c0f26e8ddc9462cc69cf75759"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "find-msvc-tools"
|
||||
version = "0.1.9"
|
||||
|
|
@ -3136,10 +3166,11 @@ dependencies = [
|
|||
"aws-smithy-types",
|
||||
"bytes",
|
||||
"futures-util",
|
||||
"proptest",
|
||||
"rstest",
|
||||
"sse-stream",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tokio-util",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -3453,6 +3484,27 @@ dependencies = [
|
|||
"veil",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-testkit"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"flate2",
|
||||
"futures-util",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
"semver",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2 0.10.9",
|
||||
"tar",
|
||||
"target-lexicon",
|
||||
"tempfile",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"toml",
|
||||
"zip",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-token-counter"
|
||||
version = "0.1.0"
|
||||
|
|
@ -5206,6 +5258,15 @@ dependencies = [
|
|||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "serde_spanned"
|
||||
version = "1.1.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6662b5879511e06e8999a8a235d848113e942c9124f211511b16466ee2995f26"
|
||||
dependencies = [
|
||||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "serde_urlencoded"
|
||||
version = "0.7.1"
|
||||
|
|
@ -5408,19 +5469,6 @@ dependencies = [
|
|||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sse-stream"
|
||||
version = "0.2.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c25ac7aff0abd1dbc474536e40416e1102c7dd9bfba0b9861c6d357f835dcfb4"
|
||||
dependencies = [
|
||||
"bytes",
|
||||
"futures-util",
|
||||
"http-body 1.1.0",
|
||||
"http-body-util",
|
||||
"pin-project-lite",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "stable_deref_trait"
|
||||
version = "1.2.1"
|
||||
|
|
@ -5537,6 +5585,17 @@ version = "0.2.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417"
|
||||
|
||||
[[package]]
|
||||
name = "tar"
|
||||
version = "0.4.46"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3f6221d9a6003c78398e3b239969f352578258df48c8eb051caadae0015bc840"
|
||||
dependencies = [
|
||||
"filetime",
|
||||
"libc",
|
||||
"xattr",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "target-lexicon"
|
||||
version = "0.13.5"
|
||||
|
|
@ -5806,6 +5865,30 @@ dependencies = [
|
|||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml"
|
||||
version = "0.9.12+spec-1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "cf92845e79fc2e2def6a5d828f0801e29a2f8acc037becc5ab08595c7d5e9863"
|
||||
dependencies = [
|
||||
"indexmap 2.14.0",
|
||||
"serde_core",
|
||||
"serde_spanned",
|
||||
"toml_datetime 0.7.5+spec-1.1.0",
|
||||
"toml_parser",
|
||||
"toml_writer",
|
||||
"winnow 0.7.15",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml_datetime"
|
||||
version = "0.7.5+spec-1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "92e1cfed4a3038bc5a127e35a2d360f145e1f4b971b551a2ba5fd7aedf7e1347"
|
||||
dependencies = [
|
||||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml_datetime"
|
||||
version = "1.1.1+spec-1.1.0"
|
||||
|
|
@ -5822,9 +5905,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "6975367e4d2ef766d86af01ffad14b622fecc8d4357a998fbc4deb6e9bacaf9b"
|
||||
dependencies = [
|
||||
"indexmap 2.14.0",
|
||||
"toml_datetime",
|
||||
"toml_datetime 1.1.1+spec-1.1.0",
|
||||
"toml_parser",
|
||||
"winnow",
|
||||
"winnow 1.0.4",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -5833,9 +5916,15 @@ version = "1.1.3+spec-1.1.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1d38ac1cf9b95face32296c0a3ede1fdc270627c9d9c02a7274dd6d960dc4d56"
|
||||
dependencies = [
|
||||
"winnow",
|
||||
"winnow 1.0.4",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml_writer"
|
||||
version = "1.1.2+spec-1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7d56353a2a665ad0f41a421187180aab746c8c325620617ad883a99a1cbe66d2"
|
||||
|
||||
[[package]]
|
||||
name = "tonic"
|
||||
version = "0.14.6"
|
||||
|
|
@ -6597,6 +6686,12 @@ version = "0.52.6"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec"
|
||||
|
||||
[[package]]
|
||||
name = "winnow"
|
||||
version = "0.7.15"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945"
|
||||
|
||||
[[package]]
|
||||
name = "winnow"
|
||||
version = "1.0.4"
|
||||
|
|
@ -6659,6 +6754,16 @@ dependencies = [
|
|||
"time",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "xattr"
|
||||
version = "1.6.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"rustix",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "xmlparser"
|
||||
version = "0.13.6"
|
||||
|
|
@ -6784,6 +6889,23 @@ dependencies = [
|
|||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "zip"
|
||||
version = "2.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fabe6324e908f85a1c52063ce7aa26b68dcb7eb6dbc83a2d148403c9bc3eba50"
|
||||
dependencies = [
|
||||
"arbitrary",
|
||||
"crc32fast",
|
||||
"crossbeam-utils",
|
||||
"displaydoc",
|
||||
"flate2",
|
||||
"indexmap 2.14.0",
|
||||
"memchr",
|
||||
"thiserror 2.0.19",
|
||||
"zopfli",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "zlib-rs"
|
||||
version = "0.6.7"
|
||||
|
|
@ -6795,3 +6917,15 @@ name = "zmij"
|
|||
version = "1.0.23"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b"
|
||||
|
||||
[[package]]
|
||||
name = "zopfli"
|
||||
version = "0.8.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f05cd8797d63865425ff89b5c4a48804f35ba0ce8d125800027ad6017d2b5249"
|
||||
dependencies = [
|
||||
"bumpalo",
|
||||
"crc32fast",
|
||||
"log",
|
||||
"simd-adler32",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -81,6 +81,12 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"]
|
|||
tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] }
|
||||
futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] }
|
||||
base64 = "0.22"
|
||||
flate2 = "1"
|
||||
semver = "1"
|
||||
tar = "0.4"
|
||||
target-lexicon = "0.13.5"
|
||||
tempfile = "3"
|
||||
zip = { version = "2", default-features = false, features = ["deflate"] }
|
||||
moka = { version = "0.12.16", features = ["future"] }
|
||||
strum = { version = "0.28.0", features = ["derive"] }
|
||||
url = "2.5.8"
|
||||
|
|
|
|||
|
|
@ -398,7 +398,7 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i
|
|||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
capabilities: capabilities.clone(),
|
||||
capabilities,
|
||||
drop_params,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
|
|
|
|||
|
|
@ -8,16 +8,17 @@ repository.workspace = true
|
|||
[features]
|
||||
default = ["aws", "sse"]
|
||||
aws = ["dep:aws-smithy-eventstream", "dep:aws-smithy-types"]
|
||||
sse = ["dep:sse-stream"]
|
||||
sse = []
|
||||
|
||||
[dependencies]
|
||||
aws-smithy-eventstream = { version = "=0.61.4", optional = true }
|
||||
aws-smithy-types = { version = "1.6.1", optional = true }
|
||||
bytes = "1"
|
||||
futures-util.workspace = true
|
||||
sse-stream = { version = "=0.2.6", optional = true }
|
||||
thiserror.workspace = true
|
||||
tokio-util = { version = "0.7", features = ["codec", "io"] }
|
||||
|
||||
[dev-dependencies]
|
||||
proptest.workspace = true
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,66 +1,47 @@
|
|||
use bytes::{Buf, Bytes, BytesMut};
|
||||
use futures_util::{Stream, StreamExt};
|
||||
use aws_smithy_eventstream::frame::{read_message_from, write_message_to};
|
||||
pub use aws_smithy_types::event_stream::{Header, HeaderValue, Message};
|
||||
use bytes::BytesMut;
|
||||
use tokio_util::codec::{Decoder, Encoder};
|
||||
|
||||
use aws_smithy_eventstream::frame::read_message_from;
|
||||
use aws_smithy_types::event_stream::Header;
|
||||
|
||||
use crate::{Error, Framer};
|
||||
use crate::EventStreamError;
|
||||
|
||||
const MIN_FRAME_BYTES: usize = 16;
|
||||
const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct AwsEventStreamFrame {
|
||||
pub headers: Vec<Header>,
|
||||
pub payload: Bytes,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
pub struct AwsEventStreamFramer;
|
||||
pub struct AwsEventStreamCodec;
|
||||
|
||||
impl Framer for AwsEventStreamFramer {
|
||||
type Frame = AwsEventStreamFrame;
|
||||
impl Decoder for AwsEventStreamCodec {
|
||||
type Item = Message;
|
||||
type Error = EventStreamError;
|
||||
|
||||
fn frame<S, B, E>(self, input: S) -> impl Stream<Item = Result<Self::Frame, Error>> + Send
|
||||
where
|
||||
S: Stream<Item = Result<B, E>> + Send,
|
||||
B: Buf + Send,
|
||||
E: std::error::Error + Send + Sync + 'static,
|
||||
{
|
||||
futures_util::stream::try_unfold(
|
||||
(Box::pin(input), BytesMut::new()),
|
||||
|(mut input, mut buffer)| async move {
|
||||
loop {
|
||||
if buffer.len() >= 4 {
|
||||
let length = (&buffer[..4]).get_u32() as usize;
|
||||
if !(16..=MAX_FRAME_BYTES).contains(&length) {
|
||||
return Err(Error::InvalidLength(length));
|
||||
}
|
||||
if buffer.len() >= length {
|
||||
let raw = buffer.split_to(length).freeze();
|
||||
let message = read_message_from(raw)?;
|
||||
let frame = AwsEventStreamFrame {
|
||||
headers: message.headers().to_vec(),
|
||||
payload: message.payload().clone(),
|
||||
};
|
||||
return Ok(Some((frame, (input, buffer))));
|
||||
}
|
||||
}
|
||||
match input.next().await {
|
||||
Some(Ok(mut chunk)) => {
|
||||
while chunk.has_remaining() {
|
||||
let bytes = chunk.chunk();
|
||||
buffer.extend_from_slice(bytes);
|
||||
let length = bytes.len();
|
||||
chunk.advance(length);
|
||||
}
|
||||
}
|
||||
Some(Err(error)) => return Err(Error::Body(Box::new(error))),
|
||||
None if buffer.is_empty() => return Ok(None),
|
||||
None => return Err(Error::Truncated),
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
.fuse()
|
||||
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Message>, EventStreamError> {
|
||||
let Some(prefix) = src.first_chunk::<4>() else {
|
||||
return Ok(None);
|
||||
};
|
||||
let length = u32::from_be_bytes(*prefix) as usize;
|
||||
if !(MIN_FRAME_BYTES..=MAX_FRAME_BYTES).contains(&length) {
|
||||
return Err(EventStreamError::InvalidLength(length));
|
||||
}
|
||||
if src.len() < length {
|
||||
return Ok(None);
|
||||
}
|
||||
Ok(Some(read_message_from(src.split_to(length).freeze())?))
|
||||
}
|
||||
|
||||
fn decode_eof(&mut self, src: &mut BytesMut) -> Result<Option<Message>, EventStreamError> {
|
||||
match self.decode(src)? {
|
||||
Some(message) => Ok(Some(message)),
|
||||
None if src.is_empty() => Ok(None),
|
||||
None => Err(EventStreamError::Truncated),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Encoder<Message> for AwsEventStreamCodec {
|
||||
type Error = EventStreamError;
|
||||
|
||||
fn encode(&mut self, message: Message, dst: &mut BytesMut) -> Result<(), EventStreamError> {
|
||||
Ok(write_message_to(&message, dst)?)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,17 +1,21 @@
|
|||
#[cfg(feature = "sse")]
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[cfg(feature = "sse")]
|
||||
#[error("SSE framing failed: {0}")]
|
||||
Sse(#[from] sse_stream::Error),
|
||||
#[cfg(feature = "aws")]
|
||||
#[error("AWS EventStream framing failed: {0}")]
|
||||
Aws(#[from] aws_smithy_eventstream::error::Error),
|
||||
pub enum SseError {
|
||||
#[error("body stream failed: {0}")]
|
||||
Body(#[source] Box<dyn std::error::Error + Send + Sync>),
|
||||
#[cfg(feature = "aws")]
|
||||
Body(#[from] std::io::Error),
|
||||
#[error("SSE field is not UTF-8: {0}")]
|
||||
InvalidUtf8(#[from] std::str::Utf8Error),
|
||||
}
|
||||
|
||||
#[cfg(feature = "aws")]
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum EventStreamError {
|
||||
#[error("body stream failed: {0}")]
|
||||
Body(#[from] std::io::Error),
|
||||
#[error("invalid AWS EventStream frame length: {0}")]
|
||||
InvalidLength(usize),
|
||||
#[cfg(feature = "aws")]
|
||||
#[error("truncated AWS EventStream frame")]
|
||||
Truncated,
|
||||
#[error("malformed AWS EventStream frame: {0}")]
|
||||
Malformed(#[from] aws_smithy_eventstream::error::Error),
|
||||
}
|
||||
|
|
|
|||
21
litellm-rust/crates/framer/src/framed.rs
Normal file
21
litellm-rust/crates/framer/src/framed.rs
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
use std::io;
|
||||
|
||||
use bytes::Buf;
|
||||
use futures_util::{Stream, StreamExt, TryStreamExt};
|
||||
use tokio_util::{
|
||||
codec::{Decoder, FramedRead},
|
||||
io::StreamReader,
|
||||
};
|
||||
|
||||
pub fn frames<S, B, E, D>(
|
||||
input: S,
|
||||
codec: D,
|
||||
) -> impl Stream<Item = Result<D::Item, D::Error>> + Send
|
||||
where
|
||||
S: Stream<Item = Result<B, E>> + Send,
|
||||
B: Buf + Send,
|
||||
E: std::error::Error + Send + Sync + 'static,
|
||||
D: Decoder + Send,
|
||||
{
|
||||
FramedRead::new(StreamReader::new(input.map_err(io::Error::other)), codec).fuse()
|
||||
}
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
mod error;
|
||||
mod framer;
|
||||
mod framed;
|
||||
|
||||
pub use error::*;
|
||||
pub use framer::*;
|
||||
pub use framed::frames;
|
||||
|
||||
#[cfg(feature = "aws")]
|
||||
pub mod aws_event_stream;
|
||||
|
|
|
|||
|
|
@ -1,43 +1,170 @@
|
|||
use futures_util::{Stream, StreamExt};
|
||||
use std::str;
|
||||
|
||||
use crate::{Error, Framer};
|
||||
use bytes::{Buf, BufMut, BytesMut};
|
||||
use tokio_util::codec::{Decoder, Encoder};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct SseFrame {
|
||||
use crate::SseError;
|
||||
|
||||
const BOM: &[u8] = b"\xEF\xBB\xBF";
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq)]
|
||||
pub struct SseEvent {
|
||||
pub event: Option<String>,
|
||||
pub data: Option<String>,
|
||||
pub data: String,
|
||||
pub id: Option<String>,
|
||||
pub retry: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
pub struct SseFramer;
|
||||
pub struct SseCodec {
|
||||
past_bom: bool,
|
||||
}
|
||||
|
||||
impl Framer for SseFramer {
|
||||
type Frame = SseFrame;
|
||||
impl Decoder for SseCodec {
|
||||
type Item = SseEvent;
|
||||
type Error = SseError;
|
||||
|
||||
fn frame<S, B, E>(self, input: S) -> impl Stream<Item = Result<SseFrame, Error>> + Send
|
||||
where
|
||||
S: Stream<Item = Result<B, E>> + Send,
|
||||
B: bytes::Buf + Send,
|
||||
E: std::error::Error + Send + Sync + 'static,
|
||||
{
|
||||
let frames = Box::pin(sse_stream::SseStream::from_bytes_stream(input));
|
||||
futures_util::stream::try_unfold(frames, |mut frames| async move {
|
||||
let Some(frame) = frames.next().await else {
|
||||
return Ok(None);
|
||||
};
|
||||
let frame = frame?;
|
||||
Ok(Some((
|
||||
SseFrame {
|
||||
event: frame.event,
|
||||
data: frame.data,
|
||||
id: frame.id,
|
||||
retry: frame.retry,
|
||||
},
|
||||
frames,
|
||||
)))
|
||||
})
|
||||
.fuse()
|
||||
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<SseEvent>, SseError> {
|
||||
if !self.skip_bom(src) {
|
||||
return Ok(None);
|
||||
}
|
||||
while let Some(end) = block_end(src) {
|
||||
let block = src.split_to(end);
|
||||
let pending = lines(&block)
|
||||
.map(|(line, _)| line)
|
||||
.take_while(|line| !line.is_empty())
|
||||
.try_fold(Pending::default(), Pending::apply)?;
|
||||
if let Some(event) = pending.dispatch() {
|
||||
return Ok(Some(event));
|
||||
}
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn decode_eof(&mut self, _pending: &mut BytesMut) -> Result<Option<SseEvent>, SseError> {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
|
||||
impl SseCodec {
|
||||
fn skip_bom(&mut self, src: &mut BytesMut) -> bool {
|
||||
if self.past_bom {
|
||||
return true;
|
||||
}
|
||||
if src.starts_with(BOM) {
|
||||
src.advance(BOM.len());
|
||||
} else if BOM.starts_with(src) {
|
||||
return false;
|
||||
}
|
||||
self.past_bom = true;
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
fn block_end(bytes: &[u8]) -> Option<usize> {
|
||||
lines(bytes)
|
||||
.find(|(line, _)| line.is_empty())
|
||||
.map(|(_, end)| end)
|
||||
}
|
||||
|
||||
fn lines(bytes: &[u8]) -> impl Iterator<Item = (&[u8], usize)> {
|
||||
let mut cursor: usize = 0;
|
||||
std::iter::from_fn(move || {
|
||||
let rest = &bytes[cursor..];
|
||||
let end = rest.iter().position(|byte| matches!(byte, b'\n' | b'\r'))?;
|
||||
cursor += end + terminator_len(&rest[end..]);
|
||||
Some((&rest[..end], cursor))
|
||||
})
|
||||
}
|
||||
|
||||
fn terminator_len(terminated: &[u8]) -> usize {
|
||||
match terminated {
|
||||
[b'\r', b'\n', ..] => 2,
|
||||
_ => 1,
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct Pending {
|
||||
event: Option<String>,
|
||||
data: Option<String>,
|
||||
id: Option<String>,
|
||||
retry: Option<u64>,
|
||||
}
|
||||
|
||||
impl Pending {
|
||||
fn apply(self, line: &[u8]) -> Result<Self, SseError> {
|
||||
let (name, value) = split_field(line);
|
||||
Ok(match name {
|
||||
b"event" => Self {
|
||||
event: Some(str::from_utf8(value)?.to_owned()),
|
||||
..self
|
||||
},
|
||||
b"data" => Self {
|
||||
data: Some(append_data(self.data, str::from_utf8(value)?)),
|
||||
..self
|
||||
},
|
||||
b"id" if !value.contains(&0) => Self {
|
||||
id: Some(str::from_utf8(value)?.to_owned()),
|
||||
..self
|
||||
},
|
||||
b"retry" => Self {
|
||||
retry: parse_retry(value).or(self.retry),
|
||||
..self
|
||||
},
|
||||
_ => self,
|
||||
})
|
||||
}
|
||||
|
||||
fn dispatch(self) -> Option<SseEvent> {
|
||||
Some(SseEvent {
|
||||
event: self.event,
|
||||
data: self.data?,
|
||||
id: self.id,
|
||||
retry: self.retry,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn split_field(line: &[u8]) -> (&[u8], &[u8]) {
|
||||
let Some(colon) = line.iter().position(|byte| *byte == b':') else {
|
||||
return (line, &[]);
|
||||
};
|
||||
let value = &line[colon + 1..];
|
||||
(&line[..colon], value.strip_prefix(b" ").unwrap_or(value))
|
||||
}
|
||||
|
||||
fn append_data(buffer: Option<String>, line: &str) -> String {
|
||||
match buffer {
|
||||
Some(existing) => format!("{existing}\n{line}"),
|
||||
None => line.to_owned(),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_retry(value: &[u8]) -> Option<u64> {
|
||||
if !value.iter().all(u8::is_ascii_digit) {
|
||||
return None;
|
||||
}
|
||||
str::from_utf8(value).ok()?.parse().ok()
|
||||
}
|
||||
|
||||
impl Encoder<SseEvent> for SseCodec {
|
||||
type Error = SseError;
|
||||
|
||||
fn encode(&mut self, event: SseEvent, dst: &mut BytesMut) -> Result<(), SseError> {
|
||||
if let Some(name) = event.event {
|
||||
dst.put_slice(format!("event: {name}\n").as_bytes());
|
||||
}
|
||||
for line in event.data.split('\n') {
|
||||
dst.put_slice(format!("data: {line}\n").as_bytes());
|
||||
}
|
||||
if let Some(id) = event.id {
|
||||
dst.put_slice(format!("id: {id}\n").as_bytes());
|
||||
}
|
||||
if let Some(retry) = event.retry {
|
||||
dst.put_slice(format!("retry: {retry}\n").as_bytes());
|
||||
}
|
||||
dst.put_u8(b'\n');
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,89 +4,174 @@ mod support;
|
|||
|
||||
use std::io;
|
||||
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_framing::aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer};
|
||||
use litellm_framing::{Error, Framer};
|
||||
use bytes::Bytes;
|
||||
use futures_util::{StreamExt, TryStreamExt, stream};
|
||||
use litellm_framing::{
|
||||
EventStreamError,
|
||||
aws_event_stream::{AwsEventStreamCodec, Header, HeaderValue, Message},
|
||||
frames,
|
||||
};
|
||||
use proptest::prelude::*;
|
||||
use rstest::{fixture, rstest};
|
||||
use support::{body_cause, cut_at, encode_all, every, input, runtime};
|
||||
|
||||
use support::encode;
|
||||
|
||||
async fn collect_aws(bytes: &[u8], chunk_size: usize) -> Result<Vec<AwsEventStreamFrame>, Error> {
|
||||
AwsEventStreamFramer
|
||||
.frame(futures_util::stream::iter(
|
||||
bytes.chunks(chunk_size).map(Ok::<_, io::Error>),
|
||||
))
|
||||
async fn collect(pieces: Vec<Bytes>) -> Result<Vec<Message>, EventStreamError> {
|
||||
frames(input(pieces), AwsEventStreamCodec)
|
||||
.try_collect()
|
||||
.await
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn two_frames() -> Vec<u8> {
|
||||
[encode(b"\xff\x00"), encode(b"second")].concat()
|
||||
fn message(payload: &[u8]) -> Message {
|
||||
Message::new(Bytes::copy_from_slice(payload))
|
||||
.add_header(Header::new(
|
||||
":event-type",
|
||||
HeaderValue::String("payload".into()),
|
||||
))
|
||||
.add_header(Header::new("sequence", HeaderValue::Int32(7)))
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn payload_frame() -> Vec<u8> {
|
||||
encode(b"payload")
|
||||
encode_all(AwsEventStreamCodec, [message(b"payload")])
|
||||
}
|
||||
|
||||
fn header_value() -> impl Strategy<Value = HeaderValue> {
|
||||
prop_oneof![
|
||||
"[a-z]{0,8}".prop_map(|text| HeaderValue::String(text.into())),
|
||||
any::<i32>().prop_map(HeaderValue::Int32),
|
||||
any::<bool>().prop_map(HeaderValue::Bool),
|
||||
proptest::collection::vec(any::<u8>(), 0..8)
|
||||
.prop_map(|bytes| HeaderValue::ByteArray(bytes.into())),
|
||||
]
|
||||
}
|
||||
|
||||
fn arbitrary_message() -> impl Strategy<Value = Message> {
|
||||
(
|
||||
proptest::collection::vec(("[a-z:-]{1,12}", header_value()), 0..3),
|
||||
proptest::collection::vec(any::<u8>(), 0..32),
|
||||
)
|
||||
.prop_map(|(headers, payload)| {
|
||||
headers.into_iter().fold(
|
||||
Message::new(Bytes::from(payload)),
|
||||
|message, (name, value)| message.add_header(Header::new(name, value)),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#[test]
|
||||
fn any_messages_survive_a_round_trip_through_any_cuts(
|
||||
messages in proptest::collection::vec(arbitrary_message(), 1..4),
|
||||
cuts in proptest::collection::vec(0_usize..512, 0..4),
|
||||
) {
|
||||
let wire = encode_all(AwsEventStreamCodec, messages.clone());
|
||||
let decoded = runtime().block_on(collect(cut_at(&wire, cuts))).unwrap();
|
||||
prop_assert_eq!(decoded, messages);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(1)]
|
||||
#[case(3)]
|
||||
#[case(12)]
|
||||
#[case(usize::MAX)]
|
||||
#[case::prelude_crc(8)]
|
||||
#[case::message_crc(usize::MAX)]
|
||||
#[tokio::test]
|
||||
async fn fragmented_and_coalesced_frames_preserve_typed_headers_and_binary_payloads(
|
||||
two_frames: Vec<u8>,
|
||||
#[case] chunk_size: usize,
|
||||
) {
|
||||
let chunk_size = chunk_size.min(two_frames.len());
|
||||
let frames = collect_aws(&two_frames, chunk_size).await.unwrap();
|
||||
assert_eq!(frames.len(), 2);
|
||||
assert_eq!(frames[0].payload, &b"\xff\x00"[..]);
|
||||
assert_eq!(frames[1].payload, "second");
|
||||
assert_eq!(
|
||||
frames[0].headers[0].value().as_string().unwrap().as_str(),
|
||||
"payload"
|
||||
);
|
||||
assert_eq!(frames[0].headers[1].value().as_int32(), Ok(7));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(8)]
|
||||
#[case(0)]
|
||||
#[tokio::test]
|
||||
async fn rejects_corrupt_crcs(payload_frame: Vec<u8>, #[case] index: usize) {
|
||||
let corrupt_index = if index == 0 {
|
||||
payload_frame.len() - 1
|
||||
} else {
|
||||
index
|
||||
};
|
||||
async fn a_corrupt_crc_is_malformed(payload_frame: Vec<u8>, #[case] index: usize) {
|
||||
let mut corrupt = payload_frame;
|
||||
corrupt[corrupt_index] ^= 1;
|
||||
assert!(matches!(collect_aws(&corrupt, 3).await, Err(Error::Aws(_))));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(0_u32)]
|
||||
#[case(15)]
|
||||
#[case(u32::MAX)]
|
||||
#[tokio::test]
|
||||
async fn rejects_invalid_lengths(#[case] length: u32) {
|
||||
let flipped = index.min(corrupt.len() - 1);
|
||||
corrupt[flipped] ^= 1;
|
||||
assert!(matches!(
|
||||
collect_aws(&length.to_be_bytes(), 1).await,
|
||||
Err(Error::InvalidLength(_))
|
||||
collect(every(&corrupt, 3)).await,
|
||||
Err(EventStreamError::Malformed(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(1)]
|
||||
#[case(3)]
|
||||
#[case(5)]
|
||||
#[case::zero(0)]
|
||||
#[case::below_minimum(15)]
|
||||
#[case::above_maximum(16 * 1024 * 1024 + 1)]
|
||||
#[case::u32_max(u32::MAX)]
|
||||
#[tokio::test]
|
||||
async fn rejects_truncation(payload_frame: Vec<u8>, #[case] end: usize) {
|
||||
async fn a_length_outside_the_frame_bounds_fails_before_buffering(#[case] length: u32) {
|
||||
assert!(matches!(
|
||||
collect_aws(&payload_frame[..end], 1).await,
|
||||
Err(Error::Truncated)
|
||||
collect(every(&length.to_be_bytes(), 1)).await,
|
||||
Err(EventStreamError::InvalidLength(seen)) if seen == length as usize
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::before_the_length(1)]
|
||||
#[case::inside_the_prelude(5)]
|
||||
#[case::one_byte_short(usize::MAX)]
|
||||
#[tokio::test]
|
||||
async fn eof_inside_a_frame_is_truncation(payload_frame: Vec<u8>, #[case] end: usize) {
|
||||
let end = end.min(payload_frame.len() - 1);
|
||||
assert!(matches!(
|
||||
collect(every(&payload_frame[..end], 1)).await,
|
||||
Err(EventStreamError::Truncated)
|
||||
));
|
||||
}
|
||||
|
||||
const FRAME_OVERHEAD_BYTES: usize = 16;
|
||||
const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024;
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_frame_at_exactly_the_maximum_length_decodes() {
|
||||
let largest = Message::new(vec![0xAB; MAX_FRAME_BYTES - FRAME_OVERHEAD_BYTES]);
|
||||
let wire = encode_all(AwsEventStreamCodec, [largest.clone()]);
|
||||
assert_eq!(wire.len(), MAX_FRAME_BYTES);
|
||||
assert_eq!(collect(every(&wire, 1 << 20)).await.unwrap(), vec![largest]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_frame_one_byte_over_the_maximum_length_is_rejected_by_its_prelude() {
|
||||
let oversized = Message::new(vec![0xAB; MAX_FRAME_BYTES - FRAME_OVERHEAD_BYTES + 1]);
|
||||
let wire = encode_all(AwsEventStreamCodec, [oversized]);
|
||||
assert!(matches!(
|
||||
collect(every(&wire[..4], 1)).await,
|
||||
Err(EventStreamError::InvalidLength(length)) if length == MAX_FRAME_BYTES + 1
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn an_empty_body_yields_nothing() {
|
||||
assert_eq!(collect(vec![]).await.unwrap(), vec![]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_complete_frame_precedes_a_truncated_following_frame() {
|
||||
let wire = encode_all(AwsEventStreamCodec, [message(b"first"), message(b"second")]);
|
||||
let mut messages = Box::pin(frames(
|
||||
input(every(&wire[..wire.len() - 1], 3)),
|
||||
AwsEventStreamCodec,
|
||||
));
|
||||
|
||||
assert_eq!(messages.next().await.unwrap().unwrap(), message(b"first"));
|
||||
assert!(matches!(
|
||||
messages.next().await,
|
||||
Some(Err(EventStreamError::Truncated))
|
||||
));
|
||||
assert!(messages.next().await.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_body_error_after_a_complete_frame_preserves_its_cause() {
|
||||
let first = encode_all(AwsEventStreamCodec, [message(b"first")]);
|
||||
let mut messages = Box::pin(frames(
|
||||
stream::iter([
|
||||
Ok(cut_at(&first, [5])[0].clone()),
|
||||
Ok(cut_at(&first, [5])[1].clone()),
|
||||
Ok(Bytes::from_static(b"\0\0\0")),
|
||||
Err(io::Error::new(io::ErrorKind::ConnectionReset, "reset")),
|
||||
]),
|
||||
AwsEventStreamCodec,
|
||||
));
|
||||
|
||||
assert_eq!(messages.next().await.unwrap().unwrap(), message(b"first"));
|
||||
let Some(Err(EventStreamError::Body(body))) = messages.next().await else {
|
||||
panic!("the body error surfaces");
|
||||
};
|
||||
assert_eq!(
|
||||
body_cause::<io::Error>(&body).unwrap().kind(),
|
||||
io::ErrorKind::ConnectionReset
|
||||
);
|
||||
assert!(messages.next().await.is_none());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,28 +2,64 @@
|
|||
|
||||
mod support;
|
||||
|
||||
use std::io;
|
||||
use bytes::Bytes;
|
||||
use futures_util::{StreamExt, TryStreamExt};
|
||||
use litellm_framing::{
|
||||
EventStreamError, SseError,
|
||||
aws_event_stream::{AwsEventStreamCodec, Message},
|
||||
frames,
|
||||
sse::{SseCodec, SseEvent},
|
||||
};
|
||||
use proptest::prelude::*;
|
||||
use support::{body_cause, cut_at, encode_all, every, input, runtime};
|
||||
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_framing::Framer;
|
||||
use litellm_framing::aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer};
|
||||
use litellm_framing::sse::SseFramer;
|
||||
fn delta(data: &str) -> SseEvent {
|
||||
SseEvent {
|
||||
event: Some("delta".into()),
|
||||
data: data.into(),
|
||||
id: Some("7".into()),
|
||||
retry: None,
|
||||
}
|
||||
}
|
||||
|
||||
use support::encode;
|
||||
fn envelopes(payloads: Vec<Bytes>) -> Vec<u8> {
|
||||
encode_all(AwsEventStreamCodec, payloads.into_iter().map(Message::new))
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#[test]
|
||||
fn an_sse_event_cut_anywhere_across_envelopes_is_reassembled(cut in 0_usize..64, chunk in 1_usize..8) {
|
||||
let sse = encode_all(SseCodec::default(), [delta("hello")]);
|
||||
let wire = envelopes(cut_at(&sse, [cut.min(sse.len())]));
|
||||
let events = runtime().block_on(async {
|
||||
let payloads = frames(input(every(&wire, chunk)), AwsEventStreamCodec)
|
||||
.map_ok(|message| message.payload().clone());
|
||||
frames(payloads, SseCodec::default()).try_collect::<Vec<_>>().await
|
||||
})
|
||||
.unwrap();
|
||||
prop_assert_eq!(events, vec![delta("hello")]);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn hosting_payloads_feed_the_same_sse_framer_across_envelope_boundaries() {
|
||||
let bytes = [encode(b"event: delta\ndata: hel"), encode(b"lo\nid: 7\n\n")].concat();
|
||||
let envelopes = AwsEventStreamFramer.frame(futures_util::stream::iter(
|
||||
bytes.chunks(3).map(Ok::<_, io::Error>),
|
||||
async fn a_truncated_envelope_after_an_sse_event_keeps_the_event_and_its_cause() {
|
||||
let complete = encode_all(SseCodec::default(), [delta("complete")]);
|
||||
let incomplete = encode_all(SseCodec::default(), [delta("incomplete")]);
|
||||
let wire = envelopes(vec![complete.into(), incomplete.into()]);
|
||||
let payloads = frames(
|
||||
input(every(&wire[..wire.len() - 1], 3)),
|
||||
AwsEventStreamCodec,
|
||||
)
|
||||
.map_ok(|message| message.payload().clone());
|
||||
let mut events = Box::pin(frames(payloads, SseCodec::default()));
|
||||
|
||||
assert_eq!(events.next().await.unwrap().unwrap(), delta("complete"));
|
||||
let Some(Err(SseError::Body(body))) = events.next().await else {
|
||||
panic!("the envelope error surfaces through the SSE layer");
|
||||
};
|
||||
assert!(matches!(
|
||||
body_cause::<EventStreamError>(&body),
|
||||
Some(EventStreamError::Truncated)
|
||||
));
|
||||
let frames = SseFramer
|
||||
.frame(envelopes.map_ok(|frame: AwsEventStreamFrame| frame.payload))
|
||||
.try_collect::<Vec<_>>()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(frames.len(), 1);
|
||||
assert_eq!(frames[0].event.as_deref(), Some("delta"));
|
||||
assert_eq!(frames[0].data.as_deref(), Some("hello"));
|
||||
assert_eq!(frames[0].id.as_deref(), Some("7"));
|
||||
assert!(events.next().await.is_none());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,67 +1,169 @@
|
|||
#![cfg(feature = "sse")]
|
||||
|
||||
mod support;
|
||||
|
||||
use std::io;
|
||||
|
||||
use futures_util::{StreamExt, TryStreamExt};
|
||||
use litellm_framing::sse::{SseFrame, SseFramer};
|
||||
use litellm_framing::{Error, Framer};
|
||||
use bytes::Bytes;
|
||||
use futures_util::{StreamExt, TryStreamExt, stream};
|
||||
use litellm_framing::{
|
||||
SseError, frames,
|
||||
sse::{SseCodec, SseEvent},
|
||||
};
|
||||
use proptest::prelude::*;
|
||||
use rstest::rstest;
|
||||
use support::{body_cause, cut_at, encode_all, every, input, runtime};
|
||||
|
||||
async fn collect_sse(chunks: &[&[u8]]) -> Result<Vec<SseFrame>, Error> {
|
||||
SseFramer
|
||||
.frame(futures_util::stream::iter(
|
||||
chunks.iter().copied().map(Ok::<_, io::Error>),
|
||||
))
|
||||
async fn collect(pieces: Vec<Bytes>) -> Result<Vec<SseEvent>, SseError> {
|
||||
frames(input(pieces), SseCodec::default())
|
||||
.try_collect()
|
||||
.await
|
||||
}
|
||||
|
||||
fn event(name: Option<&str>, data: &str) -> SseEvent {
|
||||
SseEvent {
|
||||
event: name.map(str::to_owned),
|
||||
data: data.to_owned(),
|
||||
id: None,
|
||||
retry: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn sse_event() -> impl Strategy<Value = SseEvent> {
|
||||
(
|
||||
proptest::option::of("[^\r\n\0]{0,8}"),
|
||||
"[^\r\0]{0,16}",
|
||||
proptest::option::of("[^\r\n\0]{0,8}"),
|
||||
proptest::option::of(any::<u64>()),
|
||||
)
|
||||
.prop_map(|(event, data, id, retry)| SseEvent {
|
||||
event,
|
||||
data,
|
||||
id,
|
||||
retry,
|
||||
})
|
||||
}
|
||||
|
||||
fn terminators() -> impl Strategy<Value = &'static [u8]> {
|
||||
prop_oneof![Just(&b"\n"[..]), Just(&b"\r\n"[..]), Just(&b"\r"[..])]
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#[test]
|
||||
fn any_events_survive_a_round_trip_through_any_terminator_and_any_cuts(
|
||||
events in proptest::collection::vec(sse_event(), 1..4),
|
||||
terminator in terminators(),
|
||||
cuts in proptest::collection::vec(0_usize..256, 0..4),
|
||||
bom in any::<bool>(),
|
||||
) {
|
||||
let lf_wire = encode_all(SseCodec::default(), events.clone());
|
||||
let body: Vec<u8> = lf_wire
|
||||
.iter()
|
||||
.flat_map(|byte| if *byte == b'\n' { terminator.to_vec() } else { vec![*byte] })
|
||||
.collect();
|
||||
let wire = if bom { [&b"\xEF\xBB\xBF"[..], &body].concat() } else { body };
|
||||
let decoded = runtime().block_on(collect(cut_at(&wire, cuts))).unwrap();
|
||||
prop_assert_eq!(decoded, events);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(
|
||||
&[&b":ping\r\nevent: delta\r\nid: 7\r\nretry: 10\r\ndata: \xe2"[..], &b"\x82"[..], &b"\xac\r"[..], &b"\ndata: next\r\n\r"[..], &b"\ndata: [DONE]\n\n"[..]],
|
||||
vec![
|
||||
SseFrame {
|
||||
event: Some("delta".into()),
|
||||
data: Some("€\nnext".into()),
|
||||
id: Some("7".into()),
|
||||
retry: Some(10),
|
||||
},
|
||||
SseFrame {
|
||||
event: None,
|
||||
data: Some("[DONE]".into()),
|
||||
id: None,
|
||||
retry: None,
|
||||
},
|
||||
]
|
||||
)]
|
||||
#[case::comment(b":ping\ndata: x\n\n")]
|
||||
#[case::unknown_field(b"vendor: 1\ndata: x\n\n")]
|
||||
#[case::field_without_colon(b"garbage\ndata: x\n\n")]
|
||||
#[case::retry_with_non_digits(b"retry: soon\ndata: x\n\n")]
|
||||
#[case::retry_with_a_sign(b"retry: +5\ndata: x\n\n")]
|
||||
#[case::retry_without_a_value(b"retry:\ndata: x\n\n")]
|
||||
#[case::id_with_nul(b"id: a\0b\ndata: x\n\n")]
|
||||
#[tokio::test]
|
||||
async fn fragmented_utf8_crlf_and_multiline_data_retain_metadata_and_sentinel(
|
||||
#[case] chunks: &[&[u8]],
|
||||
#[case] expected: Vec<SseFrame>,
|
||||
async fn lines_the_spec_ignores_do_not_change_the_event(#[case] wire: &[u8]) {
|
||||
assert_eq!(
|
||||
collect(every(wire, 1)).await.unwrap(),
|
||||
vec![event(None, "x")]
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_data_at_all(b"event: ping\nid: 1\n\ndata: x\n\n", vec![event(None, "x")])]
|
||||
#[case::empty_data_field(b"data:\n\n", vec![event(None, "")])]
|
||||
#[case::one_leading_space_stripped(b"data: x\n\n", vec![event(None, " x")])]
|
||||
#[case::multiline_data(b"data: a\ndata: b\ndata:\n\n", vec![event(None, "a\nb\n")])]
|
||||
#[case::last_event_name_wins(b"event: a\nevent: b\ndata: x\n\n", vec![event(Some("b"), "x")])]
|
||||
#[case::last_retry_wins(b"retry: 1\nretry: 2\ndata: x\n\n", vec![SseEvent { retry: Some(2), ..event(None, "x") }])]
|
||||
#[case::split_utf8_across_lines_is_not_joined(b"data: \xe2\x82\xac\ndata: \xe2\x82\xac\n\n", vec![event(None, "€\n€")])]
|
||||
#[tokio::test]
|
||||
async fn dispatch_follows_the_data_buffer(#[case] wire: &[u8], #[case] expected: Vec<SseEvent>) {
|
||||
assert_eq!(collect(every(wire, 1)).await.unwrap(), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::unterminated_single(b"data: partial\n", vec![])]
|
||||
#[case::unterminated_tail_after_complete(b"data: complete\n\ndata: unfinished\n", vec![event(None, "complete")])]
|
||||
#[case::lone_cr_terminates_at_eof(b"data: x\r\r", vec![event(None, "x")])]
|
||||
#[case::lone_cr_line_then_eof(b"data: x\r", vec![])]
|
||||
#[tokio::test]
|
||||
async fn eof_dispatches_only_terminated_events(
|
||||
#[case] wire: &[u8],
|
||||
#[case] expected: Vec<SseEvent>,
|
||||
) {
|
||||
assert_eq!(collect_sse(chunks).await.unwrap(), expected);
|
||||
assert_eq!(
|
||||
collect(vec![Bytes::copy_from_slice(wire)]).await.unwrap(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::inside_the_first_line(vec![&b"data: a\r"[..], &b"\ndata: b\r\n\r\n"[..]])]
|
||||
#[case::inside_the_blank_line(vec![&b"data: a\r\ndata: b\r\n\r"[..], &b"\n"[..]])]
|
||||
#[tokio::test]
|
||||
async fn a_crlf_split_across_chunks_is_one_terminator(#[case] pieces: Vec<&[u8]>) {
|
||||
let pieces = pieces.into_iter().map(Bytes::copy_from_slice).collect();
|
||||
assert_eq!(collect(pieces).await.unwrap(), vec![event(None, "a\nb")]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn eof_does_not_dispatch_an_unterminated_frame() {
|
||||
assert!(collect_sse(&[b"data: partial\n"]).await.unwrap().is_empty());
|
||||
async fn a_bom_is_stripped_only_at_the_start_of_the_stream() {
|
||||
let wire = b"\xEF\xBB\xBFdata: a\n\n\xEF\xBB\xBFdata: b\ndata: c\n\n";
|
||||
let decoded = collect(every(wire, 2)).await.unwrap();
|
||||
assert_eq!(decoded, vec![event(None, "a"), event(None, "c")]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn invalid_utf8_in_a_field_fails_after_earlier_events_and_terminates() {
|
||||
let mut events = Box::pin(frames(
|
||||
input(every(b"data: ok\n\ndata: \xff\n\n", 3)),
|
||||
SseCodec::default(),
|
||||
));
|
||||
|
||||
assert_eq!(events.next().await.unwrap().unwrap(), event(None, "ok"));
|
||||
assert!(matches!(
|
||||
events.next().await,
|
||||
Some(Err(SseError::InvalidUtf8(_)))
|
||||
));
|
||||
assert!(events.next().await.is_none());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(io::ErrorKind::ConnectionReset)]
|
||||
#[case(io::ErrorKind::UnexpectedEof)]
|
||||
#[tokio::test]
|
||||
async fn framing_errors_terminate_and_preserve_input_error_causes(#[case] kind: io::ErrorKind) {
|
||||
let mut frames = Box::pin(SseFramer.frame(futures_util::stream::iter([
|
||||
Err(io::Error::new(kind, "reset")),
|
||||
Ok(&b"data: later\n\n"[..]),
|
||||
])));
|
||||
let error = frames.next().await.unwrap().unwrap_err();
|
||||
assert!(matches!(
|
||||
error,
|
||||
Error::Sse(sse_stream::Error::Body(ref cause))
|
||||
if cause.downcast_ref::<io::Error>().unwrap().kind() == kind
|
||||
async fn a_body_error_keeps_earlier_events_and_its_cause_then_terminates(
|
||||
#[case] kind: io::ErrorKind,
|
||||
) {
|
||||
let mut events = Box::pin(frames(
|
||||
stream::iter([
|
||||
Ok(&b"data: first\n\ndata: partial"[..]),
|
||||
Err(io::Error::new(kind, "reset")),
|
||||
Ok(&b"\n\n"[..]),
|
||||
]),
|
||||
SseCodec::default(),
|
||||
));
|
||||
assert!(frames.next().await.is_none());
|
||||
assert!(frames.next().await.is_none());
|
||||
|
||||
assert_eq!(events.next().await.unwrap().unwrap(), event(None, "first"));
|
||||
let Some(Err(SseError::Body(body))) = events.next().await else {
|
||||
panic!("the body error surfaces");
|
||||
};
|
||||
assert_eq!(body_cause::<io::Error>(&body).unwrap().kind(), kind);
|
||||
assert!(events.next().await.is_none());
|
||||
assert!(events.next().await.is_none());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,15 +1,57 @@
|
|||
use aws_smithy_eventstream::frame::write_message_to;
|
||||
use aws_smithy_types::event_stream::{Header, HeaderValue, Message};
|
||||
use bytes::Bytes;
|
||||
#![allow(dead_code)]
|
||||
|
||||
pub fn encode(payload: &'static [u8]) -> Vec<u8> {
|
||||
let message = Message::new(Bytes::from_static(payload))
|
||||
.add_header(Header::new(
|
||||
":event-type",
|
||||
HeaderValue::String("payload".into()),
|
||||
))
|
||||
.add_header(Header::new("sequence", HeaderValue::Int32(7)));
|
||||
let mut bytes = Vec::new();
|
||||
write_message_to(&message, &mut bytes).unwrap();
|
||||
bytes
|
||||
use std::{error::Error, io};
|
||||
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use futures_util::{Stream, stream};
|
||||
use tokio_util::codec::Encoder;
|
||||
|
||||
pub fn encode_all<C, I>(mut codec: C, items: impl IntoIterator<Item = I>) -> Vec<u8>
|
||||
where
|
||||
C: Encoder<I>,
|
||||
C::Error: std::fmt::Debug,
|
||||
{
|
||||
let mut wire = BytesMut::new();
|
||||
for item in items {
|
||||
codec.encode(item, &mut wire).unwrap();
|
||||
}
|
||||
wire.to_vec()
|
||||
}
|
||||
|
||||
pub fn cut_at(bytes: &[u8], offsets: impl IntoIterator<Item = usize>) -> Vec<Bytes> {
|
||||
let mut sorted: Vec<usize> = offsets
|
||||
.into_iter()
|
||||
.filter(|offset| *offset <= bytes.len())
|
||||
.collect();
|
||||
sorted.sort_unstable();
|
||||
sorted.dedup();
|
||||
let bounds = std::iter::once(0)
|
||||
.chain(sorted)
|
||||
.chain(std::iter::once(bytes.len()))
|
||||
.collect::<Vec<_>>();
|
||||
bounds
|
||||
.windows(2)
|
||||
.map(|pair| Bytes::copy_from_slice(&bytes[pair[0]..pair[1]]))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn every(bytes: &[u8], size: usize) -> Vec<Bytes> {
|
||||
bytes
|
||||
.chunks(size.max(1))
|
||||
.map(Bytes::copy_from_slice)
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn input(pieces: Vec<Bytes>) -> impl Stream<Item = Result<Bytes, io::Error>> + Send {
|
||||
stream::iter(pieces.into_iter().map(Ok))
|
||||
}
|
||||
|
||||
pub fn body_cause<T: Error + 'static>(body: &io::Error) -> Option<&T> {
|
||||
body.get_ref()?.downcast_ref::<T>()
|
||||
}
|
||||
|
||||
pub fn runtime() -> tokio::runtime::Runtime {
|
||||
tokio::runtime::Builder::new_current_thread()
|
||||
.build()
|
||||
.unwrap()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,9 +2,9 @@ use base64::Engine;
|
|||
use bytes::Buf;
|
||||
use futures_util::{Stream, StreamExt};
|
||||
use litellm_framing::{
|
||||
Framer,
|
||||
aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer},
|
||||
sse::{SseFrame, SseFramer},
|
||||
aws_event_stream::{AwsEventStreamCodec, Message},
|
||||
frames,
|
||||
sse::{SseCodec, SseEvent},
|
||||
};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
|
@ -13,8 +13,6 @@ use serde_json::{Map, Value};
|
|||
pub enum Error {
|
||||
#[error("stream framing failed: {0}")]
|
||||
StreamFraming(String),
|
||||
#[error("Anthropic SSE frame has no data")]
|
||||
MissingStreamData,
|
||||
#[error("Anthropic stream event is invalid: {0}")]
|
||||
InvalidStreamEvent(String),
|
||||
#[error("Bedrock event payload is invalid: {0}")]
|
||||
|
|
@ -165,15 +163,14 @@ struct BedrockChunkPayload {
|
|||
bytes: String,
|
||||
}
|
||||
|
||||
pub fn decode_anthropic_sse_frame(frame: SseFrame) -> Result<AnthropicMessagesStreamEvent, Error> {
|
||||
let data = frame.data.ok_or(Error::MissingStreamData)?;
|
||||
serde_json::from_str(&data).map_err(|error| Error::InvalidStreamEvent(error.to_string()))
|
||||
pub fn decode_anthropic_sse_frame(event: SseEvent) -> Result<AnthropicMessagesStreamEvent, Error> {
|
||||
serde_json::from_str(&event.data).map_err(|error| Error::InvalidStreamEvent(error.to_string()))
|
||||
}
|
||||
|
||||
pub fn decode_bedrock_anthropic_frame(
|
||||
frame: AwsEventStreamFrame,
|
||||
message: Message,
|
||||
) -> Result<AnthropicMessagesStreamEvent, Error> {
|
||||
let payload: BedrockChunkPayload = serde_json::from_slice(&frame.payload)
|
||||
let payload: BedrockChunkPayload = serde_json::from_slice(message.payload())
|
||||
.map_err(|error| Error::InvalidBedrockPayload(error.to_string()))?;
|
||||
let event = base64::engine::general_purpose::STANDARD
|
||||
.decode(payload.bytes)
|
||||
|
|
@ -189,9 +186,8 @@ where
|
|||
B: Buf + Send,
|
||||
E: std::error::Error + Send + Sync + 'static,
|
||||
{
|
||||
SseFramer.frame(input).map(|frame| {
|
||||
let frame = frame.map_err(|error| Error::StreamFraming(error.to_string()))?;
|
||||
decode_anthropic_sse_frame(frame)
|
||||
frames(input, SseCodec::default()).map(|event| {
|
||||
decode_anthropic_sse_frame(event.map_err(|error| Error::StreamFraming(error.to_string()))?)
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -203,9 +199,10 @@ where
|
|||
B: Buf + Send,
|
||||
E: std::error::Error + Send + Sync + 'static,
|
||||
{
|
||||
AwsEventStreamFramer.frame(input).map(|frame| {
|
||||
let frame = frame.map_err(|error| Error::StreamFraming(error.to_string()))?;
|
||||
decode_bedrock_anthropic_frame(frame)
|
||||
frames(input, AwsEventStreamCodec).map(|message| {
|
||||
decode_bedrock_anthropic_frame(
|
||||
message.map_err(|error| Error::StreamFraming(error.to_string()))?,
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -247,12 +244,10 @@ mod tests {
|
|||
|
||||
#[test]
|
||||
fn decodes_citations_delta_events() {
|
||||
let event = decode_anthropic_sse_frame(SseFrame {
|
||||
let event = decode_anthropic_sse_frame(SseEvent {
|
||||
event: Some("content_block_delta".into()),
|
||||
data: Some(
|
||||
r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"#
|
||||
.into(),
|
||||
),
|
||||
data: r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"#
|
||||
.into(),
|
||||
id: None,
|
||||
retry: None,
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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 |
|
||||
| --- | --- |
|
||||
|
|
|
|||
32
litellm-rust/crates/testkit/Cargo.toml
Normal file
32
litellm-rust/crates/testkit/Cargo.toml
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
[package]
|
||||
name = "litellm-testkit"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
publish = false
|
||||
|
||||
[dependencies]
|
||||
flate2.workspace = true
|
||||
reqwest.workspace = true
|
||||
serde.workspace = true
|
||||
semver.workspace = true
|
||||
serde_json.workspace = true
|
||||
sha2.workspace = true
|
||||
tar.workspace = true
|
||||
target-lexicon.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio = { workspace = true, features = ["fs", "process"] }
|
||||
zip.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
flate2.workspace = true
|
||||
rstest.workspace = true
|
||||
sha2.workspace = true
|
||||
tar.workspace = true
|
||||
target-lexicon.workspace = true
|
||||
futures-util.workspace = true
|
||||
tempfile.workspace = true
|
||||
tokio.workspace = true
|
||||
toml = "0.9"
|
||||
zip.workspace = true
|
||||
181
litellm-rust/crates/testkit/src/agent/claude.rs
Normal file
181
litellm-rust/crates/testkit/src/agent/claude.rs
Normal file
|
|
@ -0,0 +1,181 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::path::Path;
|
||||
|
||||
use semver::Version;
|
||||
use serde::Deserialize;
|
||||
|
||||
use super::{
|
||||
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
|
||||
path_string,
|
||||
};
|
||||
use crate::install::release::parse;
|
||||
use crate::install::{Packaging, Release};
|
||||
use crate::{Error, Fetch, Target};
|
||||
|
||||
const RELEASES: &str = "https://downloads.claude.ai/claude-code-releases";
|
||||
|
||||
pub struct ClaudeCode;
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Manifest {
|
||||
platforms: BTreeMap<String, Platform>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Platform {
|
||||
checksum: String,
|
||||
}
|
||||
|
||||
impl Install for ClaudeCode {
|
||||
fn binary(&self) -> &'static str {
|
||||
"claude"
|
||||
}
|
||||
|
||||
async fn release(
|
||||
&self,
|
||||
fetch: &impl Fetch,
|
||||
version: &Version,
|
||||
target: Target,
|
||||
) -> Result<Release, Error> {
|
||||
let manifest_url = format!("{RELEASES}/{version}/manifest.json");
|
||||
let manifest: Manifest = parse(&manifest_url, &fetch.get(&manifest_url).await?)?;
|
||||
let key = format!(
|
||||
"{}-{}{}",
|
||||
target.os_name(),
|
||||
target.arch_name(),
|
||||
target.musl_suffix()
|
||||
);
|
||||
let platform = manifest
|
||||
.platforms
|
||||
.get(&key)
|
||||
.ok_or_else(|| Error::AssetNotFound(key.clone()))?;
|
||||
Ok(Release {
|
||||
url: format!("{RELEASES}/{version}/{key}/claude"),
|
||||
asset: key,
|
||||
sha256: platform.checksum.clone(),
|
||||
packaging: Packaging::Bare,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Configure for ClaudeCode {
|
||||
fn configure(
|
||||
&self,
|
||||
_version: &Version,
|
||||
settings: &Settings,
|
||||
home: &Path,
|
||||
) -> Result<LaunchSpec, Error> {
|
||||
if settings.wire != Wire::Messages {
|
||||
return Err(Error::UnsupportedWire {
|
||||
agent: "claude",
|
||||
wire: settings.wire,
|
||||
});
|
||||
}
|
||||
Ok(LaunchSpec {
|
||||
env: env([
|
||||
("HOME", path_string(home)),
|
||||
("CLAUDE_CONFIG_DIR", path_string(&home.join(".claude"))),
|
||||
("ANTHROPIC_BASE_URL", settings.base_url.clone()),
|
||||
("ANTHROPIC_AUTH_TOKEN", settings.api_key.clone()),
|
||||
("ANTHROPIC_MODEL", settings.model.clone()),
|
||||
("DISABLE_AUTOUPDATER", "1".to_owned()),
|
||||
]),
|
||||
files: BTreeMap::new(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
enum Event {
|
||||
Assistant {
|
||||
message: AssistantMessage,
|
||||
},
|
||||
Result(Finished),
|
||||
#[serde(other)]
|
||||
Other,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct AssistantMessage {
|
||||
content: Vec<Block>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Block {
|
||||
#[serde(rename = "type")]
|
||||
kind: String,
|
||||
name: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Finished {
|
||||
is_error: bool,
|
||||
result: Option<String>,
|
||||
usage: Option<TokenUsage>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct TokenUsage {
|
||||
input_tokens: u64,
|
||||
output_tokens: u64,
|
||||
}
|
||||
|
||||
impl Drive for ClaudeCode {
|
||||
fn args(&self, _version: &Version, settings: &Settings, prompt: &Prompt) -> Vec<String> {
|
||||
let base = [
|
||||
"-p",
|
||||
&prompt.text,
|
||||
"--output-format",
|
||||
"stream-json",
|
||||
"--verbose",
|
||||
"--model",
|
||||
&settings.model,
|
||||
];
|
||||
let tools = ["--allowedTools", "Bash,Read,Write,Edit"];
|
||||
base.into_iter()
|
||||
.chain(tools.into_iter().filter(|_| prompt.allow_tools))
|
||||
.map(str::to_owned)
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
|
||||
let events: Vec<Event> = json_lines(stdout).collect();
|
||||
let tool_calls = events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Event::Assistant { message } => Some(&message.content),
|
||||
_ => None,
|
||||
})
|
||||
.flatten()
|
||||
.filter(|block| block.kind == "tool_use")
|
||||
.filter_map(|block| block.name.clone())
|
||||
.collect();
|
||||
let finished = events.into_iter().find_map(|event| match event {
|
||||
Event::Result(finished) => Some(finished),
|
||||
_ => None,
|
||||
});
|
||||
let Some(finished) = finished else {
|
||||
return Outcome {
|
||||
tool_calls,
|
||||
..Outcome::default()
|
||||
};
|
||||
};
|
||||
let result = finished.result.unwrap_or_default();
|
||||
let (text, errors) = if finished.is_error {
|
||||
(String::new(), vec![result])
|
||||
} else {
|
||||
(result, Vec::new())
|
||||
};
|
||||
Outcome {
|
||||
text,
|
||||
tool_calls,
|
||||
usage: finished.usage.map_or_else(Usage::default, |usage| Usage {
|
||||
input_tokens: usage.input_tokens,
|
||||
output_tokens: usage.output_tokens,
|
||||
}),
|
||||
errors,
|
||||
exit_code: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
174
litellm-rust/crates/testkit/src/agent/codex.rs
Normal file
174
litellm-rust/crates/testkit/src/agent/codex.rs
Normal file
|
|
@ -0,0 +1,174 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use semver::Version;
|
||||
use serde::Deserialize;
|
||||
|
||||
use super::{
|
||||
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
|
||||
path_string, quoted, v1,
|
||||
};
|
||||
use crate::install::release::github_release;
|
||||
use crate::install::{Packaging, Release};
|
||||
use crate::target::{Arch, Os};
|
||||
use crate::{Error, Fetch, Target};
|
||||
|
||||
const RELEASES: &str = "https://api.github.com/repos/openai/codex/releases/tags";
|
||||
|
||||
pub struct Codex;
|
||||
|
||||
fn triple(target: Target) -> String {
|
||||
let arch = match target.arch {
|
||||
Arch::Aarch64 => "aarch64",
|
||||
Arch::X86_64 => "x86_64",
|
||||
};
|
||||
match target.os {
|
||||
Os::Macos => format!("{arch}-apple-darwin"),
|
||||
Os::Linux => format!("{arch}-unknown-linux-musl"),
|
||||
}
|
||||
}
|
||||
|
||||
impl Install for Codex {
|
||||
fn binary(&self) -> &'static str {
|
||||
"codex"
|
||||
}
|
||||
|
||||
async fn release(
|
||||
&self,
|
||||
fetch: &impl Fetch,
|
||||
version: &Version,
|
||||
target: Target,
|
||||
) -> Result<Release, Error> {
|
||||
let triple = triple(target);
|
||||
github_release(
|
||||
fetch,
|
||||
RELEASES,
|
||||
&format!("rust-v{version}"),
|
||||
&format!("codex-{triple}.tar.gz"),
|
||||
Packaging::TarGz {
|
||||
member: format!("codex-{triple}"),
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
impl Configure for Codex {
|
||||
fn configure(
|
||||
&self,
|
||||
_version: &Version,
|
||||
settings: &Settings,
|
||||
home: &Path,
|
||||
) -> Result<LaunchSpec, Error> {
|
||||
if settings.wire != Wire::Responses {
|
||||
return Err(Error::UnsupportedWire {
|
||||
agent: "codex",
|
||||
wire: settings.wire,
|
||||
});
|
||||
}
|
||||
let config = format!(
|
||||
"model = {model}\nmodel_provider = \"litellm\"\n\n[model_providers.litellm]\nname = \"LiteLLM\"\nbase_url = {base_url}\nenv_key = \"LITELLM_API_KEY\"\nwire_api = \"responses\"\n",
|
||||
model = quoted(&settings.model),
|
||||
base_url = quoted(&v1(settings)),
|
||||
);
|
||||
Ok(LaunchSpec {
|
||||
env: env([
|
||||
("HOME", path_string(home)),
|
||||
("CODEX_HOME", path_string(&home.join(".codex"))),
|
||||
("LITELLM_API_KEY", settings.api_key.clone()),
|
||||
]),
|
||||
files: BTreeMap::from([(PathBuf::from(".codex/config.toml"), config)]),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
enum EventKind {
|
||||
#[serde(rename = "item.completed")]
|
||||
ItemCompleted,
|
||||
#[serde(rename = "turn.completed")]
|
||||
TurnCompleted,
|
||||
#[serde(rename = "turn.failed")]
|
||||
TurnFailed,
|
||||
#[serde(other)]
|
||||
Other,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Event {
|
||||
#[serde(rename = "type")]
|
||||
kind: EventKind,
|
||||
item: Option<Item>,
|
||||
usage: Option<TokenUsage>,
|
||||
error: Option<Failure>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Item {
|
||||
#[serde(rename = "type")]
|
||||
kind: String,
|
||||
text: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct TokenUsage {
|
||||
input_tokens: u64,
|
||||
output_tokens: u64,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Failure {
|
||||
message: String,
|
||||
}
|
||||
|
||||
const NON_TOOL_ITEMS: [&str; 3] = ["agent_message", "reasoning", "error"];
|
||||
|
||||
impl Drive for Codex {
|
||||
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
|
||||
let sandbox = ["--sandbox", "workspace-write"];
|
||||
["exec", "--json", "--skip-git-repo-check"]
|
||||
.into_iter()
|
||||
.chain(sandbox.into_iter().filter(|_| prompt.allow_tools))
|
||||
.chain([prompt.text.as_str()])
|
||||
.map(str::to_owned)
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
|
||||
let events: Vec<Event> = json_lines(stdout).collect();
|
||||
let items: Vec<&Item> = events
|
||||
.iter()
|
||||
.filter(|event| matches!(event.kind, EventKind::ItemCompleted))
|
||||
.filter_map(|event| event.item.as_ref())
|
||||
.collect();
|
||||
Outcome {
|
||||
text: items
|
||||
.iter()
|
||||
.rev()
|
||||
.find(|item| item.kind == "agent_message")
|
||||
.and_then(|item| item.text.clone())
|
||||
.unwrap_or_default(),
|
||||
tool_calls: items
|
||||
.iter()
|
||||
.filter(|item| !NON_TOOL_ITEMS.contains(&item.kind.as_str()))
|
||||
.map(|item| item.kind.clone())
|
||||
.collect(),
|
||||
usage: events
|
||||
.iter()
|
||||
.filter(|event| matches!(event.kind, EventKind::TurnCompleted))
|
||||
.filter_map(|event| event.usage.as_ref())
|
||||
.map(|usage| Usage {
|
||||
input_tokens: usage.input_tokens,
|
||||
output_tokens: usage.output_tokens,
|
||||
})
|
||||
.fold(Usage::default(), |total, turn| total + turn),
|
||||
errors: events
|
||||
.iter()
|
||||
.filter(|event| matches!(event.kind, EventKind::TurnFailed))
|
||||
.filter_map(|event| event.error.as_ref())
|
||||
.map(|failure| failure.message.clone())
|
||||
.collect(),
|
||||
exit_code: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
69
litellm-rust/crates/testkit/src/agent/configure.rs
Normal file
69
litellm-rust/crates/testkit/src/agent/configure.rs
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use semver::Version;
|
||||
|
||||
use crate::Error;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Wire {
|
||||
ChatCompletions,
|
||||
Messages,
|
||||
Responses,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct Settings {
|
||||
pub base_url: String,
|
||||
pub api_key: String,
|
||||
pub model: String,
|
||||
pub wire: Wire,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct LaunchSpec {
|
||||
pub env: BTreeMap<String, String>,
|
||||
pub files: BTreeMap<PathBuf, String>,
|
||||
}
|
||||
|
||||
impl LaunchSpec {
|
||||
pub fn write_files(&self, home: &Path) -> std::io::Result<()> {
|
||||
self.files.iter().try_for_each(|(relative, contents)| {
|
||||
let path = home.join(relative);
|
||||
if let Some(parent) = path.parent() {
|
||||
std::fs::create_dir_all(parent)?;
|
||||
}
|
||||
std::fs::write(path, contents)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub trait Configure {
|
||||
fn configure(
|
||||
&self,
|
||||
version: &Version,
|
||||
settings: &Settings,
|
||||
home: &Path,
|
||||
) -> Result<LaunchSpec, Error>;
|
||||
}
|
||||
|
||||
pub(crate) fn env(
|
||||
pairs: impl IntoIterator<Item = (&'static str, String)>,
|
||||
) -> BTreeMap<String, String> {
|
||||
pairs
|
||||
.into_iter()
|
||||
.map(|(key, value)| (key.to_owned(), value))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(crate) fn path_string(path: &Path) -> String {
|
||||
path.to_string_lossy().into_owned()
|
||||
}
|
||||
|
||||
pub(crate) fn quoted(value: &str) -> String {
|
||||
serde_json::Value::from(value).to_string()
|
||||
}
|
||||
|
||||
pub(crate) fn v1(settings: &Settings) -> String {
|
||||
format!("{}/v1", settings.base_url.trim_end_matches('/'))
|
||||
}
|
||||
57
litellm-rust/crates/testkit/src/agent/drive.rs
Normal file
57
litellm-rust/crates/testkit/src/agent/drive.rs
Normal file
|
|
@ -0,0 +1,57 @@
|
|||
use std::ops::Add;
|
||||
|
||||
use semver::Version;
|
||||
|
||||
use crate::Settings;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct Prompt {
|
||||
pub text: String,
|
||||
pub allow_tools: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
|
||||
pub struct Usage {
|
||||
pub input_tokens: u64,
|
||||
pub output_tokens: u64,
|
||||
}
|
||||
|
||||
impl Add for Usage {
|
||||
type Output = Self;
|
||||
|
||||
fn add(self, other: Self) -> Self {
|
||||
Self {
|
||||
input_tokens: self.input_tokens + other.input_tokens,
|
||||
output_tokens: self.output_tokens + other.output_tokens,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq)]
|
||||
pub struct Outcome {
|
||||
pub text: String,
|
||||
pub tool_calls: Vec<String>,
|
||||
pub usage: Usage,
|
||||
pub errors: Vec<String>,
|
||||
pub exit_code: Option<i32>,
|
||||
}
|
||||
|
||||
impl Outcome {
|
||||
pub fn succeeded(&self) -> bool {
|
||||
self.exit_code == Some(0) && self.errors.is_empty()
|
||||
}
|
||||
}
|
||||
|
||||
pub trait Drive {
|
||||
fn args(&self, version: &Version, settings: &Settings, prompt: &Prompt) -> Vec<String>;
|
||||
|
||||
fn parse(&self, version: &Version, stdout: &str) -> Outcome;
|
||||
}
|
||||
|
||||
pub(crate) fn json_lines<'a, T: serde::de::DeserializeOwned + 'a>(
|
||||
stdout: &'a str,
|
||||
) -> impl Iterator<Item = T> + 'a {
|
||||
stdout
|
||||
.lines()
|
||||
.filter_map(|line| serde_json::from_str(line).ok())
|
||||
}
|
||||
17
litellm-rust/crates/testkit/src/agent/install.rs
Normal file
17
litellm-rust/crates/testkit/src/agent/install.rs
Normal file
|
|
@ -0,0 +1,17 @@
|
|||
use std::future::Future;
|
||||
|
||||
use semver::Version;
|
||||
|
||||
use crate::install::Release;
|
||||
use crate::{Error, Fetch, Target};
|
||||
|
||||
pub trait Install: Sync {
|
||||
fn binary(&self) -> &'static str;
|
||||
|
||||
fn release(
|
||||
&self,
|
||||
fetch: &impl Fetch,
|
||||
version: &Version,
|
||||
target: Target,
|
||||
) -> impl Future<Output = Result<Release, Error>> + Send;
|
||||
}
|
||||
20
litellm-rust/crates/testkit/src/agent/mod.rs
Normal file
20
litellm-rust/crates/testkit/src/agent/mod.rs
Normal file
|
|
@ -0,0 +1,20 @@
|
|||
mod claude;
|
||||
mod codex;
|
||||
mod configure;
|
||||
mod drive;
|
||||
mod install;
|
||||
mod opencode;
|
||||
|
||||
pub use claude::ClaudeCode;
|
||||
pub use codex::Codex;
|
||||
pub use configure::{Configure, LaunchSpec, Settings, Wire};
|
||||
pub use drive::{Drive, Outcome, Prompt, Usage};
|
||||
pub use install::Install;
|
||||
pub use opencode::Opencode;
|
||||
|
||||
pub(crate) use configure::{env, path_string, quoted, v1};
|
||||
pub(crate) use drive::json_lines;
|
||||
|
||||
pub trait Agent: Install + Configure + Drive {}
|
||||
|
||||
impl<T: Install + Configure + Drive> Agent for T {}
|
||||
187
litellm-rust/crates/testkit/src/agent/opencode.rs
Normal file
187
litellm-rust/crates/testkit/src/agent/opencode.rs
Normal file
|
|
@ -0,0 +1,187 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use semver::Version;
|
||||
use serde::Deserialize;
|
||||
|
||||
use super::{
|
||||
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
|
||||
path_string, v1,
|
||||
};
|
||||
use crate::install::release::github_release;
|
||||
use crate::install::{Packaging, Release};
|
||||
use crate::target::Os;
|
||||
use crate::{Error, Fetch, Target};
|
||||
|
||||
const RELEASES: &str = "https://api.github.com/repos/sst/opencode/releases/tags";
|
||||
|
||||
pub struct Opencode;
|
||||
|
||||
impl Install for Opencode {
|
||||
fn binary(&self) -> &'static str {
|
||||
"opencode"
|
||||
}
|
||||
|
||||
async fn release(
|
||||
&self,
|
||||
fetch: &impl Fetch,
|
||||
version: &Version,
|
||||
target: Target,
|
||||
) -> Result<Release, Error> {
|
||||
let stem = format!(
|
||||
"opencode-{}-{}{}",
|
||||
target.os_name(),
|
||||
target.arch_name(),
|
||||
target.musl_suffix()
|
||||
);
|
||||
let member = "opencode".to_owned();
|
||||
let (asset, packaging) = match target.os {
|
||||
Os::Macos => (format!("{stem}.zip"), Packaging::Zip { member }),
|
||||
Os::Linux => (format!("{stem}.tar.gz"), Packaging::TarGz { member }),
|
||||
};
|
||||
github_release(fetch, RELEASES, &format!("v{version}"), &asset, packaging).await
|
||||
}
|
||||
}
|
||||
|
||||
impl Configure for Opencode {
|
||||
fn configure(
|
||||
&self,
|
||||
_version: &Version,
|
||||
settings: &Settings,
|
||||
home: &Path,
|
||||
) -> Result<LaunchSpec, Error> {
|
||||
let npm = match settings.wire {
|
||||
Wire::ChatCompletions => "@ai-sdk/openai-compatible",
|
||||
Wire::Responses => "@ai-sdk/openai",
|
||||
Wire::Messages => "@ai-sdk/anthropic",
|
||||
};
|
||||
let config = serde_json::json!({
|
||||
"$schema": "https://opencode.ai/config.json",
|
||||
"model": format!("litellm/{}", settings.model),
|
||||
"provider": {
|
||||
"litellm": {
|
||||
"npm": npm,
|
||||
"name": "LiteLLM",
|
||||
"options": { "baseURL": v1(settings), "apiKey": settings.api_key },
|
||||
"models": { settings.model.clone(): { "name": settings.model } },
|
||||
}
|
||||
},
|
||||
});
|
||||
Ok(LaunchSpec {
|
||||
env: env([
|
||||
("HOME", path_string(home)),
|
||||
("XDG_CONFIG_HOME", path_string(&home.join(".config"))),
|
||||
("XDG_DATA_HOME", path_string(&home.join(".local/share"))),
|
||||
("OPENCODE_DISABLE_AUTOUPDATE", "true".to_owned()),
|
||||
]),
|
||||
files: BTreeMap::from([(
|
||||
PathBuf::from(".config/opencode/opencode.json"),
|
||||
config.to_string(),
|
||||
)]),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
enum Event {
|
||||
Text {
|
||||
part: TextPart,
|
||||
},
|
||||
ToolUse {
|
||||
part: ToolPart,
|
||||
},
|
||||
StepFinish {
|
||||
part: StepFinish,
|
||||
},
|
||||
Error {
|
||||
error: Failure,
|
||||
},
|
||||
#[serde(other)]
|
||||
Other,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct TextPart {
|
||||
text: String,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct ToolPart {
|
||||
tool: String,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct StepFinish {
|
||||
tokens: Tokens,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Tokens {
|
||||
input: u64,
|
||||
output: u64,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Failure {
|
||||
name: String,
|
||||
data: Option<FailureData>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct FailureData {
|
||||
message: Option<String>,
|
||||
}
|
||||
|
||||
impl Drive for Opencode {
|
||||
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
|
||||
["run", "--format", "json", &prompt.text]
|
||||
.map(str::to_owned)
|
||||
.to_vec()
|
||||
}
|
||||
|
||||
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
|
||||
let events: Vec<Event> = json_lines(stdout).collect();
|
||||
Outcome {
|
||||
text: events
|
||||
.iter()
|
||||
.rev()
|
||||
.find_map(|event| match event {
|
||||
Event::Text { part } => Some(part.text.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap_or_default(),
|
||||
tool_calls: events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Event::ToolUse { part } => Some(part.tool.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.collect(),
|
||||
usage: events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Event::StepFinish { part } => Some(Usage {
|
||||
input_tokens: part.tokens.input,
|
||||
output_tokens: part.tokens.output,
|
||||
}),
|
||||
_ => None,
|
||||
})
|
||||
.fold(Usage::default(), |total, step| total + step),
|
||||
errors: events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Event::Error { error } => Some(
|
||||
error
|
||||
.data
|
||||
.as_ref()
|
||||
.and_then(|data| data.message.clone())
|
||||
.unwrap_or_else(|| error.name.clone()),
|
||||
),
|
||||
_ => None,
|
||||
})
|
||||
.collect(),
|
||||
exit_code: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
56
litellm-rust/crates/testkit/src/error.rs
Normal file
56
litellm-rust/crates/testkit/src/error.rs
Normal file
|
|
@ -0,0 +1,56 @@
|
|||
use std::io;
|
||||
use std::path::PathBuf;
|
||||
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::Wire;
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum Error {
|
||||
#[error("unsupported target {0}")]
|
||||
UnsupportedTarget(String),
|
||||
#[error("{0} is not a plain x.y.z release version")]
|
||||
InvalidVersion(String),
|
||||
#[error("request to {url} failed")]
|
||||
Request {
|
||||
url: String,
|
||||
#[source]
|
||||
source: reqwest::Error,
|
||||
},
|
||||
#[error("{url} answered with status {status}")]
|
||||
Status { url: String, status: u16 },
|
||||
#[error("release metadata at {url} is malformed")]
|
||||
Metadata {
|
||||
url: String,
|
||||
#[source]
|
||||
source: serde_json::Error,
|
||||
},
|
||||
#[error("release has no asset named {0}")]
|
||||
AssetNotFound(String),
|
||||
#[error("release publishes no sha256 for {0}")]
|
||||
MissingChecksum(String),
|
||||
#[error("sha256 mismatch for {asset}: expected {expected}, got {actual}")]
|
||||
ChecksumMismatch {
|
||||
asset: String,
|
||||
expected: String,
|
||||
actual: String,
|
||||
},
|
||||
#[error("archive does not contain {0}")]
|
||||
ArchiveMemberNotFound(String),
|
||||
#[error("archive is unreadable")]
|
||||
Archive(#[source] io::Error),
|
||||
#[error("zip archive is unreadable")]
|
||||
Zip(#[from] zip::result::ZipError),
|
||||
#[error("{binary} reports version '{reported}', expected {expected}")]
|
||||
VersionMismatch {
|
||||
binary: PathBuf,
|
||||
expected: String,
|
||||
reported: String,
|
||||
},
|
||||
#[error("{agent} cannot talk to the gateway over {wire:?}")]
|
||||
UnsupportedWire { agent: &'static str, wire: Wire },
|
||||
#[error("agent did not finish within {0:?}")]
|
||||
Timeout(std::time::Duration),
|
||||
#[error("io failure")]
|
||||
Io(#[from] io::Error),
|
||||
}
|
||||
52
litellm-rust/crates/testkit/src/install/archive.rs
Normal file
52
litellm-rust/crates/testkit/src/install/archive.rs
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
use std::io::{Cursor, Read};
|
||||
|
||||
use flate2::read::GzDecoder;
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use super::release::Packaging;
|
||||
use crate::Error;
|
||||
|
||||
pub(crate) fn verify_sha256(asset: &str, expected: &str, bytes: &[u8]) -> Result<(), Error> {
|
||||
let actual = format!("{:x}", Sha256::digest(bytes));
|
||||
if actual.eq_ignore_ascii_case(expected) {
|
||||
return Ok(());
|
||||
}
|
||||
Err(Error::ChecksumMismatch {
|
||||
asset: asset.to_owned(),
|
||||
expected: expected.to_owned(),
|
||||
actual,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn extract_binary(packaging: &Packaging, bytes: &[u8]) -> Result<Vec<u8>, Error> {
|
||||
match packaging {
|
||||
Packaging::Bare => Ok(bytes.to_vec()),
|
||||
Packaging::TarGz { member } => extract_tar_gz(member, bytes),
|
||||
Packaging::Zip { member } => extract_zip(member, bytes),
|
||||
}
|
||||
}
|
||||
|
||||
fn extract_tar_gz(member: &str, bytes: &[u8]) -> Result<Vec<u8>, Error> {
|
||||
let mut archive = tar::Archive::new(GzDecoder::new(bytes));
|
||||
for entry in archive.entries().map_err(Error::Archive)? {
|
||||
let mut entry = entry.map_err(Error::Archive)?;
|
||||
let path = entry.path().map_err(Error::Archive)?;
|
||||
if path.file_name().is_some_and(|name| name == member) {
|
||||
let mut binary = Vec::new();
|
||||
entry.read_to_end(&mut binary).map_err(Error::Archive)?;
|
||||
return Ok(binary);
|
||||
}
|
||||
}
|
||||
Err(Error::ArchiveMemberNotFound(member.to_owned()))
|
||||
}
|
||||
|
||||
fn extract_zip(member: &str, bytes: &[u8]) -> Result<Vec<u8>, Error> {
|
||||
let mut archive = zip::ZipArchive::new(Cursor::new(bytes))?;
|
||||
let mut file = archive.by_name(member).map_err(|error| match error {
|
||||
zip::result::ZipError::FileNotFound => Error::ArchiveMemberNotFound(member.to_owned()),
|
||||
other => Error::Zip(other),
|
||||
})?;
|
||||
let mut binary = Vec::new();
|
||||
file.read_to_end(&mut binary).map_err(Error::Archive)?;
|
||||
Ok(binary)
|
||||
}
|
||||
55
litellm-rust/crates/testkit/src/install/fetch.rs
Normal file
55
litellm-rust/crates/testkit/src/install/fetch.rs
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
use std::future::Future;
|
||||
|
||||
use crate::Error;
|
||||
|
||||
pub trait Fetch: Sync {
|
||||
fn get(&self, url: &str) -> impl Future<Output = Result<Vec<u8>, Error>> + Send;
|
||||
}
|
||||
|
||||
pub struct HttpFetch {
|
||||
client: reqwest::Client,
|
||||
github_token: Option<String>,
|
||||
}
|
||||
|
||||
impl HttpFetch {
|
||||
pub fn new(github_token: Option<String>) -> Self {
|
||||
Self {
|
||||
client: reqwest::Client::new(),
|
||||
github_token,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn from_env() -> Self {
|
||||
Self::new(std::env::var("GITHUB_TOKEN").ok())
|
||||
}
|
||||
}
|
||||
|
||||
impl Fetch for HttpFetch {
|
||||
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
|
||||
let request = self
|
||||
.client
|
||||
.get(url)
|
||||
.header("user-agent", "litellm-testkit")
|
||||
.header("accept", "application/json, application/octet-stream");
|
||||
let request = match (
|
||||
&self.github_token,
|
||||
url.starts_with("https://api.github.com/"),
|
||||
) {
|
||||
(Some(token), true) => request.bearer_auth(token),
|
||||
_ => request,
|
||||
};
|
||||
let request_error = |source| Error::Request {
|
||||
url: url.to_owned(),
|
||||
source,
|
||||
};
|
||||
let response = request.send().await.map_err(request_error)?;
|
||||
let status = response.status();
|
||||
if !status.is_success() {
|
||||
return Err(Error::Status {
|
||||
url: url.to_owned(),
|
||||
status: status.as_u16(),
|
||||
});
|
||||
}
|
||||
Ok(response.bytes().await.map_err(request_error)?.to_vec())
|
||||
}
|
||||
}
|
||||
118
litellm-rust/crates/testkit/src/install/mod.rs
Normal file
118
litellm-rust/crates/testkit/src/install/mod.rs
Normal file
|
|
@ -0,0 +1,118 @@
|
|||
mod archive;
|
||||
mod fetch;
|
||||
pub(crate) mod release;
|
||||
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::process::Stdio;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
use semver::Version;
|
||||
use tokio::fs;
|
||||
use tokio::process::Command;
|
||||
|
||||
use crate::{Error, Install, Target};
|
||||
use archive::{extract_binary, verify_sha256};
|
||||
|
||||
static STAGING_COUNTER: AtomicU64 = AtomicU64::new(0);
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct Installed {
|
||||
pub version: Version,
|
||||
pub binary: PathBuf,
|
||||
}
|
||||
|
||||
pub struct Installer<F> {
|
||||
fetch: F,
|
||||
cache_root: PathBuf,
|
||||
target: Target,
|
||||
}
|
||||
|
||||
impl<F: Fetch> Installer<F> {
|
||||
pub fn new(fetch: F, cache_root: impl Into<PathBuf>, target: Target) -> Self {
|
||||
Self {
|
||||
fetch,
|
||||
cache_root: cache_root.into(),
|
||||
target,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn install(
|
||||
&self,
|
||||
agent: &impl Install,
|
||||
version: &Version,
|
||||
) -> Result<Installed, Error> {
|
||||
validate_release(version)?;
|
||||
let dir = self
|
||||
.cache_root
|
||||
.join(agent.binary())
|
||||
.join(version.to_string());
|
||||
let binary = dir.join(agent.binary());
|
||||
let installed = Installed {
|
||||
version: version.clone(),
|
||||
binary: binary.clone(),
|
||||
};
|
||||
if fs::try_exists(&binary).await? && probe_version(&binary, version).await.is_ok() {
|
||||
return Ok(installed);
|
||||
}
|
||||
|
||||
let release = agent.release(&self.fetch, version, self.target).await?;
|
||||
let archive = self.fetch.get(&release.url).await?;
|
||||
verify_sha256(&release.asset, &release.sha256, &archive)?;
|
||||
let contents = extract_binary(&release.packaging, &archive)?;
|
||||
|
||||
fs::create_dir_all(&dir).await?;
|
||||
let staging = dir.join(format!(
|
||||
".{}.{}.{}.partial",
|
||||
agent.binary(),
|
||||
std::process::id(),
|
||||
STAGING_COUNTER.fetch_add(1, Ordering::Relaxed)
|
||||
));
|
||||
fs::write(&staging, contents).await?;
|
||||
fs::set_permissions(&staging, std::fs::Permissions::from_mode(0o755)).await?;
|
||||
fs::rename(&staging, &binary).await?;
|
||||
|
||||
match probe_version(&binary, version).await {
|
||||
Ok(()) => Ok(installed),
|
||||
Err(error) => {
|
||||
fs::remove_file(&binary).await?;
|
||||
Err(error)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_release(version: &Version) -> Result<(), Error> {
|
||||
if version.pre.is_empty() && version.build.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
Err(Error::InvalidVersion(version.to_string()))
|
||||
}
|
||||
|
||||
async fn probe_version(binary: &Path, expected: &Version) -> Result<(), Error> {
|
||||
let home = std::env::temp_dir();
|
||||
let output = Command::new(binary)
|
||||
.arg("--version")
|
||||
.env_clear()
|
||||
.env("HOME", home)
|
||||
.env("DISABLE_AUTOUPDATER", "1")
|
||||
.stdin(Stdio::null())
|
||||
.output()
|
||||
.await?;
|
||||
let stdout = String::from_utf8_lossy(&output.stdout);
|
||||
if stdout
|
||||
.split_whitespace()
|
||||
.filter_map(|token| Version::parse(token).ok())
|
||||
.any(|reported| &reported == expected)
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
Err(Error::VersionMismatch {
|
||||
binary: binary.to_owned(),
|
||||
expected: expected.to_string(),
|
||||
reported: stdout.trim().to_owned(),
|
||||
})
|
||||
}
|
||||
|
||||
pub use fetch::{Fetch, HttpFetch};
|
||||
pub use release::{Packaging, Release};
|
||||
65
litellm-rust/crates/testkit/src/install/release.rs
Normal file
65
litellm-rust/crates/testkit/src/install/release.rs
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
use serde::Deserialize;
|
||||
|
||||
use crate::{Error, Fetch};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum Packaging {
|
||||
Bare,
|
||||
TarGz { member: String },
|
||||
Zip { member: String },
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct Release {
|
||||
pub asset: String,
|
||||
pub url: String,
|
||||
pub sha256: String,
|
||||
pub packaging: Packaging,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct GithubRelease {
|
||||
assets: Vec<GithubAsset>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct GithubAsset {
|
||||
name: String,
|
||||
digest: Option<String>,
|
||||
browser_download_url: String,
|
||||
}
|
||||
|
||||
pub(crate) async fn github_release(
|
||||
fetch: &impl Fetch,
|
||||
releases_url: &str,
|
||||
tag: &str,
|
||||
asset_name: &str,
|
||||
packaging: Packaging,
|
||||
) -> Result<Release, Error> {
|
||||
let url = format!("{releases_url}/{tag}");
|
||||
let release: GithubRelease = parse(&url, &fetch.get(&url).await?)?;
|
||||
let asset = release
|
||||
.assets
|
||||
.into_iter()
|
||||
.find(|asset| asset.name == asset_name)
|
||||
.ok_or_else(|| Error::AssetNotFound(asset_name.to_owned()))?;
|
||||
let sha256 = asset
|
||||
.digest
|
||||
.as_deref()
|
||||
.and_then(|digest| digest.strip_prefix("sha256:"))
|
||||
.ok_or_else(|| Error::MissingChecksum(asset_name.to_owned()))?
|
||||
.to_owned();
|
||||
Ok(Release {
|
||||
asset: asset.name,
|
||||
url: asset.browser_download_url,
|
||||
sha256,
|
||||
packaging,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn parse<T: for<'de> Deserialize<'de>>(url: &str, body: &[u8]) -> Result<T, Error> {
|
||||
serde_json::from_slice(body).map_err(|source| Error::Metadata {
|
||||
url: url.to_owned(),
|
||||
source,
|
||||
})
|
||||
}
|
||||
15
litellm-rust/crates/testkit/src/lib.rs
Normal file
15
litellm-rust/crates/testkit/src/lib.rs
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
mod agent;
|
||||
mod error;
|
||||
mod install;
|
||||
mod session;
|
||||
mod target;
|
||||
|
||||
pub use agent::{
|
||||
Agent, ClaudeCode, Codex, Configure, Drive, Install, LaunchSpec, Opencode, Outcome, Prompt,
|
||||
Settings, Usage, Wire,
|
||||
};
|
||||
pub use error::Error;
|
||||
pub use install::{Fetch, HttpFetch, Installed, Installer, Packaging, Release};
|
||||
pub use semver::Version;
|
||||
pub use session::Session;
|
||||
pub use target::{Arch, Os, Target};
|
||||
76
litellm-rust/crates/testkit/src/session.rs
Normal file
76
litellm-rust/crates/testkit/src/session.rs
Normal file
|
|
@ -0,0 +1,76 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::path::PathBuf;
|
||||
use std::process::Stdio;
|
||||
use std::time::Duration;
|
||||
|
||||
use semver::Version;
|
||||
use tokio::process::Command;
|
||||
use tokio::time::timeout;
|
||||
|
||||
use crate::{Configure, Drive, Error, Installed, Outcome, Prompt, Settings};
|
||||
|
||||
const STDERR_LIMIT_CHARS: usize = 2000;
|
||||
|
||||
pub struct Session {
|
||||
binary: PathBuf,
|
||||
home: PathBuf,
|
||||
version: Version,
|
||||
settings: Settings,
|
||||
env: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl Session {
|
||||
pub fn prepare(
|
||||
agent: &impl Configure,
|
||||
installed: &Installed,
|
||||
settings: Settings,
|
||||
home: impl Into<PathBuf>,
|
||||
) -> Result<Self, Error> {
|
||||
let home = home.into();
|
||||
let spec = agent.configure(&installed.version, &settings, &home)?;
|
||||
spec.write_files(&home)?;
|
||||
Ok(Self {
|
||||
binary: installed.binary.clone(),
|
||||
home,
|
||||
version: installed.version.clone(),
|
||||
settings,
|
||||
env: spec.env,
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn run(
|
||||
&self,
|
||||
agent: &impl Drive,
|
||||
prompt: &Prompt,
|
||||
limit: Duration,
|
||||
) -> Result<Outcome, Error> {
|
||||
let child = Command::new(&self.binary)
|
||||
.args(agent.args(&self.version, &self.settings, prompt))
|
||||
.env_clear()
|
||||
.env("PATH", "/usr/bin:/bin")
|
||||
.envs(&self.env)
|
||||
.current_dir(&self.home)
|
||||
.stdin(Stdio::null())
|
||||
.kill_on_drop(true)
|
||||
.output();
|
||||
let output = timeout(limit, child)
|
||||
.await
|
||||
.map_err(|_| Error::Timeout(limit))??;
|
||||
let parsed = agent.parse(&self.version, &String::from_utf8_lossy(&output.stdout));
|
||||
let failed_silently = !output.status.success() && parsed.errors.is_empty();
|
||||
Ok(Outcome {
|
||||
errors: if failed_silently {
|
||||
vec![
|
||||
String::from_utf8_lossy(&output.stderr)
|
||||
.chars()
|
||||
.take(STDERR_LIMIT_CHARS)
|
||||
.collect(),
|
||||
]
|
||||
} else {
|
||||
parsed.errors
|
||||
},
|
||||
exit_code: output.status.code(),
|
||||
..parsed
|
||||
})
|
||||
}
|
||||
}
|
||||
69
litellm-rust/crates/testkit/src/target.rs
Normal file
69
litellm-rust/crates/testkit/src/target.rs
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
use target_lexicon::{Architecture, Environment, OperatingSystem, Triple};
|
||||
|
||||
use crate::Error;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Os {
|
||||
Macos,
|
||||
Linux,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Arch {
|
||||
Aarch64,
|
||||
X86_64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub struct Target {
|
||||
pub os: Os,
|
||||
pub arch: Arch,
|
||||
pub musl: bool,
|
||||
}
|
||||
|
||||
impl Target {
|
||||
pub fn host() -> Result<Self, Error> {
|
||||
Self::try_from(&Triple::host())
|
||||
}
|
||||
|
||||
pub(crate) const fn os_name(self) -> &'static str {
|
||||
match self.os {
|
||||
Os::Macos => "darwin",
|
||||
Os::Linux => "linux",
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) const fn arch_name(self) -> &'static str {
|
||||
match self.arch {
|
||||
Arch::Aarch64 => "arm64",
|
||||
Arch::X86_64 => "x64",
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) const fn musl_suffix(self) -> &'static str {
|
||||
if self.musl { "-musl" } else { "" }
|
||||
}
|
||||
}
|
||||
|
||||
impl TryFrom<&Triple> for Target {
|
||||
type Error = Error;
|
||||
|
||||
fn try_from(triple: &Triple) -> Result<Self, Error> {
|
||||
let unsupported = || Error::UnsupportedTarget(triple.to_string());
|
||||
let os = match triple.operating_system {
|
||||
OperatingSystem::Darwin(_) | OperatingSystem::MacOSX(_) => Os::Macos,
|
||||
OperatingSystem::Linux => Os::Linux,
|
||||
_ => return Err(unsupported()),
|
||||
};
|
||||
let arch = match triple.architecture {
|
||||
Architecture::Aarch64(_) => Arch::Aarch64,
|
||||
Architecture::X86_64 => Arch::X86_64,
|
||||
_ => return Err(unsupported()),
|
||||
};
|
||||
Ok(Self {
|
||||
os,
|
||||
arch,
|
||||
musl: triple.environment == Environment::Musl,
|
||||
})
|
||||
}
|
||||
}
|
||||
133
litellm-rust/crates/testkit/tests/configure.rs
Normal file
133
litellm-rust/crates/testkit/tests/configure.rs
Normal file
|
|
@ -0,0 +1,133 @@
|
|||
use std::path::Path;
|
||||
|
||||
use litellm_testkit::{ClaudeCode, Codex, Configure, Error, Opencode, Settings, Version, Wire};
|
||||
use rstest::rstest;
|
||||
|
||||
fn settings(wire: Wire) -> Settings {
|
||||
Settings {
|
||||
base_url: "http://localhost:4000/".to_owned(),
|
||||
api_key: "sk-test \"quoted\"".to_owned(),
|
||||
model: "some-model".to_owned(),
|
||||
wire,
|
||||
}
|
||||
}
|
||||
|
||||
fn version() -> Version {
|
||||
Version::new(1, 2, 3)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(&ClaudeCode, Wire::Messages)]
|
||||
#[case(&Codex, Wire::Responses)]
|
||||
#[case(&Opencode, Wire::ChatCompletions)]
|
||||
fn every_agent_runs_inside_the_given_home(#[case] agent: &impl Configure, #[case] wire: Wire) {
|
||||
let home = Path::new("/scratch/home");
|
||||
|
||||
let spec = agent.configure(&version(), &settings(wire), home).unwrap();
|
||||
|
||||
assert_eq!(spec.env["HOME"], "/scratch/home");
|
||||
assert!(
|
||||
spec.env
|
||||
.iter()
|
||||
.filter(|(key, _)| key.ends_with("_HOME") || key.as_str() == "CLAUDE_CONFIG_DIR")
|
||||
.all(|(_, value)| value.starts_with("/scratch/home"))
|
||||
);
|
||||
assert!(spec.files.keys().all(|path| path.is_relative()));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::claude_code(&ClaudeCode, &[Wire::ChatCompletions, Wire::Responses])]
|
||||
#[case::codex(&Codex, &[Wire::ChatCompletions, Wire::Messages])]
|
||||
fn wires_an_agent_cannot_speak_are_refused(
|
||||
#[case] agent: &impl Configure,
|
||||
#[case] refused: &[Wire],
|
||||
) {
|
||||
refused.iter().for_each(|wire| {
|
||||
let result = agent.configure(&version(), &settings(*wire), Path::new("/h"));
|
||||
|
||||
assert!(matches!(result, Err(Error::UnsupportedWire { wire: got, .. }) if got == *wire));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_code_points_at_the_gateway_root_with_the_key_and_model() {
|
||||
let spec = ClaudeCode
|
||||
.configure(&version(), &settings(Wire::Messages), Path::new("/h"))
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(spec.env["ANTHROPIC_BASE_URL"], "http://localhost:4000/");
|
||||
assert_eq!(spec.env["ANTHROPIC_AUTH_TOKEN"], "sk-test \"quoted\"");
|
||||
assert_eq!(spec.env["ANTHROPIC_MODEL"], "some-model");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_config_is_valid_toml_routing_the_responses_api_to_the_gateway() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let spec = Codex
|
||||
.configure(&version(), &settings(Wire::Responses), dir.path())
|
||||
.unwrap();
|
||||
spec.write_files(dir.path()).unwrap();
|
||||
|
||||
let config: toml::Table =
|
||||
toml::from_str(&std::fs::read_to_string(dir.path().join(".codex/config.toml")).unwrap())
|
||||
.unwrap();
|
||||
let provider = &config["model_providers"]["litellm"];
|
||||
|
||||
assert_eq!(config["model"].as_str(), Some("some-model"));
|
||||
assert_eq!(config["model_provider"].as_str(), Some("litellm"));
|
||||
assert_eq!(
|
||||
provider["base_url"].as_str(),
|
||||
Some("http://localhost:4000/v1")
|
||||
);
|
||||
assert_eq!(provider["wire_api"].as_str(), Some("responses"));
|
||||
let key_var = provider["env_key"].as_str().unwrap();
|
||||
assert_eq!(spec.env[key_var], "sk-test \"quoted\"");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(Wire::ChatCompletions)]
|
||||
#[case(Wire::Responses)]
|
||||
#[case(Wire::Messages)]
|
||||
fn opencode_config_is_valid_json_registering_the_gateway_model(#[case] wire: Wire) {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let spec = Opencode
|
||||
.configure(&version(), &settings(wire), dir.path())
|
||||
.unwrap();
|
||||
spec.write_files(dir.path()).unwrap();
|
||||
|
||||
let config: serde_json::Value = serde_json::from_str(
|
||||
&std::fs::read_to_string(dir.path().join(".config/opencode/opencode.json")).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
let provider = &config["provider"]["litellm"];
|
||||
|
||||
assert_eq!(config["model"], "litellm/some-model");
|
||||
assert_eq!(provider["options"]["baseURL"], "http://localhost:4000/v1");
|
||||
assert_eq!(provider["options"]["apiKey"], "sk-test \"quoted\"");
|
||||
assert!(provider["models"]["some-model"].is_object());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn opencode_uses_a_different_provider_package_for_every_wire() {
|
||||
let package = |wire| {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let spec = Opencode
|
||||
.configure(&version(), &settings(wire), dir.path())
|
||||
.unwrap();
|
||||
let config: serde_json::Value =
|
||||
serde_json::from_str(spec.files.values().next().unwrap()).unwrap();
|
||||
config["provider"]["litellm"]["npm"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.to_owned()
|
||||
};
|
||||
let packages = [Wire::ChatCompletions, Wire::Responses, Wire::Messages].map(package);
|
||||
|
||||
assert_eq!(
|
||||
packages
|
||||
.iter()
|
||||
.collect::<std::collections::BTreeSet<_>>()
|
||||
.len(),
|
||||
packages.len()
|
||||
);
|
||||
}
|
||||
262
litellm-rust/crates/testkit/tests/install.rs
Normal file
262
litellm-rust/crates/testkit/tests/install.rs
Normal file
|
|
@ -0,0 +1,262 @@
|
|||
mod support;
|
||||
|
||||
use std::str::FromStr;
|
||||
|
||||
use litellm_testkit::{ClaudeCode, Codex, Error, Installer, Opencode, Target, Version};
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
use support::{FakeFetch, script_printing, sha256, tar_gz, zip_archive};
|
||||
use target_lexicon::Triple;
|
||||
|
||||
fn target(triple: &str) -> Target {
|
||||
Target::try_from(&Triple::from_str(triple).unwrap()).unwrap()
|
||||
}
|
||||
|
||||
fn linux() -> Target {
|
||||
target("x86_64-unknown-linux-gnu")
|
||||
}
|
||||
fn version() -> Version {
|
||||
Version::new(9, 8, 7)
|
||||
}
|
||||
|
||||
fn github_release(asset: &str, download_url: &str, digest: Option<String>) -> Vec<u8> {
|
||||
json!({
|
||||
"assets": [
|
||||
{ "name": "unrelated.txt", "digest": "sha256:00", "browser_download_url": "https://example.test/unrelated" },
|
||||
{ "name": asset, "digest": digest, "browser_download_url": download_url },
|
||||
]
|
||||
})
|
||||
.to_string()
|
||||
.into_bytes()
|
||||
}
|
||||
|
||||
fn claude_routes(binary: &[u8], checksum: &str) -> Vec<(String, Vec<u8>)> {
|
||||
let base = "https://downloads.claude.ai/claude-code-releases/9.8.7";
|
||||
let manifest = json!({ "platforms": { "linux-x64": { "checksum": checksum } } });
|
||||
vec![
|
||||
(
|
||||
format!("{base}/manifest.json"),
|
||||
manifest.to_string().into_bytes(),
|
||||
),
|
||||
(format!("{base}/linux-x64/claude"), binary.to_vec()),
|
||||
]
|
||||
}
|
||||
|
||||
fn codex_routes(archive: Vec<u8>, digest: Option<String>) -> Vec<(String, Vec<u8>)> {
|
||||
let release = github_release(
|
||||
"codex-x86_64-unknown-linux-musl.tar.gz",
|
||||
"https://example.test/codex.tar.gz",
|
||||
digest,
|
||||
);
|
||||
vec![
|
||||
(
|
||||
"https://api.github.com/repos/openai/codex/releases/tags/rust-v9.8.7".to_owned(),
|
||||
release,
|
||||
),
|
||||
("https://example.test/codex.tar.gz".to_owned(), archive),
|
||||
]
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn claude_bare_binary_is_installed_and_runnable() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let installed = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&ClaudeCode, &version())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(installed.binary, cache.path().join("claude/9.8.7/claude"));
|
||||
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn codex_binary_is_extracted_from_the_tarball_under_its_own_name() {
|
||||
let binary = script_printing("codex-cli 9.8.7");
|
||||
let archive = tar_gz("codex-x86_64-unknown-linux-musl", &binary);
|
||||
let fetch = FakeFetch::new(codex_routes(
|
||||
archive.clone(),
|
||||
Some(format!("sha256:{}", sha256(&archive))),
|
||||
));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let installed = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&Codex, &version())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
|
||||
assert_eq!(installed.binary, cache.path().join("codex/9.8.7/codex"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn opencode_binary_is_extracted_from_the_darwin_zip() {
|
||||
let binary = script_printing("9.8.7");
|
||||
let archive = zip_archive("opencode", &binary);
|
||||
let release = github_release(
|
||||
"opencode-darwin-arm64.zip",
|
||||
"https://example.test/opencode.zip",
|
||||
Some(format!("sha256:{}", sha256(&archive))),
|
||||
);
|
||||
let fetch = FakeFetch::new([
|
||||
(
|
||||
"https://api.github.com/repos/sst/opencode/releases/tags/v9.8.7".to_owned(),
|
||||
release,
|
||||
),
|
||||
("https://example.test/opencode.zip".to_owned(), archive),
|
||||
]);
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let installed = Installer::new(&fetch, cache.path(), target("aarch64-apple-darwin"))
|
||||
.install(&Opencode, &version())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tampered_download_is_rejected_and_nothing_is_left_behind() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(b"what the vendor signed")));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let result = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&ClaudeCode, &version())
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::ChecksumMismatch { .. })));
|
||||
assert!(!cache.path().join("claude/9.8.7").exists());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn github_asset_without_a_digest_is_refused() {
|
||||
let archive = tar_gz(
|
||||
"codex-x86_64-unknown-linux-musl",
|
||||
&script_printing("codex-cli 9.8.7"),
|
||||
);
|
||||
let fetch = FakeFetch::new(codex_routes(archive, None));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let result = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&Codex, &version())
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::MissingChecksum(_))));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn binary_reporting_a_different_version_is_removed() {
|
||||
let binary = script_printing("1.0.0 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let result = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&ClaudeCode, &version())
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::VersionMismatch { .. })));
|
||||
assert!(!cache.path().join("claude/9.8.7/claude").exists());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn second_install_reuses_the_cached_binary_without_downloading() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
let installer = Installer::new(&fetch, cache.path(), linux());
|
||||
|
||||
let first = installer.install(&ClaudeCode, &version()).await.unwrap();
|
||||
let calls_after_first = fetch.calls();
|
||||
let second = installer.install(&ClaudeCode, &version()).await.unwrap();
|
||||
|
||||
assert_eq!(first, second);
|
||||
assert_eq!(fetch.calls(), calls_after_first);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn corrupted_cache_entry_is_replaced_by_a_fresh_download() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
let installer = Installer::new(&fetch, cache.path(), linux());
|
||||
let installed = installer.install(&ClaudeCode, &version()).await.unwrap();
|
||||
std::fs::write(&installed.binary, script_printing("0.0.1")).unwrap();
|
||||
|
||||
installer.install(&ClaudeCode, &version()).await.unwrap();
|
||||
|
||||
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case("9.8.7-beta.1")]
|
||||
#[case("9.8.7+build.5")]
|
||||
#[tokio::test]
|
||||
async fn pre_releases_never_reach_the_network_or_the_filesystem(#[case] version: &str) {
|
||||
let fetch = FakeFetch::new([]);
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let result = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&ClaudeCode, &Version::parse(version).unwrap())
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::InvalidVersion(_))));
|
||||
assert_eq!(fetch.calls(), 0);
|
||||
assert_eq!(std::fs::read_dir(cache.path()).unwrap().count(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn musl_linux_picks_the_musl_claude_build() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let base = "https://downloads.claude.ai/claude-code-releases/9.8.7";
|
||||
let manifest = json!({ "platforms": {
|
||||
"linux-x64": { "checksum": sha256(b"glibc build") },
|
||||
"linux-x64-musl": { "checksum": sha256(&binary) },
|
||||
} });
|
||||
let fetch = FakeFetch::new([
|
||||
(
|
||||
format!("{base}/manifest.json"),
|
||||
manifest.to_string().into_bytes(),
|
||||
),
|
||||
(format!("{base}/linux-x64-musl/claude"), binary.clone()),
|
||||
]);
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let installed = Installer::new(&fetch, cache.path(), target("x86_64-unknown-linux-musl"))
|
||||
.install(&ClaudeCode, &version())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case("x86_64-pc-windows-msvc")]
|
||||
#[case("riscv64gc-unknown-linux-gnu")]
|
||||
#[case("wasm32-unknown-unknown")]
|
||||
fn targets_no_agent_ships_for_are_rejected(#[case] triple: &str) {
|
||||
let result = Target::try_from(&Triple::from_str(triple).unwrap());
|
||||
|
||||
assert!(matches!(result, Err(Error::UnsupportedTarget(_))));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_installs_of_the_same_version_both_succeed() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
let installer = Installer::new(&fetch, cache.path(), linux());
|
||||
|
||||
let wanted = version();
|
||||
let installs =
|
||||
futures_util::future::join_all((0..8).map(|_| installer.install(&ClaudeCode, &wanted)))
|
||||
.await;
|
||||
|
||||
assert!(installs.iter().all(Result::is_ok));
|
||||
assert_eq!(
|
||||
std::fs::read(&installs[0].as_ref().unwrap().binary).unwrap(),
|
||||
binary
|
||||
);
|
||||
}
|
||||
133
litellm-rust/crates/testkit/tests/live.rs
Normal file
133
litellm-rust/crates/testkit/tests/live.rs
Normal file
|
|
@ -0,0 +1,133 @@
|
|||
//! Drives the real agents through a real gateway. Run with `cargo test -p litellm-testkit --test live -- --ignored`
|
||||
//! after exporting `TESTKIT_GATEWAY_URL`, `TESTKIT_GATEWAY_KEY`, one `TESTKIT_MODEL_<WIRE>` per wire
|
||||
//! (`MESSAGES`, `RESPONSES`, `CHAT_COMPLETIONS`) and one `TESTKIT_<AGENT>_VERSION` per agent
|
||||
//! (`CLAUDE`, `CODEX`, `OPENCODE`). `TESTKIT_CACHE_DIR` and `GITHUB_TOKEN` are optional.
|
||||
|
||||
use std::path::PathBuf;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_testkit::{
|
||||
Agent, ClaudeCode, Codex, HttpFetch, Installer, Opencode, Outcome, Prompt, Session, Settings,
|
||||
Target, Version, Wire,
|
||||
};
|
||||
use rstest::rstest;
|
||||
|
||||
const LIMIT: Duration = Duration::from_secs(180);
|
||||
|
||||
fn required(name: &str) -> String {
|
||||
std::env::var(name).unwrap_or_else(|_| panic!("{name} must be set to run the live tests"))
|
||||
}
|
||||
|
||||
fn model_var(wire: Wire) -> &'static str {
|
||||
match wire {
|
||||
Wire::Messages => "TESTKIT_MODEL_MESSAGES",
|
||||
Wire::Responses => "TESTKIT_MODEL_RESPONSES",
|
||||
Wire::ChatCompletions => "TESTKIT_MODEL_CHAT_COMPLETIONS",
|
||||
}
|
||||
}
|
||||
|
||||
async fn drive(
|
||||
agent: &impl Agent,
|
||||
version_var: &str,
|
||||
wire: Wire,
|
||||
model: Option<&str>,
|
||||
prompt: Prompt,
|
||||
) -> Outcome {
|
||||
let cache = std::env::var("TESTKIT_CACHE_DIR")
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|_| std::env::temp_dir().join("litellm-testkit-cache"));
|
||||
let installer = Installer::new(HttpFetch::from_env(), cache, Target::host().unwrap());
|
||||
let installed = installer
|
||||
.install(agent, &Version::parse(&required(version_var)).unwrap())
|
||||
.await
|
||||
.unwrap();
|
||||
let settings = Settings {
|
||||
base_url: required("TESTKIT_GATEWAY_URL"),
|
||||
api_key: required("TESTKIT_GATEWAY_KEY"),
|
||||
model: model.map_or_else(|| required(model_var(wire)), str::to_owned),
|
||||
wire,
|
||||
};
|
||||
let home = tempfile::tempdir().unwrap();
|
||||
let session = Session::prepare(agent, &installed, settings, home.path()).unwrap();
|
||||
session.run(agent, &prompt, LIMIT).await.unwrap()
|
||||
}
|
||||
|
||||
fn text_prompt() -> Prompt {
|
||||
Prompt {
|
||||
text: "Reply with the single word: pong".to_owned(),
|
||||
allow_tools: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn tool_prompt() -> Prompt {
|
||||
Prompt {
|
||||
text: "Run the shell command 'echo tool-ok' and reply with exactly its output.".to_owned(),
|
||||
allow_tools: true,
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
|
||||
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
|
||||
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
|
||||
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
|
||||
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
|
||||
#[ignore = "needs a live gateway, see the module docs"]
|
||||
#[tokio::test]
|
||||
async fn plain_prompt_gets_an_answer_and_token_usage(
|
||||
#[case] agent: &impl Agent,
|
||||
#[case] version_var: &str,
|
||||
#[case] wire: Wire,
|
||||
) {
|
||||
let outcome = drive(agent, version_var, wire, None, text_prompt()).await;
|
||||
|
||||
assert!(outcome.succeeded(), "{outcome:?}");
|
||||
assert!(outcome.text.to_lowercase().contains("pong"), "{outcome:?}");
|
||||
assert!(outcome.usage.output_tokens > 0, "{outcome:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
|
||||
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
|
||||
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
|
||||
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
|
||||
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
|
||||
#[ignore = "needs a live gateway, see the module docs"]
|
||||
#[tokio::test]
|
||||
async fn tool_use_is_reported_and_its_result_reaches_the_answer(
|
||||
#[case] agent: &impl Agent,
|
||||
#[case] version_var: &str,
|
||||
#[case] wire: Wire,
|
||||
) {
|
||||
let outcome = drive(agent, version_var, wire, None, tool_prompt()).await;
|
||||
|
||||
assert!(outcome.succeeded(), "{outcome:?}");
|
||||
assert!(!outcome.tool_calls.is_empty(), "{outcome:?}");
|
||||
assert!(outcome.text.contains("tool-ok"), "{outcome:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
|
||||
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
|
||||
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
|
||||
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
|
||||
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
|
||||
#[ignore = "needs a live gateway, see the module docs"]
|
||||
#[tokio::test]
|
||||
async fn model_the_gateway_rejects_is_reported_as_an_error(
|
||||
#[case] agent: &impl Agent,
|
||||
#[case] version_var: &str,
|
||||
#[case] wire: Wire,
|
||||
) {
|
||||
let outcome = drive(
|
||||
agent,
|
||||
version_var,
|
||||
wire,
|
||||
Some("testkit-no-such-model"),
|
||||
text_prompt(),
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(!outcome.succeeded(), "{outcome:?}");
|
||||
assert!(!outcome.errors.is_empty(), "{outcome:?}");
|
||||
}
|
||||
155
litellm-rust/crates/testkit/tests/session.rs
Normal file
155
litellm-rust/crates/testkit/tests/session.rs
Normal file
|
|
@ -0,0 +1,155 @@
|
|||
use std::os::unix::fs::PermissionsExt;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_testkit::{
|
||||
Configure, Drive, Error, Installed, LaunchSpec, Outcome, Prompt, Session, Settings, Version,
|
||||
Wire,
|
||||
};
|
||||
|
||||
struct Scripted;
|
||||
|
||||
impl Configure for Scripted {
|
||||
fn configure(
|
||||
&self,
|
||||
version: &Version,
|
||||
_settings: &Settings,
|
||||
home: &Path,
|
||||
) -> Result<LaunchSpec, Error> {
|
||||
Ok(LaunchSpec {
|
||||
env: [
|
||||
("AGENT_HOME".to_owned(), home.to_string_lossy().into_owned()),
|
||||
("AGENT_SAW_VERSION".to_owned(), version.to_string()),
|
||||
]
|
||||
.into(),
|
||||
files: [(
|
||||
PathBuf::from("conf/agent.toml"),
|
||||
"configured = true\n".to_owned(),
|
||||
)]
|
||||
.into(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Drive for Scripted {
|
||||
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
|
||||
vec!["--prompt".to_owned(), prompt.text.clone()]
|
||||
}
|
||||
|
||||
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
|
||||
Outcome {
|
||||
text: stdout.to_owned(),
|
||||
..Outcome::default()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn settings() -> Settings {
|
||||
Settings {
|
||||
base_url: "http://gateway.test".to_owned(),
|
||||
api_key: "sk-test".to_owned(),
|
||||
model: "some-model".to_owned(),
|
||||
wire: Wire::Messages,
|
||||
}
|
||||
}
|
||||
|
||||
fn prompt(text: &str) -> Prompt {
|
||||
Prompt {
|
||||
text: text.to_owned(),
|
||||
allow_tools: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn session(script: &str) -> (Session, tempfile::TempDir) {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let binary = dir.path().join("agent");
|
||||
std::fs::write(&binary, format!("#!/bin/sh\n{script}\n")).unwrap();
|
||||
std::fs::set_permissions(&binary, std::fs::Permissions::from_mode(0o755)).unwrap();
|
||||
let home = dir.path().join("home");
|
||||
std::fs::create_dir(&home).unwrap();
|
||||
let installed = Installed {
|
||||
version: Version::new(4, 5, 6),
|
||||
binary,
|
||||
};
|
||||
(
|
||||
Session::prepare(&Scripted, &installed, settings(), home).unwrap(),
|
||||
dir,
|
||||
)
|
||||
}
|
||||
|
||||
const LIMIT: Duration = Duration::from_secs(20);
|
||||
|
||||
#[tokio::test]
|
||||
async fn prepare_writes_the_config_files_under_home() {
|
||||
let (_session, dir) = session("true");
|
||||
|
||||
let written = std::fs::read_to_string(dir.path().join("home/conf/agent.toml")).unwrap();
|
||||
|
||||
assert_eq!(written, "configured = true\n");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn configure_and_drive_are_given_the_installed_version() {
|
||||
let (session, _dir) = session("echo \"$AGENT_SAW_VERSION\"");
|
||||
|
||||
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
|
||||
|
||||
assert_eq!(outcome.text.trim(), "4.5.6");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn agent_runs_in_home_with_only_its_own_environment() {
|
||||
let (session, dir) = session("pwd -P; env");
|
||||
|
||||
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
|
||||
|
||||
let home = dir.path().join("home").canonicalize().unwrap();
|
||||
assert_eq!(outcome.text.lines().next().unwrap(), home.to_string_lossy());
|
||||
assert!(outcome.text.contains("AGENT_HOME="));
|
||||
assert!(
|
||||
!outcome.text.contains("CARGO_"),
|
||||
"test runner environment leaked into the agent"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn prompt_reaches_the_agent_as_one_untouched_argument() {
|
||||
let (session, _dir) = session("printf '%s|' \"$@\"");
|
||||
let text = "two spaces; $(echo injected) 'quoted'";
|
||||
|
||||
let outcome = session.run(&Scripted, &prompt(text), LIMIT).await.unwrap();
|
||||
|
||||
assert_eq!(outcome.text, format!("--prompt|{text}|"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn clean_exit_is_a_success() {
|
||||
let (session, _dir) = session("echo done");
|
||||
|
||||
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
|
||||
|
||||
assert_eq!(outcome.exit_code, Some(0));
|
||||
assert!(outcome.succeeded());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn failing_exit_without_a_parsed_error_reports_stderr() {
|
||||
let (session, _dir) = session("echo boom >&2; exit 3");
|
||||
|
||||
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
|
||||
|
||||
assert_eq!(outcome.exit_code, Some(3));
|
||||
assert!(!outcome.succeeded());
|
||||
assert_eq!(outcome.errors, ["boom\n"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn agent_that_outlives_the_limit_is_stopped() {
|
||||
let (session, _dir) = session("sleep 30");
|
||||
|
||||
let result = session
|
||||
.run(&Scripted, &prompt("hi"), Duration::from_millis(200))
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::Timeout(_))));
|
||||
}
|
||||
70
litellm-rust/crates/testkit/tests/support/mod.rs
Normal file
70
litellm-rust/crates/testkit/tests/support/mod.rs
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
use std::collections::HashMap;
|
||||
use std::io::Write;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use litellm_testkit::{Error, Fetch};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
pub struct FakeFetch {
|
||||
routes: HashMap<String, Vec<u8>>,
|
||||
calls: AtomicUsize,
|
||||
}
|
||||
|
||||
impl FakeFetch {
|
||||
pub fn new(routes: impl IntoIterator<Item = (String, Vec<u8>)>) -> Self {
|
||||
Self {
|
||||
routes: routes.into_iter().collect(),
|
||||
calls: AtomicUsize::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn calls(&self) -> usize {
|
||||
self.calls.load(Ordering::SeqCst)
|
||||
}
|
||||
}
|
||||
|
||||
impl Fetch for FakeFetch {
|
||||
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
|
||||
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
self.routes.get(url).cloned().ok_or_else(|| Error::Status {
|
||||
url: url.to_owned(),
|
||||
status: 404,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Fetch for &FakeFetch {
|
||||
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
|
||||
(*self).get(url).await
|
||||
}
|
||||
}
|
||||
|
||||
pub fn sha256(bytes: &[u8]) -> String {
|
||||
format!("{:x}", Sha256::digest(bytes))
|
||||
}
|
||||
|
||||
pub fn script_printing(output: &str) -> Vec<u8> {
|
||||
format!("#!/bin/sh\necho '{output}'\n").into_bytes()
|
||||
}
|
||||
|
||||
pub fn tar_gz(member: &str, contents: &[u8]) -> Vec<u8> {
|
||||
let mut builder = tar::Builder::new(Vec::new());
|
||||
let mut header = tar::Header::new_gnu();
|
||||
header.set_size(contents.len() as u64);
|
||||
header.set_mode(0o755);
|
||||
header.set_cksum();
|
||||
builder.append_data(&mut header, member, contents).unwrap();
|
||||
let tarball = builder.into_inner().unwrap();
|
||||
let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
|
||||
encoder.write_all(&tarball).unwrap();
|
||||
encoder.finish().unwrap()
|
||||
}
|
||||
|
||||
pub fn zip_archive(member: &str, contents: &[u8]) -> Vec<u8> {
|
||||
let mut writer = zip::ZipWriter::new(std::io::Cursor::new(Vec::new()));
|
||||
writer
|
||||
.start_file(member, zip::write::SimpleFileOptions::default())
|
||||
.unwrap();
|
||||
writer.write_all(contents).unwrap();
|
||||
writer.finish().unwrap().into_inner()
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from litellm.litellm_core_utils.core_helpers import process_response_headers
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.llms.anthropic.common_utils import ANTHROPIC_ERROR_STATUS_CODE_MAP
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import INCOMPLETE_STREAM_ERROR_MESSAGE
|
||||
from litellm.proxy.pass_through_endpoints.success_handler import (
|
||||
PassThroughEndpointLogging,
|
||||
)
|
||||
|
|
@ -28,11 +29,6 @@ GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ: Final = PassThroughEndpointLogging()
|
|||
_UPSTREAM_PUMP_TASKS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: stdlib strong-ref set for pump tasks
|
||||
_DETACHED_STREAM_DRAINS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: bounded strong-ref set, detached drains
|
||||
|
||||
INCOMPLETE_STREAM_ERROR_MESSAGE: Final = (
|
||||
"Provider stream ended before emitting a message_stop event; "
|
||||
"the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated."
|
||||
)
|
||||
|
||||
|
||||
def _is_message_stop_chunk(chunk: object) -> bool:
|
||||
if isinstance(chunk, dict):
|
||||
|
|
|
|||
|
|
@ -15,6 +15,12 @@ if TYPE_CHECKING:
|
|||
from litellm.exceptions import ContentPolicyViolationError
|
||||
|
||||
|
||||
INCOMPLETE_STREAM_ERROR_MESSAGE: Final = (
|
||||
"Provider stream ended before emitting a message_stop event; "
|
||||
"the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated."
|
||||
)
|
||||
|
||||
|
||||
def get_safeguard_refusal_stop_details(response: object) -> Mapping[str, Any] | None:
|
||||
"""
|
||||
Return the ``stop_details`` of an Anthropic Messages response refused by a
|
||||
|
|
|
|||
|
|
@ -2,20 +2,25 @@
|
|||
## Translates OpenAI call to Anthropic `/v1/messages` format
|
||||
import asyncio
|
||||
import json
|
||||
import traceback
|
||||
from collections import deque
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, field_validator
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm._logging import redact_internal_details_from_client_message
|
||||
from litellm._uuid import uuid
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
encrypted_reasoning_signature,
|
||||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
|
||||
INCOMPLETE_STREAM_ERROR_MESSAGE,
|
||||
refusal_stop_details,
|
||||
responses_output_refusal_text,
|
||||
)
|
||||
from litellm.responses.streaming_iterator import stream_error_status_and_message
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage
|
||||
|
||||
from .transformation import (
|
||||
|
|
@ -27,6 +32,72 @@ if TYPE_CHECKING:
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
|
||||
|
||||
|
||||
class _UpstreamFailure(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
status_code: int | None = None
|
||||
message: str | None = None
|
||||
|
||||
@field_validator("status_code", mode="before")
|
||||
@classmethod
|
||||
def http_error_status_or_none(cls, value: object) -> int | None:
|
||||
candidate: Final = (
|
||||
value
|
||||
if isinstance(value, int) and not isinstance(value, bool)
|
||||
else int(value)
|
||||
if isinstance(value, str) and value.isdecimal()
|
||||
else None
|
||||
)
|
||||
return candidate if candidate is not None and 400 <= candidate <= 599 else None
|
||||
|
||||
@field_validator("message", mode="before")
|
||||
@classmethod
|
||||
def str_or_none(cls, value: object) -> str | None:
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
class _FailedResponse(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, from_attributes=True)
|
||||
|
||||
error: object | None = None
|
||||
|
||||
|
||||
class _FailedResponseEvent(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, from_attributes=True)
|
||||
|
||||
response: _FailedResponse | None = None
|
||||
|
||||
|
||||
def _original_failure(exception: Exception) -> Exception:
|
||||
failure = exception # rebind-ok: walks the MidStreamFallbackError chain down to the provider failure
|
||||
while isinstance(failure, MidStreamFallbackError) and failure.original_exception is not None:
|
||||
failure = failure.original_exception
|
||||
return failure
|
||||
|
||||
|
||||
def _failure_status_and_message(exception: Exception) -> tuple[int, str]:
|
||||
original: Final = _original_failure(exception)
|
||||
failure: Final = _UpstreamFailure.model_validate(
|
||||
{"status_code": getattr(original, "status_code", None), "message": getattr(original, "message", None)}
|
||||
)
|
||||
status_code: Final = failure.status_code if failure.status_code is not None else 500
|
||||
message: Final = failure.message or str(original) or INCOMPLETE_STREAM_ERROR_MESSAGE
|
||||
return status_code, message
|
||||
|
||||
|
||||
def _anthropic_error_chunk(status_code: int, message: str) -> dict[str, object]:
|
||||
from litellm.anthropic_interface.exceptions.exception_mapping_utils import (
|
||||
AnthropicExceptionMapping,
|
||||
)
|
||||
|
||||
return dict(
|
||||
AnthropicExceptionMapping.transform_to_anthropic_error(
|
||||
status_code=status_code,
|
||||
raw_message=redact_internal_details_from_client_message(message),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class AnthropicResponsesStreamWrapper:
|
||||
"""
|
||||
Wraps a Responses API streaming iterator and re-emits events in Anthropic SSE format.
|
||||
|
|
@ -40,6 +111,7 @@ class AnthropicResponsesStreamWrapper:
|
|||
response.function_call_arguments.delta -> content_block_delta (input_json_delta)
|
||||
response.output_item.done -> content_block_delta (signature_delta) + content_block_stop
|
||||
response.completed -> message_delta + message_stop
|
||||
response.failed -> error (the stream ends without message_stop)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
|
|
@ -60,6 +132,7 @@ class AnthropicResponsesStreamWrapper:
|
|||
self._pending_tool_ids: dict[str, str] = {} # item_id -> call_id / name accumulator
|
||||
self._sent_message_start = False
|
||||
self._sent_message_stop = False
|
||||
self._stream_failed = False
|
||||
self._chunk_queue: deque[dict[str, object]] = deque()
|
||||
self._refusal_text: str = ""
|
||||
self._sync_responses_iterator: Iterator[object] | None = None
|
||||
|
|
@ -293,10 +366,23 @@ class AnthropicResponsesStreamWrapper:
|
|||
)
|
||||
return
|
||||
|
||||
if event_type == "response.failed":
|
||||
failed: Final = _FailedResponseEvent.model_validate(event)
|
||||
status_code, message = stream_error_status_and_message(
|
||||
failed.response.error if failed.response is not None else None
|
||||
)
|
||||
verbose_logger.error(
|
||||
"AnthropicResponsesStreamWrapper: upstream Responses stream for %s failed (%s): %s",
|
||||
self.model,
|
||||
status_code,
|
||||
message,
|
||||
)
|
||||
self._fail_stream(status_code, message)
|
||||
return
|
||||
|
||||
# ---- response completed -> message_delta + message_stop ----
|
||||
if event_type in (
|
||||
"response.completed",
|
||||
"response.failed",
|
||||
"response.incomplete",
|
||||
):
|
||||
response_obj: Final = getattr(event, "response", None) or (
|
||||
|
|
@ -350,21 +436,24 @@ class AnthropicResponsesStreamWrapper:
|
|||
self._sent_message_stop = True
|
||||
return
|
||||
|
||||
def _fail_stream(self, status_code: int, message: str) -> None:
|
||||
self._stream_failed = True
|
||||
self._chunk_queue.append(_anthropic_error_chunk(status_code, message))
|
||||
|
||||
def __aiter__(self) -> "AnthropicResponsesStreamWrapper":
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> dict[str, object]:
|
||||
# Return any queued chunks first
|
||||
if self._chunk_queue:
|
||||
return self._chunk_queue.popleft()
|
||||
if self._stream_failed:
|
||||
raise StopAsyncIteration
|
||||
|
||||
# Emit message_start if not yet done (fallback if response.created wasn't fired)
|
||||
if not self._sent_message_start:
|
||||
self._sent_message_start = True
|
||||
self._chunk_queue.append(self._make_message_start())
|
||||
return self._chunk_queue.popleft()
|
||||
|
||||
# Consume the upstream stream
|
||||
try:
|
||||
if hasattr(self.responses_stream, "__aiter__"):
|
||||
async for event in self.responses_stream:
|
||||
|
|
@ -382,10 +471,19 @@ class AnthropicResponsesStreamWrapper:
|
|||
return self._chunk_queue.popleft()
|
||||
except StopAsyncIteration:
|
||||
pass
|
||||
except Exception as e:
|
||||
verbose_logger.error("AnthropicResponsesStreamWrapper error: %s\n%s", e, traceback.format_exc())
|
||||
except Exception as e: # noqa: BLE001 # every upstream failure becomes a client error event
|
||||
verbose_logger.exception(
|
||||
"AnthropicResponsesStreamWrapper: upstream Responses stream for %s failed", self.model
|
||||
)
|
||||
self._fail_stream(*_failure_status_and_message(e))
|
||||
|
||||
if not self._chunk_queue and not self._sent_message_stop and not self._stream_failed:
|
||||
verbose_logger.error(
|
||||
"AnthropicResponsesStreamWrapper: upstream Responses stream for %s ended without a terminal event",
|
||||
self.model,
|
||||
)
|
||||
self._fail_stream(500, INCOMPLETE_STREAM_ERROR_MESSAGE)
|
||||
|
||||
# Drain any remaining queued chunks
|
||||
if self._chunk_queue:
|
||||
return self._chunk_queue.popleft()
|
||||
|
||||
|
|
|
|||
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.",
|
||||
|
|
|
|||
|
|
@ -2781,15 +2781,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 +3002,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,
|
||||
|
|
|
|||
|
|
@ -4618,6 +4618,13 @@ class GoogleSSOHandler:
|
|||
return result or {}
|
||||
|
||||
|
||||
def _raise_if_sso_debug_disabled() -> None:
|
||||
"""The debug routes run the browser-redirect SSO flow, so they cannot carry a
|
||||
bearer credential; an explicit opt-in flag is the only way to gate them."""
|
||||
if get_secret_bool("ENABLE_SSO_DEBUG") is not True:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Not Found")
|
||||
|
||||
|
||||
@router.get("/sso/debug/login", tags=["experimental"], include_in_schema=False)
|
||||
async def debug_sso_login(request: Request):
|
||||
"""
|
||||
|
|
@ -4625,6 +4632,8 @@ async def debug_sso_login(request: Request):
|
|||
PROXY_BASE_URL should be the your deployed proxy endpoint, e.g. PROXY_BASE_URL="https://litellm-production-7002.up.railway.app/"
|
||||
Example:
|
||||
"""
|
||||
_raise_if_sso_debug_disabled()
|
||||
|
||||
from litellm.proxy.proxy_server import premium_user
|
||||
|
||||
microsoft_client_id: Final = os.getenv("MICROSOFT_CLIENT_ID", None)
|
||||
|
|
@ -4670,6 +4679,8 @@ async def debug_sso_callback(request: Request):
|
|||
"""
|
||||
Returns the OpenID object returned by the SSO provider
|
||||
"""
|
||||
_raise_if_sso_debug_disabled()
|
||||
|
||||
import json
|
||||
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
|
|
|||
|
|
@ -6342,6 +6342,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:
|
||||
|
|
|
|||
|
|
@ -230,6 +230,11 @@ def _status_code_for_error_fields(error_type: str | None, error_code: str | None
|
|||
return next((status for status in map(_status_code_for_error_field, fields) if status is not None), 500)
|
||||
|
||||
|
||||
def stream_error_status_and_message(error_obj: object) -> tuple[int, str]:
|
||||
message, error_type, error_code = _error_event_fields(error_obj)
|
||||
return _status_code_for_error_fields(error_type, error_code), message
|
||||
|
||||
|
||||
def _map_stream_error_to_exception(error_obj: object, model: str, custom_llm_provider: str) -> Exception:
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -52,7 +52,7 @@ from tests._vcr_redis_persister import (
|
|||
# network call entirely, so skip tests record nothing (NOOP) and passing tests
|
||||
# stop carrying a volatile github episode. This matches the established idiom in
|
||||
# the unit-test suite, which sets the same flag (see e.g.
|
||||
# tests/test_litellm/test_cost_calculator.py). ``setdefault`` so an explicit
|
||||
# tests/unit/test_cost_calculator.py). ``setdefault`` so an explicit
|
||||
# override still wins.
|
||||
os.environ.setdefault("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
|
||||
|
|
|
|||
|
|
@ -13,15 +13,16 @@ def check_for_litellm_module_deletion(base_dir):
|
|||
del sys.modules[module]
|
||||
"""
|
||||
problematic_files = []
|
||||
test_dir = os.path.join(base_dir, "test_litellm")
|
||||
candidate_dirs = [os.path.join(base_dir, name) for name in ("test_litellm", "unit")]
|
||||
test_dirs = [test_dir for test_dir in candidate_dirs if os.path.exists(test_dir)]
|
||||
|
||||
if not os.path.exists(test_dir):
|
||||
print(f"Warning: Directory {test_dir} does not exist.")
|
||||
if not test_dirs:
|
||||
print(f"Warning: None of {candidate_dirs} exist.")
|
||||
return []
|
||||
|
||||
print(f"Checking directory: {test_dir}")
|
||||
print(f"Checking directories: {test_dirs}")
|
||||
|
||||
for root, _, files in os.walk(test_dir):
|
||||
for root, _, files in (entry for test_dir in test_dirs for entry in os.walk(test_dir)):
|
||||
for file in files:
|
||||
if file.endswith(".py"):
|
||||
file_path = os.path.join(root, file)
|
||||
|
|
@ -173,7 +174,7 @@ def main():
|
|||
f"This can cause import issues and test failures. Files: {problematic_files}"
|
||||
)
|
||||
else:
|
||||
print("✓ No litellm module deletion patterns found in test_litellm directory.")
|
||||
print("✓ No litellm module deletion patterns found in tests/test_litellm or tests/unit.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ def get_all_functions_called_in_tests(base_dir):
|
|||
specifically in files containing the word 'router'.
|
||||
"""
|
||||
called_functions = set()
|
||||
test_dirs = ["local_testing", "router_unit_tests", "test_litellm"]
|
||||
test_dirs = ["local_testing", "router_unit_tests", "test_litellm", "unit"]
|
||||
|
||||
for test_dir in test_dirs:
|
||||
dir_path = os.path.join(base_dir, test_dir)
|
||||
|
|
|
|||
|
|
@ -85,6 +85,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"}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -22,7 +22,7 @@ class TestBedrockGPTOSS(BaseLLMChatTest):
|
|||
"""Bedrock GPT-OSS intermittently emits truncated toolUse.input deltas on
|
||||
the live endpoint, which makes the inherited live integration test flaky.
|
||||
The accumulation side is covered deterministically by
|
||||
tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py::test_transform_tool_calls_index;
|
||||
tests/unit/llms/bedrock/chat/test_invoke_handler.py::test_transform_tool_calls_index;
|
||||
the GPT-OSS-specific request-body transformation is covered by
|
||||
test_function_calling_request_body_gpt_oss below.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -277,4 +277,4 @@ class BaseSkillsAPITest(ABC):
|
|||
#
|
||||
# Transformation logic (URL construction, headers, request/response parsing) is
|
||||
# covered by unit tests in:
|
||||
# tests/test_litellm/test_anthropic_skills_transformation.py
|
||||
# tests/unit/test_anthropic_skills_transformation.py
|
||||
|
|
|
|||
|
|
@ -324,7 +324,7 @@ def test_parallel_function_call_anthropic_error_msg(model, messages):
|
|||
Anthropic (and Bedrock Invoke via ``AnthropicConfig.transform_request``)
|
||||
inject a dummy tool so CLIs work with ``modify_params`` left off. Bedrock
|
||||
Converse's no-raise behavior is covered offline in
|
||||
``tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py``
|
||||
``tests/unit/llms/bedrock/chat/test_converse_transformation.py``
|
||||
(see #24158, #27138), which needs no live credentials.
|
||||
"""
|
||||
# Force modify_params off as a clean baseline: it exercises the Anthropic
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ body can still arrive, released once the caller is done with the response.
|
|||
|
||||
Nothing here re-tests the shapes ``_handler_may_close_client`` covers -- a
|
||||
borrowed ``handler.client``, a caller-supplied client, an evicted-but-held
|
||||
client. Those are pinned in ``tests/test_litellm/llms/custom_httpx/
|
||||
client. Those are pinned in ``tests/unit/llms/custom_httpx/
|
||||
test_http_handler.py``. What is uncovered there is the in-flight response, so no
|
||||
test here may keep the client in a local: that inflates the very refcount under
|
||||
test, and the test then passes on a broken handler. They hold weak references
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Integration tests for SageMaker Nova provider.
|
|||
These tests require a live SageMaker Nova endpoint and AWS credentials.
|
||||
They are skipped by default — run manually with:
|
||||
|
||||
pytest tests/test_litellm/llms/sagemaker/test_sagemaker_nova_integration.py -v --no-header -rN
|
||||
pytest tests/local_testing/test_sagemaker_nova_integration.py -v --no-header -rN
|
||||
|
||||
Prerequisites:
|
||||
export AWS_PROFILE=<your-profile> # or set AWS_ACCESS_KEY_ID / AWS_SECRET_ACCESS_KEY
|
||||
|
|
@ -251,7 +251,7 @@ class TestSagemakerNova2LiteIntegration:
|
|||
|
||||
Run with:
|
||||
export SAGEMAKER_NOVA2_LITE_ENDPOINT=<your-nova-2-lite-endpoint>
|
||||
pytest tests/test_litellm/llms/sagemaker/test_sagemaker_nova_integration.py::TestSagemakerNova2LiteIntegration -v
|
||||
pytest tests/local_testing/test_sagemaker_nova_integration.py::TestSagemakerNova2LiteIntegration -v
|
||||
"""
|
||||
|
||||
def test_should_accept_reasoning_effort_low(self):
|
||||
|
|
|
|||
|
|
@ -85,7 +85,7 @@ class TestBingGroundingSearch(BaseSearchTest):
|
|||
class TestBingGroundingSearchTransformation:
|
||||
"""
|
||||
Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked.
|
||||
Transformation details are unit-tested in tests/test_litellm/llms/azure/search/.
|
||||
Transformation details are unit-tested in tests/unit/llms/azure/search/.
|
||||
"""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
|
|
|
|||
|
|
@ -58,7 +58,7 @@ class TestNimbleSearch(BaseSearchTest):
|
|||
class TestNimbleSearchTransformation:
|
||||
"""
|
||||
Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked.
|
||||
Transformation details are unit-tested in tests/test_litellm/llms/nimble/search/.
|
||||
Transformation details are unit-tested in tests/unit/llms/nimble/search/.
|
||||
"""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
|
|
|
|||
|
|
@ -1,387 +0,0 @@
|
|||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
import litellm.batches.batch_utils as bu
|
||||
from litellm.types.llms.openai import Batch
|
||||
|
||||
GROUNDED_USAGE_METADATA = {
|
||||
"promptTokenCount": 19,
|
||||
"candidatesTokenCount": 59,
|
||||
"thoughtsTokenCount": 406,
|
||||
"toolUsePromptTokenCount": 73,
|
||||
"totalTokenCount": 557,
|
||||
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 19}],
|
||||
"candidatesTokensDetails": [{"modality": "TEXT", "tokenCount": 59}],
|
||||
"toolUsePromptTokensDetails": [{"modality": "TEXT", "tokenCount": 73}],
|
||||
"trafficType": "ON_DEMAND",
|
||||
}
|
||||
PASSTHROUGH_OUTPUT_URI = (
|
||||
"gs://litellm-bucket/litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/u/"
|
||||
"predictions.jsonl"
|
||||
)
|
||||
UNGROUNDED_USAGE_METADATA = {
|
||||
"promptTokenCount": 20,
|
||||
"candidatesTokenCount": 48,
|
||||
"thoughtsTokenCount": 195,
|
||||
"toolUsePromptTokenCount": 73,
|
||||
"totalTokenCount": 336,
|
||||
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 20}],
|
||||
"trafficType": "ON_DEMAND",
|
||||
}
|
||||
|
||||
|
||||
def _batch(output_file_id: str) -> Batch:
|
||||
return Batch(
|
||||
id="b",
|
||||
completion_window="24h",
|
||||
created_at=1,
|
||||
endpoint="/v1/chat/completions",
|
||||
input_file_id="f",
|
||||
object="batch",
|
||||
status="completed",
|
||||
output_file_id=output_file_id,
|
||||
)
|
||||
|
||||
|
||||
def _vertex_jsonl(rows: list[dict]) -> bytes:
|
||||
return "\n".join(json.dumps(row) for row in rows).encode()
|
||||
|
||||
|
||||
def _vertex_openai_row(custom_id: str, model: str, prompt_tokens: int, completion_tokens: int) -> dict:
|
||||
return {
|
||||
"id": f"batch_req_{custom_id}",
|
||||
"custom_id": custom_id,
|
||||
"response": {
|
||||
"status_code": 200,
|
||||
"request_id": custom_id,
|
||||
"body": {
|
||||
"id": f"chatcmpl-{custom_id}",
|
||||
"object": "chat.completion",
|
||||
"model": model,
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
|
||||
"usage": {
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"completion_tokens": completion_tokens,
|
||||
"total_tokens": prompt_tokens + completion_tokens,
|
||||
},
|
||||
},
|
||||
},
|
||||
"error": None,
|
||||
}
|
||||
|
||||
|
||||
def _native_vertex_row(usage_metadata: dict, *, grounded: bool, model_version: str | None = "gemini-2.5-flash"):
|
||||
candidate = {"content": {"role": "model", "parts": [{"text": "ok"}]}, "finishReason": "STOP"}
|
||||
grounding = {"groundingMetadata": {"webSearchQueries": ["q"]}} if grounded else {}
|
||||
response = {"candidates": [{**candidate, **grounding}], "usageMetadata": usage_metadata}
|
||||
return {
|
||||
"request": {"contents": [{"role": "user", "parts": [{"text": "q"}]}], "tools": [{"googleSearch": {}}]},
|
||||
"status": "",
|
||||
"response": {**response, **({"modelVersion": model_version} if model_version else {})},
|
||||
"processed_time": "2026-09-23T19:02:00.000+00:00",
|
||||
}
|
||||
|
||||
|
||||
def _capture_cost_calls(monkeypatch, prompt_cost=0.5, completion_cost=0.25) -> list:
|
||||
import litellm.cost_calculator as cc
|
||||
|
||||
calls: list = []
|
||||
|
||||
def _calc(**kw):
|
||||
calls.append(kw)
|
||||
return (prompt_cost, completion_cost)
|
||||
|
||||
monkeypatch.setattr(cc, "batch_cost_calculator", _calc)
|
||||
return calls
|
||||
|
||||
|
||||
def test_vertex_native_cost_bills_embedding_rows(monkeypatch):
|
||||
monkeypatch.setitem(litellm.model_cost, "vertex_ai/gemini-embedding-2", {"input_cost_per_token_batches": 1e-7})
|
||||
rows = [
|
||||
{
|
||||
"key": "id_1",
|
||||
"status": "",
|
||||
"request": {"content": {"parts": [{"text": "hello world"}]}},
|
||||
"response": {"embedding": {"values": [0.1, 0.2]}, "usageMetadata": {"promptTokenCount": 2}},
|
||||
},
|
||||
{
|
||||
"key": "id_2",
|
||||
"status": "",
|
||||
"request": {"content": {"parts": [{"text": "hello"}]}},
|
||||
"response": {"embedding": {"values": [0.3]}, "tokenCount": "3"},
|
||||
},
|
||||
{"key": "id_3", "status": "INVALID_ARGUMENT", "request": {"content": {"parts": [{"text": ""}]}}},
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-embedding-2")
|
||||
|
||||
assert (result.successful_requests, result.failed_requests) == (2, 1)
|
||||
assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (5, 0, 5)
|
||||
assert result.cost == pytest.approx(5 * 1e-7)
|
||||
assert result.models == ["gemini-embedding-2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_vertex_rows_route_to_vertex_cost_path_without_flag(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
|
||||
monkeypatch.setattr(
|
||||
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
|
||||
)
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
|
||||
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False),
|
||||
]
|
||||
|
||||
result = await bu.calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
|
||||
)
|
||||
|
||||
assert result.cost == pytest.approx(1.5)
|
||||
assert (result.successful_requests, result.failed_requests) == (2, 0)
|
||||
assert result.models == ["gemini-2.5-flash"]
|
||||
assert {(call["model"], call["custom_llm_provider"]) for call in calls} == {("gemini-2.5-flash", "vertex_ai")}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_shaped_vertex_rows_keep_the_generic_path_without_flag(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
|
||||
monkeypatch.setattr(
|
||||
bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run")
|
||||
)
|
||||
_capture_cost_calls(monkeypatch)
|
||||
rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)]
|
||||
|
||||
result = await bu.calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
|
||||
)
|
||||
|
||||
assert result.successful_requests == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_vertex_rows_on_another_provider_keep_the_generic_path(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run")
|
||||
)
|
||||
_capture_cost_calls(monkeypatch)
|
||||
|
||||
result = await bu.calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
assert result.successful_requests == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_completed_batch_routes_native_rows_without_flag(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
|
||||
raw_rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)]
|
||||
|
||||
async def fake_fetch(batch, custom_llm_provider, litellm_params=None):
|
||||
return _vertex_jsonl(raw_rows)
|
||||
|
||||
monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch)
|
||||
monkeypatch.setattr(
|
||||
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
|
||||
)
|
||||
calls = _capture_cost_calls(monkeypatch, prompt_cost=0.7, completion_cost=0.3)
|
||||
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
|
||||
|
||||
result = await bu._handle_completed_batch(
|
||||
_batch(PASSTHROUGH_OUTPUT_URI),
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_name="gemini-2.5-flash",
|
||||
model_info=deployment_model_info,
|
||||
)
|
||||
|
||||
assert result.cost == pytest.approx(1.0)
|
||||
assert result.usage.total_tokens == 557
|
||||
assert [call["model_info"] for call in calls] == [deployment_model_info]
|
||||
|
||||
|
||||
def test_native_vertex_usage_is_billed_like_the_online_path(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
grounded = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)
|
||||
ungrounded = _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False)
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage([grounded, ungrounded], "gemini-2.5-flash")
|
||||
|
||||
grounded_usage, ungrounded_usage = (call["usage"] for call in calls)
|
||||
assert grounded_usage.prompt_tokens == 19
|
||||
assert grounded_usage.completion_tokens == 59 + 406
|
||||
assert grounded_usage.completion_tokens_details.reasoning_tokens == 406
|
||||
assert ungrounded_usage.prompt_tokens == 20 + 73
|
||||
assert ungrounded_usage.completion_tokens == 48 + 195
|
||||
assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (
|
||||
19 + 93,
|
||||
465 + 243,
|
||||
557 + 336,
|
||||
)
|
||||
|
||||
|
||||
def test_native_vertex_rows_are_priced_by_model_version_without_a_model_name(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"),
|
||||
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-pro"),
|
||||
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None),
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows)
|
||||
|
||||
assert [call["model"] for call in calls] == ["gemini-2.5-flash", "gemini-2.5-pro"]
|
||||
assert result.models == ["gemini-2.5-flash", "gemini-2.5-pro"]
|
||||
assert result.cost == pytest.approx(1.5)
|
||||
assert result.successful_requests == 3
|
||||
assert result.usage.total_tokens == 557 + 336 + 336
|
||||
|
||||
|
||||
def test_native_vertex_rows_without_usage_metadata_count_as_failed(monkeypatch):
|
||||
_capture_cost_calls(monkeypatch)
|
||||
rows = [
|
||||
{"request": {"contents": []}, "status": "Error: bad request", "processed_time": "t"},
|
||||
{"request": {"contents": []}, "response": {"candidates": []}},
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
|
||||
|
||||
assert (result.successful_requests, result.failed_requests) == (1, 2)
|
||||
assert result.usage.total_tokens == 557
|
||||
|
||||
|
||||
def test_native_vertex_batch_whose_rows_all_failed_still_names_the_deployment_model(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [{"request": {"contents": []}, "status": "Error: quota exceeded", "processed_time": "t"}] * 2
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
|
||||
|
||||
assert result.models == ["gemini-2.5-flash"]
|
||||
assert (result.successful_requests, result.failed_requests, result.cost) == (0, 2, 0.0)
|
||||
assert calls == []
|
||||
|
||||
|
||||
def test_native_vertex_rows_are_priced_with_the_deployment_model_info(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
|
||||
|
||||
bu.calculate_vertex_ai_batch_cost_and_usage(
|
||||
[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
|
||||
"gemini-2.5-flash",
|
||||
model_info=deployment_model_info,
|
||||
)
|
||||
|
||||
assert [call["model_info"] for call in calls] == [deployment_model_info]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_vertex_rows_keep_the_deployment_model_info_through_the_batch_entrypoint(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
deployment_model_info = {"input_cost_per_token_batches": 1e-6}
|
||||
|
||||
await bu.calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_name="gemini-2.5-flash",
|
||||
model_info=deployment_model_info,
|
||||
)
|
||||
|
||||
assert [call["model_info"] for call in calls] == [deployment_model_info]
|
||||
|
||||
|
||||
def test_native_vertex_rows_are_priced_by_the_deployment_model_over_model_version(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-pro")]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
|
||||
|
||||
assert [call["model"] for call in calls] == ["gemini-2.5-flash"]
|
||||
assert result.models == ["gemini-2.5-flash"]
|
||||
|
||||
|
||||
def test_native_vertex_rows_that_fail_response_validation_count_as_failed(monkeypatch):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [
|
||||
{"request": {"contents": []}, "response": {"candidates": "nope", "usageMetadata": GROUNDED_USAGE_METADATA}},
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
|
||||
|
||||
assert (result.successful_requests, result.failed_requests) == (1, 1)
|
||||
assert result.usage.total_tokens == 557
|
||||
assert len(calls) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("wildcard_model", ["*", "vertex_ai/*"])
|
||||
def test_native_vertex_rows_under_a_wildcard_deployment_are_priced_by_model_version(monkeypatch, wildcard_model):
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"),
|
||||
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None),
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, wildcard_model)
|
||||
|
||||
assert [call["model"] for call in calls] == ["gemini-2.5-flash", wildcard_model]
|
||||
assert result.cost == pytest.approx(1.5)
|
||||
assert (result.successful_requests, result.failed_requests) == (2, 0)
|
||||
assert result.usage.total_tokens == 557 + 336
|
||||
|
||||
|
||||
def test_native_vertex_row_without_model_version_under_a_wildcard_deployment_bills_its_explicit_prices():
|
||||
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
|
||||
with_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash")
|
||||
without_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version=None)
|
||||
|
||||
twin = bu.calculate_vertex_ai_batch_cost_and_usage([with_version], "vertex_ai/*", model_info=deployment_model_info)
|
||||
both = bu.calculate_vertex_ai_batch_cost_and_usage(
|
||||
[with_version, without_version], "vertex_ai/*", model_info=deployment_model_info
|
||||
)
|
||||
|
||||
assert twin.cost > 0
|
||||
assert both.cost == pytest.approx(2 * twin.cost)
|
||||
assert (both.successful_requests, both.failed_requests) == (2, 0)
|
||||
|
||||
|
||||
def test_native_vertex_row_the_cost_map_cannot_price_is_billed_at_zero_and_the_rest_still_bills(monkeypatch):
|
||||
import litellm.cost_calculator as cc
|
||||
|
||||
def _calc(**kw):
|
||||
if kw["model"] == "gemini-unpriced":
|
||||
raise ValueError("no pricing")
|
||||
return (0.5, 0.25)
|
||||
|
||||
monkeypatch.setattr(cc, "batch_cost_calculator", _calc)
|
||||
rows = [
|
||||
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-unpriced"),
|
||||
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-flash"),
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows)
|
||||
|
||||
assert result.cost == pytest.approx(0.75)
|
||||
assert (result.successful_requests, result.failed_requests) == (2, 0)
|
||||
assert result.usage.total_tokens == 557 + 336
|
||||
assert result.models == ["gemini-unpriced", "gemini-2.5-flash"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_sends_every_vertex_row_down_the_native_path_when_a_model_is_known(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False)
|
||||
monkeypatch.setattr(
|
||||
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
|
||||
)
|
||||
calls = _capture_cost_calls(monkeypatch)
|
||||
rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)]
|
||||
|
||||
result = await bu.calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
|
||||
)
|
||||
|
||||
assert calls == []
|
||||
assert (result.successful_requests, result.failed_requests) == (0, 1)
|
||||
|
|
@ -1,117 +0,0 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.chat_completions import dispatch
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.catalog import Route, RouteRule, Rules
|
||||
from litellm.rust_bridge.chat_completions.entrypoints import (
|
||||
LiteLLMChatCompletionsRequest,
|
||||
NativeAcompletion,
|
||||
NativeCompletion,
|
||||
)
|
||||
from litellm.rust_bridge.configuration import Rollout
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
MESSAGES: Final = [{"role": "user", "content": "hi"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_public_completion_calls_keep_the_python_result() -> None:
|
||||
sync_response: Final = litellm.completion(model="openai/test-model", messages=MESSAGES, mock_response="ok")
|
||||
async_response: Final = await litellm.acompletion(model="openai/test-model", messages=MESSAGES, mock_response="ok")
|
||||
|
||||
assert isinstance(sync_response, ModelResponse)
|
||||
assert isinstance(async_response, ModelResponse)
|
||||
assert sync_response.choices[0].message.content == "ok"
|
||||
assert async_response.choices[0].message.content == "ok"
|
||||
|
||||
|
||||
def test_sync_completion_request_projects_public_arguments() -> None:
|
||||
rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),)
|
||||
expected: Final = ModelResponse()
|
||||
|
||||
def native(
|
||||
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
|
||||
) -> ModelResponse:
|
||||
assert request.model == "test-model"
|
||||
assert request.messages == MESSAGES
|
||||
assert request.custom_llm_provider == "openai"
|
||||
assert request.stream is True
|
||||
return expected
|
||||
|
||||
binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None)
|
||||
binding.override(native)
|
||||
response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
|
||||
("test-model", MESSAGES),
|
||||
{"custom_llm_provider": "openai", "stream": True},
|
||||
python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"),
|
||||
binding=binding,
|
||||
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
|
||||
rules=rules,
|
||||
)
|
||||
|
||||
assert response is expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_completion_falls_back_after_native_declines() -> None:
|
||||
from litellm.rust_bridge.bindings import native_exception_types
|
||||
|
||||
native_types: Final = native_exception_types()
|
||||
if native_types is None:
|
||||
pytest.skip("native bridge is unavailable")
|
||||
declined, _ = native_types
|
||||
expected: Final = ModelResponse()
|
||||
rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_OPT_OUT),)
|
||||
|
||||
async def native(
|
||||
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
|
||||
) -> ModelResponse:
|
||||
raise declined("unsupported")
|
||||
|
||||
async def python(*args: object, **kwargs: object) -> ModelResponse:
|
||||
return expected
|
||||
|
||||
binding: Final[NativeBinding[NativeAcompletion]] = NativeBinding("acompletion", validate=lambda _: None)
|
||||
binding.override(native)
|
||||
response: Final = await dispatch._ADISPATCH.arun( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
|
||||
("test-model", MESSAGES),
|
||||
{},
|
||||
python=python,
|
||||
binding=binding,
|
||||
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
|
||||
rules=rules,
|
||||
)
|
||||
|
||||
assert response is expected
|
||||
|
||||
|
||||
def test_internal_acompletion_marker_bypasses_native() -> None:
|
||||
rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),)
|
||||
expected: Final = ModelResponse()
|
||||
|
||||
def python(*args: object, **kwargs: object) -> ModelResponse:
|
||||
return expected
|
||||
|
||||
def native(
|
||||
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
|
||||
) -> ModelResponse:
|
||||
pytest.fail("acompletion's inner completion call must stay on Python")
|
||||
|
||||
binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None)
|
||||
binding.override(native)
|
||||
response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
|
||||
("test-model", MESSAGES),
|
||||
{"custom_llm_provider": "openai", "acompletion": True},
|
||||
python=python,
|
||||
binding=binding,
|
||||
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
|
||||
rules=rules,
|
||||
)
|
||||
|
||||
assert response is expected
|
||||
|
|
@ -14,6 +14,7 @@ from pathlib import Path
|
|||
from types import SimpleNamespace
|
||||
import httpx
|
||||
import pytest
|
||||
from pytest_socket import _remove_restrictions
|
||||
|
||||
import asyncio
|
||||
|
||||
|
|
@ -509,6 +510,14 @@ def setup_and_teardown():
|
|||
print(f"[conftest] Module teardown complete (worker: {worker_id or 'master'})")
|
||||
|
||||
|
||||
def pytest_collectstart():
|
||||
_remove_restrictions()
|
||||
|
||||
|
||||
def pytest_runtest_setup():
|
||||
_remove_restrictions()
|
||||
|
||||
|
||||
def pytest_collection_modifyitems(config, items):
|
||||
"""
|
||||
Customize test collection order.
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
# Levo integration tests
|
||||
|
|
@ -7,10 +7,6 @@ the litellm_responses bridge provider, which calls litellm.responses() internall
|
|||
|
||||
import os
|
||||
|
||||
from litellm.interactions.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesInteractionsConfig,
|
||||
)
|
||||
from litellm.types.interactions import Turn
|
||||
from tests.test_litellm.interactions.base_interactions_test import (
|
||||
BaseInteractionsTest,
|
||||
)
|
||||
|
|
@ -30,71 +26,3 @@ class TestLiteLLMResponsesBridge(BaseInteractionsTest):
|
|||
def get_api_key(self) -> str:
|
||||
"""Return the OpenAI API key from environment."""
|
||||
return os.getenv("OPENAI_API_KEY", "")
|
||||
|
||||
|
||||
class TestBridgeInputTransformation:
|
||||
"""Regression tests for translating Interactions input into Responses API input.
|
||||
|
||||
The bridge used to pass Google content parts through raw ({"type": "text"}),
|
||||
which the Responses API rejects with a 400, and it dropped the role encoded
|
||||
in step types and in the legacy "model" turn role.
|
||||
"""
|
||||
|
||||
def test_step_input_maps_roles_and_content_types(self):
|
||||
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input(
|
||||
[
|
||||
{"type": "user_input", "content": [{"type": "text", "text": "I like apples."}]},
|
||||
{"type": "model_output", "content": [{"type": "text", "text": "I like oranges."}]},
|
||||
{"type": "user_input", "content": [{"type": "text", "text": "What did you say?"}]},
|
||||
]
|
||||
)
|
||||
assert transformed == [
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "I like apples."}]},
|
||||
{"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]},
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "What did you say?"}]},
|
||||
]
|
||||
|
||||
def test_legacy_turn_input_maps_model_role_to_assistant(self):
|
||||
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input(
|
||||
[
|
||||
{"role": "user", "content": [{"type": "text", "text": "I like apples."}]},
|
||||
{"role": "model", "content": [{"type": "text", "text": "I like oranges."}]},
|
||||
]
|
||||
)
|
||||
assert transformed == [
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "I like apples."}]},
|
||||
{"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]},
|
||||
]
|
||||
|
||||
def test_turn_pydantic_model_with_string_content(self):
|
||||
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input(
|
||||
[Turn(role="model", content="I like oranges.")]
|
||||
)
|
||||
assert transformed == [
|
||||
{"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]}
|
||||
]
|
||||
|
||||
def test_string_input_passes_through(self):
|
||||
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input("Hello")
|
||||
assert transformed == "Hello"
|
||||
|
||||
def test_content_list_input_becomes_single_user_message(self):
|
||||
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input(
|
||||
[{"type": "text", "text": "Hello"}, "world"]
|
||||
)
|
||||
assert transformed == [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "input_text", "text": "Hello"},
|
||||
{"type": "input_text", "text": "world"},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
def test_non_text_content_passes_through_unchanged(self):
|
||||
image_part = {"type": "image", "data": "base64data", "mime_type": "image/png"}
|
||||
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input(
|
||||
[{"type": "user_input", "content": [image_part]}]
|
||||
)
|
||||
assert transformed == [{"role": "user", "content": [image_part]}]
|
||||
|
|
|
|||
|
|
@ -9,171 +9,6 @@ import os
|
|||
import pytest
|
||||
|
||||
|
||||
from litellm.llms.cometapi.chat.transformation import (
|
||||
CometAPIChatCompletionStreamingHandler,
|
||||
CometAPIConfig,
|
||||
)
|
||||
from litellm.llms.cometapi.common_utils import CometAPIException
|
||||
|
||||
|
||||
class TestCometAPIChatCompletionStreamingHandler:
|
||||
def test_chunk_parser_successful(self):
|
||||
handler = CometAPIChatCompletionStreamingHandler(
|
||||
streaming_response=None, sync_stream=True
|
||||
)
|
||||
|
||||
# Test input chunk
|
||||
chunk = {
|
||||
"id": "test_id",
|
||||
"created": 1234567890,
|
||||
"model": "gpt-3.5-turbo",
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
|
||||
"choices": [
|
||||
{"delta": {"content": "test content", "reasoning": "test reasoning"}}
|
||||
],
|
||||
}
|
||||
|
||||
# Parse chunk
|
||||
result = handler.chunk_parser(chunk)
|
||||
|
||||
# Verify response
|
||||
assert result.id == "test_id"
|
||||
assert result.object == "chat.completion.chunk"
|
||||
assert result.created == 1234567890
|
||||
assert result.model == "gpt-3.5-turbo"
|
||||
assert result.usage.prompt_tokens == chunk["usage"]["prompt_tokens"]
|
||||
assert result.usage.completion_tokens == chunk["usage"]["completion_tokens"]
|
||||
assert result.usage.total_tokens == chunk["usage"]["total_tokens"]
|
||||
assert len(result.choices) == 1
|
||||
assert result.choices[0]["delta"]["reasoning_content"] == "test reasoning"
|
||||
|
||||
def test_chunk_parser_error_response(self):
|
||||
handler = CometAPIChatCompletionStreamingHandler(
|
||||
streaming_response=None, sync_stream=True
|
||||
)
|
||||
|
||||
# Test error chunk
|
||||
error_chunk = {
|
||||
"error": {
|
||||
"message": "test error",
|
||||
"code": 400,
|
||||
}
|
||||
}
|
||||
|
||||
# Verify error handling
|
||||
with pytest.raises(CometAPIException) as exc_info:
|
||||
handler.chunk_parser(error_chunk)
|
||||
|
||||
assert "CometAPI Error: test error" in str(exc_info.value)
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
def test_chunk_parser_key_error(self):
|
||||
handler = CometAPIChatCompletionStreamingHandler(
|
||||
streaming_response=None, sync_stream=True
|
||||
)
|
||||
|
||||
# Test invalid chunk missing required fields
|
||||
invalid_chunk = {"incomplete": "data"}
|
||||
|
||||
# Verify KeyError handling
|
||||
with pytest.raises(CometAPIException) as exc_info:
|
||||
handler.chunk_parser(invalid_chunk)
|
||||
|
||||
assert "KeyError" in str(exc_info.value)
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
class TestCometAPIConfig:
|
||||
def test_transform_request_basic(self):
|
||||
"""Test basic request transformation"""
|
||||
config = CometAPIConfig()
|
||||
|
||||
transformed_request = config.transform_request(
|
||||
model="cometapi/gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert transformed_request["model"] == "cometapi/gpt-3.5-turbo"
|
||||
assert transformed_request["messages"] == [
|
||||
{"role": "user", "content": "Hello, world!"}
|
||||
]
|
||||
|
||||
def test_transform_request_with_extra_body(self):
|
||||
"""Test request transformation with extra_body parameters"""
|
||||
config = CometAPIConfig()
|
||||
|
||||
transformed_request = config.transform_request(
|
||||
model="cometapi/gpt-4",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
optional_params={"extra_body": {"custom_param": "custom_value"}},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
# Validate that extra_body parameters are merged into the request
|
||||
assert transformed_request["custom_param"] == "custom_value"
|
||||
assert transformed_request["messages"] == [
|
||||
{"role": "user", "content": "Hello, world!"}
|
||||
]
|
||||
|
||||
def test_cache_control_flag_removal(self):
|
||||
"""Test cache control flag removal from messages"""
|
||||
config = CometAPIConfig()
|
||||
|
||||
transformed_request = config.transform_request(
|
||||
model="cometapi/gpt-3.5-turbo",
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hello, world!",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
# CometAPI should remove cache_control flags by default
|
||||
assert transformed_request["messages"][0].get("cache_control") is None
|
||||
|
||||
def test_map_openai_params(self):
|
||||
"""Test OpenAI parameter mapping"""
|
||||
config = CometAPIConfig()
|
||||
|
||||
non_default_params = {
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 100,
|
||||
"top_p": 0.9,
|
||||
}
|
||||
|
||||
mapped_params = config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params={},
|
||||
model="cometapi/gpt-3.5-turbo",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert mapped_params["temperature"] == 0.7
|
||||
assert mapped_params["max_tokens"] == 100
|
||||
assert mapped_params["top_p"] == 0.9
|
||||
|
||||
def test_get_error_class(self):
|
||||
"""Test error class creation"""
|
||||
config = CometAPIConfig()
|
||||
|
||||
error = config.get_error_class(
|
||||
error_message="Test error",
|
||||
status_code=400,
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
|
||||
assert isinstance(error, CometAPIException)
|
||||
assert error.message == "Test error"
|
||||
assert error.status_code == 400
|
||||
|
||||
|
||||
# Integration test example (requires real API key)
|
||||
|
|
|
|||
|
|
@ -1,79 +0,0 @@
|
|||
import json
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
def test_completion_merges_leading_system_and_developer_messages_for_chat_template_models(
|
||||
respx_mock: respx.MockRouter,
|
||||
):
|
||||
upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock(
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"id": "chatcmpl-123",
|
||||
"object": "chat.completion",
|
||||
"created": 1677652288,
|
||||
"model": "my-custom-model",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10},
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
response: Final = litellm.completion(
|
||||
model="databricks/my-custom-model",
|
||||
messages=[
|
||||
{"role": "system", "content": "You are terse."},
|
||||
{"role": "developer", "content": "Skills: none."},
|
||||
{"role": "user", "content": "Hello"},
|
||||
],
|
||||
api_base="https://example.databricks.test/serving-endpoints",
|
||||
api_key="fake-databricks-api-key",
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
assert upstream.call_count == 1
|
||||
request_body: Final = json.loads(upstream.calls[0].request.read())
|
||||
assert request_body["messages"] == [
|
||||
{"role": "system", "content": "You are terse.\n\nSkills: none."},
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
assert response.choices[0].message.content == "Answer"
|
||||
|
||||
|
||||
def test_completion_merges_system_messages_when_one_has_empty_content(respx_mock: respx.MockRouter):
|
||||
upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock(
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"id": "chatcmpl-123",
|
||||
"object": "chat.completion",
|
||||
"created": 1677652288,
|
||||
"model": "my-custom-model",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10},
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
litellm.completion(
|
||||
model="databricks/my-custom-model",
|
||||
messages=[
|
||||
{"role": "system", "content": "You are terse."},
|
||||
{"role": "system", "content": ""},
|
||||
{"role": "user", "content": "Hello"},
|
||||
],
|
||||
api_base="https://example.databricks.test/serving-endpoints",
|
||||
api_key="fake-databricks-api-key",
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
request_body: Final = json.loads(upstream.calls[0].request.read())
|
||||
assert request_body["messages"] == [
|
||||
{"role": "system", "content": "You are terse."},
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
|
|
@ -1,433 +0,0 @@
|
|||
"""
|
||||
Integration tests for DeepInfra rerank functionality.
|
||||
Tests the full rerank flow following the repository patterns.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
def assert_response_shape(response, custom_llm_provider):
|
||||
"""Helper function to validate response structure specific to DeepInfra."""
|
||||
assert hasattr(response, "id")
|
||||
assert hasattr(response, "results")
|
||||
assert hasattr(response, "meta")
|
||||
assert isinstance(response.results, list)
|
||||
|
||||
for result in response.results:
|
||||
assert "index" in result
|
||||
assert "relevance_score" in result
|
||||
assert isinstance(result["index"], int)
|
||||
assert isinstance(result["relevance_score"], (int, float))
|
||||
|
||||
# Check meta structure
|
||||
assert "tokens" in response.meta
|
||||
assert "billed_units" in response.meta
|
||||
assert "input_tokens" in response.meta["tokens"]
|
||||
assert "total_tokens" in response.meta["billed_units"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post")
|
||||
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
|
||||
def test_basic_rerank_deepinfra(mock_sync_post, mock_async_post, sync_mode):
|
||||
"""Test basic DeepInfra rerank functionality."""
|
||||
# Mock response data that matches DeepInfra API format
|
||||
mock_response_data = {
|
||||
"scores": [0.9, 0.1],
|
||||
"input_tokens": 25,
|
||||
"request_id": "deepinfra-request-123",
|
||||
"inference_status": {
|
||||
"status": "success",
|
||||
"runtime_ms": 150,
|
||||
"cost": 0.0001,
|
||||
"tokens_generated": 0,
|
||||
"tokens_input": 25,
|
||||
},
|
||||
}
|
||||
|
||||
def return_val():
|
||||
return mock_response_data
|
||||
|
||||
api_key = "test_deepinfra_api_key"
|
||||
api_base = "https://api.deepinfra.com"
|
||||
|
||||
if sync_mode:
|
||||
# Create mock response object for sync
|
||||
mock_response = MagicMock()
|
||||
mock_response.json = return_val
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.text = json.dumps(mock_response_data)
|
||||
mock_sync_post.return_value = mock_response
|
||||
|
||||
response = litellm.rerank(
|
||||
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
|
||||
query="hello",
|
||||
documents=["hello", "world"],
|
||||
top_n=2,
|
||||
custom_llm_provider="deepinfra",
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
mock_sync_post.assert_called_once()
|
||||
else:
|
||||
# Create mock response object for async
|
||||
mock_response = AsyncMock()
|
||||
|
||||
def return_val():
|
||||
return mock_response_data
|
||||
|
||||
mock_response.json = return_val
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.text = json.dumps(mock_response_data)
|
||||
mock_async_post.return_value = mock_response
|
||||
|
||||
response = asyncio.run(
|
||||
litellm.arerank(
|
||||
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
|
||||
query="hello",
|
||||
documents=["hello", "world"],
|
||||
top_n=2,
|
||||
custom_llm_provider="deepinfra",
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
)
|
||||
mock_async_post.assert_called_once()
|
||||
|
||||
# Verify response structure
|
||||
assert response.id == "deepinfra-request-123"
|
||||
assert response.results is not None
|
||||
assert len(response.results) == 2
|
||||
assert response.results[0]["index"] == 0
|
||||
assert response.results[0]["relevance_score"] == 0.9
|
||||
assert response.results[1]["index"] == 1
|
||||
assert response.results[1]["relevance_score"] == 0.1
|
||||
|
||||
# Verify metadata
|
||||
assert response.meta["tokens"]["input_tokens"] == 25
|
||||
assert response.meta["billed_units"]["total_tokens"] == 25
|
||||
|
||||
# Verify hidden params specific to DeepInfra
|
||||
assert response._hidden_params["status"] == "success"
|
||||
assert response._hidden_params["runtime_ms"] == 150
|
||||
assert response._hidden_params["cost"] == 0.0001
|
||||
# Note: The model name is processed and the 'deepinfra/' prefix is removed
|
||||
assert response._hidden_params["model"] == "Qwen/Qwen3-Reranker-0.6B"
|
||||
|
||||
assert_response_shape(response, custom_llm_provider="deepinfra")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post")
|
||||
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
|
||||
def test_deepinfra_rerank_with_queries_param(
|
||||
mock_sync_post, mock_async_post, sync_mode
|
||||
):
|
||||
"""Test DeepInfra rerank with multiple queries parameter."""
|
||||
mock_response_data = {
|
||||
"scores": [0.8, 0.6, 0.2],
|
||||
"input_tokens": 35,
|
||||
"request_id": "deepinfra-multi-query-123",
|
||||
"inference_status": {"status": "success", "runtime_ms": 200},
|
||||
}
|
||||
|
||||
def return_val():
|
||||
return mock_response_data
|
||||
|
||||
if sync_mode:
|
||||
mock_response = MagicMock()
|
||||
mock_response.json = return_val
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.text = json.dumps(mock_response_data)
|
||||
mock_sync_post.return_value = mock_response
|
||||
|
||||
response = litellm.rerank(
|
||||
model="deepinfra/Qwen/Qwen3-Reranker-4B",
|
||||
query="hello",
|
||||
documents=["hello", "world", "test"],
|
||||
queries=["hello", "hi there"], # DeepInfra specific param
|
||||
custom_llm_provider="deepinfra",
|
||||
api_key="test_key",
|
||||
api_base="https://api.deepinfra.com",
|
||||
)
|
||||
|
||||
mock_sync_post.assert_called_once()
|
||||
# Verify that queries parameter was passed in request
|
||||
call_data = json.loads(mock_sync_post.call_args.kwargs["data"])
|
||||
assert "queries" in call_data
|
||||
assert call_data["queries"] == ["hello", "hi there"]
|
||||
else:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.json = return_val
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.text = json.dumps(mock_response_data)
|
||||
mock_async_post.return_value = mock_response
|
||||
|
||||
response = asyncio.run(
|
||||
litellm.arerank(
|
||||
model="deepinfra/Qwen/Qwen3-Reranker-4B",
|
||||
query="hello",
|
||||
documents=["hello", "world", "test"],
|
||||
queries=["hello", "hi there"],
|
||||
custom_llm_provider="deepinfra",
|
||||
api_key="test_key",
|
||||
api_base="https://api.deepinfra.com",
|
||||
)
|
||||
)
|
||||
|
||||
mock_async_post.assert_called_once()
|
||||
call_data = json.loads(mock_async_post.call_args.kwargs["data"])
|
||||
assert "queries" in call_data
|
||||
assert call_data["queries"] == ["hello", "hi there"]
|
||||
|
||||
assert response.results is not None
|
||||
assert len(response.results) == 3
|
||||
|
||||
|
||||
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
|
||||
def test_deepinfra_rerank_with_service_tier(mock_post):
|
||||
"""Test DeepInfra rerank with service_tier parameter."""
|
||||
mock_response_data = {
|
||||
"scores": [0.95, 0.75],
|
||||
"input_tokens": 30,
|
||||
"request_id": "deepinfra-premium-123",
|
||||
}
|
||||
|
||||
def return_val():
|
||||
return mock_response_data
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json = return_val
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.text = json.dumps(mock_response_data)
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
response = litellm.rerank(
|
||||
model="deepinfra/Qwen/Qwen3-Reranker-8B",
|
||||
query="premium search",
|
||||
documents=["doc1", "doc2"],
|
||||
service_tier="premium", # DeepInfra specific param
|
||||
custom_llm_provider="deepinfra",
|
||||
api_key="test_key",
|
||||
api_base="https://api.deepinfra.com",
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
|
||||
# Verify URL
|
||||
call_url = mock_post.call_args.kwargs["url"]
|
||||
assert "api.deepinfra.com/inference/Qwen/Qwen3-Reranker-8B" in call_url
|
||||
|
||||
# Verify request contains service_tier
|
||||
call_data = json.loads(mock_post.call_args.kwargs["data"])
|
||||
assert call_data["service_tier"] == "premium"
|
||||
|
||||
assert response.results is not None
|
||||
|
||||
|
||||
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
|
||||
def test_deepinfra_rerank_with_env_vars(mock_post, monkeypatch):
|
||||
"""Test DeepInfra rerank with environment variable configuration."""
|
||||
monkeypatch.setenv("DEEPINFRA_API_KEY", "env_test_key")
|
||||
monkeypatch.setenv("DEEPINFRA_API_BASE", "https://custom-deepinfra.com")
|
||||
|
||||
mock_response_data = {
|
||||
"scores": [0.88, 0.22],
|
||||
"input_tokens": 28,
|
||||
"request_id": "env-test-123",
|
||||
}
|
||||
|
||||
def return_val():
|
||||
return mock_response_data
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json = return_val
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.text = json.dumps(mock_response_data)
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
response = litellm.rerank(
|
||||
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
|
||||
query="hello",
|
||||
documents=["hello", "world"],
|
||||
custom_llm_provider="deepinfra",
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
|
||||
# Verify headers contain env API key
|
||||
headers = mock_post.call_args.kwargs.get("headers", {})
|
||||
assert "Bearer env_test_key" in headers.get("Authorization", "")
|
||||
|
||||
assert response.results is not None
|
||||
|
||||
|
||||
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
|
||||
def test_deepinfra_rerank_error_handling(mock_post):
|
||||
"""Test DeepInfra rerank error handling."""
|
||||
error_response = {"detail": {"error": "Invalid API key"}}
|
||||
|
||||
def return_val():
|
||||
return error_response
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 401
|
||||
mock_response.json = return_val
|
||||
mock_response.text = json.dumps(error_response)
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
# The current implementation handles errors gracefully, so we expect a successful response
|
||||
# with the error information in the hidden params
|
||||
response = litellm.rerank(
|
||||
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
|
||||
query="hello",
|
||||
documents=["hello", "world"],
|
||||
custom_llm_provider="deepinfra",
|
||||
api_key="invalid_key",
|
||||
api_base="https://api.deepinfra.com",
|
||||
)
|
||||
|
||||
# Verify that the response contains error information
|
||||
assert (
|
||||
response._hidden_params["status"] == "unknown"
|
||||
) # Default status when error occurs
|
||||
|
||||
|
||||
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
|
||||
def test_deepinfra_rerank_defaults_api_base_when_missing(mock_post, monkeypatch):
|
||||
"""With no api_base anywhere, the call still goes out against DeepInfra's own base."""
|
||||
monkeypatch.delenv("DEEPINFRA_API_BASE", raising=False)
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json = lambda: {"scores": [0.9, 0.1], "input_tokens": 20}
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
response = litellm.rerank(
|
||||
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
|
||||
query="hello",
|
||||
documents=["hello", "world"],
|
||||
custom_llm_provider="deepinfra",
|
||||
api_key="test_key",
|
||||
# api_base is intentionally missing
|
||||
)
|
||||
|
||||
assert "api.deepinfra.com" in mock_post.call_args.kwargs["url"]
|
||||
assert [result["relevance_score"] for result in response.results] == [0.9, 0.1]
|
||||
|
||||
|
||||
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
|
||||
def test_deepinfra_rerank_request_format(mock_post):
|
||||
"""Test that the request is properly formatted for DeepInfra API."""
|
||||
mock_response_data = {"scores": [0.9, 0.1], "input_tokens": 20}
|
||||
|
||||
def return_val():
|
||||
return mock_response_data
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json = return_val
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.text = json.dumps(mock_response_data)
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
response = litellm.rerank(
|
||||
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
|
||||
query="test query",
|
||||
documents=["doc1", "doc2"],
|
||||
custom_llm_provider="deepinfra",
|
||||
api_key="test_key",
|
||||
api_base="https://api.deepinfra.com",
|
||||
instruction="custom instruction",
|
||||
webhook="https://webhook.example.com",
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
|
||||
# Verify URL format
|
||||
call_url = mock_post.call_args.kwargs["url"]
|
||||
assert call_url == "https://api.deepinfra.com/inference/Qwen/Qwen3-Reranker-0.6B"
|
||||
|
||||
# Verify headers
|
||||
headers = mock_post.call_args.kwargs["headers"]
|
||||
assert headers["Authorization"] == "Bearer test_key"
|
||||
assert headers["accept"] == "application/json"
|
||||
assert headers["content-type"] == "application/json"
|
||||
|
||||
# Verify request body format
|
||||
request_data = json.loads(mock_post.call_args.kwargs["data"])
|
||||
assert request_data["queries"] == [
|
||||
"test query",
|
||||
"test query",
|
||||
] # DeepInfra requires queries to match documents length
|
||||
assert request_data["documents"] == ["doc1", "doc2"]
|
||||
assert request_data["instruction"] == "custom instruction"
|
||||
assert request_data["webhook"] == "https://webhook.example.com"
|
||||
|
||||
assert response.results is not None
|
||||
|
||||
|
||||
def test_deepinfra_rerank_models():
|
||||
"""Test that DeepInfra Qwen rerank models are recognized."""
|
||||
# These should not raise errors during model validation
|
||||
models = [
|
||||
"deepinfra/Qwen/Qwen3-Reranker-0.6B",
|
||||
"deepinfra/Qwen/Qwen3-Reranker-4B",
|
||||
"deepinfra/Qwen/Qwen3-Reranker-8B",
|
||||
]
|
||||
|
||||
for model in models:
|
||||
resolved_model, provider, _, api_base = litellm.get_llm_provider(model=model)
|
||||
assert provider == "deepinfra"
|
||||
assert resolved_model == model.removeprefix("deepinfra/")
|
||||
assert api_base == "https://api.deepinfra.com/v1/openai"
|
||||
|
||||
|
||||
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
|
||||
def test_deepinfra_rerank_minimal_response(mock_post):
|
||||
"""Test handling of minimal DeepInfra response."""
|
||||
# Minimal response with just scores
|
||||
mock_response_data = {"scores": [0.7, 0.3]}
|
||||
|
||||
def return_val():
|
||||
return mock_response_data
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json = return_val
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.text = json.dumps(mock_response_data)
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
response = litellm.rerank(
|
||||
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
|
||||
query="hello",
|
||||
documents=["hello", "world"],
|
||||
custom_llm_provider="deepinfra",
|
||||
api_key="test_key",
|
||||
api_base="https://api.deepinfra.com",
|
||||
)
|
||||
|
||||
# Should handle minimal response gracefully
|
||||
assert response.results is not None
|
||||
assert len(response.results) == 2
|
||||
assert response.results[0]["relevance_score"] == 0.7
|
||||
assert response.results[1]["relevance_score"] == 0.3
|
||||
|
||||
# Should have default values for missing fields
|
||||
assert response.meta["tokens"]["input_tokens"] == 0 # Default when missing
|
||||
assert response._hidden_params["status"] == "unknown" # Default when missing
|
||||
|
|
@ -1 +0,0 @@
|
|||
"""Tests for Gemini files functionality"""
|
||||
|
|
@ -1 +0,0 @@
|
|||
# Gemini Video Generation Tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# Manus provider tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# Manus Responses API tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# MiniMax tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# MiniMax chat tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# MiniMax messages tests
|
||||
|
|
@ -1,19 +1,9 @@
|
|||
import os
|
||||
from typing import Dict
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import litellm
|
||||
import pytest
|
||||
|
||||
from litellm.llms.base_llm.audio_transcription.transformation import (
|
||||
BaseAudioTranscriptionConfig,
|
||||
)
|
||||
from litellm.llms.mistral.audio_transcription.transformation import (
|
||||
MistralAudioTranscriptionConfig,
|
||||
)
|
||||
from litellm.types.utils import TranscriptionResponse
|
||||
from litellm.utils import ProviderConfigManager
|
||||
from tests.llm_translation.base_audio_transcription_unit_tests import (
|
||||
BaseLLMAudioTranscriptionTest,
|
||||
)
|
||||
|
|
@ -37,184 +27,3 @@ class TestMistralAudioTranscription(BaseLLMAudioTranscriptionTest):
|
|||
"Async audio transcription test for Mistral is skipped in this suite; "
|
||||
"async test plugins (e.g. pytest-asyncio/anyio) are not configured here."
|
||||
)
|
||||
|
||||
|
||||
def test_mistral_audio_transcription_config_installed():
|
||||
"""Ensure Mistral audio transcription config is registered with ProviderConfigManager."""
|
||||
config = ProviderConfigManager.get_provider_audio_transcription_config(
|
||||
model="mistral/voxtral-mini-latest",
|
||||
provider=litellm.LlmProviders.MISTRAL,
|
||||
)
|
||||
assert config is not None
|
||||
assert isinstance(config, BaseAudioTranscriptionConfig)
|
||||
assert isinstance(config, MistralAudioTranscriptionConfig)
|
||||
|
||||
|
||||
def test_mistral_audio_transcription_get_complete_url():
|
||||
config = MistralAudioTranscriptionConfig()
|
||||
url = config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key="fake-key",
|
||||
model="voxtral-mini-latest",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
assert url == "https://api.mistral.ai/v1/audio/transcriptions"
|
||||
|
||||
|
||||
def test_mistral_audio_transcription_get_complete_url_custom_base():
|
||||
config = MistralAudioTranscriptionConfig()
|
||||
url = config.get_complete_url(
|
||||
api_base="https://custom.api.example.com/v1/",
|
||||
api_key="fake-key",
|
||||
model="voxtral-mini-latest",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
assert url == "https://custom.api.example.com/v1/audio/transcriptions"
|
||||
|
||||
|
||||
def test_mistral_audio_transcription_validate_environment():
|
||||
config = MistralAudioTranscriptionConfig()
|
||||
headers = config.validate_environment(
|
||||
headers={},
|
||||
model="voxtral-mini-latest",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key="test-key-123",
|
||||
)
|
||||
assert headers["Authorization"] == "Bearer test-key-123"
|
||||
assert headers["accept"] == "application/json"
|
||||
|
||||
|
||||
def test_mistral_audio_transcription_supported_params():
|
||||
config = MistralAudioTranscriptionConfig()
|
||||
params = config.get_supported_openai_params("voxtral-mini-latest")
|
||||
assert "language" in params
|
||||
assert "temperature" in params
|
||||
assert "response_format" in params
|
||||
assert "timestamp_granularities" in params
|
||||
|
||||
|
||||
def test_mistral_audio_transcription_request_transform():
|
||||
config = MistralAudioTranscriptionConfig()
|
||||
|
||||
wav_path = os.path.join(
|
||||
os.path.dirname(__file__),
|
||||
"../../../../..",
|
||||
"tests",
|
||||
"llm_translation",
|
||||
"gettysburg.wav",
|
||||
)
|
||||
audio_file = open(wav_path, "rb")
|
||||
|
||||
result = config.transform_audio_transcription_request(
|
||||
model="voxtral-mini-latest",
|
||||
audio_file=audio_file,
|
||||
optional_params={"language": "en", "temperature": 0.0},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
audio_file.close()
|
||||
|
||||
assert isinstance(result.data, dict)
|
||||
assert result.data["model"] == "voxtral-mini-latest"
|
||||
assert result.data["language"] == "en"
|
||||
assert result.data["temperature"] == 0.0
|
||||
assert result.files is not None
|
||||
assert "file" in result.files
|
||||
|
||||
|
||||
def test_mistral_audio_transcription_request_with_diarize():
|
||||
"""Test that Mistral-specific params like diarize are passed through."""
|
||||
config = MistralAudioTranscriptionConfig()
|
||||
|
||||
wav_path = os.path.join(
|
||||
os.path.dirname(__file__),
|
||||
"../../../../..",
|
||||
"tests",
|
||||
"llm_translation",
|
||||
"gettysburg.wav",
|
||||
)
|
||||
audio_file = open(wav_path, "rb")
|
||||
|
||||
result = config.transform_audio_transcription_request(
|
||||
model="voxtral-mini-latest",
|
||||
audio_file=audio_file,
|
||||
optional_params={"diarize": True},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
audio_file.close()
|
||||
|
||||
assert isinstance(result.data, dict)
|
||||
assert result.data["diarize"] == "true"
|
||||
|
||||
|
||||
def test_mistral_audio_transcription_response_transform():
|
||||
config = MistralAudioTranscriptionConfig()
|
||||
|
||||
mock_response = MagicMock(spec=httpx.Response)
|
||||
mock_response.json.return_value = {"text": "Four score and seven years ago..."}
|
||||
|
||||
response = config.transform_audio_transcription_response(mock_response)
|
||||
|
||||
assert isinstance(response, TranscriptionResponse)
|
||||
assert response.text == "Four score and seven years ago..."
|
||||
|
||||
|
||||
def test_mistral_audio_transcription_response_transform_diarized():
|
||||
"""Test that diarized responses preserve segments and language."""
|
||||
config = MistralAudioTranscriptionConfig()
|
||||
|
||||
mock_response = MagicMock(spec=httpx.Response)
|
||||
mock_response.json.return_value = {
|
||||
"model": "voxtral-mini-latest",
|
||||
"text": "Hello, how are you? I am fine.",
|
||||
"language": None,
|
||||
"segments": [
|
||||
{
|
||||
"text": "Hello, how are you?",
|
||||
"start": 0.3,
|
||||
"end": 2.1,
|
||||
"speaker_id": "speaker_1",
|
||||
"type": "transcription_segment",
|
||||
},
|
||||
{
|
||||
"text": "I am fine.",
|
||||
"start": 2.5,
|
||||
"end": 3.8,
|
||||
"speaker_id": "speaker_2",
|
||||
"type": "transcription_segment",
|
||||
},
|
||||
],
|
||||
"usage": {
|
||||
"prompt_audio_seconds": 4,
|
||||
"prompt_tokens": 5,
|
||||
"total_tokens": 50,
|
||||
"completion_tokens": 20,
|
||||
},
|
||||
}
|
||||
|
||||
response = config.transform_audio_transcription_response(mock_response)
|
||||
|
||||
assert isinstance(response, TranscriptionResponse)
|
||||
assert response.text == "Hello, how are you? I am fine."
|
||||
assert response["segments"] is not None
|
||||
assert len(response["segments"]) == 2
|
||||
assert response["segments"][0]["speaker_id"] == "speaker_1"
|
||||
assert response["segments"][1]["speaker_id"] == "speaker_2"
|
||||
assert response["language"] is None
|
||||
|
||||
|
||||
def test_mistral_audio_transcription_response_transform_empty():
|
||||
config = MistralAudioTranscriptionConfig()
|
||||
|
||||
mock_response = MagicMock(spec=httpx.Response)
|
||||
mock_response.json.return_value = {}
|
||||
|
||||
response = config.transform_audio_transcription_response(mock_response)
|
||||
|
||||
assert isinstance(response, TranscriptionResponse)
|
||||
assert response.text == ""
|
||||
|
|
|
|||
|
|
@ -3,321 +3,12 @@ Tests for JSON-based provider configuration system.
|
|||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import patch
|
||||
|
||||
try:
|
||||
import pytest
|
||||
except ImportError:
|
||||
# pytest not available, will run as standalone script
|
||||
pytest = None
|
||||
|
||||
# Add workspace to path
|
||||
workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))
|
||||
sys.path.insert(0, workspace_path)
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
class TestJSONProviderLoader:
|
||||
"""Test JSON provider loading and configuration"""
|
||||
|
||||
def test_load_json_providers(self):
|
||||
"""Test that JSON providers load correctly"""
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
# Verify publicai is loaded
|
||||
assert JSONProviderRegistry.exists("publicai")
|
||||
|
||||
# Get publicai config
|
||||
publicai = JSONProviderRegistry.get("publicai")
|
||||
assert publicai is not None
|
||||
assert publicai.base_url == "https://api.publicai.co/v1"
|
||||
assert publicai.api_key_env == "PUBLICAI_API_KEY"
|
||||
assert publicai.api_base_env == "PUBLICAI_API_BASE"
|
||||
assert publicai.param_mappings.get("max_completion_tokens") == "max_tokens"
|
||||
|
||||
def test_dynamic_config_generation(self):
|
||||
"""Test dynamic config class creation"""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("publicai")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
# Test API info resolution
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
|
||||
assert api_base == "https://api.publicai.co/v1"
|
||||
|
||||
# Test with custom base
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(
|
||||
"https://custom.api.com", "test-key"
|
||||
)
|
||||
assert api_base == "https://custom.api.com"
|
||||
assert api_key == "test-key"
|
||||
|
||||
def test_parameter_mapping(self):
|
||||
"""Test parameter mapping works"""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("publicai")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
# Test parameter mapping
|
||||
optional_params = {}
|
||||
non_default_params = {"max_completion_tokens": 100, "temperature": 0.7}
|
||||
result = config.map_openai_params(
|
||||
non_default_params, optional_params, "gpt-4", False
|
||||
)
|
||||
|
||||
# max_completion_tokens should be mapped to max_tokens
|
||||
assert "max_tokens" in result
|
||||
assert result["max_tokens"] == 100
|
||||
assert "max_completion_tokens" not in result
|
||||
|
||||
# temperature should be passed through
|
||||
assert result["temperature"] == 0.7
|
||||
|
||||
def test_supported_params(self):
|
||||
"""Test that config returns supported params"""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("publicai")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
# Get supported params
|
||||
supported = config.get_supported_openai_params("gpt-4")
|
||||
|
||||
# Should have standard OpenAI params
|
||||
assert isinstance(supported, list)
|
||||
assert len(supported) > 0
|
||||
|
||||
def test_tool_params_excluded_when_function_calling_not_supported(self):
|
||||
"""Test that tool-related params are excluded for models that don't support
|
||||
function calling. Regression test for https://github.com/BerriAI/litellm/issues/21125
|
||||
"""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("publicai")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
# Mock supports_function_calling to return False
|
||||
with patch("litellm.utils.supports_function_calling", return_value=False):
|
||||
supported = config.get_supported_openai_params("some-model-without-fc")
|
||||
|
||||
tool_params = [
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"function_call",
|
||||
"functions",
|
||||
"parallel_tool_calls",
|
||||
]
|
||||
for param in tool_params:
|
||||
assert (
|
||||
param not in supported
|
||||
), f"'{param}' should not be in supported params when function calling is not supported"
|
||||
|
||||
# Non-tool params should still be present
|
||||
assert "temperature" in supported
|
||||
assert "max_tokens" in supported
|
||||
assert "stop" in supported
|
||||
|
||||
def test_tool_params_included_when_function_calling_supported(self):
|
||||
"""Test that tool-related params are included for models that support function calling."""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("publicai")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
# Mock supports_function_calling to return True
|
||||
with patch("litellm.utils.supports_function_calling", return_value=True):
|
||||
supported = config.get_supported_openai_params("some-model-with-fc")
|
||||
|
||||
assert "tools" in supported
|
||||
assert "tool_choice" in supported
|
||||
|
||||
def test_provider_resolution(self):
|
||||
"""Test that provider resolution finds JSON providers"""
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import (
|
||||
get_llm_provider,
|
||||
)
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="publicai/gpt-4",
|
||||
custom_llm_provider=None,
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert model == "gpt-4"
|
||||
assert provider == "publicai"
|
||||
assert api_base == "https://api.publicai.co/v1"
|
||||
|
||||
def test_provider_config_manager(self):
|
||||
"""Test that ProviderConfigManager returns JSON-based configs"""
|
||||
from litellm import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
config = ProviderConfigManager.get_provider_chat_config(
|
||||
model="gpt-4", provider=LlmProviders.PUBLICAI
|
||||
)
|
||||
|
||||
assert config is not None
|
||||
assert config.custom_llm_provider == "publicai"
|
||||
|
||||
|
||||
class TestPinstripes:
|
||||
"""Tests for Pinstripes JSON-configured provider"""
|
||||
|
||||
def test_pinstripes_json_config_exists(self):
|
||||
"""Test that pinstripes is configured in providers.json"""
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
assert JSONProviderRegistry.exists("pinstripes")
|
||||
|
||||
pinstripes = JSONProviderRegistry.get("pinstripes")
|
||||
assert pinstripes is not None
|
||||
assert pinstripes.base_url == "https://pinstripes.io/v1"
|
||||
assert pinstripes.api_key_env == "PINSTRIPES_API_KEY"
|
||||
assert pinstripes.param_mappings.get("max_completion_tokens") == "max_tokens"
|
||||
|
||||
def test_pinstripes_provider_resolution(self):
|
||||
"""Test that provider resolution finds pinstripes and returns the default base URL"""
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="pinstripes/ps/glm-4.5-air",
|
||||
custom_llm_provider=None,
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert model == "ps/glm-4.5-air"
|
||||
assert provider == "pinstripes"
|
||||
assert api_base == "https://pinstripes.io/v1"
|
||||
|
||||
def test_pinstripes_dynamic_config(self):
|
||||
"""Test dynamic config class creation for pinstripes"""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("pinstripes")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
|
||||
assert api_base == "https://pinstripes.io/v1"
|
||||
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(
|
||||
"https://custom.pinstripes.io/v1", "test-key"
|
||||
)
|
||||
assert api_base == "https://custom.pinstripes.io/v1"
|
||||
assert api_key == "test-key"
|
||||
|
||||
def test_pinstripes_parameter_mapping(self):
|
||||
"""Test that max_completion_tokens is mapped to max_tokens for pinstripes"""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("pinstripes")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
optional_params = {}
|
||||
non_default_params = {"max_completion_tokens": 100, "temperature": 0.7}
|
||||
result = config.map_openai_params(
|
||||
non_default_params, optional_params, "ps/glm-4.5-air", False
|
||||
)
|
||||
|
||||
assert "max_tokens" in result
|
||||
assert result["max_tokens"] == 100
|
||||
assert "max_completion_tokens" not in result
|
||||
assert result["temperature"] == 0.7
|
||||
|
||||
|
||||
class TestDarkbloom:
|
||||
def test_darkbloom_json_config_exists(self):
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
darkbloom = JSONProviderRegistry.get("darkbloom")
|
||||
assert darkbloom is not None
|
||||
assert darkbloom.base_url == "https://api.darkbloom.dev/v1"
|
||||
assert darkbloom.api_key_env == "DARKBLOOM_API_KEY"
|
||||
assert darkbloom.api_base_env == "DARKBLOOM_API_BASE"
|
||||
assert darkbloom.param_mappings.get("max_completion_tokens") == "max_tokens"
|
||||
|
||||
def test_darkbloom_provider_resolution(self):
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="darkbloom/gemma-4-26b",
|
||||
custom_llm_provider=None,
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert model == "gemma-4-26b"
|
||||
assert provider == "darkbloom"
|
||||
assert api_key is None
|
||||
assert api_base == "https://api.darkbloom.dev/v1"
|
||||
|
||||
def test_darkbloom_dynamic_config(self):
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("darkbloom")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
|
||||
assert api_base == "https://api.darkbloom.dev/v1"
|
||||
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(
|
||||
"https://custom.darkbloom.dev/v1", "test-key"
|
||||
)
|
||||
assert api_base == "https://custom.darkbloom.dev/v1"
|
||||
assert api_key == "test-key"
|
||||
|
||||
def test_darkbloom_complete_url_appends_endpoint(self):
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("darkbloom")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
url = config.get_complete_url(
|
||||
api_base="https://api.darkbloom.dev/v1",
|
||||
api_key="test-key",
|
||||
model="darkbloom/gemma-4-26b",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert url == "https://api.darkbloom.dev/v1/chat/completions"
|
||||
|
||||
def test_darkbloom_provider_config_manager(self):
|
||||
from litellm import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
config = ProviderConfigManager.get_provider_chat_config(
|
||||
model="gemma-4-26b", provider=LlmProviders.DARKBLOOM
|
||||
)
|
||||
|
||||
assert config is not None
|
||||
assert config.custom_llm_provider == "darkbloom"
|
||||
|
||||
|
||||
class TestPublicAIIntegration:
|
||||
"""Integration tests for PublicAI provider"""
|
||||
|
||||
|
|
@ -457,55 +148,3 @@ class TestPublicAIIntegration:
|
|||
pytest.fail(f"Content list conversion test failed: {str(e)}")
|
||||
else:
|
||||
raise
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run basic tests
|
||||
print("Testing JSON Provider System...")
|
||||
|
||||
test_loader = TestJSONProviderLoader()
|
||||
print("\n1. Testing JSON provider loading...")
|
||||
test_loader.test_load_json_providers()
|
||||
print(" ✓ JSON providers loaded")
|
||||
|
||||
print("\n2. Testing dynamic config generation...")
|
||||
test_loader.test_dynamic_config_generation()
|
||||
print(" ✓ Dynamic config works")
|
||||
|
||||
print("\n3. Testing parameter mapping...")
|
||||
test_loader.test_parameter_mapping()
|
||||
print(" ✓ Parameter mapping works")
|
||||
|
||||
print("\n4. Testing excluded params...")
|
||||
test_loader.test_excluded_params()
|
||||
print(" ✓ Excluded params work")
|
||||
|
||||
print("\n5. Testing provider resolution...")
|
||||
test_loader.test_provider_resolution()
|
||||
print(" ✓ Provider resolution works")
|
||||
|
||||
print("\n6. Testing provider config manager...")
|
||||
test_loader.test_provider_config_manager()
|
||||
print(" ✓ Config manager works")
|
||||
|
||||
print("\n" + "=" * 50)
|
||||
print("PublicAI Integration Tests...")
|
||||
print("=" * 50)
|
||||
|
||||
test_integration = TestPublicAIIntegration()
|
||||
|
||||
print("\n7. Testing basic completion...")
|
||||
test_integration.test_publicai_completion_basic()
|
||||
|
||||
print("\n8. Testing streaming...")
|
||||
test_integration.test_publicai_completion_with_streaming()
|
||||
|
||||
print("\n9. Testing parameter mapping...")
|
||||
test_integration.test_publicai_parameter_mapping()
|
||||
|
||||
print("\n10. Testing content list conversion...")
|
||||
test_integration.test_publicai_content_list_conversion()
|
||||
|
||||
print("\n" + "=" * 50)
|
||||
print("✓ All tests passed!")
|
||||
print("=" * 50)
|
||||
|
|
|
|||
|
|
@ -4,86 +4,12 @@ Related to issue #18794
|
|||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
try:
|
||||
import pytest
|
||||
except ImportError:
|
||||
pytest = None
|
||||
|
||||
# Add workspace to path
|
||||
workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))
|
||||
sys.path.insert(0, workspace_path)
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
class TestXiaomiMiMoProviderConfig:
|
||||
"""Test Xiaomi MiMo provider configuration"""
|
||||
|
||||
def test_xiaomi_mimo_in_provider_list(self):
|
||||
"""Test that xiaomi_mimo is in the provider list (fixes #18794)"""
|
||||
from litellm import LlmProviders
|
||||
|
||||
# Verify xiaomi_mimo is in the enum
|
||||
assert hasattr(LlmProviders, "XIAOMI_MIMO")
|
||||
assert LlmProviders.XIAOMI_MIMO.value == "xiaomi_mimo"
|
||||
|
||||
# Verify it's in the provider list
|
||||
assert "xiaomi_mimo" in litellm.provider_list
|
||||
|
||||
def test_xiaomi_mimo_json_config_exists(self):
|
||||
"""Test that xiaomi_mimo is configured in providers.json"""
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
# Verify xiaomi_mimo is loaded
|
||||
assert JSONProviderRegistry.exists("xiaomi_mimo")
|
||||
|
||||
# Get xiaomi_mimo config
|
||||
xiaomi_mimo = JSONProviderRegistry.get("xiaomi_mimo")
|
||||
assert xiaomi_mimo is not None
|
||||
assert xiaomi_mimo.base_url == "https://api.xiaomimimo.com/v1"
|
||||
assert xiaomi_mimo.api_key_env == "XIAOMI_MIMO_API_KEY"
|
||||
assert xiaomi_mimo.param_mappings.get("max_completion_tokens") == "max_tokens"
|
||||
|
||||
def test_xiaomi_mimo_provider_resolution(self):
|
||||
"""Test that provider resolution finds xiaomi_mimo"""
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="xiaomi_mimo/mimo-v2-flash",
|
||||
custom_llm_provider=None,
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert model == "mimo-v2-flash"
|
||||
assert provider == "xiaomi_mimo"
|
||||
assert api_base == "https://api.xiaomimimo.com/v1"
|
||||
|
||||
def test_xiaomi_mimo_router_config(self):
|
||||
"""Test that xiaomi_mimo can be used in Router configuration (fixes #18794)"""
|
||||
from litellm import Router
|
||||
|
||||
# This should not raise "Unsupported provider - xiaomi_mimo"
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "mimo-v2-flash",
|
||||
"litellm_params": {
|
||||
"model": "xiaomi_mimo/mimo-v2-flash",
|
||||
"api_key": "test-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
# Verify the deployment was created successfully
|
||||
assert len(router.model_list) == 1
|
||||
assert router.model_list[0]["model_name"] == "mimo-v2-flash"
|
||||
|
||||
|
||||
class TestXiaomiMiMoIntegration:
|
||||
"""Integration tests for Xiaomi MiMo provider"""
|
||||
|
||||
|
|
@ -128,30 +54,3 @@ class TestXiaomiMiMoIntegration:
|
|||
pytest.fail(f"Xiaomi MiMo completion failed: {str(e)}")
|
||||
else:
|
||||
raise
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run basic tests
|
||||
print("Testing Xiaomi MiMo Provider...")
|
||||
|
||||
test_config = TestXiaomiMiMoProviderConfig()
|
||||
|
||||
print("\n1. Testing provider in list...")
|
||||
test_config.test_xiaomi_mimo_in_provider_list()
|
||||
print(" ✓ xiaomi_mimo in provider list")
|
||||
|
||||
print("\n2. Testing JSON config...")
|
||||
test_config.test_xiaomi_mimo_json_config_exists()
|
||||
print(" ✓ xiaomi_mimo JSON config loaded")
|
||||
|
||||
print("\n3. Testing provider resolution...")
|
||||
test_config.test_xiaomi_mimo_provider_resolution()
|
||||
print(" ✓ Provider resolution works")
|
||||
|
||||
print("\n4. Testing router configuration...")
|
||||
test_config.test_xiaomi_mimo_router_config()
|
||||
print(" ✓ Router configuration works (issue #18794 fixed)")
|
||||
|
||||
print("\n" + "=" * 50)
|
||||
print("✓ All configuration tests passed!")
|
||||
print("=" * 50)
|
||||
|
|
|
|||
|
|
@ -54,61 +54,3 @@ def test_ovhcloud_audio_transcription_config_installed():
|
|||
|
||||
assert config is not None
|
||||
assert isinstance(config, BaseAudioTranscriptionConfig)
|
||||
|
||||
|
||||
|
||||
class TestOVHCloudDurationFieldMigration:
|
||||
"""Tests for OVHCloud duration -> seconds field migration."""
|
||||
|
||||
def test_seconds_field_mapped_to_duration(self):
|
||||
"""New `seconds` field should be normalized to `duration`."""
|
||||
from litellm.llms.ovhcloud.audio_transcription.transformation import (
|
||||
OVHCloudAudioTranscriptionConfig,
|
||||
)
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
config = OVHCloudAudioTranscriptionConfig()
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"text": "Hello world",
|
||||
"seconds": 3.14,
|
||||
}
|
||||
|
||||
result = config.transform_audio_transcription_response(mock_response)
|
||||
|
||||
assert result.text == "Hello world"
|
||||
assert result._hidden_params["duration"] == 3.14
|
||||
|
||||
def test_legacy_duration_field_still_works(self):
|
||||
"""Legacy `duration` field should still be accepted."""
|
||||
from litellm.llms.ovhcloud.audio_transcription.transformation import (
|
||||
OVHCloudAudioTranscriptionConfig,
|
||||
)
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
config = OVHCloudAudioTranscriptionConfig()
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"text": "Hello world",
|
||||
"duration": 2.71,
|
||||
}
|
||||
|
||||
result = config.transform_audio_transcription_response(mock_response)
|
||||
|
||||
assert result.text == "Hello world"
|
||||
assert result._hidden_params["duration"] == 2.71
|
||||
|
||||
|
||||
|
||||
def test_seconds_zero_mapped_to_duration(self):
|
||||
"""seconds=0.0 must not be treated as falsy and lost."""
|
||||
from litellm.llms.ovhcloud.audio_transcription.transformation import (
|
||||
OVHCloudAudioTranscriptionConfig,
|
||||
)
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
config = OVHCloudAudioTranscriptionConfig()
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"text": "silence", "seconds": 0.0}
|
||||
result = config.transform_audio_transcription_response(mock_response)
|
||||
assert result._hidden_params["duration"] == 0.0
|
||||
|
|
@ -6,174 +6,12 @@ import os
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.ovhcloud.utils import OVHCloudException
|
||||
from litellm.utils import get_optional_params
|
||||
|
||||
|
||||
from litellm.llms.ovhcloud.chat.transformation import (
|
||||
OVHCloudChatCompletionStreamingHandler,
|
||||
OVHCloudChatConfig,
|
||||
)
|
||||
|
||||
config = OVHCloudChatConfig()
|
||||
model = "ovhcloud/Mistral-7B-Instruct-v0.3"
|
||||
|
||||
|
||||
class TestOvhCloudChatCompletionStreamingHandler:
|
||||
def test_chunk_parser_successful(self):
|
||||
handler = OVHCloudChatCompletionStreamingHandler(
|
||||
streaming_response=None, sync_stream=True
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"id": "test_id",
|
||||
"created": 1234567890,
|
||||
"model": "gpt-oss-20b",
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
|
||||
"choices": [
|
||||
{"delta": {"content": "test content", "reasoning": "test reasoning"}}
|
||||
],
|
||||
}
|
||||
|
||||
result = handler.chunk_parser(chunk)
|
||||
|
||||
assert result.id == "test_id"
|
||||
assert result.object == "chat.completion.chunk"
|
||||
assert result.created == 1234567890
|
||||
assert result.model == "gpt-oss-20b"
|
||||
assert result.usage.prompt_tokens == chunk["usage"]["prompt_tokens"]
|
||||
assert result.usage.completion_tokens == chunk["usage"]["completion_tokens"]
|
||||
assert result.usage.total_tokens == chunk["usage"]["total_tokens"]
|
||||
assert len(result.choices) == 1
|
||||
assert result.choices[0]["delta"]["reasoning_content"] == "test reasoning"
|
||||
|
||||
def test_chunk_parser_error_response(self):
|
||||
handler = OVHCloudChatCompletionStreamingHandler(
|
||||
streaming_response=None, sync_stream=True
|
||||
)
|
||||
|
||||
error_chunk = {
|
||||
"error": {
|
||||
"message": "test error",
|
||||
"code": 400,
|
||||
}
|
||||
}
|
||||
|
||||
with pytest.raises(OVHCloudException) as exc_info:
|
||||
handler.chunk_parser(error_chunk)
|
||||
|
||||
assert "OVHCloud Error: test error" in str(exc_info.value)
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
def test_chunk_parser_key_error(self):
|
||||
handler = OVHCloudChatCompletionStreamingHandler(
|
||||
streaming_response=None, sync_stream=True
|
||||
)
|
||||
|
||||
invalid_chunk = {"incomplete": "data"}
|
||||
|
||||
with pytest.raises(OVHCloudException) as exc_info:
|
||||
handler.chunk_parser(invalid_chunk)
|
||||
|
||||
assert "KeyError" in str(exc_info.value)
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
class TestOVHCloudConfig:
|
||||
def test_transform_request_basic(self):
|
||||
"""Test basic request transformation"""
|
||||
transformed_request = config.transform_request(
|
||||
model,
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert transformed_request["model"] == model
|
||||
assert transformed_request["messages"] == [
|
||||
{"role": "user", "content": "Hello, world!"}
|
||||
]
|
||||
|
||||
def test_transform_request_with_extra_body(self):
|
||||
"""Test request transformation with extra_body parameters"""
|
||||
transformed_request = config.transform_request(
|
||||
model,
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
optional_params={"extra_body": {"custom_param": "custom_value"}},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert transformed_request["custom_param"] == "custom_value"
|
||||
assert transformed_request["messages"] == [
|
||||
{"role": "user", "content": "Hello, world!"}
|
||||
]
|
||||
|
||||
def test_map_openai_params(self):
|
||||
"""Test OpenAI parameter mapping"""
|
||||
non_default_params = {
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 100,
|
||||
"top_p": 0.9,
|
||||
}
|
||||
|
||||
mapped_params = config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params={},
|
||||
model=model,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert mapped_params["temperature"] == 0.7
|
||||
assert mapped_params["max_tokens"] == 100
|
||||
assert mapped_params["top_p"] == 0.9
|
||||
|
||||
def test_get_error_class(self):
|
||||
"""Test error class creation"""
|
||||
error = config.get_error_class(
|
||||
error_message="Test error",
|
||||
status_code=400,
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
|
||||
assert isinstance(error, OVHCloudException)
|
||||
assert error.message == "Test error"
|
||||
assert error.status_code == 400
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"Meta-Llama-3_3-70B-Instruct",
|
||||
"Meta-Llama-3_1-70B-Instruct",
|
||||
"Mixtral-8x7B-Instruct-v0.1",
|
||||
"gpt-oss-120b",
|
||||
"some-model-not-in-the-cost-map",
|
||||
],
|
||||
)
|
||||
def test_tools_not_filtered_by_static_model_map(self, model):
|
||||
"""
|
||||
OVHCloud AI Endpoints are OpenAI-compatible; tools/tool_choice must pass
|
||||
through for any model. The server is responsible for rejecting unsupported
|
||||
tool calls — LiteLLM must not strip them based on a stale static catalog.
|
||||
"""
|
||||
|
||||
params = get_optional_params(
|
||||
model=model,
|
||||
custom_llm_provider="ovhcloud",
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {"name": "x", "parameters": {}},
|
||||
}
|
||||
],
|
||||
tool_choice="auto",
|
||||
)
|
||||
|
||||
assert "tools" in params
|
||||
assert "tool_choice" in params
|
||||
|
||||
|
||||
def test_ovhcloud_integration():
|
||||
from litellm import completion
|
||||
|
||||
|
|
@ -285,78 +123,3 @@ def test_ovhcloud_with_custom_base_url():
|
|||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
|
||||
|
||||
class TestOVHCloudReasoningFieldMigration:
|
||||
"""Tests for OVHCloud reasoning_content -> reasoning field migration."""
|
||||
|
||||
def test_streaming_new_reasoning_field(self):
|
||||
"""New `reasoning` field should be mapped to `reasoning_content`."""
|
||||
handler = OVHCloudChatCompletionStreamingHandler(
|
||||
streaming_response=iter([]),
|
||||
sync_stream=True,
|
||||
)
|
||||
chunk = {
|
||||
"id": "test-id",
|
||||
"created": 1234567890,
|
||||
"model": "test-model",
|
||||
"choices": [
|
||||
{
|
||||
"delta": {
|
||||
"role": "assistant",
|
||||
"reasoning": "Let me think...",
|
||||
},
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
}
|
||||
result = handler.chunk_parser(chunk)
|
||||
assert result.choices[0]["delta"]["reasoning_content"] == "Let me think..."
|
||||
|
||||
def test_streaming_legacy_reasoning_content_unchanged(self):
|
||||
"""Legacy `reasoning_content` field should pass through untouched."""
|
||||
handler = OVHCloudChatCompletionStreamingHandler(
|
||||
streaming_response=iter([]),
|
||||
sync_stream=True,
|
||||
)
|
||||
chunk = {
|
||||
"id": "test-id",
|
||||
"created": 1234567890,
|
||||
"model": "test-model",
|
||||
"choices": [
|
||||
{
|
||||
"delta": {
|
||||
"role": "assistant",
|
||||
"reasoning_content": "Already correct field.",
|
||||
},
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
}
|
||||
result = handler.chunk_parser(chunk)
|
||||
assert result.choices[0]["delta"]["reasoning_content"] == "Already correct field."
|
||||
|
||||
def test_streaming_both_fields_legacy_wins(self):
|
||||
"""When both fields present, existing `reasoning_content` is not overwritten."""
|
||||
handler = OVHCloudChatCompletionStreamingHandler(
|
||||
streaming_response=iter([]),
|
||||
sync_stream=True,
|
||||
)
|
||||
chunk = {
|
||||
"id": "test-id",
|
||||
"created": 1234567890,
|
||||
"model": "test-model",
|
||||
"choices": [
|
||||
{
|
||||
"delta": {
|
||||
"reasoning": "new field",
|
||||
"reasoning_content": "legacy field",
|
||||
},
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
}
|
||||
result = handler.chunk_parser(chunk)
|
||||
assert result.choices[0]["delta"]["reasoning_content"] == "legacy field"
|
||||
|
||||
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue