mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Merge remote-tracking branch 'origin/main' into litellm_lit6852_cli_session_spend_key
This commit is contained in:
commit
b28c55450a
189 changed files with 2540 additions and 973 deletions
|
|
@ -31,7 +31,7 @@ while IFS= read -r file || [ -n "$file" ]; do
|
|||
case "$file" in
|
||||
model_prices_and_context_window.json | litellm/model_prices_and_context_window_backup.json | model_prices_and_context_window.schema.json)
|
||||
has_cost_map=true ;;
|
||||
tests/test_litellm/* | tests/proxy_unit_tests/*) : ;;
|
||||
tests/test_litellm/* | tests/proxy_unit_tests/* | tests/unit/proxy/*) : ;;
|
||||
*) outside_cost_map_set=true ;;
|
||||
esac
|
||||
done
|
||||
|
|
|
|||
135
.circleci/scripts/unit_selection.sh
Executable file
135
.circleci/scripts/unit_selection.sh
Executable file
|
|
@ -0,0 +1,135 @@
|
|||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
flag="${1:?usage: unit_selection.sh <codecov flag>}"
|
||||
|
||||
legacy_flags=(
|
||||
caching-local
|
||||
enterprise-package
|
||||
enterprise-routing
|
||||
proxy-db-auth-checks
|
||||
proxy-db-budgets
|
||||
proxy-db-custom-logging
|
||||
proxy-db-db-and-spend
|
||||
proxy-db-endpoints-and-responses
|
||||
proxy-db-guardrails-hooks
|
||||
proxy-db-jwt-and-keys
|
||||
proxy-db-key-generation
|
||||
proxy-db-logging-misc
|
||||
proxy-db-proxy-runtime
|
||||
proxy-db-proxy-server-core
|
||||
proxy-db-proxy-utils
|
||||
proxy-extras
|
||||
proxy-infra
|
||||
)
|
||||
|
||||
legacy_paths() {
|
||||
case "$1" in
|
||||
caching-local) echo tests/unit/caching ;;
|
||||
enterprise-package)
|
||||
echo tests/unit/enterprise/integrations
|
||||
echo tests/unit/enterprise/proxy/auth
|
||||
echo tests/unit/enterprise/proxy/guardrails
|
||||
echo tests/unit/enterprise/proxy/hooks
|
||||
echo tests/unit/enterprise/proxy/management_endpoints
|
||||
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/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
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py
|
||||
echo tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py
|
||||
echo tests/unit/enterprise/proxy/test_enterprise_routes.py
|
||||
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 ;;
|
||||
proxy-db-auth-checks)
|
||||
echo tests/unit/proxy/auth/test_auth_checks.py
|
||||
echo tests/unit/proxy/auth/test_user_api_key_auth.py
|
||||
echo tests/unit/proxy/test_deprecated_key_grace_period.py ;;
|
||||
proxy-db-budgets)
|
||||
echo tests/unit/proxy/auth/test_default_end_user_budget_simple.py
|
||||
echo tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py
|
||||
echo tests/unit/proxy/test_zero_cost_model_budget_bypass.py ;;
|
||||
proxy-db-custom-logging)
|
||||
echo tests/unit/proxy/test_custom_callback_input.py
|
||||
echo tests/unit/proxy/test_custom_logger_s3_gcs.py ;;
|
||||
proxy-db-db-and-spend)
|
||||
echo tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py
|
||||
echo tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py
|
||||
echo tests/unit/proxy/db/test_update_daily_tag_spend.py
|
||||
echo tests/unit/proxy/test_db_schema_changes.py
|
||||
echo tests/unit/proxy/test_prisma_client_backoff_retry.py
|
||||
echo tests/unit/proxy/test_update_spend.py
|
||||
echo tests/unit/skills/test_skills_db.py ;;
|
||||
proxy-db-endpoints-and-responses)
|
||||
echo tests/unit/proxy/auth/test_models_fallback_endpoint.py
|
||||
echo tests/unit/proxy/common_utils/test_check_batch_cost.py
|
||||
echo tests/unit/proxy/common_utils/test_check_responses_cost.py
|
||||
echo tests/unit/proxy/common_utils/test_realtime_cache.py
|
||||
echo tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py
|
||||
echo tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py
|
||||
echo tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py
|
||||
echo tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py
|
||||
echo tests/unit/proxy/response_polling/test_response_polling_handler.py
|
||||
echo tests/unit/proxy/test_custom_tokenizer_bug.py
|
||||
echo tests/unit/proxy/test_get_favicon.py
|
||||
echo tests/unit/proxy/test_get_image.py
|
||||
echo tests/unit/proxy/test_prompt_test_endpoint.py
|
||||
echo tests/unit/proxy/test_reducto_ocr_route.py
|
||||
echo tests/unit/proxy/test_response_polling_pre_call_checks.py
|
||||
echo tests/unit/proxy/test_ui_path_detection.py ;;
|
||||
proxy-db-guardrails-hooks)
|
||||
echo tests/unit/proxy/hooks/test_banned_keyword_list.py
|
||||
echo tests/unit/proxy/test_proxy_setting_guardrails.py
|
||||
echo tests/unit/proxy/test_unit_test_proxy_hooks.py ;;
|
||||
proxy-db-jwt-and-keys)
|
||||
echo tests/unit/proxy/auth/test_jwt.py
|
||||
echo tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py
|
||||
echo tests/unit/proxy/test_proxy_custom_auth.py ;;
|
||||
proxy-db-key-generation) echo tests/unit/proxy/management_endpoints/test_key_generate_prisma.py ;;
|
||||
proxy-db-logging-misc)
|
||||
echo tests/unit/proxy/management_helpers/test_audit_logs_proxy.py
|
||||
echo tests/unit/proxy/spend_tracking/test_search_api_logging.py
|
||||
echo tests/unit/proxy/test_proxy_reject_logging.py ;;
|
||||
proxy-db-proxy-runtime)
|
||||
echo tests/unit/proxy/auth/test_multipart_bypass_repro.py
|
||||
echo tests/unit/proxy/auth/test_proxy_routes.py
|
||||
echo tests/unit/proxy/middleware/test_request_size_limit_middleware.py
|
||||
echo tests/unit/proxy/test_proxy_config_unit_test.py
|
||||
echo tests/unit/proxy/test_proxy_token_counter.py
|
||||
echo tests/unit/proxy/test_server_root_path.py ;;
|
||||
proxy-db-proxy-server-core)
|
||||
echo tests/unit/proxy/test_aproxy_startup.py
|
||||
echo tests/unit/proxy/test_proxy_server.py ;;
|
||||
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 ;;
|
||||
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
|
||||
esac
|
||||
}
|
||||
|
||||
expand() {
|
||||
while read -r path; do
|
||||
if [ -d "$path" ]; then
|
||||
find "$path" -name 'test_*.py'
|
||||
elif [ -f "$path" ]; then
|
||||
echo "$path"
|
||||
else
|
||||
echo "unit_selection.sh: $path does not exist" >&2
|
||||
exit 1
|
||||
fi
|
||||
done
|
||||
}
|
||||
|
||||
if [ "$flag" = unit ]; then
|
||||
comm -23 \
|
||||
<(find tests/unit -name 'test_*.py' | sort) \
|
||||
<(for legacy in "${legacy_flags[@]}"; do legacy_paths "$legacy"; done | expand | sort)
|
||||
exit 0
|
||||
fi
|
||||
|
||||
legacy_paths "$flag" | expand | sort
|
||||
|
|
@ -74,6 +74,7 @@ commands:
|
|||
steps:
|
||||
- run:
|
||||
name: Install Codecov CLI (pinned v11.3.1)
|
||||
when: always
|
||||
command: |
|
||||
curl -sSLf -o /tmp/codecov https://cli.codecov.io/v11.3.1/linux/codecov
|
||||
curl -sSLf -o /tmp/codecov.SHA256SUM https://cli.codecov.io/v11.3.1/linux/codecov.SHA256SUM
|
||||
|
|
@ -90,7 +91,6 @@ commands:
|
|||
uv run --no-sync python -c "import litellm_enterprise; print('litellm-enterprise OK:', litellm_enterprise.__file__)"
|
||||
setup_test_deps:
|
||||
steps:
|
||||
- checkout
|
||||
- install_uv
|
||||
- install_rust
|
||||
- restore_cache:
|
||||
|
|
@ -165,42 +165,59 @@ commands:
|
|||
jobs:
|
||||
unit:
|
||||
parameters:
|
||||
tests_path:
|
||||
type: string
|
||||
default: tests/unit
|
||||
flag:
|
||||
type: string
|
||||
default: unit
|
||||
shards:
|
||||
type: integer
|
||||
default: 6
|
||||
workers:
|
||||
type: integer
|
||||
default: 4
|
||||
dist:
|
||||
type: string
|
||||
default: loadscope
|
||||
base_ref:
|
||||
type: string
|
||||
default: ""
|
||||
pull_request_url:
|
||||
type: string
|
||||
default: ""
|
||||
reruns:
|
||||
type: integer
|
||||
default: 0
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
resource_class: large
|
||||
working_directory: ~/project
|
||||
parallelism: << parameters.shards >>
|
||||
environment:
|
||||
COVERAGE_CORE: sysmon
|
||||
LITELLM_LOCAL_MODEL_COST_MAP: "True"
|
||||
steps:
|
||||
- setup_test_deps
|
||||
- checkout
|
||||
- skip_unless_relevant:
|
||||
base_ref: << parameters.base_ref >>
|
||||
pull_request_url: << parameters.pull_request_url >>
|
||||
- setup_test_deps
|
||||
- run:
|
||||
name: "Run << parameters.tests_path >> shard"
|
||||
name: "Run << parameters.flag >> shard"
|
||||
no_output_timeout: 20m
|
||||
command: |
|
||||
mkdir -p test-results/<< parameters.flag >>
|
||||
mapfile -t files < <(find << parameters.tests_path >> -name 'test_*.py' | sort | circleci tests split --split-by=timings --timings-type=filename)
|
||||
if [ "${#files[@]}" -eq 0 ]; then echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.tests_path >> files; nothing to run"; exit 0; fi
|
||||
selection="$(bash .circleci/scripts/unit_selection.sh << parameters.flag >>)" || { echo "unit_selection.sh failed for << parameters.flag >>"; exit 1; }
|
||||
[ -n "${selection}" ] || { echo "unit_selection.sh produced no files for << parameters.flag >>"; exit 1; }
|
||||
shard="$(printf '%s\n' "${selection}" | circleci tests split --split-by=timings --timings-type=filename)" || { echo "circleci tests split failed for << parameters.flag >>"; exit 1; }
|
||||
[ -n "${shard}" ] || { echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.flag >> files; nothing to run"; exit 0; }
|
||||
mapfile -t files < <(printf '%s\n' "${shard}")
|
||||
xdist_args=()
|
||||
if [ "<< parameters.workers >>" -gt 0 ]; then xdist_args=(-n << parameters.workers >> --dist=<< parameters.dist >>); fi
|
||||
rerun_args=(-p no:rerunfailures)
|
||||
if [ "<< parameters.reruns >>" -gt 0 ]; then rerun_args=(--reruns << parameters.reruns >> --reruns-delay 1 --rerun-except "from pytest-timeout"); fi
|
||||
test_env=(PATH="$PATH" HOME="$HOME" CI=true COVERAGE_CORE="$COVERAGE_CORE" LITELLM_LOCAL_MODEL_COST_MAP="$LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
set +e
|
||||
uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml
|
||||
env -i "${test_env[@]}" \
|
||||
uv run --no-sync pytest "${files[@]}" "${rerun_args[@]}" -p no:pytest-retry --timeout=90 "${xdist_args[@]}" --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml
|
||||
status=$?
|
||||
set -e
|
||||
if [ "$status" -eq 5 ]; then echo "pytest collected no tests from the shard; passing"; exit 0; fi
|
||||
|
|
@ -224,6 +241,7 @@ jobs:
|
|||
resource_class: large
|
||||
working_directory: ~/project
|
||||
steps:
|
||||
- checkout
|
||||
- setup_test_deps
|
||||
- run:
|
||||
name: Checkout litellm-docs
|
||||
|
|
@ -250,16 +268,17 @@ jobs:
|
|||
resource_class: large
|
||||
working_directory: ~/project
|
||||
steps:
|
||||
- setup_test_deps
|
||||
- checkout
|
||||
- skip_unless_relevant:
|
||||
base_ref: << parameters.base_ref >>
|
||||
pull_request_url: << parameters.pull_request_url >>
|
||||
- setup_test_deps
|
||||
- start_postgres:
|
||||
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
|
||||
- start_redis
|
||||
- run:
|
||||
name: Run owned integration contracts
|
||||
command: bash .circleci/scripts/run_integration.sh << parameters.suite >>
|
||||
command: env -i PATH="$PATH" HOME="$HOME" CIRCLE_SHA1="$CIRCLE_SHA1" CIRCLE_WORKFLOW_ID="$CIRCLE_WORKFLOW_ID" bash .circleci/scripts/run_integration.sh << parameters.suite >>
|
||||
no_output_timeout: 15m
|
||||
- run:
|
||||
name: Stop owned database and Redis
|
||||
|
|
@ -282,6 +301,53 @@ workflows:
|
|||
- unit:
|
||||
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-<< matrix.flag >>
|
||||
shards: 1
|
||||
workers: 2
|
||||
reruns: 2
|
||||
matrix:
|
||||
parameters:
|
||||
flag: [caching-local, proxy-extras, enterprise-routing]
|
||||
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-<< matrix.flag >>
|
||||
shards: 1
|
||||
reruns: 2
|
||||
matrix:
|
||||
parameters:
|
||||
flag:
|
||||
- enterprise-package
|
||||
- proxy-infra
|
||||
- proxy-db-auth-checks
|
||||
- proxy-db-jwt-and-keys
|
||||
- proxy-db-proxy-server-core
|
||||
- proxy-db-proxy-runtime
|
||||
- proxy-db-custom-logging
|
||||
- proxy-db-logging-misc
|
||||
- proxy-db-db-and-spend
|
||||
- proxy-db-guardrails-hooks
|
||||
- proxy-db-budgets
|
||||
- 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-proxy-db-proxy-utils
|
||||
flag: proxy-db-proxy-utils
|
||||
shards: 1
|
||||
reruns: 2
|
||||
dist: worksteal
|
||||
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-key-generation
|
||||
flag: proxy-db-key-generation
|
||||
shards: 1
|
||||
workers: 0
|
||||
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 "" >>
|
||||
- documentation
|
||||
- integration:
|
||||
name: integration-<< matrix.suite >>
|
||||
|
|
|
|||
6
.github/pull_request_template.md
vendored
6
.github/pull_request_template.md
vendored
|
|
@ -6,7 +6,8 @@
|
|||
|
||||
## TLDR
|
||||
|
||||
<!-- Fill in the bullets below and keep each one short and concrete: one line per bullet, roughly 10 words max -->
|
||||
<!-- Fill in the bullets below and keep each one short and concrete: one line per bullet, roughly 10 words max
|
||||
If the PR intentionally changes what existing users see or how a screen behaves, add a line under the bullets that starts "Intentional product change:" describing what changes, why, and what users lose. Reviewers must never have to infer a deliberate UX change from the diff -->
|
||||
|
||||
Problem this solves:
|
||||
|
||||
|
|
@ -28,7 +29,8 @@ How it solves it:
|
|||
No LiteLLM internals: never name functions, files, DB tables, config classes, hooks, callbacks, or code paths. "The upload hands back an ID that looks like OpenAI's own `file-abc123` instead of the scrambled one the gateway returned" is right, "no managed-file row was registered" is wrong
|
||||
Keep the two lists step-for-step identical until they diverge, so the changed step is obvious
|
||||
If the bug had a security or authorization consequence, end each list with what another user could or could no longer do
|
||||
Regenerate this section whenever new commits change the PR's behavior, so it never describes an older revision
|
||||
Regenerate this section, screenshots included, whenever new commits change the PR's behavior, so it never describes an older revision
|
||||
If the PR changes what an Admin UI page shows, embed a before and an after screenshot of that page right after its list, taken at the same URL on the same data, with the rows, fields, or controls that changed boxed in red so a reader spots the difference without reading the steps. These are the UI screenshots for Screenshots / Proof of Fix too: embed them once here and have that section's Before and After steps point back to them instead of repeating the images
|
||||
|
||||
Example:
|
||||
|
||||
|
|
|
|||
13
.github/scripts/assert_ci_coverage.py
vendored
13
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -34,7 +34,6 @@ GLOB_CHARS = frozenset("*?")
|
|||
# tests has to be named by some shard or it runs nowhere. A child listed here is
|
||||
# itself decomposed one level deeper and is checked through its own entry.
|
||||
SHARDED_ROOTS: tuple[str, ...] = (
|
||||
"tests/proxy_unit_tests",
|
||||
"tests/test_litellm",
|
||||
"tests/test_litellm/proxy",
|
||||
)
|
||||
|
|
@ -120,6 +119,13 @@ def _invoked_test_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
|
|||
)
|
||||
|
||||
|
||||
def _unit_selection_tokens(repo_root: pathlib.Path = REPO_ROOT) -> frozenset[str]:
|
||||
script: Final = repo_root / ".circleci/scripts/unit_selection.sh"
|
||||
if not script.is_file():
|
||||
return frozenset()
|
||||
return frozenset(match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(_uncommented(script.read_text())))
|
||||
|
||||
|
||||
def _built_dockerfile_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
|
||||
return frozenset(
|
||||
match.group(0)
|
||||
|
|
@ -611,7 +617,10 @@ def main() -> int:
|
|||
scalars = _all_scalars()
|
||||
|
||||
integration_paths, ownership_findings = _integration_ownership()
|
||||
test_findings = _uncovered_tests(allowlist, _invoked_test_tokens(scalars) | integration_paths) + ownership_findings
|
||||
test_findings = (
|
||||
_uncovered_tests(allowlist, _invoked_test_tokens(scalars) | _unit_selection_tokens() | integration_paths)
|
||||
+ ownership_findings
|
||||
)
|
||||
dockerfile_findings = _uncovered_dockerfiles(allowlist, _built_dockerfile_tokens(scalars))
|
||||
stale_findings = _stale_allowlist_paths(allowlist, test_files=_test_files(), dockerfiles=_dockerfiles())
|
||||
|
||||
|
|
|
|||
33
.github/workflows/_test-unit-base.yml
vendored
33
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -13,6 +13,15 @@ on:
|
|||
have its path existence-checked like any other token.
|
||||
required: true
|
||||
type: string
|
||||
fork-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.
|
||||
required: false
|
||||
type: string
|
||||
default: ""
|
||||
workers:
|
||||
description: "Number of pytest-xdist workers"
|
||||
required: false
|
||||
|
|
@ -92,6 +101,7 @@ jobs:
|
|||
pull-requests: read
|
||||
outputs:
|
||||
decision: ${{ steps.changes.outputs.decision }}
|
||||
has-coverage: ${{ steps.tests.outputs.has-coverage }}
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
|
|
@ -160,10 +170,13 @@ jobs:
|
|||
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Run tests
|
||||
id: tests
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
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 }}
|
||||
MAX_FAILURES: ${{ inputs.max-failures }}
|
||||
WORKERS: ${{ inputs.workers }}
|
||||
RERUNS: ${{ inputs.reruns }}
|
||||
|
|
@ -171,9 +184,18 @@ jobs:
|
|||
DIST: ${{ inputs.dist }}
|
||||
COVERAGE_CORE: sysmon
|
||||
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' ' ')"
|
||||
fi
|
||||
if [ -z "${selection// /}" ]; then
|
||||
echo "shard selection is empty on this event (CircleCI flag ${FORK_FLAG:-none} owns it); nothing to run"
|
||||
exit 0
|
||||
fi
|
||||
pytest_args=()
|
||||
existing_paths=0
|
||||
for token in ${TEST_PATH:?}; do
|
||||
for token in ${selection}; do
|
||||
case "${token}" in
|
||||
-*) pytest_args+=("${token}") ;;
|
||||
*)
|
||||
|
|
@ -187,7 +209,7 @@ jobs:
|
|||
esac
|
||||
done
|
||||
if [ "${existing_paths}" -eq 0 ]; then
|
||||
echo "No path in TEST_PATH exists (${TEST_PATH}); nothing to run"
|
||||
echo "No path in the selection exists (${selection}); nothing to run"
|
||||
exit 0
|
||||
fi
|
||||
xdist_args=()
|
||||
|
|
@ -209,8 +231,11 @@ jobs:
|
|||
--cov-config=pyproject.toml
|
||||
status=$?
|
||||
set -e
|
||||
if [ -f coverage.xml ]; then
|
||||
echo "has-coverage=true" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
if [ "$status" -eq 5 ]; then
|
||||
echo "pytest collected no tests from ${TEST_PATH}; passing"
|
||||
echo "pytest collected no tests from ${selection}; passing"
|
||||
exit 0
|
||||
fi
|
||||
exit "$status"
|
||||
|
|
@ -226,7 +251,7 @@ jobs:
|
|||
upload-coverage:
|
||||
name: Upload coverage to Codecov
|
||||
needs: run
|
||||
if: always() && needs.run.outputs.decision != 'skip'
|
||||
if: always() && needs.run.outputs.decision != 'skip' && needs.run.outputs.has-coverage == 'true'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
13
.github/workflows/compat-matrix-image.yml
vendored
13
.github/workflows/compat-matrix-image.yml
vendored
|
|
@ -4,6 +4,7 @@ on:
|
|||
pull_request:
|
||||
paths:
|
||||
- tests/e2e/claude_code/cron_vm/**
|
||||
- tests/e2e/claude_code/pr_gate_version_resolver.py
|
||||
- .github/workflows/compat-matrix-image.yml
|
||||
workflow_dispatch:
|
||||
|
||||
|
|
@ -28,6 +29,14 @@ jobs:
|
|||
- name: Build the Render cron image
|
||||
run: docker build -f tests/e2e/claude_code/cron_vm/Dockerfile -t compat-matrix:${{ github.sha }} tests/e2e
|
||||
|
||||
- name: Run the pinned binaries as the cron user
|
||||
- name: Resolve and install the Claude Code CLI as the cron user
|
||||
run: |
|
||||
docker run --rm compat-matrix:${{ github.sha }} bash -c 'set -e; whoami; claude --version; gh --version; uv --version'
|
||||
docker run --rm compat-matrix:${{ github.sha }} bash -c '
|
||||
set -euo pipefail
|
||||
whoami
|
||||
gh --version
|
||||
uv --version
|
||||
version="$(uv run --no-project --python 3.12 python /opt/litellm/tests/e2e/claude_code/pr_gate_version_resolver.py)"
|
||||
/opt/litellm/tests/e2e/claude_code/cron_vm/install_claude_code.sh "${version}" /tmp/claude-cli
|
||||
/tmp/claude-cli/claude --version
|
||||
'
|
||||
|
|
|
|||
97
.github/workflows/test-unit-proxy-db.yml
vendored
97
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -20,6 +20,12 @@ concurrency:
|
|||
# rather than alphabetical letter ranges. Adding a new test file means adding it
|
||||
# to whichever group it belongs to, not reshuffling slices.
|
||||
#
|
||||
# `.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.
|
||||
#
|
||||
# Design targets:
|
||||
# * Every shard runs in <= 7 minutes of wall-clock on the default runner.
|
||||
# Most of a shard's time is pytest plugin load + xdist worker imports +
|
||||
|
|
@ -58,7 +64,7 @@ jobs:
|
|||
proxy-db:
|
||||
needs: assert-shard-coverage
|
||||
# Display only the semantic shard name in the checks UI instead of GHA's
|
||||
# default "proxy-db (key-generation, tests/proxy_unit_tests/…, 0, loadscope, 20)"
|
||||
# default "proxy-db (key-generation, tests/unit/proxy/…, 0, loadscope, 20)"
|
||||
# which includes every matrix field and gets truncated past the test-path.
|
||||
name: ${{ matrix.test-group }}
|
||||
permissions:
|
||||
|
|
@ -71,132 +77,93 @@ jobs:
|
|||
include:
|
||||
# Must run serially — event-loop conflict with the logging worker.
|
||||
- test-group: key-generation
|
||||
test-path: "tests/proxy_unit_tests/test_key_generate_prisma.py"
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-key-generation
|
||||
workers: 0
|
||||
dist: loadscope
|
||||
timeout: 20
|
||||
|
||||
# ---- auth: split into 2 shards ----
|
||||
- test-group: auth-checks
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_auth_checks.py
|
||||
tests/proxy_unit_tests/test_user_api_key_auth.py
|
||||
tests/proxy_unit_tests/test_deprecated_key_grace_period.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-auth-checks
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: jwt-and-keys
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_jwt.py
|
||||
tests/proxy_unit_tests/test_jwt_key_mapping.py
|
||||
tests/proxy_unit_tests/test_proxy_custom_auth.py
|
||||
tests/proxy_unit_tests/test_key_generate_dynamodb.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-jwt-and-keys
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
# ---- test_proxy_utils.py, single shard, worksteal distribution ----
|
||||
- test-group: proxy-utils
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_utils.py"
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-proxy-utils
|
||||
workers: 4
|
||||
dist: worksteal
|
||||
timeout: 15
|
||||
|
||||
# ---- proxy server: split into 2 shards ----
|
||||
- test-group: proxy-server-core
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_proxy_server.py
|
||||
tests/proxy_unit_tests/test_aproxy_startup.py
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py"
|
||||
fork-flag: proxy-db-proxy-server-core
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: proxy-runtime
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_proxy_config_unit_test.py
|
||||
tests/proxy_unit_tests/test_proxy_routes.py
|
||||
tests/proxy_unit_tests/test_server_root_path.py
|
||||
tests/proxy_unit_tests/test_proxy_pass_user_config.py
|
||||
tests/proxy_unit_tests/test_proxy_token_counter.py
|
||||
tests/proxy_unit_tests/test_request_size_limit_middleware.py
|
||||
tests/proxy_unit_tests/test_multipart_bypass_repro.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-proxy-runtime
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
# ---- logging: split into 2 shards ----
|
||||
- test-group: custom-logging
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_custom_callback_input.py
|
||||
tests/proxy_unit_tests/test_custom_logger_s3_gcs.py
|
||||
tests/proxy_unit_tests/test_proxy_custom_logger.py
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_custom_logger.py"
|
||||
fork-flag: proxy-db-custom-logging
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: logging-misc
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_proxy_reject_logging.py
|
||||
tests/proxy_unit_tests/test_audit_logs_proxy.py
|
||||
tests/proxy_unit_tests/test_search_api_logging.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-logging-misc
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
- test-group: db-and-spend
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_prisma_client_backoff_retry.py
|
||||
tests/proxy_unit_tests/test_db_schema_changes.py
|
||||
tests/proxy_unit_tests/test_e2e_pod_lock_manager.py
|
||||
tests/proxy_unit_tests/test_skills_db.py
|
||||
tests/proxy_unit_tests/test_update_daily_tag_spend.py
|
||||
tests/proxy_unit_tests/test_update_spend.py
|
||||
tests/proxy_unit_tests/test_proxy_encrypt_decrypt.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-db-and-spend
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
# ---- guardrails + budget + hooks: split into 2 ----
|
||||
- test-group: guardrails-hooks
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_proxy_setting_guardrails.py
|
||||
tests/proxy_unit_tests/test_banned_keyword_list.py
|
||||
tests/proxy_unit_tests/test_unit_test_proxy_hooks.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-guardrails-hooks
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: budgets
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_default_end_user_budget_simple.py
|
||||
tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py
|
||||
tests/proxy_unit_tests/test_zero_cost_model_budget_bypass.py
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-budgets
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
- test-group: endpoints-and-responses
|
||||
test-path: >-
|
||||
tests/proxy_unit_tests/test_blog_posts_endpoint.py
|
||||
tests/proxy_unit_tests/test_models_fallback_endpoint.py
|
||||
tests/proxy_unit_tests/test_google_endpoint_routing.py
|
||||
tests/proxy_unit_tests/test_google_gemini_proxy_request.py
|
||||
tests/proxy_unit_tests/test_gemini_agents_endpoints.py
|
||||
tests/proxy_unit_tests/test_get_favicon.py
|
||||
tests/proxy_unit_tests/test_get_image.py
|
||||
tests/proxy_unit_tests/test_reducto_ocr_route.py
|
||||
tests/proxy_unit_tests/test_ui_path_detection.py
|
||||
tests/proxy_unit_tests/test_prompt_test_endpoint.py
|
||||
tests/proxy_unit_tests/test_check_batch_cost.py
|
||||
tests/proxy_unit_tests/test_check_responses_cost.py
|
||||
tests/proxy_unit_tests/test_response_polling_handler.py
|
||||
tests/proxy_unit_tests/test_response_polling_pre_call_checks.py
|
||||
tests/proxy_unit_tests/test_realtime_cache.py
|
||||
tests/proxy_unit_tests/test_proxy_exception_mapping.py
|
||||
tests/proxy_unit_tests/test_custom_tokenizer_bug.py
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_exception_mapping.py"
|
||||
fork-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 }}
|
||||
workers: ${{ matrix.workers }}
|
||||
reruns: 2
|
||||
timeout-minutes: ${{ matrix.timeout }}
|
||||
|
|
|
|||
30
.github/workflows/test-unit.yml
vendored
30
.github/workflows/test-unit.yml
vendored
|
|
@ -31,10 +31,14 @@ concurrency:
|
|||
# number, so a partially-specified entry would fail the call rather than fall
|
||||
# back to the default.
|
||||
#
|
||||
# tests/proxy_unit_tests keeps its own caller (test-unit-proxy-db.yml): it is
|
||||
# already a matrix and carries a shard-coverage guard that reads that file by
|
||||
# name. Folding it in here is a follow-up, together with generalising that guard
|
||||
# into assert_ci_coverage.py.
|
||||
# tests/unit/proxy keeps its own caller (test-unit-proxy-db.yml): it is already
|
||||
# a matrix and carries a shard-coverage guard that reads that file by name.
|
||||
# Folding it in here is a follow-up, together with generalising that guard into
|
||||
# assert_ci_coverage.py.
|
||||
#
|
||||
# `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.
|
||||
jobs:
|
||||
unit:
|
||||
name: ${{ matrix.shard }}
|
||||
|
|
@ -65,10 +69,10 @@ jobs:
|
|||
- shard: enterprise-routing
|
||||
artifact-name: enterprise-routing
|
||||
test-path: >-
|
||||
tests/test_litellm/enterprise
|
||||
tests/test_litellm/google_genai
|
||||
tests/test_litellm/router_utils
|
||||
tests/test_litellm/router_strategy
|
||||
fork-flag: enterprise-routing
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -200,7 +204,7 @@ jobs:
|
|||
tests/test_litellm/proxy/types_utils
|
||||
tests/test_litellm/proxy/logging_endpoints
|
||||
tests/test_litellm/proxy/test_*.py
|
||||
tests/test_gateway
|
||||
fork-flag: proxy-infra
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -208,11 +212,8 @@ jobs:
|
|||
|
||||
- shard: caching-local
|
||||
artifact-name: caching-local
|
||||
test-path: >-
|
||||
tests/local_testing/test_cache_preset_key.py
|
||||
tests/local_testing/test_caching_handler.py
|
||||
tests/local_testing/test_responses_stream_cache_keys.py
|
||||
tests/local_testing/test_unit_test_caching.py
|
||||
test-path: ""
|
||||
fork-flag: caching-local
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -220,7 +221,8 @@ jobs:
|
|||
|
||||
- shard: proxy-extras
|
||||
artifact-name: proxy-extras
|
||||
test-path: "tests/litellm-proxy-extras"
|
||||
test-path: ""
|
||||
fork-flag: proxy-extras
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -228,7 +230,8 @@ jobs:
|
|||
|
||||
- shard: enterprise-package
|
||||
artifact-name: enterprise-package
|
||||
test-path: "tests/enterprise"
|
||||
test-path: ""
|
||||
fork-flag: enterprise-package
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -247,6 +250,7 @@ jobs:
|
|||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: ${{ matrix.test-path }}
|
||||
fork-flag: ${{ matrix.fork-flag || '' }}
|
||||
workers: ${{ matrix.workers }}
|
||||
reruns: ${{ matrix.reruns }}
|
||||
timeout-minutes: ${{ matrix.timeout-minutes }}
|
||||
|
|
|
|||
12
Makefile
12
Makefile
|
|
@ -51,8 +51,8 @@ help:
|
|||
@echo " make test-unit-core-utils - Run core utils tests (~32 files)"
|
||||
@echo " make test-unit-other - Run other tests (caching, responses, etc., ~69 files)"
|
||||
@echo " make test-unit-root - Run root-level tests (~34 files)"
|
||||
@echo " make test-proxy-unit-a - Run proxy_unit_tests (a-o, ~20 files)"
|
||||
@echo " make test-proxy-unit-b - Run proxy_unit_tests (p-z, ~28 files)"
|
||||
@echo " make test-proxy-unit-a - Run tests/unit/proxy (a-o)"
|
||||
@echo " make test-proxy-unit-b - Run tests/unit/proxy (p-z)"
|
||||
@echo " make test-integration - Run integration tests"
|
||||
@echo " make test-unit-helm - Run helm unit tests"
|
||||
@echo " make test-rust-extension - Build the Rust extension and run its public Python tests"
|
||||
|
|
@ -332,17 +332,17 @@ 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/test_litellm/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/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
|
||||
|
||||
test-unit-root: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
# Proxy unit tests (tests/proxy_unit_tests split alphabetically)
|
||||
# Proxy unit tests (tests/unit/proxy split alphabetically)
|
||||
test-proxy-unit-a: install-test-deps
|
||||
$(UV_RUN) pytest tests/proxy_unit_tests/test_[a-o]*.py --tb=short -vv -n 2 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/proxy --ignore-glob='tests/unit/proxy/test_[p-z]*.py' --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-proxy-unit-b: install-test-deps
|
||||
$(UV_RUN) pytest tests/proxy_unit_tests/test_[p-z]*.py --tb=short -vv -n 2 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/proxy/test_[p-z]*.py tests/unit/skills --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-integration: install-test-deps
|
||||
$(UV_RUN) pytest tests/ -k "not test_litellm"
|
||||
|
|
|
|||
|
|
@ -421,6 +421,24 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
): # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm
|
||||
pass
|
||||
|
||||
async def async_filter_listed_models(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
model_names: Sequence[str],
|
||||
) -> Sequence[str]:
|
||||
"""Runs on the model listing routes (`/v1/models`, `/v1/models/{id}`, `/model/info`,
|
||||
`/model_group/info`) with the public model names the route would otherwise return, so a
|
||||
lookup of one model may offer just that name: decide per name, never by position in the
|
||||
sequence. Return the names to keep as a sequence of strings; a name left out disappears
|
||||
from every listing, any alias of it offered in the same call goes with it, and
|
||||
`/v1/models/{id}` answers 404 for it, exactly as for a model that does not exist. Names
|
||||
outside `model_names` are ignored, so a callback can only narrow the listing, never widen
|
||||
it. Under `use_team_public_model_name: false`, `/v1/models` and `/model_group/info` list a
|
||||
team model by its internal routing name while `/model/info` keeps its public name, so hide
|
||||
both names to hide it on every route.
|
||||
"""
|
||||
return model_names
|
||||
|
||||
async def async_post_call_response_headers_hook(
|
||||
self,
|
||||
data: dict,
|
||||
|
|
|
|||
|
|
@ -369,7 +369,7 @@ model LiteLLM_SkillsTable {
|
|||
Run the tests:
|
||||
|
||||
```bash
|
||||
pytest tests/proxy_unit_tests/test_skills_db.py -v
|
||||
pytest tests/unit/skills/test_skills_db.py -v
|
||||
```
|
||||
|
||||
Tests cover:
|
||||
|
|
|
|||
|
|
@ -6092,7 +6092,7 @@
|
|||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"deprecation_date": "2027-06-25",
|
||||
"deprecation_date": "2027-07-31",
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_image_token": 8e-07,
|
||||
"input_cost_per_token": 6e-07,
|
||||
|
|
@ -26212,7 +26212,6 @@
|
|||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
"input_cost_per_token_flex": 1.5e-07,
|
||||
"input_cost_per_token_priority": 5.4e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 32768,
|
||||
|
|
@ -49416,7 +49415,6 @@
|
|||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
"input_cost_per_token_flex": 1.5e-07,
|
||||
"input_cost_per_token_priority": 5.4e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 32768,
|
||||
|
|
|
|||
|
|
@ -210,6 +210,25 @@ class ProxyInitializationHelpers:
|
|||
response: Final = httpx.get(url=f"http://{host}:{port}/health")
|
||||
print(json.dumps(response.json(), indent=4))
|
||||
|
||||
@staticmethod
|
||||
def _run_config_validation(config: str | None) -> None:
|
||||
if config is None:
|
||||
raise click.UsageError("--validate_config requires --config <path>")
|
||||
import asyncio
|
||||
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
async def _load() -> int:
|
||||
_, model_list, _ = await ProxyConfig().load_config(router=None, config_file_path=config)
|
||||
return len(model_list)
|
||||
|
||||
try:
|
||||
model_count: Final = asyncio.run(_load())
|
||||
except Exception as error:
|
||||
click.echo(f"LiteLLM: config validation failed: {error}", err=True)
|
||||
raise click.exceptions.Exit(1) from error
|
||||
click.echo(f"LiteLLM: config OK ({model_count} models)")
|
||||
|
||||
@staticmethod
|
||||
def _run_test_chat_completion(
|
||||
host: str,
|
||||
|
|
@ -887,6 +906,12 @@ class ProxyInitializationHelpers:
|
|||
default=False,
|
||||
help="Skip starting the server after setup (useful for migrations only)",
|
||||
)
|
||||
@click.option(
|
||||
"--validate_config",
|
||||
is_flag=True,
|
||||
default=False,
|
||||
help="Load and validate the config file (including mcp_servers) without starting the server, then exit. Exit code 1 on any config error.",
|
||||
)
|
||||
@click.option(
|
||||
"--keepalive_timeout",
|
||||
default=None,
|
||||
|
|
@ -1027,6 +1052,7 @@ def run_server(
|
|||
log_config,
|
||||
use_prisma_db_push: bool,
|
||||
skip_server_startup,
|
||||
validate_config: bool,
|
||||
keepalive_timeout,
|
||||
timeout_worker_healthcheck,
|
||||
max_requests_before_restart,
|
||||
|
|
@ -1069,6 +1095,9 @@ def run_server(
|
|||
if version is True:
|
||||
ProxyInitializationHelpers._echo_litellm_version()
|
||||
return
|
||||
if validate_config is True:
|
||||
ProxyInitializationHelpers._run_config_validation(config)
|
||||
return
|
||||
if model and "ollama" in model and api_base is None:
|
||||
ProxyInitializationHelpers._run_ollama_serve()
|
||||
if health is True:
|
||||
|
|
|
|||
|
|
@ -152,6 +152,7 @@ from litellm.router_utils.auto_router_tuning_baseline import (
|
|||
snapshot_tuning_baselines,
|
||||
tuning_limit_violation,
|
||||
)
|
||||
from litellm.router_utils.common_utils import resolve_model_group_alias
|
||||
from litellm.router_utils.routing_groups import parse_routing_groups
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.types.utils import (
|
||||
|
|
@ -11104,6 +11105,40 @@ class ProxyStartupEvent:
|
|||
|
||||
|
||||
#### API ENDPOINTS ####
|
||||
async def _names_hidden_by_listing_callbacks(
|
||||
user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
|
||||
) -> frozenset[str]:
|
||||
hidden: Final = await proxy_logging_obj.hidden_by_listing_callbacks(user_api_key_dict, model_names)
|
||||
if not hidden or llm_router is None:
|
||||
return hidden
|
||||
aliases: Final = llm_router.model_group_alias
|
||||
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, general_settings)
|
||||
return hidden | frozenset(
|
||||
alias
|
||||
for alias in aliases
|
||||
if (target := resolve_model_group_alias(aliases, alias)) is not None
|
||||
and internal_to_public.get(target, target) in hidden
|
||||
)
|
||||
|
||||
|
||||
async def _entries_kept_by_listing_callbacks(
|
||||
entries: Sequence[tuple[str, str]], user_api_key_dict: UserAPIKeyAuth
|
||||
) -> tuple[tuple[str, str], ...]:
|
||||
hidden: Final = await _names_hidden_by_listing_callbacks(
|
||||
user_api_key_dict, tuple(response_id for response_id, _ in entries)
|
||||
)
|
||||
if not hidden:
|
||||
return tuple(entries)
|
||||
return tuple(entry for entry in entries if entry[0] not in hidden)
|
||||
|
||||
|
||||
async def _deployment_hidden_by_listing_callbacks(deployment: Deployment, user_api_key_dict: UserAPIKeyAuth) -> bool:
|
||||
listed_name: Final = _translate_model_name_for_response(deployment.model_dump(exclude_none=True)).get("model_name")
|
||||
if not isinstance(listed_name, str):
|
||||
return False
|
||||
return listed_name in await _names_hidden_by_listing_callbacks(user_api_key_dict, (listed_name,))
|
||||
|
||||
|
||||
@router.get("/v1/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"])
|
||||
@router.get(
|
||||
"/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"]
|
||||
|
|
@ -11254,7 +11289,9 @@ async def model_list(
|
|||
# The internal routing key drives the metadata/fallback lookup, while the
|
||||
# public name is what the client sees as the model id.
|
||||
model_data = []
|
||||
admin_entries: Final = TeamModelNameTranslator.listing_entries(all_models, llm_router, settings)
|
||||
admin_entries: Final = await _entries_kept_by_listing_callbacks(
|
||||
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), user_api_key_dict
|
||||
)
|
||||
for response_id, lookup_id in admin_entries:
|
||||
model_info = create_model_info_response(
|
||||
model_id=lookup_id,
|
||||
|
|
@ -11310,7 +11347,10 @@ async def model_list(
|
|||
# public name is what the client sees as the model id.
|
||||
model_data = []
|
||||
entries: Final = alias_listing_entries(
|
||||
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), caller_aliases
|
||||
await _entries_kept_by_listing_callbacks(
|
||||
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), user_api_key_dict
|
||||
),
|
||||
caller_aliases,
|
||||
)
|
||||
for response_id, lookup_id in entries:
|
||||
model_info = create_model_info_response(
|
||||
|
|
@ -11404,13 +11444,24 @@ async def model_info(
|
|||
llm_router=llm_router,
|
||||
)
|
||||
hidden_names: Final = blocked_names | unhealthy_names
|
||||
if hidden_names:
|
||||
all_models = [m for m in all_models if m not in hidden_names]
|
||||
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, settings)
|
||||
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(
|
||||
user_api_key_dict,
|
||||
tuple(
|
||||
response_id
|
||||
for response_id, _ in TeamModelNameTranslator.listing_entries(
|
||||
tuple(m for m in all_models if m not in hidden_names), llm_router, settings
|
||||
)
|
||||
),
|
||||
)
|
||||
if hidden_names or callback_hidden_names:
|
||||
all_models = [
|
||||
m for m in all_models if m not in hidden_names and internal_to_public.get(m, m) not in callback_hidden_names
|
||||
]
|
||||
undiscoverable_names: Final = undiscoverable_model_names(
|
||||
all_models, llm_router, user_api_key_dict, team_id or user_api_key_dict.team_id
|
||||
)
|
||||
|
||||
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, settings)
|
||||
aliased_model_id: Final = alias_target(
|
||||
model_id,
|
||||
caller_alias_maps(
|
||||
|
|
@ -15730,7 +15781,7 @@ async def model_info_v1(
|
|||
if litellm_model_id is not None:
|
||||
# user is trying to get specific model from litellm router
|
||||
deployment_info: Final = llm_router.get_deployment(model_id=litellm_model_id)
|
||||
if deployment_info is None:
|
||||
if deployment_info is None or await _deployment_hidden_by_listing_callbacks(deployment_info, user_api_key_dict):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": f"Model id = {litellm_model_id} not found on litellm proxy"},
|
||||
|
|
@ -15819,10 +15870,17 @@ async def model_info_v1(
|
|||
general_settings=general_settings,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
visible_models: Final = discoverable_rows(
|
||||
servable_rows: Final = discoverable_rows(
|
||||
(model for model in all_models if model.get("model_name") not in hidden_names),
|
||||
user_api_key_dict,
|
||||
)
|
||||
listed_names: Final = tuple(
|
||||
dict.fromkeys(name for model in servable_rows if isinstance(name := model.get("model_name"), str))
|
||||
)
|
||||
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(user_api_key_dict, listed_names)
|
||||
visible_models: Final = tuple(
|
||||
model for model in servable_rows if model.get("model_name") not in callback_hidden_names
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug("all_models: %s", visible_models)
|
||||
return _model_info_json_response(visible_models)
|
||||
|
|
@ -15871,7 +15929,7 @@ async def model_deprecations(
|
|||
|
||||
|
||||
def _get_model_group_info(
|
||||
llm_router: Router, all_models_str: list[str], model_group: str | None
|
||||
llm_router: Router, all_models_str: Sequence[str], model_group: str | None
|
||||
) -> list[ModelGroupInfoProxy]:
|
||||
model_groups: Final[list[ModelGroupInfoProxy]] = []
|
||||
|
||||
|
|
@ -16104,23 +16162,34 @@ async def model_group_info(
|
|||
undiscoverable_group_names: Final = undiscoverable_model_names(
|
||||
all_models_str, llm_router, user_api_key_dict, user_api_key_dict.team_id
|
||||
)
|
||||
model_groups: list[ModelGroupInfoProxy] = _get_model_group_info(
|
||||
llm_router=llm_router,
|
||||
all_models_str=[name for name in all_models_str if name not in undiscoverable_group_names],
|
||||
model_group=model_group,
|
||||
)
|
||||
listed_group_names: Final = tuple(name for name in all_models_str if name not in undiscoverable_group_names)
|
||||
|
||||
# Append A2A agents to model groups
|
||||
from litellm.proxy.agent_endpoints.model_list_helpers import (
|
||||
append_agents_to_model_group,
|
||||
)
|
||||
|
||||
model_groups = await append_agents_to_model_group(
|
||||
model_groups=model_groups,
|
||||
model_groups: Final = await append_agents_to_model_group(
|
||||
model_groups=_get_model_group_info(
|
||||
llm_router=llm_router, all_models_str=listed_group_names, model_group=model_group
|
||||
),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, general_settings)
|
||||
public_group_names: Final = tuple(
|
||||
internal_to_public.get(group.model_group, group.model_group) for group in model_groups
|
||||
)
|
||||
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(
|
||||
user_api_key_dict, tuple(dict.fromkeys(public_group_names))
|
||||
)
|
||||
|
||||
return {"data": model_groups}
|
||||
return {
|
||||
"data": [
|
||||
group
|
||||
for group, public_name in zip(model_groups, public_group_names, strict=True)
|
||||
if public_name not in callback_hidden_names
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ from typing import (
|
|||
Final,
|
||||
Generic,
|
||||
Literal,
|
||||
NoReturn,
|
||||
Optional,
|
||||
Protocol,
|
||||
TypeAlias,
|
||||
|
|
@ -106,7 +107,7 @@ except ImportError:
|
|||
raise ImportError("backoff is not installed. Please install it via 'pip install backoff'")
|
||||
|
||||
from fastapi import HTTPException, status
|
||||
from pydantic import TypeAdapter
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
import litellm.litellm_core_utils
|
||||
|
|
@ -1136,11 +1137,51 @@ class _CallbackCapabilities:
|
|||
# avoids the per-request ``get_custom_logger_compatible_class`` walk for
|
||||
# every string entry in ``litellm.callbacks``.
|
||||
resolved_callbacks: tuple[object, ...] = field(default_factory=tuple)
|
||||
listed_models_filters: tuple[CustomLogger, ...] = field(default_factory=tuple)
|
||||
|
||||
|
||||
def _overrides_hook(callback: CustomLogger, hook_name: str) -> bool:
|
||||
leaf_to_base: Final = takewhile(lambda klass: klass is not CustomLogger, type(callback).__mro__)
|
||||
return any(hook_name in klass.__dict__ for klass in leaf_to_base)
|
||||
|
||||
|
||||
def _overrides_moderation_hook(callback: CustomLogger) -> bool:
|
||||
leaf_to_base: Final = takewhile(lambda klass: klass is not CustomLogger, type(callback).__mro__)
|
||||
return any("async_moderation_hook" in klass.__dict__ for klass in leaf_to_base)
|
||||
return _overrides_hook(callback, "async_moderation_hook")
|
||||
|
||||
|
||||
_LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...])
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MalformedListingFilterReturn:
|
||||
callback: str
|
||||
tag: Literal["malformed_listing_filter_return"] = "malformed_listing_filter_return"
|
||||
|
||||
|
||||
async def _names_kept_by_listing_callbacks(
|
||||
callbacks: Sequence[CustomLogger],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
model_names: tuple[str, ...],
|
||||
) -> tuple[str, ...] | MalformedListingFilterReturn:
|
||||
if not callbacks or not model_names:
|
||||
return model_names
|
||||
returned: Final = await callbacks[0].async_filter_listed_models(user_api_key_dict, model_names)
|
||||
try:
|
||||
kept: Final = frozenset(_LISTED_MODEL_NAMES.validate_python(returned))
|
||||
except ValidationError:
|
||||
return MalformedListingFilterReturn(callback=type(callbacks[0]).__name__)
|
||||
return await _names_kept_by_listing_callbacks(
|
||||
callbacks[1:], user_api_key_dict, tuple(name for name in model_names if name in kept)
|
||||
)
|
||||
|
||||
|
||||
def _raise_malformed_listing_filter_return(error: MalformedListingFilterReturn) -> NoReturn:
|
||||
raise ProxyException(
|
||||
message=f"{error.callback}.async_filter_listed_models must return a sequence of model names",
|
||||
type=ProxyErrorTypes.internal_server_error,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
|
||||
|
||||
class ProxyLogging:
|
||||
|
|
@ -2808,6 +2849,9 @@ class ProxyLogging:
|
|||
has_moderation_override=has_moderation_override,
|
||||
iterator_overrides=tuple(iterator_overrides),
|
||||
resolved_callbacks=tuple(resolved_callbacks),
|
||||
listed_models_filters=tuple(
|
||||
callback for callback in resolved_callbacks if _overrides_hook(callback, "async_filter_listed_models")
|
||||
),
|
||||
)
|
||||
# Limit cache to handle test churn without leaking; production
|
||||
# callback lists are stable so this rarely grows past 1 entry.
|
||||
|
|
@ -3715,6 +3759,18 @@ class ProxyLogging:
|
|||
verbose_proxy_logger.exception("Error in post_call_response_headers_hook: %s", str(e))
|
||||
return merged_headers
|
||||
|
||||
async def hidden_by_listing_callbacks(
|
||||
self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
|
||||
) -> frozenset[str]:
|
||||
filters: Final = ProxyLogging._callback_capabilities().listed_models_filters
|
||||
if not filters:
|
||||
return frozenset()
|
||||
candidates: Final = tuple(model_names)
|
||||
kept: Final = await _names_kept_by_listing_callbacks(filters, user_api_key_dict, candidates)
|
||||
if isinstance(kept, MalformedListingFilterReturn):
|
||||
_raise_malformed_listing_filter_return(kept)
|
||||
return frozenset(candidates).difference(kept)
|
||||
|
||||
@staticmethod
|
||||
def _build_litellm_call_info(data: dict, response: object) -> dict[str, object]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -6092,7 +6092,7 @@
|
|||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"deprecation_date": "2027-06-25",
|
||||
"deprecation_date": "2027-07-31",
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_image_token": 8e-07,
|
||||
"input_cost_per_token": 6e-07,
|
||||
|
|
@ -26212,7 +26212,6 @@
|
|||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
"input_cost_per_token_flex": 1.5e-07,
|
||||
"input_cost_per_token_priority": 5.4e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 32768,
|
||||
|
|
@ -49416,7 +49415,6 @@
|
|||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
"input_cost_per_token_flex": 1.5e-07,
|
||||
"input_cost_per_token_priority": 5.4e-07,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 32768,
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ longer signal it.
|
|||
|
||||
### Added
|
||||
|
||||
- **model**: Optional `display_name` argument on `litellm_model`, sent as `model_info.display_name` and returned as `display_name` by `/v1/models`, so client model pickers show a readable name; changes are persisted through `/model/{id}/update` since `/model/update` ignores `model_info`; also exported by the `litellm_model` and `litellm_models` data sources
|
||||
- **key**: Computed `server_metadata` attribute on `litellm_key` exposing every metadata entry the proxy stores, so metadata created outside Terraform is visible in state and drift on it shows on refresh, while `metadata` keeps tracking only the declared entries and updates keep preserving undeclared ones
|
||||
- **team_member_add**: `tpm_limit`, `rpm_limit`, `budget_duration`, and `allowed_models` attributes on `litellm_team_member_add`, applied to every member of the resource; `budget_duration` and `allowed_models` ride on `/team/member_add`, while the limits are sent through `/team/member_update`, which is where the proxy accepts them
|
||||
- **team**: Optional `team_id` argument on `litellm_team`, so teams can be created with a stable, human-readable ID instead of a provider-generated UUID; changing it forces replacement
|
||||
|
|
@ -48,6 +49,7 @@ longer signal it.
|
|||
|
||||
### Fixed
|
||||
|
||||
- **model**: `litellm_model` refresh now reads the `{"data": [...]}` envelope `/model/info` returns, so `model_info` fields changed outside Terraform show up as drift instead of silently keeping the previous state
|
||||
- **key**: An update that changes `team_id` and fails because the key was already cascade-deleted along with its previous team now recovers by recreating the key under the new team, instead of aborting the apply. The key's absence is confirmed against the proxy first, so an unrelated failure still errors out, and a `team_id` change between two teams that both still exist stays a plain in-place update
|
||||
- **credential**: create now reports a `credential_name` collision as a clear error naming the `terraform import` command that adopts the existing credential, instead of surfacing the proxy's raw 500 with a Prisma `Unique constraint failed` message. New `adopt_existing` argument (default `false`) opts into taking the existing credential over during create, which makes `apply` idempotent again once state loses track of a credential that still exists on the proxy. Requires a proxy that answers 409 on the collision; older proxies are still detected by their 500 message
|
||||
- **credential**: credential names and `model_id` are now percent-encoded in request URLs, so a name containing `/`, `?`, `#` or spaces reaches the proxy intact instead of being cut at the first reserved character and read, updated or deleted as a different credential
|
||||
|
|
|
|||
|
|
@ -43,6 +43,7 @@ In addition to all arguments above, the following attributes are exported:
|
|||
* `tier` - Model tier (`free` or `paid`).
|
||||
* `mode` - Model mode, e.g. `chat` or `embedding`.
|
||||
* `team_id` - Team the deployment is scoped to, if any.
|
||||
* `display_name` - Human-readable name returned by `/v1/models`, if configured.
|
||||
* `db_model` - Whether the deployment is stored in the database (as opposed to config).
|
||||
|
||||
## Security Note
|
||||
|
|
|
|||
|
|
@ -41,4 +41,5 @@ In addition to all arguments above, the following attributes are exported:
|
|||
* `tier` - Model tier (`free` or `paid`).
|
||||
* `mode` - Model mode, e.g. `chat` or `embedding`.
|
||||
* `team_id` - Team the deployment is scoped to, if any.
|
||||
* `display_name` - Human-readable name returned by `/v1/models`, if configured.
|
||||
* `db_model` - Whether the deployment is stored in the database.
|
||||
|
|
|
|||
|
|
@ -126,6 +126,8 @@ The following arguments are supported:
|
|||
|
||||
* `team_id` - (Optional) string. Associate the model with a specific team.
|
||||
|
||||
* `display_name` - (Optional) string. Human-readable name stored in `model_info.display_name` and returned as `display_name` by `/v1/models`, so clients such as Claude Code and Claude Desktop show it in their model picker instead of `model_name`. When unset, clients fall back to `model_name`.
|
||||
|
||||
* `mode` - (Optional) string. The intended use of the model. Valid values are:
|
||||
* `completion`
|
||||
* `embedding`
|
||||
|
|
|
|||
|
|
@ -23,14 +23,15 @@ type modelInfoParams struct {
|
|||
}
|
||||
|
||||
type modelInfoMeta struct {
|
||||
ID string `json:"id"`
|
||||
DBModel bool `json:"db_model"`
|
||||
BaseModel string `json:"base_model"`
|
||||
Tier string `json:"tier"`
|
||||
Mode string `json:"mode"`
|
||||
TeamID string `json:"team_id"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
ID string `json:"id"`
|
||||
DBModel bool `json:"db_model"`
|
||||
BaseModel string `json:"base_model"`
|
||||
Tier string `json:"tier"`
|
||||
Mode string `json:"mode"`
|
||||
TeamID string `json:"team_id"`
|
||||
DisplayName string `json:"display_name"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
type modelInfoEntry struct {
|
||||
|
|
@ -112,6 +113,10 @@ func dataSourceLiteLLMModel() *schema.Resource {
|
|||
Type: schema.TypeString,
|
||||
Computed: true,
|
||||
},
|
||||
"display_name": {
|
||||
Type: schema.TypeString,
|
||||
Computed: true,
|
||||
},
|
||||
"db_model": {
|
||||
Type: schema.TypeBool,
|
||||
Computed: true,
|
||||
|
|
@ -161,6 +166,7 @@ func dataSourceLiteLLMModelRead(d *schema.ResourceData, m interface{}) error {
|
|||
d.Set("tier", entry.ModelInfo.Tier)
|
||||
d.Set("mode", entry.ModelInfo.Mode)
|
||||
d.Set("team_id", entry.ModelInfo.TeamID)
|
||||
d.Set("display_name", entry.ModelInfo.DisplayName)
|
||||
d.Set("db_model", entry.ModelInfo.DBModel)
|
||||
|
||||
log.Printf("[INFO] Successfully read model with ID: %s", modelID)
|
||||
|
|
@ -197,6 +203,7 @@ func dataSourceLiteLLMModels() *schema.Resource {
|
|||
"tier": {Type: schema.TypeString, Computed: true},
|
||||
"mode": {Type: schema.TypeString, Computed: true},
|
||||
"team_id": {Type: schema.TypeString, Computed: true},
|
||||
"display_name": {Type: schema.TypeString, Computed: true},
|
||||
"db_model": {Type: schema.TypeBool, Computed: true},
|
||||
},
|
||||
},
|
||||
|
|
@ -247,6 +254,7 @@ func dataSourceLiteLLMModelsRead(d *schema.ResourceData, m interface{}) error {
|
|||
"tier": entry.ModelInfo.Tier,
|
||||
"mode": entry.ModelInfo.Mode,
|
||||
"team_id": entry.ModelInfo.TeamID,
|
||||
"display_name": entry.ModelInfo.DisplayName,
|
||||
"db_model": entry.ModelInfo.DBModel,
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -35,7 +35,8 @@ func TestDataSourceModelReadSingleObject(t *testing.T) {
|
|||
"base_model": "gpt-4o",
|
||||
"tier": "paid",
|
||||
"mode": "chat",
|
||||
"team_id": "team-1"
|
||||
"team_id": "team-1",
|
||||
"display_name": "GPT-4o"
|
||||
}
|
||||
}
|
||||
}`))
|
||||
|
|
@ -66,6 +67,7 @@ func TestDataSourceModelReadSingleObject(t *testing.T) {
|
|||
"tier": "paid",
|
||||
"mode": "chat",
|
||||
"team_id": "team-1",
|
||||
"display_name": "GPT-4o",
|
||||
"db_model": true,
|
||||
}
|
||||
for attr, want := range checks {
|
||||
|
|
@ -115,7 +117,7 @@ func TestDataSourceModelsRead(t *testing.T) {
|
|||
w.Header().Set("Content-Type", "application/json")
|
||||
w.Write([]byte(`{
|
||||
"data": [
|
||||
{"model_name": "a", "litellm_params": {"model": "openai/a", "custom_llm_provider": "openai"}, "model_info": {"id": "id-1", "db_model": true}},
|
||||
{"model_name": "a", "litellm_params": {"model": "openai/a", "custom_llm_provider": "openai"}, "model_info": {"id": "id-1", "db_model": true, "display_name": "Model A"}},
|
||||
{"model_name": "b", "litellm_params": {"model": "anthropic/b", "custom_llm_provider": "anthropic"}, "model_info": {"id": "id-2"}}
|
||||
]
|
||||
}`))
|
||||
|
|
@ -143,7 +145,11 @@ func TestDataSourceModelsRead(t *testing.T) {
|
|||
t.Fatalf("expected 2 models, got %d", len(models))
|
||||
}
|
||||
first := models[0].(map[string]interface{})
|
||||
if first["model_name"] != "a" || first["custom_llm_provider"] != "openai" || first["db_model"] != true {
|
||||
if first["model_name"] != "a" || first["custom_llm_provider"] != "openai" || first["db_model"] != true || first["display_name"] != "Model A" {
|
||||
t.Errorf("unexpected first model: %v", first)
|
||||
}
|
||||
second := models[1].(map[string]interface{})
|
||||
if second["display_name"] != "" {
|
||||
t.Errorf("expected empty display_name for model without one, got %v", second["display_name"])
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -93,6 +93,11 @@ func resourceLiteLLMModel() *schema.Resource {
|
|||
Type: schema.TypeString,
|
||||
Optional: true,
|
||||
},
|
||||
"display_name": {
|
||||
Type: schema.TypeString,
|
||||
Optional: true,
|
||||
Description: "Human-readable name returned as display_name by /v1/models, shown in client model pickers instead of model_name",
|
||||
},
|
||||
"mode": {
|
||||
Type: schema.TypeString,
|
||||
Optional: true,
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import (
|
|||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
|
@ -53,6 +54,7 @@ func retryModelRead(d *schema.ResourceData, m interface{}, maxRetries int) error
|
|||
const (
|
||||
endpointModelNew = "/model/new"
|
||||
endpointModelUpdate = "/model/update"
|
||||
endpointModelPatch = "/model/%s/update"
|
||||
endpointModelInfo = "/model/info"
|
||||
endpointModelDelete = "/model/delete"
|
||||
)
|
||||
|
|
@ -246,12 +248,13 @@ func createOrUpdateModel(d *schema.ResourceData, m interface{}, isUpdate bool) e
|
|||
ModelName: d.Get("model_name").(string),
|
||||
LiteLLMParams: litellmParams,
|
||||
ModelInfo: ModelInfo{
|
||||
ID: modelID,
|
||||
DBModel: true,
|
||||
BaseModel: pricingBaseModel,
|
||||
Tier: d.Get("tier").(string),
|
||||
Mode: d.Get("mode").(string),
|
||||
TeamID: d.Get("team_id").(string),
|
||||
ID: modelID,
|
||||
DBModel: true,
|
||||
BaseModel: pricingBaseModel,
|
||||
Tier: d.Get("tier").(string),
|
||||
Mode: d.Get("mode").(string),
|
||||
TeamID: d.Get("team_id").(string),
|
||||
DisplayName: d.Get("display_name").(string),
|
||||
},
|
||||
Additional: make(map[string]interface{}),
|
||||
}
|
||||
|
|
@ -275,6 +278,12 @@ func createOrUpdateModel(d *schema.ResourceData, m interface{}, isUpdate bool) e
|
|||
return fmt.Errorf("failed to %s model: %w", map[bool]string{true: "update", false: "create"}[isUpdate], err)
|
||||
}
|
||||
|
||||
if isUpdate && d.HasChange("display_name") {
|
||||
if err := patchModelDisplayName(client, modelID, d.Get("display_name").(string)); err != nil {
|
||||
return fmt.Errorf("failed to update model display_name: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
d.SetId(modelID)
|
||||
|
||||
log.Printf("[INFO] Model created with ID %s. Starting retry mechanism to read the model...", modelID)
|
||||
|
|
@ -282,6 +291,19 @@ func createOrUpdateModel(d *schema.ResourceData, m interface{}, isUpdate bool) e
|
|||
return retryModelRead(d, m, 5)
|
||||
}
|
||||
|
||||
// /model/update only merges litellm_params, so model_info changes go through the PATCH endpoint.
|
||||
func patchModelDisplayName(client *Client, modelID, displayName string) error {
|
||||
resp, err := MakeRequest(client, "PATCH", fmt.Sprintf(endpointModelPatch, url.PathEscape(modelID)), ModelInfoPatch{
|
||||
ModelInfo: ModelInfoPatchFields{ID: modelID, DisplayName: displayName},
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
_, err = handleAPIResponse(resp, nil, client)
|
||||
return err
|
||||
}
|
||||
|
||||
func resourceLiteLLMModelCreate(d *schema.ResourceData, m interface{}) error {
|
||||
return createOrUpdateModel(d, m, false)
|
||||
}
|
||||
|
|
@ -327,6 +349,7 @@ func resourceLiteLLMModelRead(d *schema.ResourceData, m interface{}) error {
|
|||
d.Set("tier", GetStringValue(modelResp.ModelInfo.Tier, d.Get("tier").(string)))
|
||||
d.Set("mode", GetStringValue(modelResp.ModelInfo.Mode, d.Get("mode").(string)))
|
||||
d.Set("team_id", GetStringValue(modelResp.ModelInfo.TeamID, d.Get("team_id").(string)))
|
||||
d.Set("display_name", modelResp.ModelInfo.DisplayName)
|
||||
|
||||
// Preserve credential name from state since it might not be returned by API
|
||||
d.Set("litellm_credential_name", d.Get("litellm_credential_name").(string))
|
||||
|
|
|
|||
244
terraform/provider/litellm/resource_model_test.go
Normal file
244
terraform/provider/litellm/resource_model_test.go
Normal file
|
|
@ -0,0 +1,244 @@
|
|||
package litellm
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/helper/schema"
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/terraform"
|
||||
)
|
||||
|
||||
func modelInfoBody(displayName string) string {
|
||||
modelInfo := map[string]interface{}{
|
||||
"id": "model-123",
|
||||
"db_model": true,
|
||||
"base_model": "claude-sonnet-4-5",
|
||||
"tier": "free",
|
||||
"mode": "chat",
|
||||
}
|
||||
if displayName != "" {
|
||||
modelInfo["display_name"] = displayName
|
||||
}
|
||||
body, _ := json.Marshal(map[string]interface{}{
|
||||
"model_name": "sonnet-4-5-anthropic",
|
||||
"litellm_params": map[string]interface{}{"model": "anthropic/claude-sonnet-4-5", "custom_llm_provider": "anthropic"},
|
||||
"model_info": modelInfo,
|
||||
})
|
||||
return string(body)
|
||||
}
|
||||
|
||||
func modelInfoDataEnvelope(displayName string) string {
|
||||
return `{"data": [` + modelInfoBody(displayName) + `]}`
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMModelCreateSendsDisplayName(t *testing.T) {
|
||||
var createPayload map[string]interface{}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.URL.Path {
|
||||
case "/model/new":
|
||||
if err := json.NewDecoder(r.Body).Decode(&createPayload); err != nil {
|
||||
t.Errorf("failed to decode create payload: %v", err)
|
||||
}
|
||||
w.Write([]byte(modelInfoBody("Claude Sonnet 4.5")))
|
||||
case "/model/info":
|
||||
w.Write([]byte(modelInfoBody("Claude Sonnet 4.5")))
|
||||
default:
|
||||
t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path)
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMModel().Schema, map[string]interface{}{
|
||||
"model_name": "sonnet-4-5-anthropic",
|
||||
"custom_llm_provider": "anthropic",
|
||||
"base_model": "claude-sonnet-4-5",
|
||||
"model_api_key": "sk-ant-test",
|
||||
"mode": "chat",
|
||||
"display_name": "Claude Sonnet 4.5",
|
||||
})
|
||||
|
||||
if err := resourceLiteLLMModelCreate(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("create failed: %v", err)
|
||||
}
|
||||
|
||||
modelInfo, ok := createPayload["model_info"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected model_info object in create payload, got %v", createPayload["model_info"])
|
||||
}
|
||||
if modelInfo["display_name"] != "Claude Sonnet 4.5" {
|
||||
t.Errorf("expected model_info.display_name 'Claude Sonnet 4.5', got %v", modelInfo["display_name"])
|
||||
}
|
||||
if got := d.Get("display_name").(string); got != "Claude Sonnet 4.5" {
|
||||
t.Errorf("expected state display_name 'Claude Sonnet 4.5', got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMModelCreateOmitsUnsetDisplayName(t *testing.T) {
|
||||
var createPayload map[string]interface{}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.URL.Path {
|
||||
case "/model/new":
|
||||
if err := json.NewDecoder(r.Body).Decode(&createPayload); err != nil {
|
||||
t.Errorf("failed to decode create payload: %v", err)
|
||||
}
|
||||
w.Write([]byte(modelInfoBody("")))
|
||||
case "/model/info":
|
||||
w.Write([]byte(modelInfoBody("")))
|
||||
default:
|
||||
t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path)
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMModel().Schema, map[string]interface{}{
|
||||
"model_name": "sonnet-4-5-anthropic",
|
||||
"custom_llm_provider": "anthropic",
|
||||
"base_model": "claude-sonnet-4-5",
|
||||
"model_api_key": "sk-ant-test",
|
||||
})
|
||||
|
||||
if err := resourceLiteLLMModelCreate(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("create failed: %v", err)
|
||||
}
|
||||
|
||||
modelInfo := createPayload["model_info"].(map[string]interface{})
|
||||
if _, present := modelInfo["display_name"]; present {
|
||||
t.Errorf("expected display_name to be omitted from model_info when unset, got %v", modelInfo["display_name"])
|
||||
}
|
||||
if got := d.Get("display_name").(string); got != "" {
|
||||
t.Errorf("expected empty state display_name, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMModelReadDisplayName(t *testing.T) {
|
||||
cases := map[string]struct {
|
||||
serverBody string
|
||||
want string
|
||||
}{
|
||||
"server value wins inside data envelope": {serverBody: modelInfoDataEnvelope("Renamed In Admin UI"), want: "Renamed In Admin UI"},
|
||||
"server value wins unwrapped": {serverBody: modelInfoBody("Renamed In Admin UI"), want: "Renamed In Admin UI"},
|
||||
"external removal clears state": {serverBody: modelInfoDataEnvelope(""), want: ""},
|
||||
}
|
||||
for name, tc := range cases {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/model/info" {
|
||||
t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path)
|
||||
}
|
||||
w.Write([]byte(tc.serverBody))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMModel().Schema, map[string]interface{}{
|
||||
"model_name": "sonnet-4-5-anthropic",
|
||||
"custom_llm_provider": "anthropic",
|
||||
"base_model": "claude-sonnet-4-5",
|
||||
"display_name": "Claude Sonnet 4.5",
|
||||
})
|
||||
d.SetId("model-123")
|
||||
|
||||
if err := resourceLiteLLMModelRead(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("read failed: %v", err)
|
||||
}
|
||||
if got := d.Get("display_name").(string); got != tc.want {
|
||||
t.Errorf("expected display_name %q, got %q", tc.want, got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func updateResourceData(t *testing.T, oldDisplayName, newDisplayName string) *schema.ResourceData {
|
||||
t.Helper()
|
||||
res := resourceLiteLLMModel()
|
||||
attrs := map[string]string{
|
||||
"model_name": "sonnet-4-5-anthropic",
|
||||
"custom_llm_provider": "anthropic",
|
||||
"base_model": "claude-sonnet-4-5",
|
||||
}
|
||||
if oldDisplayName != "" {
|
||||
attrs["display_name"] = oldDisplayName
|
||||
}
|
||||
state := &terraform.InstanceState{ID: "model-123", Attributes: attrs}
|
||||
diff, err := res.Diff(context.Background(), state, &terraform.ResourceConfig{Config: map[string]interface{}{
|
||||
"model_name": "sonnet-4-5-anthropic",
|
||||
"custom_llm_provider": "anthropic",
|
||||
"base_model": "claude-sonnet-4-5",
|
||||
"display_name": newDisplayName,
|
||||
}}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("diff failed: %v", err)
|
||||
}
|
||||
d, err := schema.InternalMap(res.Schema).Data(state, diff)
|
||||
if err != nil {
|
||||
t.Fatalf("data failed: %v", err)
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMModelUpdatePatchesDisplayName(t *testing.T) {
|
||||
cases := map[string]struct {
|
||||
newName string
|
||||
}{
|
||||
"changed name is patched": {newName: "Claude Sonnet 4.5 v2"},
|
||||
"cleared name is patched": {newName: ""},
|
||||
}
|
||||
for name, tc := range cases {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
var patchPayload map[string]interface{}
|
||||
var patchPath string
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch {
|
||||
case r.Method == http.MethodPost && r.URL.Path == "/model/update":
|
||||
w.Write([]byte(modelInfoBody("Claude Sonnet 4.5")))
|
||||
case r.Method == http.MethodPatch:
|
||||
patchPath = r.URL.Path
|
||||
if err := json.NewDecoder(r.Body).Decode(&patchPayload); err != nil {
|
||||
t.Errorf("failed to decode patch payload: %v", err)
|
||||
}
|
||||
w.Write([]byte(modelInfoBody(tc.newName)))
|
||||
case r.URL.Path == "/model/info":
|
||||
w.Write([]byte(modelInfoDataEnvelope(tc.newName)))
|
||||
default:
|
||||
t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path)
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
d := updateResourceData(t, "Claude Sonnet 4.5", tc.newName)
|
||||
if err := resourceLiteLLMModelUpdate(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("update failed: %v", err)
|
||||
}
|
||||
if patchPath != "/model/model-123/update" {
|
||||
t.Fatalf("expected PATCH /model/model-123/update, got %q", patchPath)
|
||||
}
|
||||
modelInfo := patchPayload["model_info"].(map[string]interface{})
|
||||
if modelInfo["display_name"] != tc.newName {
|
||||
t.Errorf("expected patched display_name %q, got %v", tc.newName, modelInfo["display_name"])
|
||||
}
|
||||
if got := d.Get("display_name").(string); got != tc.newName {
|
||||
t.Errorf("expected state display_name %q, got %q", tc.newName, got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMModelUpdateSkipsPatchWhenDisplayNameUnchanged(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodPatch {
|
||||
t.Errorf("unexpected PATCH %s", r.URL.Path)
|
||||
}
|
||||
w.Write([]byte(modelInfoDataEnvelope("Claude Sonnet 4.5")))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
d := updateResourceData(t, "Claude Sonnet 4.5", "Claude Sonnet 4.5")
|
||||
if err := resourceLiteLLMModelUpdate(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("update failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
|
@ -25,6 +25,16 @@ type ModelResponse struct {
|
|||
Additional map[string]interface{} `json:"additional"`
|
||||
}
|
||||
|
||||
// ModelInfoPatch is the body for PATCH /model/{id}/update; display_name is sent even when empty so it can be cleared.
|
||||
type ModelInfoPatch struct {
|
||||
ModelInfo ModelInfoPatchFields `json:"model_info"`
|
||||
}
|
||||
|
||||
type ModelInfoPatchFields struct {
|
||||
ID string `json:"id"`
|
||||
DisplayName string `json:"display_name"`
|
||||
}
|
||||
|
||||
// ModelRequest represents a request to create or update a model.
|
||||
type ModelRequest struct {
|
||||
ModelName string `json:"model_name"`
|
||||
|
|
@ -108,12 +118,13 @@ type LiteLLMParams struct {
|
|||
|
||||
// ModelInfo represents information about a model.
|
||||
type ModelInfo struct {
|
||||
ID string `json:"id"`
|
||||
DBModel bool `json:"db_model"`
|
||||
BaseModel string `json:"base_model"`
|
||||
Tier string `json:"tier"`
|
||||
Mode string `json:"mode"`
|
||||
TeamID string `json:"team_id,omitempty"`
|
||||
ID string `json:"id"`
|
||||
DBModel bool `json:"db_model"`
|
||||
BaseModel string `json:"base_model"`
|
||||
Tier string `json:"tier"`
|
||||
Mode string `json:"mode"`
|
||||
TeamID string `json:"team_id,omitempty"`
|
||||
DisplayName string `json:"display_name,omitempty"`
|
||||
}
|
||||
|
||||
// Key represents a LiteLLM API key.
|
||||
|
|
|
|||
|
|
@ -55,6 +55,13 @@ func handleAPIResponse(resp *http.Response, reqBody interface{}, client *Client)
|
|||
resp.Status, client.redactSensitiveData(string(bodyBytes)), client.redactSensitiveData(string(reqBodyBytes)))
|
||||
}
|
||||
|
||||
var envelope struct {
|
||||
Data []json.RawMessage `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(bodyBytes, &envelope); err == nil && len(envelope.Data) > 0 {
|
||||
bodyBytes = envelope.Data[0]
|
||||
}
|
||||
|
||||
var modelResp ModelResponse
|
||||
if err := json.Unmarshal(bodyBytes, &modelResp); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse response: %v", err)
|
||||
|
|
|
|||
|
|
@ -97,7 +97,6 @@ GET /guardrails/{guardrail_id}
|
|||
GET /prompts/{prompt_id}
|
||||
GET /prompts/{prompt_id}/versions
|
||||
PATCH /guardrails/{guardrail_id}
|
||||
PATCH /model/{model_id}/update
|
||||
PATCH /prompts/{prompt_id}
|
||||
PATCH /team/{team_id}
|
||||
POST /team/model/add
|
||||
|
|
|
|||
|
|
@ -31,11 +31,11 @@ def get_function_names_from_file(file_path):
|
|||
def get_all_functions_called_in_tests(base_dir):
|
||||
"""
|
||||
Returns a set of function names that are called in test functions
|
||||
inside 'local_testing' and 'proxy_unit_tests' directories,
|
||||
inside 'local_testing' and 'unit/proxy' directories,
|
||||
specifically in files containing the word 'router'.
|
||||
"""
|
||||
called_functions = set()
|
||||
test_dirs = ["local_testing", "proxy_unit_tests"]
|
||||
test_dirs = ["local_testing", "unit/proxy"]
|
||||
|
||||
for test_dir in test_dirs:
|
||||
dir_path = os.path.join(base_dir, test_dir)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,60 @@
|
|||
"""Unit tests for the Claude Code PR-gate version resolver.
|
||||
|
||||
Markerless harness tests: they feed the resolver a hand-built packument and a
|
||||
fixed clock, so they run without a proxy, never reach the npm registry, and
|
||||
carry no `e2e` marker.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final, Mapping
|
||||
|
||||
import pytest
|
||||
|
||||
from claude_code.pr_gate_version_resolver import NoEligibleVersionError, resolve_pr_gate_version
|
||||
|
||||
NOW: Final = datetime(2026, 4, 25, 12, 0, tzinfo=timezone.utc)
|
||||
INSIDE_THE_2_1_88_WINDOW: Final = datetime(2026, 4, 3, 12, 0, tzinfo=timezone.utc)
|
||||
|
||||
|
||||
def _packument(times: Mapping[str, str], unpublished: frozenset[str] = frozenset()) -> dict[str, object]:
|
||||
return {
|
||||
"name": "@anthropic-ai/claude-code",
|
||||
"time": {"created": "2024-01-01T00:00:00.000Z", "modified": "2026-04-25T00:00:00.000Z", **times},
|
||||
"versions": {version: {"version": version} for version in times if version not in unpublished},
|
||||
}
|
||||
|
||||
|
||||
def test_skips_a_version_npm_has_unpublished() -> None:
|
||||
metadata: Final = _packument(
|
||||
{
|
||||
"2.1.87": "2026-03-28T20:00:00.000Z",
|
||||
"2.1.88": "2026-03-30T22:36:48.424Z",
|
||||
"2.1.89": "2026-03-31T23:32:40.000Z",
|
||||
},
|
||||
unpublished=frozenset({"2.1.88"}),
|
||||
)
|
||||
assert resolve_pr_gate_version(metadata=metadata, as_of=INSIDE_THE_2_1_88_WINDOW) == "2.1.87"
|
||||
|
||||
|
||||
def test_raises_when_the_only_old_enough_version_is_unpublished() -> None:
|
||||
metadata: Final = _packument(
|
||||
{"2.1.88": "2026-03-30T22:36:48.424Z", "2.1.89": "2026-03-31T23:32:40.000Z"},
|
||||
unpublished=frozenset({"2.1.88"}),
|
||||
)
|
||||
with pytest.raises(NoEligibleVersionError):
|
||||
resolve_pr_gate_version(metadata=metadata, as_of=INSIDE_THE_2_1_88_WINDOW)
|
||||
|
||||
|
||||
def test_picks_the_newest_published_version_at_least_min_age_old() -> None:
|
||||
metadata: Final = _packument(
|
||||
{
|
||||
"2.1.118": "2026-04-15T10:00:00.000Z",
|
||||
"2.1.119": "2026-04-21T10:00:00.000Z",
|
||||
"2.2.0-rc.1": "2026-04-22T10:00:00.000Z",
|
||||
"2.1.120": "2026-04-23T10:00:00.000Z",
|
||||
"2.1.121": "2026-04-25T11:00:00.000Z",
|
||||
}
|
||||
)
|
||||
assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "2.1.119"
|
||||
|
|
@ -297,6 +297,7 @@ def run_claude(
|
|||
}
|
||||
env["ANTHROPIC_BASE_URL"] = base_url
|
||||
env["ANTHROPIC_AUTH_TOKEN"] = api_key
|
||||
env["DISABLE_AUTOUPDATER"] = "1"
|
||||
# Hand the CLI a fresh empty HOME so a compromised claude package
|
||||
# or a model-directed Read tool call can't see the runtime user's
|
||||
# real dotfiles. Created here, removed in the `finally` below
|
||||
|
|
|
|||
|
|
@ -4,8 +4,6 @@ ARG GH_VERSION=2.101.0
|
|||
ARG GH_SHA256=9bca2d1c16825f109907a23307628a2f0698fbf99662b73a5cf0b020293072b8
|
||||
ARG UV_VERSION=0.10.9
|
||||
ARG UV_SHA256=20d79708222611fa540b5c9ed84f352bcd3937740e51aacc0f8b15b271c57594
|
||||
ARG CLAUDE_CODE_VERSION=2.1.228
|
||||
ARG CLAUDE_CODE_SHA256=d535985e6941a3eb00179ccd7f52ceb0c6623a0305a518ebc4e6514f84a94c99
|
||||
|
||||
SHELL ["/bin/bash", "-o", "pipefail", "-c"]
|
||||
|
||||
|
|
@ -23,11 +21,6 @@ RUN curl -fsSLo /tmp/uv.tar.gz "https://github.com/astral-sh/uv/releases/downloa
|
|||
&& tar -xzf /tmp/uv.tar.gz -C /usr/local/bin --strip-components=1 uv-x86_64-unknown-linux-gnu/uv \
|
||||
&& rm /tmp/uv.tar.gz
|
||||
|
||||
RUN curl -fsSLo /tmp/claude "https://downloads.claude.ai/claude-code-releases/${CLAUDE_CODE_VERSION}/linux-x64/claude" \
|
||||
&& echo "${CLAUDE_CODE_SHA256} /tmp/claude" | sha256sum -c - \
|
||||
&& install -m 0755 /tmp/claude /usr/local/bin/claude \
|
||||
&& rm /tmp/claude
|
||||
|
||||
RUN groupadd --gid 1000 populator && useradd --uid 1000 --gid 1000 --create-home populator
|
||||
|
||||
ENV HOME=/home/populator \
|
||||
|
|
|
|||
|
|
@ -16,16 +16,18 @@ than as a GitHub Action or on a dedicated VM. Trade-offs:
|
|||
clone of litellm plus a cold `uv sync`. That adds a few minutes on
|
||||
top of the ~10 minute test run; the job's 12 hour ceiling is nowhere
|
||||
near.
|
||||
- ⚠️ The Claude Code CLI version under test is pinned in the
|
||||
`Dockerfile` (`CLAUDE_CODE_VERSION` + its checksum). Bumping it is a
|
||||
PR, see the gotchas below.
|
||||
- ✅ The Claude Code CLI under test is chosen on every run (the newest
|
||||
npm release published at least 3 days ago) and downloaded
|
||||
checksum-verified, so the matrix follows CLI releases without a PR;
|
||||
see the gotchas for pinning a run.
|
||||
|
||||
## Layout
|
||||
|
||||
| File | Purpose |
|
||||
| --- | --- |
|
||||
| `Dockerfile` | The image Render builds: Debian bookworm-slim plus pinned, checksum-verified `gh`, `uv`, and the Claude Code CLI, with this `tests/e2e/` tree copied to `/opt/litellm/tests/e2e/`. Runs as the non-root user `populator` (uid/gid 1000, which is what Render's secret files are readable by). |
|
||||
| `run_daily.sh` | The actual cron job. Resolves versions, clones the worktree, boots the proxy, runs pytest, builds the JSON, opens (or updates) a docs PR, sweeps stale compat-matrix PRs. |
|
||||
| `Dockerfile` | The image Render builds: Debian bookworm-slim plus pinned, checksum-verified `gh` and `uv`, with this `tests/e2e/` tree copied to `/opt/litellm/tests/e2e/`. Runs as the non-root user `populator` (uid/gid 1000, which is what Render's secret files are readable by). |
|
||||
| `run_daily.sh` | The actual cron job. Resolves versions, clones the worktree, installs the Claude Code CLI under test, boots the proxy, runs pytest, builds the JSON, opens (or updates) a docs PR, sweeps stale compat-matrix PRs. |
|
||||
| `install_claude_code.sh` | Downloads one Claude Code release (`<version> <dest-dir>`) from the vendor's native release channel, verifies it against the sha256 in that release's `manifest.json`, and refuses a binary whose `--version` disagrees. Run by the cron and by the `compat-matrix-image` GitHub workflow. |
|
||||
| `build_matrix.py` | Tiny Python CLI that wraps `claude_code.matrix_builder.build_from_paths`. Exists only because the bash script needs *some* way to render the per-cell aggregation, and the builder is already Python. |
|
||||
| `check_regressions.py` | Tiny Python CLI that wraps `claude_code.matrix_builder.find_regressions`. Diffs the freshly built matrix against the currently-published one and exits `3` if any cell flipped green→red, which gates auto-merge. |
|
||||
| `litellm-compat-matrix.env.example` | The service's env vars, one per line, with what each is for. |
|
||||
|
|
@ -35,10 +37,7 @@ than as a GitHub Action or on a dedicated VM. Trade-offs:
|
|||
1. **Resolves the latest LiteLLM final release tag** (newest bare
|
||||
`vX.Y.Z`, skipping `-rc.N`/`-dev.N` pre-releases) by paging the
|
||||
GitHub Releases API (`curl | jq`).
|
||||
2. **Reads the Claude Code CLI version** via `claude --version`. That
|
||||
is whatever the `Dockerfile` pins; the job never upgrades it on its
|
||||
own.
|
||||
3. **Clones the worktree** at `~/litellm-cron-worktree/` (a
|
||||
2. **Clones the worktree** at `~/litellm-cron-worktree/` (a
|
||||
`--filter=blob:none` clone, so only the checked-out tag's blobs are
|
||||
fetched), `git checkout --force <tag>`, then `uv sync --frozen
|
||||
--no-install-project` against a uv-managed CPython 3.12 followed by
|
||||
|
|
@ -52,12 +51,20 @@ than as a GitHub Action or on a dedicated VM. Trade-offs:
|
|||
`claude_code/` so the tree's EKS-harness `conftest.py` (whose imports
|
||||
the stable venv doesn't install) is never loaded. The tag's own
|
||||
`tests/e2e/` is deliberately not used.
|
||||
3. **Resolves and installs the Claude Code CLI under test**:
|
||||
`pr_gate_version_resolver.py` (run on the venv, the image has no
|
||||
Python of its own) picks the newest `@anthropic-ai/claude-code` npm
|
||||
release published at least 3 days ago, the same buffer the PR gate
|
||||
uses, unless `CLAUDE_CODE_VERSION` pins one, and
|
||||
`install_claude_code.sh` downloads that release's `linux-x64` binary
|
||||
into the run's scratch dir, verified against the release manifest.
|
||||
4. **Boots the proxy** as a `setsid` background process on port `4100`
|
||||
bound to loopback, then polls `/health/liveliness` until it's up.
|
||||
5. **Runs pytest** on `tests/e2e/claude_code/` with `LITELLM_PROXY_URL`
|
||||
pointed at the proxy and `COMPAT_RESULTS_PATH` set so the conftest
|
||||
hook writes the per-test results artifact. Test failures become
|
||||
`fail` cells in the JSON, not script errors.
|
||||
5. **Runs pytest** on `tests/e2e/claude_code/` with that CLI first on
|
||||
`PATH`, `LITELLM_PROXY_URL` pointed at the proxy, and
|
||||
`COMPAT_RESULTS_PATH` set so the conftest hook writes the per-test
|
||||
results artifact. Test failures become `fail` cells in the JSON, not
|
||||
script errors.
|
||||
6. **Builds `compatibility-matrix.json`** by handing the artifact +
|
||||
manifest to `build_matrix.py`.
|
||||
7. **Opens or updates a docs PR**: `gh repo clone` of `litellm-docs`
|
||||
|
|
@ -70,7 +77,8 @@ than as a GitHub Action or on a dedicated VM. Trade-offs:
|
|||
branch ... already exists" is treated as success). If the JSON is
|
||||
byte-identical to what `main` already publishes, the push is skipped
|
||||
entirely. These PRs are not gated on a second human review.
|
||||
8. **Gates auto-merge on a regression check**: before enabling
|
||||
|
||||
**Auto-merge is gated on a regression check**: before enabling
|
||||
auto-merge, `check_regressions.py` diffs the new matrix against the
|
||||
one currently on `main`. Auto-merge (`gh pr merge --auto --squash`)
|
||||
is only enabled when **no cell flipped green→red** — i.e. every
|
||||
|
|
@ -82,7 +90,7 @@ than as a GitHub Action or on a dedicated VM. Trade-offs:
|
|||
auto-merge a prior same-day run enabled is explicitly disabled — so a
|
||||
human reviews before it lands on the public table. The check fails
|
||||
*closed*: if it errors, auto-merge is withheld.
|
||||
9. **Sweeps stale compat-matrix PRs**: once today's PR exists, every
|
||||
8. **Sweeps stale compat-matrix PRs**: once today's PR exists, every
|
||||
other open `compat-matrix/*` PR that the publishing account opened
|
||||
from a branch on the docs repo itself is closed (and its bot-owned
|
||||
branch deleted), so at most one compat-matrix PR is ever open — the
|
||||
|
|
@ -165,7 +173,8 @@ curl -fsS -X POST "https://api.render.com/v1/services/${CRON_ID}/deploys" \
|
|||
curl -fsS "https://api.render.com/v1/services/${CRON_ID}/deploys?limit=1" \
|
||||
-H "Authorization: Bearer ${RENDER_API_KEY}"
|
||||
|
||||
# A run that does NOT open a PR (first-time validation, CLI bumps):
|
||||
# A run that does NOT open a PR (first-time validation, a CLI pinned
|
||||
# with CLAUDE_CODE_VERSION):
|
||||
# set SKIP_PUBLISH=1 on the service, trigger a run, then remove it.
|
||||
# The matrix JSON is printed at the end of the run's log (nothing on
|
||||
# the container's disk outlives the run) and saved to
|
||||
|
|
@ -217,21 +226,25 @@ docker run --rm --platform linux/amd64 \
|
|||
or fine-grained Contents:RW + Pull requests:RW). It is delivered as
|
||||
a file, not an env var, so pytest, the proxy, and the claude CLI
|
||||
never inherit it; manual runs export `GITHUB_TOKEN` instead.
|
||||
- **Bumping the Claude Code CLI is a PR.** Change `CLAUDE_CODE_VERSION`
|
||||
in the `Dockerfile` and set `CLAUDE_CODE_SHA256` to the `linux-x64`
|
||||
checksum from
|
||||
`https://downloads.claude.ai/claude-code-releases/<version>/manifest.json`.
|
||||
The first run on a new CLI is the riskiest one: if the new CLI
|
||||
changes its wire format the matrix run can produce systematic
|
||||
failures, so trigger a `SKIP_PUBLISH=1` run before the next scheduled
|
||||
fire. `gh` and `uv` bump the same way, with the checksum from the
|
||||
release's `gh_<version>_checksums.txt` and the tarball's `.sha256`
|
||||
sidecar respectively.
|
||||
- **The Claude Code CLI is chosen per run, not pinned.** Each run
|
||||
tests the newest `@anthropic-ai/claude-code` npm release published
|
||||
at least 3 days ago, downloaded from
|
||||
`https://downloads.claude.ai/claude-code-releases/<version>/linux-x64/claude`
|
||||
and verified against the sha256 in that release's `manifest.json`.
|
||||
A CLI release that breaks a cell shows up as a green→red flip, which
|
||||
withholds auto-merge on that day's docs PR for review. To rerun the
|
||||
matrix on one specific CLI, set `CLAUDE_CODE_VERSION` on the run.
|
||||
`gh` and `uv` stay pinned in the `Dockerfile`; bump them in a PR with
|
||||
the checksum from the release's `gh_<version>_checksums.txt` and the
|
||||
tarball's `.sha256` sidecar respectively.
|
||||
- **A local build on Apple silicon only proves the image assembles.**
|
||||
Under QEMU the Claude Code binary (a Bun executable) dies with
|
||||
`CPU lacks AVX support` and `gh` panics in the Go runtime, so
|
||||
`claude --version` and a full run are verified with a
|
||||
`SKIP_PUBLISH=1` run on Render, not locally.
|
||||
`CPU lacks AVX support` and `gh` panics in the Go runtime, so the CLI
|
||||
download and `claude --version` are verified by the
|
||||
`compat-matrix-image` GitHub workflow (an x86 runner that builds the
|
||||
image and runs `install_claude_code.sh` in it on every PR touching
|
||||
this directory) and a full run with a `SKIP_PUBLISH=1` run on Render,
|
||||
not locally.
|
||||
- **Nothing persists between runs.** A failed run leaves no
|
||||
half-installed venv behind, but also no cache: don't expect a rerun
|
||||
to be faster than the first one.
|
||||
|
|
|
|||
38
tests/e2e/claude_code/cron_vm/install_claude_code.sh
Executable file
38
tests/e2e/claude_code/cron_vm/install_claude_code.sh
Executable file
|
|
@ -0,0 +1,38 @@
|
|||
#!/usr/bin/env bash
|
||||
|
||||
set -Eeuo pipefail
|
||||
|
||||
RELEASES_URL="https://downloads.claude.ai/claude-code-releases"
|
||||
|
||||
log() { printf '==> %s\n' "$*" >&2; }
|
||||
die() { printf 'ERROR: %s\n' "$*" >&2; exit 1; }
|
||||
|
||||
[[ $# -eq 2 ]] || die "usage: $(basename "$0") <version> <dest-dir>"
|
||||
VERSION="$1"
|
||||
DEST_DIR="$2"
|
||||
[[ "${VERSION}" =~ ^[0-9]+\.[0-9]+\.[0-9]+$ ]] || die "not a Claude Code release version: '${VERSION}'"
|
||||
|
||||
mkdir -p "${DEST_DIR}"
|
||||
MANIFEST="${DEST_DIR}/manifest.json"
|
||||
curl -fsSL --retry 3 --retry-all-errors --output "${MANIFEST}" "${RELEASES_URL}/${VERSION}/manifest.json" \
|
||||
|| die "no release manifest for claude code ${VERSION} at ${RELEASES_URL}"
|
||||
CHECKSUM="$(jq -r '.platforms["linux-x64"].checksum // empty' "${MANIFEST}")"
|
||||
[[ "${CHECKSUM}" =~ ^[0-9a-f]{64}$ ]] || die "manifest for claude code ${VERSION} carries no linux-x64 sha256"
|
||||
|
||||
log "downloading claude code ${VERSION} (linux-x64)"
|
||||
DOWNLOAD="${DEST_DIR}/claude.download"
|
||||
curl -fsSL --retry 3 --retry-all-errors --output "${DOWNLOAD}" "${RELEASES_URL}/${VERSION}/linux-x64/claude"
|
||||
echo "${CHECKSUM} ${DOWNLOAD}" | sha256sum -c - >/dev/null \
|
||||
|| die "claude code ${VERSION} sha256 mismatch; refusing to install"
|
||||
chmod 0755 "${DOWNLOAD}"
|
||||
mv "${DOWNLOAD}" "${DEST_DIR}/claude"
|
||||
|
||||
PROBE_HOME="$(mktemp -d -t claude-probe-home.XXXXXX)"
|
||||
trap 'rm -rf "${PROBE_HOME}"' EXIT
|
||||
REPORTED="$(
|
||||
env -i HOME="${PROBE_HOME}" PATH="${PATH}" DISABLE_AUTOUPDATER=1 \
|
||||
"${DEST_DIR}/claude" --version | awk '{print $1}'
|
||||
)" || die "claude code ${VERSION} could not run --version"
|
||||
[[ "${REPORTED}" == "${VERSION}" ]] \
|
||||
|| die "installed claude code reports '${REPORTED}', expected ${VERSION}"
|
||||
log "installed claude code ${VERSION} at ${DEST_DIR}/claude"
|
||||
|
|
@ -7,15 +7,21 @@
|
|||
# 1. Resolve the latest LiteLLM final release tag from the GitHub
|
||||
# Releases API.
|
||||
# 2. Update a long-lived worktree at $WORKTREE to that tag and `uv sync` it.
|
||||
# 3. Boot the proxy as a background subprocess on $PROXY_PORT (default
|
||||
# 3. Resolve the Claude Code CLI version under test (the newest npm
|
||||
# release at least 3 days old, via pr_gate_version_resolver.py, or
|
||||
# $CLAUDE_CODE_VERSION when set) and download that release's
|
||||
# linux-x64 binary into the run's scratch dir, checksum-verified
|
||||
# against the vendor's release manifest (install_claude_code.sh).
|
||||
# 4. Boot the proxy as a background subprocess on $PROXY_PORT (default
|
||||
# 4100; a separate port from the human-tended :4000 proxy).
|
||||
# 4. Run `pytest tests/e2e/claude_code/` against the proxy. Test
|
||||
# failures become `fail` cells in the JSON, not script errors.
|
||||
# 5. Hand the per-test results artifact + manifest to a small Python
|
||||
# 5. Run `pytest tests/e2e/claude_code/` against the proxy with that
|
||||
# CLI first on PATH. Test failures become `fail` cells in the JSON,
|
||||
# not script errors.
|
||||
# 6. Hand the per-test results artifact + manifest to a small Python
|
||||
# CLI (`build_matrix.py`) that wraps the existing
|
||||
# `matrix_builder.build_from_paths` to produce the published
|
||||
# compatibility-matrix.json.
|
||||
# 6. `gh repo clone` litellm-docs, write the JSON to a deterministic
|
||||
# 7. `gh repo clone` litellm-docs, write the JSON to a deterministic
|
||||
# branch (`compat-matrix/<litellm>-<claude>-<UTC-date>`), commit,
|
||||
# push the branch straight to BerriAI/litellm-docs (mateo-berri has
|
||||
# write access), `gh pr create`, then — *only if no cell regressed
|
||||
|
|
@ -23,7 +29,7 @@
|
|||
# auto-merge so the PR merges itself once required checks pass. A
|
||||
# green→red regression leaves auto-merge off for human review; an
|
||||
# already-red cell (red→red) does not block.
|
||||
# 7. Sweep stale compat-matrix PRs: once today's PR exists, close any
|
||||
# 8. Sweep stale compat-matrix PRs: once today's PR exists, close any
|
||||
# other open `compat-matrix/*` PR (and delete its bot-owned branch)
|
||||
# so at most ONE compat-matrix PR is ever open — the newest. A
|
||||
# gate-withheld PR that nobody triages is superseded by the next
|
||||
|
|
@ -33,7 +39,7 @@
|
|||
# rather than spawning a new one. If the JSON is byte-identical to the
|
||||
# docs branch, we skip the push entirely.
|
||||
#
|
||||
# Required commands on $PATH: git, uv, gh, jq, curl, claude.
|
||||
# Required commands on $PATH: git, uv, gh, jq, curl.
|
||||
# Required state: a litellm checkout at $LITELLM_REPO (this file lives in
|
||||
# it); $WORKTREE is created on first run.
|
||||
#
|
||||
|
|
@ -51,6 +57,9 @@ DOCS_BRANCH="${DOCS_BRANCH:-main}"
|
|||
DOCS_TARGET_PATH="${DOCS_TARGET_PATH:-src/data/compatibility-matrix.json}"
|
||||
SKIP_PUBLISH="${SKIP_PUBLISH:-0}"
|
||||
PYTEST_K="${PYTEST_K:-}"
|
||||
# Empty means "resolve it": the newest @anthropic-ai/claude-code npm
|
||||
# release published at least 3 days ago. Set it to pin a manual run.
|
||||
CLAUDE_CODE_VERSION="${CLAUDE_CODE_VERSION:-}"
|
||||
# The e2e suite uses PEP 695 `type` aliases, so the venv needs Python
|
||||
# >= 3.12 (also what repo CI runs) even when the host's system python is
|
||||
# older. uv fetches a managed CPython of this version on first use --
|
||||
|
|
@ -108,7 +117,7 @@ trap cleanup EXIT INT TERM
|
|||
log() { printf '==> %s\n' "$*" >&2; }
|
||||
die() { printf 'ERROR: %s\n' "$*" >&2; exit 1; }
|
||||
|
||||
for cmd in git uv gh jq curl claude; do
|
||||
for cmd in git uv gh jq curl; do
|
||||
command -v "${cmd}" >/dev/null 2>&1 || die "missing required command: ${cmd}"
|
||||
done
|
||||
|
||||
|
|
@ -199,10 +208,6 @@ LITELLM_VERSION="$(
|
|||
[[ -n "${LITELLM_VERSION}" ]] || die "could not resolve latest PEP 440 final release (vX.Y.Z) in 5 pages of releases"
|
||||
log "resolved litellm: ${LITELLM_VERSION}"
|
||||
|
||||
CLAUDE_CODE_VERSION="$(claude --version 2>/dev/null | awk '{print $1}')"
|
||||
[[ -n "${CLAUDE_CODE_VERSION}" ]] || die "could not read 'claude --version'"
|
||||
log "local claude code: ${CLAUDE_CODE_VERSION}"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. Update the worktree to that tag
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -314,7 +319,27 @@ PROXY_CONFIG="${WORKTREE}/tests/e2e/claude_code/test_config.yaml"
|
|||
[[ -f "${PROXY_CONFIG}" ]] || die "proxy config not found at ${PROXY_CONFIG} (shim incomplete?)"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. Boot the proxy
|
||||
# 3. Resolve and install the Claude Code CLI under test
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# The resolver is stdlib-only, but the image ships no python of its
|
||||
# own, so it runs on the venv the sync above just built. Its 3-day
|
||||
# publish-age buffer (PRD #26476) keeps a release that gets pulled or
|
||||
# patched within days from ever driving the published matrix.
|
||||
if [[ -z "${CLAUDE_CODE_VERSION}" ]]; then
|
||||
CLAUDE_CODE_VERSION="$(
|
||||
cd "${WORKTREE}" \
|
||||
&& "${WORKTREE_UV}" run --no-sync python "${POPULATOR_DIR}/../pr_gate_version_resolver.py"
|
||||
)" || die "could not resolve the Claude Code version to test"
|
||||
log "resolved claude code: ${CLAUDE_CODE_VERSION}"
|
||||
else
|
||||
log "CLAUDE_CODE_VERSION set; testing claude code ${CLAUDE_CODE_VERSION}"
|
||||
fi
|
||||
CLAUDE_CLI_DIR="${WORKDIR}/claude-cli"
|
||||
"${POPULATOR_DIR}/install_claude_code.sh" "${CLAUDE_CODE_VERSION}" "${CLAUDE_CLI_DIR}"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 4. Boot the proxy
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
log "starting proxy on 127.0.0.1:${PROXY_PORT}"
|
||||
|
|
@ -350,7 +375,7 @@ curl -fsS "${HEALTH_URL}" >/dev/null \
|
|||
|| { tail -50 "${WORKDIR}/proxy.log" >&2; die "proxy did not become healthy"; }
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 4. Run pytest
|
||||
# 5. Run pytest
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
RESULTS_JSON="${WORKDIR}/compat-results.json"
|
||||
|
|
@ -374,6 +399,7 @@ set +e
|
|||
&& LITELLM_PROXY_URL="http://127.0.0.1:${PROXY_PORT}" \
|
||||
LITELLM_MASTER_KEY="${PROXY_API_KEY}" \
|
||||
COMPAT_RESULTS_PATH="${RESULTS_JSON}" \
|
||||
PATH="${CLAUDE_CLI_DIR}:${PATH}" \
|
||||
"${WORKTREE_UV}" run --no-sync pytest "${PYTEST_ARGS[@]}"
|
||||
)
|
||||
PYTEST_EXIT=$?
|
||||
|
|
@ -386,7 +412,7 @@ log "pytest exit code: ${PYTEST_EXIT} (failures become 'fail' cells, not script
|
|||
[[ -f "${RESULTS_JSON}" ]] || die "pytest did not produce ${RESULTS_JSON}"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 5. Build the matrix JSON
|
||||
# 6. Build the matrix JSON
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
MATRIX_JSON="${WORKDIR}/compatibility-matrix.json"
|
||||
|
|
@ -402,7 +428,7 @@ log "building ${MATRIX_JSON}"
|
|||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 6. Open a docs-repo PR
|
||||
# 7. Open a docs-repo PR
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
if [[ "${SKIP_PUBLISH}" == "1" ]]; then
|
||||
|
|
@ -634,7 +660,9 @@ else
|
|||
|| die "auto-merge still armed on ${BRANCH_NAME} (enabled ${AUTOMERGE_ARMED}) after --disable-auto"
|
||||
fi
|
||||
|
||||
# --- Stale-PR sweep ----------------------------------------------------------
|
||||
# ---------------------------------------------------------------------------
|
||||
# 8. Sweep stale compat-matrix PRs
|
||||
# ---------------------------------------------------------------------------
|
||||
# Keep at most ONE compat-matrix PR open: today's. Any other open
|
||||
# `compat-matrix/*` PR is a leftover from a day whose regression gate
|
||||
# withheld auto-merge and nobody triaged it; the PR we just opened or
|
||||
|
|
|
|||
|
|
@ -80,7 +80,9 @@ def resolve_pr_gate_version(
|
|||
|
||||
"Newest" means newest by **publish time**, not semver string order —
|
||||
if a patch lands on an older major after a newer release, the
|
||||
patched line is the eligible one.
|
||||
patched line is the eligible one. A version npm has unpublished keeps
|
||||
its ``time`` entry but drops out of ``versions``, so only versions
|
||||
still present in ``versions`` are candidates.
|
||||
|
||||
Args:
|
||||
metadata: Pre-fetched npm packument (skips the HTTP call). Useful
|
||||
|
|
@ -101,6 +103,7 @@ def resolve_pr_gate_version(
|
|||
metadata = fetch(package_name)
|
||||
|
||||
times = metadata.get("time") or {}
|
||||
versions = metadata.get("versions") or {}
|
||||
if as_of is None:
|
||||
as_of = datetime.now(timezone.utc)
|
||||
cutoff = as_of - min_age
|
||||
|
|
@ -109,6 +112,8 @@ def resolve_pr_gate_version(
|
|||
for version, raw_ts in times.items():
|
||||
if version in _TIME_META_KEYS:
|
||||
continue
|
||||
if version not in versions:
|
||||
continue
|
||||
if not isinstance(raw_ts, str):
|
||||
continue
|
||||
if "-" in version:
|
||||
|
|
|
|||
|
|
@ -0,0 +1,78 @@
|
|||
"""Live e2e: `/utils/token_counter?call_endpoint=true` counts Gemini `contents` upstream.
|
||||
|
||||
Google's countTokens API is the only tokenizer that knows Gemini's real token
|
||||
boundaries, so the proxy must forward `contents` to it for both the AI Studio and
|
||||
Vertex deployments and hand back the provider's `promptTokensDetails`. Claude on
|
||||
Vertex is covered by `/v1/messages/count_tokens`; this is the Gemini `contents`
|
||||
route the claude_code rows never reach
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from proxy_client import ProxyClient
|
||||
from pydantic import BaseModel
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
GEMINI_DEPLOYMENTS = ("gemini-2.5-flash", "gemini-2.5-flash-vertex")
|
||||
|
||||
|
||||
class _Part(BaseModel):
|
||||
text: str
|
||||
|
||||
|
||||
class _Content(BaseModel):
|
||||
parts: tuple[_Part, ...]
|
||||
|
||||
|
||||
class _TokenCountBody(BaseModel):
|
||||
model: str
|
||||
contents: tuple[_Content, ...]
|
||||
|
||||
|
||||
class _CallEndpoint(BaseModel):
|
||||
call_endpoint: bool = True
|
||||
|
||||
|
||||
class _ModalityTokens(BaseModel):
|
||||
modality: str
|
||||
tokenCount: int
|
||||
|
||||
|
||||
class _CountTokensUpstream(BaseModel):
|
||||
totalTokens: int
|
||||
promptTokensDetails: tuple[_ModalityTokens, ...]
|
||||
|
||||
|
||||
class _TokenCountResponse(BaseModel):
|
||||
total_tokens: int
|
||||
request_model: str
|
||||
model_used: str
|
||||
tokenizer_type: str
|
||||
original_response: _CountTokensUpstream
|
||||
|
||||
|
||||
class TestGeminiContentsTokenCounting:
|
||||
@pytest.mark.parametrize("model", GEMINI_DEPLOYMENTS)
|
||||
def test_contents_are_counted_by_the_provider_endpoint(
|
||||
self, proxy: ProxyClient, scoped_key: str, model: str
|
||||
) -> None:
|
||||
text = f"Hello world, how are you doing today? {unique_marker()}"
|
||||
body = _TokenCountBody(model=model, contents=(_Content(parts=(_Part(text=text),)),))
|
||||
|
||||
result = proxy.transport.send(
|
||||
"/utils/token_counter",
|
||||
headers=proxy.transport.bearer(scoped_key),
|
||||
json=body,
|
||||
params=_CallEndpoint(),
|
||||
)
|
||||
|
||||
require_successful_call(result)
|
||||
counted = _TokenCountResponse.model_validate_json(result.body)
|
||||
assert counted.request_model == model, counted
|
||||
assert counted.original_response.totalTokens == counted.total_tokens > 0, counted
|
||||
assert counted.original_response.promptTokensDetails, counted
|
||||
assert all(detail.tokenCount > 0 for detail in counted.original_response.promptTokensDetails), counted
|
||||
|
|
@ -2,6 +2,7 @@ import { chromium, expect, request } from "@playwright/test";
|
|||
import { users, Role, STORAGE_PATHS } from "./fixtures/users";
|
||||
import { ARTIFACT_DIR, UI_BASE_URL } from "./constants";
|
||||
import { expectUnrestrictedDashboard, setInvitedUserPassword } from "./helpers/userOnboarding";
|
||||
import { hideLiteAdmin } from "./helpers/navigation";
|
||||
import * as fs from "fs";
|
||||
import * as path from "path";
|
||||
|
||||
|
|
@ -75,6 +76,9 @@ async function globalSetup() {
|
|||
if (await dismiss.isVisible({ timeout: 1_500 }).catch(() => false)) {
|
||||
await dismiss.click();
|
||||
}
|
||||
if (role === Role.ProxyAdmin) {
|
||||
await hideLiteAdmin(page);
|
||||
}
|
||||
// The login flow stores a post-login return URL in the litellm_return_url
|
||||
// cookie. If the snapshot captures it before the app consumes it, every
|
||||
// test inheriting this storageState gets yanked to that stale URL the
|
||||
|
|
|
|||
|
|
@ -62,6 +62,20 @@ export async function dismissFeedbackPopup(page: PlaywrightPage): Promise<void>
|
|||
}
|
||||
}
|
||||
|
||||
export async function hideLiteAdmin(page: PlaywrightPage): Promise<void> {
|
||||
await page.getByRole("button", { name: /Account menu/i }).click();
|
||||
const panel = page.getByTestId("sidebar-account-menu-panel");
|
||||
await expect(panel).toBeVisible({ timeout: 5_000 });
|
||||
const toggle = panel.getByRole("switch", { name: "Toggle hide LiteAdmin" });
|
||||
if ((await toggle.getAttribute("aria-checked")) !== "true") {
|
||||
await toggle.click();
|
||||
}
|
||||
await expect(toggle).toHaveAttribute("aria-checked", "true");
|
||||
await page.keyboard.press("Escape");
|
||||
await expect(panel).toBeHidden();
|
||||
await expect(page.getByRole("button", { name: "LiteAdmin", exact: true })).toBeHidden();
|
||||
}
|
||||
|
||||
/**
|
||||
* Click on a team ID in the table. Team IDs are rendered differently depending
|
||||
* on the component version — try button first (Tremor Button), fall back to
|
||||
|
|
|
|||
103
tests/integration/routing/test_user_config_routing.py
Normal file
103
tests/integration/routing/test_user_config_routing.py
Normal file
|
|
@ -0,0 +1,103 @@
|
|||
import json
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
from integration._support.client import Gateway
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
|
||||
USER_KEY: Final = "sk-user-supplied-" + uuid.uuid4().hex
|
||||
|
||||
|
||||
def _completion(request: Request) -> Reply:
|
||||
if request.target != "/v1/chat/completions":
|
||||
return Reply(status=404, body=b"{}")
|
||||
body: Final = json.loads(request.body)
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "chatcmpl-user-config",
|
||||
"object": "chat.completion",
|
||||
"created": 0,
|
||||
"model": body["model"],
|
||||
"choices": [
|
||||
{"index": 0, "message": {"role": "assistant", "content": "routed"}, "finish_reason": "stop"}
|
||||
],
|
||||
"usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _user_config(upstream_url: str) -> dict[str, object]:
|
||||
return {
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "user-config-deployment",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4.1-mini",
|
||||
"api_base": upstream_url + "/v1",
|
||||
"api_key": USER_KEY,
|
||||
},
|
||||
}
|
||||
],
|
||||
"num_retries": 0,
|
||||
}
|
||||
|
||||
|
||||
def _opt_in_config(directory: Path, upstream_url: str) -> Path:
|
||||
config: Final = directory / "allow_client_side_credentials_config.yaml"
|
||||
config.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "admin-deployment",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4.1-mini",
|
||||
"api_base": upstream_url + "/v1",
|
||||
"api_key": "sk-admin-configured",
|
||||
},
|
||||
}
|
||||
],
|
||||
"general_settings": {
|
||||
"master_key": "os.environ/LITELLM_MASTER_KEY",
|
||||
"database_url": "os.environ/DATABASE_URL",
|
||||
"store_model_in_db": True,
|
||||
"allow_client_side_credentials": True,
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
return config
|
||||
|
||||
|
||||
def _request_body(upstream_url: str) -> dict[str, object]:
|
||||
return {
|
||||
"model": "user-config-deployment",
|
||||
"messages": [{"role": "user", "content": "user config control"}],
|
||||
"user_config": _user_config(upstream_url),
|
||||
}
|
||||
|
||||
|
||||
def test_user_config_routes_to_the_user_supplied_deployment_when_opted_in(gateway: Gateway, tmp_path: Path) -> None:
|
||||
with wire_server(_completion) as upstream:
|
||||
config: Final = _opt_in_config(tmp_path, upstream.url)
|
||||
with owned_proxy(gateway, tmp_path, {}, config=config) as candidate:
|
||||
response: Final = candidate.request("POST", "/v1/chat/completions", _request_body(upstream.url))
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json()["choices"][0]["message"]["content"] == "routed"
|
||||
outbound: Final = tuple(upstream.received.get_nowait() for _ in range(upstream.received.qsize()))
|
||||
completions: Final = tuple(request for request in outbound if request.target == "/v1/chat/completions")
|
||||
assert len(completions) == 1, outbound
|
||||
assert completions[0].headers["authorization"] == f"Bearer {USER_KEY}"
|
||||
assert json.loads(completions[0].body)["model"] == "gpt-4.1-mini"
|
||||
|
||||
|
||||
def test_user_config_is_rejected_without_the_opt_in(gateway: Gateway) -> None:
|
||||
with wire_server(_completion) as upstream:
|
||||
response: Final = gateway.request("POST", "/v1/chat/completions", _request_body(upstream.url))
|
||||
assert response.status_code == 401, response.text
|
||||
assert "user_config is not allowed in request body" in response.text
|
||||
assert upstream.received.empty()
|
||||
|
|
@ -1,5 +1,3 @@
|
|||
import logging
|
||||
import os
|
||||
import pytest
|
||||
from mcp.types import Tool as MCPTool
|
||||
from typing import List, Any, cast
|
||||
|
|
@ -846,161 +844,6 @@ async def test_streaming_mcp_events_validation():
|
|||
assert mock_get_tools.called, "MCP tools should have been fetched"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
pytest.param("gpt-4o-mini", id="openai"),
|
||||
pytest.param("claude-haiku-4-5", id="anthropic"),
|
||||
],
|
||||
)
|
||||
async def test_streaming_responses_api_with_mcp_tools(
|
||||
model: str, caplog: pytest.LogCaptureFixture
|
||||
):
|
||||
"""
|
||||
Test the streaming responses API with MCP tools when using server_url="litellm_proxy"
|
||||
|
||||
Under the hood the follow occurs
|
||||
|
||||
- MCP: responses called litellm MCP manager.list_tools (MOCKED)
|
||||
- Request 1: Made to model under test with fetched tools (REAL LLM CALL)
|
||||
- MCP: Execute tool call from request 1 and returns result (MOCKED)
|
||||
- Request 2: Made to model under test with fetched tools and tool results (REAL LLM CALL)
|
||||
|
||||
Return the user the result of request 2
|
||||
"""
|
||||
# Skip test if API keys are not set for the respective models
|
||||
if ("claude" in model.lower() or "anthropic" in model.lower()) and not os.getenv(
|
||||
"ANTHROPIC_API_KEY"
|
||||
):
|
||||
pytest.skip("ANTHROPIC_API_KEY not set, skipping anthropic model test")
|
||||
if ("gpt" in model.lower() or "openai" in model.lower()) and not os.getenv(
|
||||
"OPENAI_API_KEY"
|
||||
):
|
||||
pytest.skip("OPENAI_API_KEY not set, skipping openai model test")
|
||||
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
print("🧪 Testing basic streaming with MCP tools...")
|
||||
|
||||
# Mock MCP tools that would be returned from the manager
|
||||
mock_mcp_tools = [
|
||||
MCPTool.model_validate({
|
||||
"name": "search_repo",
|
||||
"description": "Search BerriAI/litellm repository for information",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {"type": "string", "description": "Search query"}
|
||||
},
|
||||
"required": ["query"],
|
||||
},
|
||||
}, by_name=False)
|
||||
]
|
||||
|
||||
# Only mock the MCP-specific operations, let LLM responses be real
|
||||
with caplog.at_level(logging.ERROR):
|
||||
with (
|
||||
patch.object(
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
"_get_mcp_tools_from_manager",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_tools,
|
||||
patch.object(
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
"_execute_tool_calls",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_execute_tools,
|
||||
):
|
||||
# Setup MCP mocks only
|
||||
mock_get_tools.return_value = (mock_mcp_tools, ["litellm_proxy"])
|
||||
|
||||
# Create a dynamic mock that will match the actual tool call ID from the LLM response
|
||||
def mock_execute_tool_calls_side_effect(
|
||||
tool_calls, user_api_key_auth, **kwargs
|
||||
):
|
||||
"""Mock function that returns results matching the actual tool call IDs from the LLM"""
|
||||
results = []
|
||||
for tool_call in tool_calls:
|
||||
# Extract call_id from the tool call
|
||||
call_id = None
|
||||
if isinstance(tool_call, dict):
|
||||
call_id = tool_call.get("call_id") or tool_call.get("id")
|
||||
elif hasattr(tool_call, "call_id"):
|
||||
call_id = tool_call.call_id
|
||||
elif hasattr(tool_call, "id"):
|
||||
call_id = tool_call.id
|
||||
|
||||
if call_id:
|
||||
results.append(
|
||||
{
|
||||
"tool_call_id": call_id,
|
||||
"result": "LiteLLM is a unified interface for 100+ LLMs that translates inputs to provider-specific completion endpoints and provides consistent OpenAI-format output.",
|
||||
}
|
||||
)
|
||||
return results
|
||||
|
||||
mock_execute_tools.side_effect = mock_execute_tool_calls_side_effect
|
||||
|
||||
# Make the actual call - LLM responses will be real
|
||||
mcp_tool_config = cast(
|
||||
Any,
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_url": "litellm_proxy",
|
||||
"require_approval": "never",
|
||||
},
|
||||
)
|
||||
response = await litellm.aresponses(
|
||||
model=model,
|
||||
tools=[mcp_tool_config],
|
||||
tool_choice="required",
|
||||
input=[
|
||||
{
|
||||
"role": "user",
|
||||
"type": "message",
|
||||
"content": "give me a TLDR of what BerriAI/litellm is about",
|
||||
}
|
||||
],
|
||||
stream=True,
|
||||
)
|
||||
|
||||
print(f"📋 Response type: {type(response)}")
|
||||
assert hasattr(
|
||||
response, "__aiter__"
|
||||
), "Response should be an async streaming response"
|
||||
|
||||
# Collect streaming chunks
|
||||
chunks = []
|
||||
async for chunk in response:
|
||||
chunks.append(chunk)
|
||||
print(f"📦 Chunk type: {getattr(chunk, 'type', 'unknown')}")
|
||||
|
||||
print(f"📊 Total chunks received: {len(chunks)}")
|
||||
|
||||
# Verify MCP mocks were called (may be called multiple times in streaming)
|
||||
assert (
|
||||
mock_get_tools.call_count >= 1
|
||||
), f"Expected MCP tools to be fetched at least once, got {mock_get_tools.call_count}"
|
||||
print(f"MCP tools fetched: {len(mock_mcp_tools)}")
|
||||
|
||||
# Verify we got a response
|
||||
assert response is not None
|
||||
assert len(chunks) > 0, "Should have received streaming chunks"
|
||||
|
||||
print("Basic streaming responses API with MCP tools test passed!")
|
||||
|
||||
lite_errors = [
|
||||
record
|
||||
for record in caplog.records
|
||||
if record.levelno >= logging.ERROR
|
||||
and ("LiteLLM" in record.name or "LiteLLM" in record.getMessage())
|
||||
]
|
||||
assert not lite_errors, "Unexpected LiteLLM errors: " + ", ".join(
|
||||
record.getMessage() for record in lite_errors
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_parameter_preparation_helpers():
|
||||
"""
|
||||
|
|
@ -1215,7 +1058,7 @@ async def test_no_duplicate_mcp_tools_in_streaming_e2e():
|
|||
The test mocks the MCP manager response but validates the actual tools
|
||||
sent to the LLM to ensure no duplication occurs.
|
||||
"""
|
||||
from unittest.mock import AsyncMock, patch, call
|
||||
from unittest.mock import AsyncMock, patch
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
)
|
||||
|
|
@ -1432,221 +1275,3 @@ async def test_no_duplicate_mcp_tools_in_streaming_e2e():
|
|||
"tools_per_call": [len(tools) for tools in llm_call_tools],
|
||||
"duplicate_tools_found": False,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", ["gpt-4o-mini"])
|
||||
async def test_streaming_mcp_event_order_and_response_id_consistency(
|
||||
model: str, caplog: pytest.LogCaptureFixture
|
||||
):
|
||||
"""
|
||||
Test that:
|
||||
1. Streaming events are emitted in correct order (response.created, response.in_progress, response.output_item.added before MCP events)
|
||||
2. All response lifecycle events share the same response ID within a cycle
|
||||
"""
|
||||
if ("gpt" in model.lower() or "openai" in model.lower()) and not os.getenv(
|
||||
"OPENAI_API_KEY"
|
||||
):
|
||||
pytest.skip("OPENAI_API_KEY not set, skipping openai model test")
|
||||
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
mock_mcp_tools = [
|
||||
MCPTool.model_validate({
|
||||
"name": "get_weather",
|
||||
"description": "Get weather for a city",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"city": {"type": "string", "description": "City name"}
|
||||
},
|
||||
"required": ["city"],
|
||||
},
|
||||
}, by_name=False)
|
||||
]
|
||||
|
||||
with caplog.at_level(logging.ERROR):
|
||||
with (
|
||||
patch.object(
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
"_get_mcp_tools_from_manager",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_tools,
|
||||
patch.object(
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
"_execute_tool_calls",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_execute_tools,
|
||||
):
|
||||
mock_get_tools.return_value = (mock_mcp_tools, ["litellm_proxy"])
|
||||
|
||||
def mock_execute_side_effect(tool_calls, user_api_key_auth, **kwargs):
|
||||
results = []
|
||||
for tool_call in tool_calls:
|
||||
call_id = None
|
||||
if isinstance(tool_call, dict):
|
||||
call_id = tool_call.get("call_id") or tool_call.get("id")
|
||||
elif hasattr(tool_call, "call_id"):
|
||||
call_id = tool_call.call_id
|
||||
elif hasattr(tool_call, "id"):
|
||||
call_id = tool_call.id
|
||||
if call_id:
|
||||
results.append(
|
||||
{
|
||||
"tool_call_id": call_id,
|
||||
"result": "Sunny, 72°F",
|
||||
}
|
||||
)
|
||||
return results
|
||||
|
||||
mock_execute_tools.side_effect = mock_execute_side_effect
|
||||
|
||||
mcp_tool_config = cast(
|
||||
Any,
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_url": "litellm_proxy",
|
||||
"require_approval": "never",
|
||||
},
|
||||
)
|
||||
|
||||
response = await litellm.aresponses(
|
||||
model=model,
|
||||
tools=[mcp_tool_config],
|
||||
input=[
|
||||
{
|
||||
"role": "user",
|
||||
"type": "message",
|
||||
"content": "What's the weather in San Francisco?",
|
||||
}
|
||||
],
|
||||
stream=True,
|
||||
)
|
||||
|
||||
events = []
|
||||
async for chunk in response:
|
||||
events.append(chunk)
|
||||
|
||||
assert len(events) > 0, "Should receive streaming events"
|
||||
|
||||
created_idx = next(
|
||||
(
|
||||
i
|
||||
for i, e in enumerate(events)
|
||||
if getattr(e, "type", None) == "response.created"
|
||||
),
|
||||
None,
|
||||
)
|
||||
in_progress_idx = next(
|
||||
(
|
||||
i
|
||||
for i, e in enumerate(events)
|
||||
if getattr(e, "type", None) == "response.in_progress"
|
||||
),
|
||||
None,
|
||||
)
|
||||
output_item_added_idx = next(
|
||||
(
|
||||
i
|
||||
for i, e in enumerate(events)
|
||||
if getattr(e, "type", None) == "response.output_item.added"
|
||||
),
|
||||
None,
|
||||
)
|
||||
mcp_in_progress_idx = next(
|
||||
(
|
||||
i
|
||||
for i, e in enumerate(events)
|
||||
if "mcp_list_tools.in_progress" in str(getattr(e, "type", ""))
|
||||
),
|
||||
None,
|
||||
)
|
||||
completed_idx = next(
|
||||
(
|
||||
i
|
||||
for i, e in enumerate(events)
|
||||
if getattr(e, "type", None) == "response.completed"
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
assert created_idx is not None, "response.created event should be present"
|
||||
assert (
|
||||
in_progress_idx is not None
|
||||
), "response.in_progress event should be present"
|
||||
assert (
|
||||
output_item_added_idx is not None
|
||||
), "response.output_item.added event should be present"
|
||||
|
||||
assert (
|
||||
created_idx < in_progress_idx
|
||||
), "response.created should come before response.in_progress"
|
||||
assert (
|
||||
in_progress_idx < output_item_added_idx
|
||||
), "response.in_progress should come before response.output_item.added"
|
||||
|
||||
if mcp_in_progress_idx is not None:
|
||||
assert (
|
||||
output_item_added_idx < mcp_in_progress_idx
|
||||
), "response.output_item.added should come before response.mcp_list_tools.in_progress"
|
||||
|
||||
response_ids = []
|
||||
for i, event in enumerate(events):
|
||||
event_type = getattr(event, "type", None)
|
||||
if hasattr(event, "response"):
|
||||
response_obj = getattr(event, "response", None)
|
||||
if response_obj and hasattr(response_obj, "id"):
|
||||
event_type_value = (
|
||||
event_type.value
|
||||
if hasattr(event_type, "value")
|
||||
else str(event_type)
|
||||
)
|
||||
if any(
|
||||
x in event_type_value
|
||||
for x in [
|
||||
"response.created",
|
||||
"response.in_progress",
|
||||
"response.completed",
|
||||
]
|
||||
):
|
||||
response_ids.append((i, event_type_value, response_obj.id))
|
||||
|
||||
assert (
|
||||
len(response_ids) >= 2
|
||||
), f"Should have at least 2 response lifecycle events. Found {len(response_ids)}"
|
||||
|
||||
cycles = []
|
||||
current_cycle = []
|
||||
current_id = None
|
||||
|
||||
for idx, event_type, resp_id in response_ids:
|
||||
if current_id is None or resp_id == current_id:
|
||||
current_cycle.append((idx, event_type, resp_id))
|
||||
current_id = resp_id
|
||||
else:
|
||||
if current_cycle:
|
||||
cycles.append(current_cycle)
|
||||
current_cycle = [(idx, event_type, resp_id)]
|
||||
current_id = resp_id
|
||||
if current_cycle:
|
||||
cycles.append(current_cycle)
|
||||
|
||||
for cycle_num, cycle in enumerate(cycles):
|
||||
cycle_ids = set(resp_id for _, _, resp_id in cycle)
|
||||
assert (
|
||||
len(cycle_ids) == 1
|
||||
), f"Cycle {cycle_num + 1} should have consistent response ID. Found {len(cycle_ids)} unique IDs"
|
||||
|
||||
assert (
|
||||
completed_idx is not None
|
||||
), "response.completed event should be present"
|
||||
|
||||
lite_errors = [
|
||||
record
|
||||
for record in caplog.records
|
||||
if record.levelno >= logging.ERROR
|
||||
and ("LiteLLM" in record.name or "LiteLLM" in record.getMessage())
|
||||
]
|
||||
assert not lite_errors, "Unexpected LiteLLM errors: " + ", ".join(
|
||||
record.getMessage() for record in lite_errors
|
||||
)
|
||||
|
|
|
|||
389
tests/mcp_tests/test_aresponses_api_with_mcp_providers.py
Normal file
389
tests/mcp_tests/test_aresponses_api_with_mcp_providers.py
Normal file
|
|
@ -0,0 +1,389 @@
|
|||
import logging
|
||||
import os
|
||||
import pytest
|
||||
from mcp.types import Tool as MCPTool
|
||||
from typing import Any, cast
|
||||
|
||||
import litellm
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import LiteLLM_Proxy_MCP_Handler
|
||||
|
||||
|
||||
class MockUserAPIKeyAuth:
|
||||
"""Mock UserAPIKeyAuth for testing"""
|
||||
|
||||
def __init__(self):
|
||||
self.api_key = "test_key"
|
||||
self.user_id = "test_user"
|
||||
self.team_id = "test_team"
|
||||
self.user_email = "test@example.com"
|
||||
self.max_budget = 100.0
|
||||
self.spend = 0.0
|
||||
self.models = []
|
||||
self.aliases = {}
|
||||
self.config = {}
|
||||
self.permissions = {}
|
||||
self.metadata = {}
|
||||
self.object_permission_id = "test_permission_id"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
pytest.param("gpt-4o-mini", id="openai"),
|
||||
pytest.param("claude-haiku-4-5", id="anthropic"),
|
||||
],
|
||||
)
|
||||
async def test_streaming_responses_api_with_mcp_tools(
|
||||
model: str, caplog: pytest.LogCaptureFixture
|
||||
):
|
||||
"""
|
||||
Test the streaming responses API with MCP tools when using server_url="litellm_proxy"
|
||||
|
||||
Under the hood the follow occurs
|
||||
|
||||
- MCP: responses called litellm MCP manager.list_tools (MOCKED)
|
||||
- Request 1: Made to model under test with fetched tools (REAL LLM CALL)
|
||||
- MCP: Execute tool call from request 1 and returns result (MOCKED)
|
||||
- Request 2: Made to model under test with fetched tools and tool results (REAL LLM CALL)
|
||||
|
||||
Return the user the result of request 2
|
||||
"""
|
||||
if ("claude" in model.lower() or "anthropic" in model.lower()) and not os.getenv(
|
||||
"ANTHROPIC_API_KEY"
|
||||
):
|
||||
pytest.skip("ANTHROPIC_API_KEY not set, skipping anthropic model test")
|
||||
if ("gpt" in model.lower() or "openai" in model.lower()) and not os.getenv(
|
||||
"OPENAI_API_KEY"
|
||||
):
|
||||
pytest.skip("OPENAI_API_KEY not set, skipping openai model test")
|
||||
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
print("🧪 Testing basic streaming with MCP tools...")
|
||||
|
||||
mock_mcp_tools = [
|
||||
MCPTool.model_validate({
|
||||
"name": "search_repo",
|
||||
"description": "Search BerriAI/litellm repository for information",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {"type": "string", "description": "Search query"}
|
||||
},
|
||||
"required": ["query"],
|
||||
},
|
||||
}, by_name=False)
|
||||
]
|
||||
|
||||
with caplog.at_level(logging.ERROR):
|
||||
with (
|
||||
patch.object(
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
"_get_mcp_tools_from_manager",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_tools,
|
||||
patch.object(
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
"_execute_tool_calls",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_execute_tools,
|
||||
):
|
||||
mock_get_tools.return_value = (mock_mcp_tools, ["litellm_proxy"])
|
||||
|
||||
def mock_execute_tool_calls_side_effect(
|
||||
tool_calls, user_api_key_auth, **kwargs
|
||||
):
|
||||
"""Mock function that returns results matching the actual tool call IDs from the LLM"""
|
||||
results = []
|
||||
for tool_call in tool_calls:
|
||||
call_id = None
|
||||
if isinstance(tool_call, dict):
|
||||
call_id = tool_call.get("call_id") or tool_call.get("id")
|
||||
elif hasattr(tool_call, "call_id"):
|
||||
call_id = tool_call.call_id
|
||||
elif hasattr(tool_call, "id"):
|
||||
call_id = tool_call.id
|
||||
|
||||
if call_id:
|
||||
results.append(
|
||||
{
|
||||
"tool_call_id": call_id,
|
||||
"result": "LiteLLM is a unified interface for 100+ LLMs that translates inputs to provider-specific completion endpoints and provides consistent OpenAI-format output.",
|
||||
}
|
||||
)
|
||||
return results
|
||||
|
||||
mock_execute_tools.side_effect = mock_execute_tool_calls_side_effect
|
||||
|
||||
mcp_tool_config = cast(
|
||||
Any,
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_url": "litellm_proxy",
|
||||
"require_approval": "never",
|
||||
},
|
||||
)
|
||||
response = await litellm.aresponses(
|
||||
model=model,
|
||||
tools=[mcp_tool_config],
|
||||
tool_choice="required",
|
||||
input=[
|
||||
{
|
||||
"role": "user",
|
||||
"type": "message",
|
||||
"content": "give me a TLDR of what BerriAI/litellm is about",
|
||||
}
|
||||
],
|
||||
stream=True,
|
||||
)
|
||||
|
||||
print(f"📋 Response type: {type(response)}")
|
||||
assert hasattr(
|
||||
response, "__aiter__"
|
||||
), "Response should be an async streaming response"
|
||||
|
||||
chunks = []
|
||||
async for chunk in response:
|
||||
chunks.append(chunk)
|
||||
print(f"📦 Chunk type: {getattr(chunk, 'type', 'unknown')}")
|
||||
|
||||
print(f"📊 Total chunks received: {len(chunks)}")
|
||||
|
||||
assert (
|
||||
mock_get_tools.call_count >= 1
|
||||
), f"Expected MCP tools to be fetched at least once, got {mock_get_tools.call_count}"
|
||||
print(f"MCP tools fetched: {len(mock_mcp_tools)}")
|
||||
|
||||
assert response is not None
|
||||
assert len(chunks) > 0, "Should have received streaming chunks"
|
||||
|
||||
print("Basic streaming responses API with MCP tools test passed!")
|
||||
|
||||
lite_errors = [
|
||||
record
|
||||
for record in caplog.records
|
||||
if record.levelno >= logging.ERROR
|
||||
and ("LiteLLM" in record.name or "LiteLLM" in record.getMessage())
|
||||
]
|
||||
assert not lite_errors, "Unexpected LiteLLM errors: " + ", ".join(
|
||||
record.getMessage() for record in lite_errors
|
||||
)
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", ["gpt-4o-mini"])
|
||||
async def test_streaming_mcp_event_order_and_response_id_consistency(
|
||||
model: str, caplog: pytest.LogCaptureFixture
|
||||
):
|
||||
"""
|
||||
Test that:
|
||||
1. Streaming events are emitted in correct order (response.created, response.in_progress, response.output_item.added before MCP events)
|
||||
2. All response lifecycle events share the same response ID within a cycle
|
||||
"""
|
||||
if ("gpt" in model.lower() or "openai" in model.lower()) and not os.getenv(
|
||||
"OPENAI_API_KEY"
|
||||
):
|
||||
pytest.skip("OPENAI_API_KEY not set, skipping openai model test")
|
||||
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
mock_mcp_tools = [
|
||||
MCPTool.model_validate({
|
||||
"name": "get_weather",
|
||||
"description": "Get weather for a city",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"city": {"type": "string", "description": "City name"}
|
||||
},
|
||||
"required": ["city"],
|
||||
},
|
||||
}, by_name=False)
|
||||
]
|
||||
|
||||
with caplog.at_level(logging.ERROR):
|
||||
with (
|
||||
patch.object(
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
"_get_mcp_tools_from_manager",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_tools,
|
||||
patch.object(
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
"_execute_tool_calls",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_execute_tools,
|
||||
):
|
||||
mock_get_tools.return_value = (mock_mcp_tools, ["litellm_proxy"])
|
||||
|
||||
def mock_execute_side_effect(tool_calls, user_api_key_auth, **kwargs):
|
||||
results = []
|
||||
for tool_call in tool_calls:
|
||||
call_id = None
|
||||
if isinstance(tool_call, dict):
|
||||
call_id = tool_call.get("call_id") or tool_call.get("id")
|
||||
elif hasattr(tool_call, "call_id"):
|
||||
call_id = tool_call.call_id
|
||||
elif hasattr(tool_call, "id"):
|
||||
call_id = tool_call.id
|
||||
if call_id:
|
||||
results.append(
|
||||
{
|
||||
"tool_call_id": call_id,
|
||||
"result": "Sunny, 72°F",
|
||||
}
|
||||
)
|
||||
return results
|
||||
|
||||
mock_execute_tools.side_effect = mock_execute_side_effect
|
||||
|
||||
mcp_tool_config = cast(
|
||||
Any,
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_url": "litellm_proxy",
|
||||
"require_approval": "never",
|
||||
},
|
||||
)
|
||||
|
||||
response = await litellm.aresponses(
|
||||
model=model,
|
||||
tools=[mcp_tool_config],
|
||||
input=[
|
||||
{
|
||||
"role": "user",
|
||||
"type": "message",
|
||||
"content": "What's the weather in San Francisco?",
|
||||
}
|
||||
],
|
||||
stream=True,
|
||||
)
|
||||
|
||||
events = []
|
||||
async for chunk in response:
|
||||
events.append(chunk)
|
||||
|
||||
assert len(events) > 0, "Should receive streaming events"
|
||||
|
||||
created_idx = next(
|
||||
(
|
||||
i
|
||||
for i, e in enumerate(events)
|
||||
if getattr(e, "type", None) == "response.created"
|
||||
),
|
||||
None,
|
||||
)
|
||||
in_progress_idx = next(
|
||||
(
|
||||
i
|
||||
for i, e in enumerate(events)
|
||||
if getattr(e, "type", None) == "response.in_progress"
|
||||
),
|
||||
None,
|
||||
)
|
||||
output_item_added_idx = next(
|
||||
(
|
||||
i
|
||||
for i, e in enumerate(events)
|
||||
if getattr(e, "type", None) == "response.output_item.added"
|
||||
),
|
||||
None,
|
||||
)
|
||||
mcp_in_progress_idx = next(
|
||||
(
|
||||
i
|
||||
for i, e in enumerate(events)
|
||||
if "mcp_list_tools.in_progress" in str(getattr(e, "type", ""))
|
||||
),
|
||||
None,
|
||||
)
|
||||
completed_idx = next(
|
||||
(
|
||||
i
|
||||
for i, e in enumerate(events)
|
||||
if getattr(e, "type", None) == "response.completed"
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
assert created_idx is not None, "response.created event should be present"
|
||||
assert (
|
||||
in_progress_idx is not None
|
||||
), "response.in_progress event should be present"
|
||||
assert (
|
||||
output_item_added_idx is not None
|
||||
), "response.output_item.added event should be present"
|
||||
|
||||
assert (
|
||||
created_idx < in_progress_idx
|
||||
), "response.created should come before response.in_progress"
|
||||
assert (
|
||||
in_progress_idx < output_item_added_idx
|
||||
), "response.in_progress should come before response.output_item.added"
|
||||
|
||||
if mcp_in_progress_idx is not None:
|
||||
assert (
|
||||
output_item_added_idx < mcp_in_progress_idx
|
||||
), "response.output_item.added should come before response.mcp_list_tools.in_progress"
|
||||
|
||||
response_ids = []
|
||||
for i, event in enumerate(events):
|
||||
event_type = getattr(event, "type", None)
|
||||
if hasattr(event, "response"):
|
||||
response_obj = getattr(event, "response", None)
|
||||
if response_obj and hasattr(response_obj, "id"):
|
||||
event_type_value = (
|
||||
event_type.value
|
||||
if hasattr(event_type, "value")
|
||||
else str(event_type)
|
||||
)
|
||||
if any(
|
||||
x in event_type_value
|
||||
for x in [
|
||||
"response.created",
|
||||
"response.in_progress",
|
||||
"response.completed",
|
||||
]
|
||||
):
|
||||
response_ids.append((i, event_type_value, response_obj.id))
|
||||
|
||||
assert (
|
||||
len(response_ids) >= 2
|
||||
), f"Should have at least 2 response lifecycle events. Found {len(response_ids)}"
|
||||
|
||||
cycles = []
|
||||
current_cycle = []
|
||||
current_id = None
|
||||
|
||||
for idx, event_type, resp_id in response_ids:
|
||||
if current_id is None or resp_id == current_id:
|
||||
current_cycle.append((idx, event_type, resp_id))
|
||||
current_id = resp_id
|
||||
else:
|
||||
if current_cycle:
|
||||
cycles.append(current_cycle)
|
||||
current_cycle = [(idx, event_type, resp_id)]
|
||||
current_id = resp_id
|
||||
if current_cycle:
|
||||
cycles.append(current_cycle)
|
||||
|
||||
for cycle_num, cycle in enumerate(cycles):
|
||||
cycle_ids = set(resp_id for _, _, resp_id in cycle)
|
||||
assert (
|
||||
len(cycle_ids) == 1
|
||||
), f"Cycle {cycle_num + 1} should have consistent response ID. Found {len(cycle_ids)} unique IDs"
|
||||
|
||||
assert (
|
||||
completed_idx is not None
|
||||
), "response.completed event should be present"
|
||||
|
||||
lite_errors = [
|
||||
record
|
||||
for record in caplog.records
|
||||
if record.levelno >= logging.ERROR
|
||||
and ("LiteLLM" in record.name or "LiteLLM" in record.getMessage())
|
||||
]
|
||||
assert not lite_errors, "Unexpected LiteLLM errors: " + ", ".join(
|
||||
record.getMessage() for record in lite_errors
|
||||
)
|
||||
|
|
@ -1,114 +0,0 @@
|
|||
import sys, os
|
||||
import traceback
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
import io
|
||||
|
||||
# this file is to test litellm/proxy
|
||||
|
||||
import pytest, logging, asyncio
|
||||
import litellm
|
||||
from litellm import embedding, completion, completion_cost, Timeout
|
||||
from litellm import RateLimitError
|
||||
|
||||
# Configure logging
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG, # Set the desired logging level
|
||||
format="%(asctime)s - %(levelname)s - %(message)s",
|
||||
)
|
||||
|
||||
# test /chat/completion request to the proxy
|
||||
from fastapi.testclient import TestClient
|
||||
from fastapi import FastAPI
|
||||
from litellm.proxy.proxy_server import (
|
||||
router,
|
||||
save_worker_config,
|
||||
initialize,
|
||||
) # Replace with the actual module where your FastAPI router is defined
|
||||
|
||||
# Your bearer token
|
||||
token = "sk-1234"
|
||||
|
||||
headers = {"Authorization": f"Bearer {token}"}
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
def client_no_auth():
|
||||
# Assuming litellm.proxy.proxy_server is an object
|
||||
from litellm.proxy.proxy_server import cleanup_router_config_variables
|
||||
|
||||
cleanup_router_config_variables()
|
||||
filepath = os.path.dirname(os.path.abspath(__file__))
|
||||
config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml"
|
||||
# initialize can get run in parallel, it sets specific variables for the fast api app, sinc eit gets run in parallel different tests use the wrong variables
|
||||
asyncio.run(initialize(config=config_fp, debug=True))
|
||||
app = FastAPI()
|
||||
app.include_router(router) # Include your router in the test app
|
||||
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
os.environ.get("AZURE_AI_API_KEY") is None
|
||||
or os.environ.get("OPENAI_API_KEY") is None,
|
||||
reason="AZURE_AI_API_KEY or OPENAI_API_KEY not set - skipping integration test",
|
||||
)
|
||||
def test_chat_completion(client_no_auth):
|
||||
global headers
|
||||
|
||||
from litellm.types.router import RouterConfig, ModelConfig
|
||||
from litellm.types.completion import CompletionRequest
|
||||
|
||||
user_config = RouterConfig(
|
||||
model_list=[
|
||||
ModelConfig(
|
||||
model_name="user-azure-instance",
|
||||
litellm_params=CompletionRequest(
|
||||
model="azure/gpt-4.1-mini",
|
||||
api_key=os.getenv("AZURE_AI_API_KEY"),
|
||||
api_version=os.getenv("AZURE_API_VERSION"),
|
||||
api_base=os.getenv("AZURE_AI_API_BASE"),
|
||||
timeout=10,
|
||||
),
|
||||
tpm=240000,
|
||||
rpm=1800,
|
||||
),
|
||||
ModelConfig(
|
||||
model_name="user-openai-instance",
|
||||
litellm_params=CompletionRequest(
|
||||
model="gpt-3.5-turbo",
|
||||
api_key=os.getenv("OPENAI_API_KEY"),
|
||||
timeout=10,
|
||||
),
|
||||
tpm=240000,
|
||||
rpm=1800,
|
||||
),
|
||||
],
|
||||
num_retries=2,
|
||||
allowed_fails=3,
|
||||
fallbacks=[{"user-azure-instance": ["user-openai-instance"]}],
|
||||
).dict()
|
||||
|
||||
try:
|
||||
# Your test data
|
||||
test_data = {
|
||||
"model": "user-azure-instance",
|
||||
"messages": [
|
||||
{"role": "user", "content": "hi"},
|
||||
],
|
||||
"max_tokens": 10,
|
||||
"user_config": user_config,
|
||||
}
|
||||
|
||||
print("testing proxy server with chat completions")
|
||||
response = client_no_auth.post("/v1/chat/completions", json=test_data)
|
||||
print(f"response - {response.text}")
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
print(f"Received response: {result}")
|
||||
except Exception as e:
|
||||
pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}")
|
||||
|
||||
|
||||
# Run the test
|
||||
|
|
@ -0,0 +1,51 @@
|
|||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not os.getenv("GEMINI_API_KEY") and not os.getenv("GOOGLE_API_KEY"),
|
||||
reason="Requires GEMINI_API_KEY or GOOGLE_API_KEY.",
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_pass_through_endpoint():
|
||||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
Request,
|
||||
Response,
|
||||
gemini_proxy_route,
|
||||
)
|
||||
|
||||
body = b"""
|
||||
{
|
||||
"contents": [{
|
||||
"parts":[{
|
||||
"text": "The quick brown fox jumps over the lazy dog."
|
||||
}]
|
||||
}]
|
||||
}
|
||||
"""
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/gemini/v1beta/models/gemini-2.5-flash:countTokens",
|
||||
"query_string": b"key=sk-1234",
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
],
|
||||
}
|
||||
|
||||
async def async_receive():
|
||||
return {"type": "http.request", "body": body, "more_body": False}
|
||||
|
||||
request = Request(
|
||||
scope=scope,
|
||||
receive=async_receive,
|
||||
)
|
||||
|
||||
await gemini_proxy_route(
|
||||
endpoint="v1beta/models/gemini-2.5-flash:countTokens?key=sk-1234",
|
||||
request=request,
|
||||
fastapi_response=Response(),
|
||||
)
|
||||
|
||||
|
|
@ -53,7 +53,8 @@ def client_no_auth():
|
|||
config_fp = (
|
||||
repo_root
|
||||
/ "tests"
|
||||
/ "proxy_unit_tests"
|
||||
/ "unit"
|
||||
/ "proxy"
|
||||
/ "test_configs"
|
||||
/ "test_config_no_auth.yaml"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -322,7 +322,9 @@ def test_get_proxy_model_info_shows_litellm_params_pricing_and_names_it_as_an_ov
|
|||
def test_get_proxy_model_info_names_config_model_info_pricing_as_an_override(monkeypatch, local_model_cost_map):
|
||||
"""Pricing declared under ``model_info`` in config.yaml overrides the cost map too."""
|
||||
info = _enriched_model_info(
|
||||
monkeypatch, {"model": "openai/gpt-5.6"}, {"id": "dep-config", "db_model": False, "output_cost_per_token": 7e-06}
|
||||
monkeypatch,
|
||||
{"model": "openai/gpt-5.6"},
|
||||
{"id": "dep-config", "db_model": False, "output_cost_per_token": 7e-06},
|
||||
)
|
||||
assert info["pricing_overrides"] == ("output_cost_per_token",)
|
||||
assert info["output_cost_per_token"] == 7e-06
|
||||
|
|
@ -399,7 +401,9 @@ def test_model_info_reports_null_cost_for_unpriced_deployment_and_zero_for_decla
|
|||
|
||||
def enriched_cost(model_name: str) -> tuple:
|
||||
deployment = router.get_model_list(model_name=model_name)[0]
|
||||
info = proxy_server._enrich_model_info_with_litellm_data({**deployment, "model_info": dict(deployment["model_info"])})["model_info"]
|
||||
info = proxy_server._enrich_model_info_with_litellm_data(
|
||||
{**deployment, "model_info": dict(deployment["model_info"])}
|
||||
)["model_info"]
|
||||
return info.get("input_cost_per_token"), info.get("output_cost_per_token")
|
||||
|
||||
assert enriched_cost("vllm-unpriced") == (None, None)
|
||||
|
|
@ -643,7 +647,6 @@ def model_group_info_router(monkeypatch):
|
|||
monkeypatch.setattr(proxy_server, "user_model", None)
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
monkeypatch.setattr(proxy_server, "proxy_logging_obj", None)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", None)
|
||||
monkeypatch.setattr(proxy_server, "_get_model_group_info", model_group_info)
|
||||
|
||||
|
|
@ -671,7 +674,9 @@ def test_model_group_info_proxy_admin_ignores_key_model_restriction(
|
|||
|
||||
|
||||
@pytest.mark.parametrize("admin_role", ["proxy_admin", "proxy_admin_viewer"])
|
||||
def test_model_group_info_proxy_admin_expands_wildcard_deployments(client, auth_as, model_group_info_router, admin_role):
|
||||
def test_model_group_info_proxy_admin_expands_wildcard_deployments(
|
||||
client, auth_as, model_group_info_router, admin_role
|
||||
):
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.proxy.auth.model_checks import get_known_models_from_wildcard
|
||||
|
||||
|
|
|
|||
425
tests/test_litellm/proxy/test_model_list_callback_filter.py
Normal file
425
tests/test_litellm/proxy/test_model_list_callback_filter.py
Normal file
|
|
@ -0,0 +1,425 @@
|
|||
"""
|
||||
Tests for `CustomLogger.async_filter_listed_models` on the model listing endpoints:
|
||||
GET /v1/models (`model_list`, OpenAI and Anthropic shapes), GET /v1/models/{id}
|
||||
(`model_info`), GET /v1/model/info (`model_info_v1`) and GET /model_group/info
|
||||
(`model_group_info`). A registered callback that overrides the hook decides per
|
||||
caller which of the names the route would list are kept; the rest disappear and
|
||||
`/v1/models/{id}` answers 404 for them.
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from starlette.requests import Request
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
|
||||
class _Gate(CustomLogger):
|
||||
def __init__(self, hidden: frozenset[str] = frozenset(), extra: tuple[str, ...] = ()) -> None:
|
||||
super().__init__()
|
||||
self.hidden = hidden
|
||||
self.extra = extra
|
||||
self.seen: list[tuple[str, ...]] = []
|
||||
|
||||
async def async_filter_listed_models(
|
||||
self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
|
||||
) -> Sequence[str]:
|
||||
self.seen.append(tuple(model_names))
|
||||
return [*(name for name in model_names if name not in self.hidden), *self.extra]
|
||||
|
||||
|
||||
class _InferenceOnlyGate(CustomLogger):
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
if data.get("model") == "restricted-model":
|
||||
raise HTTPException(status_code=403, detail="not entitled to this model")
|
||||
return data
|
||||
|
||||
|
||||
class _RaisingGate(CustomLogger):
|
||||
async def async_filter_listed_models(
|
||||
self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
|
||||
) -> Sequence[str]:
|
||||
raise HTTPException(status_code=503, detail="entitlement service down")
|
||||
|
||||
|
||||
class _ReversingGate(CustomLogger):
|
||||
async def async_filter_listed_models(
|
||||
self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
|
||||
) -> Sequence[str]:
|
||||
return list(reversed(model_names))
|
||||
|
||||
|
||||
class _StringReturningGate(CustomLogger):
|
||||
async def async_filter_listed_models(self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]) -> str:
|
||||
return "open-model"
|
||||
|
||||
|
||||
def _deployment(model_name: str, model: str = "openai/gpt-4o", **model_info):
|
||||
return {
|
||||
"model_name": model_name,
|
||||
"litellm_params": {"model": model, "api_key": "sk-fake"},
|
||||
"model_info": {"id": f"{model_name}-id", **model_info},
|
||||
}
|
||||
|
||||
|
||||
def _install_router(monkeypatch, *deployments, **router_kwargs) -> Router:
|
||||
router = Router(model_list=list(deployments), **router_kwargs)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
monkeypatch.setattr(proxy_server, "llm_model_list", router.model_list)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
monkeypatch.setattr(proxy_server, "user_model", None)
|
||||
return router
|
||||
|
||||
|
||||
def _register(monkeypatch, *callbacks: CustomLogger) -> None:
|
||||
monkeypatch.setattr(litellm, "callbacks", list(callbacks))
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def two_model_router(monkeypatch) -> Router:
|
||||
return _install_router(monkeypatch, _deployment("open-model"), _deployment("restricted-model"))
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def team_router(monkeypatch) -> Router:
|
||||
return _install_router(
|
||||
monkeypatch,
|
||||
_deployment("gpt-4"),
|
||||
_deployment("model_name_team1_abc", team_id="team1", team_public_model_name="team-gpt"),
|
||||
_deployment("model_name_team1_def", team_id="team1", team_public_model_name="team-chat"),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def team_admin_privileges(monkeypatch) -> None:
|
||||
from litellm.proxy.management_endpoints import common_utils
|
||||
|
||||
async def _is_team_admin(**kwargs) -> bool:
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _is_team_admin)
|
||||
|
||||
|
||||
def _non_admin(**kwargs) -> UserAPIKeyAuth:
|
||||
return UserAPIKeyAuth(api_key="sk-test", user_role=LitellmUserRoles.INTERNAL_USER, **kwargs)
|
||||
|
||||
|
||||
def _admin() -> UserAPIKeyAuth:
|
||||
return UserAPIKeyAuth(api_key="sk-test", user_id="u", user_role=LitellmUserRoles.PROXY_ADMIN, team_models=[])
|
||||
|
||||
|
||||
def _team_member() -> UserAPIKeyAuth:
|
||||
return UserAPIKeyAuth(
|
||||
api_key="sk-test",
|
||||
user_id="u",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
team_id="team1",
|
||||
team_models=["model_name_team1_abc", "model_name_team1_def"],
|
||||
models=["model_name_team1_abc", "model_name_team1_def"],
|
||||
)
|
||||
|
||||
|
||||
def _anthropic_request() -> Request:
|
||||
return Request(
|
||||
scope={
|
||||
"type": "http",
|
||||
"method": "GET",
|
||||
"path": "/v1/models",
|
||||
"query_string": b"",
|
||||
"headers": [(b"anthropic-version", b"2023-06-01")],
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
async def _v1_models(user_api_key_dict: UserAPIKeyAuth, **kwargs) -> list[str]:
|
||||
response = await proxy_server.model_list(user_api_key_dict=user_api_key_dict, **kwargs)
|
||||
return [m["id"] for m in response["data"]]
|
||||
|
||||
|
||||
async def _v1_model_info_names(user_api_key_dict: UserAPIKeyAuth) -> list[str]:
|
||||
response = await proxy_server.model_info_v1(user_api_key_dict=user_api_key_dict)
|
||||
return [row["model_name"] for row in json.loads(response.body)["data"]]
|
||||
|
||||
|
||||
async def _model_groups(user_api_key_dict: UserAPIKeyAuth) -> list[str]:
|
||||
response = await proxy_server.model_group_info(user_api_key_dict=user_api_key_dict)
|
||||
return [group.model_group for group in response["data"]]
|
||||
|
||||
|
||||
async def _model_by_id_status(model_id: str, user_api_key_dict: UserAPIKeyAuth) -> int:
|
||||
try:
|
||||
response = await proxy_server.model_info(model_id=model_id, user_api_key_dict=user_api_key_dict)
|
||||
except HTTPException as error:
|
||||
return error.status_code
|
||||
assert response["id"] == model_id
|
||||
return 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v1_models_lists_only_the_names_the_callback_keeps(two_model_router, monkeypatch):
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
|
||||
|
||||
assert await _v1_models(_non_admin()) == ["open-model"]
|
||||
assert await _v1_models(_admin()) == ["open-model"]
|
||||
assert await _v1_models(_non_admin(), request=_anthropic_request()) == ["open-model"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v1_models_scope_expand_applies_the_callback(two_model_router, team_admin_privileges, monkeypatch):
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
|
||||
|
||||
assert await _v1_models(_non_admin(), scope="expand") == ["open-model"]
|
||||
assert await _v1_models(_admin(), scope="expand") == ["open-model"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v1_models_by_id_answers_404_for_a_name_the_callback_leaves_out(two_model_router, monkeypatch):
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
|
||||
|
||||
assert await _model_by_id_status("restricted-model", _non_admin()) == 404
|
||||
assert await _model_by_id_status("open-model", _non_admin()) == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v1_model_info_lists_only_the_rows_the_callback_keeps(two_model_router, monkeypatch):
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
|
||||
|
||||
assert await _v1_model_info_names(_non_admin()) == ["open-model"]
|
||||
assert await _v1_model_info_names(_admin()) == ["open-model"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_group_info_lists_only_the_groups_the_callback_keeps(two_model_router, monkeypatch):
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
|
||||
|
||||
assert await _model_groups(_non_admin()) == ["open-model"]
|
||||
assert await _model_groups(_admin()) == ["open-model"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_callback_without_the_hook_changes_no_listing(two_model_router, monkeypatch):
|
||||
_register(monkeypatch, _InferenceOnlyGate())
|
||||
|
||||
assert await _v1_models(_non_admin()) == ["open-model", "restricted-model"]
|
||||
assert await _model_by_id_status("restricted-model", _non_admin()) == 200
|
||||
assert await _v1_model_info_names(_non_admin()) == ["open-model", "restricted-model"]
|
||||
assert await _model_groups(_non_admin()) == ["open-model", "restricted-model"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_cannot_add_a_name_it_was_not_offered(two_model_router, monkeypatch):
|
||||
_register(monkeypatch, _Gate(extra=("ghost-model",)))
|
||||
|
||||
assert await _v1_models(_non_admin()) == ["open-model", "restricted-model"]
|
||||
assert await _model_by_id_status("ghost-model", _non_admin()) == 404
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callbacks_narrow_in_registration_order(monkeypatch):
|
||||
_install_router(monkeypatch, _deployment("a"), _deployment("b"), _deployment("c"))
|
||||
first: _Gate = _Gate(hidden=frozenset({"a"}))
|
||||
second: _Gate = _Gate(hidden=frozenset({"b"}))
|
||||
_register(monkeypatch, first, second)
|
||||
|
||||
assert await _v1_models(_non_admin()) == ["c"]
|
||||
assert first.seen == [("a", "b", "c")]
|
||||
assert second.seen == [("b", "c")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_sees_and_filters_team_models_by_their_public_name(team_router, monkeypatch):
|
||||
gate: _Gate = _Gate(hidden=frozenset({"team-gpt"}))
|
||||
_register(monkeypatch, gate)
|
||||
|
||||
assert await _v1_models(_team_member()) == ["team-chat"]
|
||||
assert await _model_by_id_status("team-gpt", _team_member()) == 404
|
||||
assert await _model_by_id_status("team-chat", _team_member()) == 200
|
||||
assert all("team-gpt" in seen and "model_name_team1_abc" not in seen for seen in gate.seen)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_sees_public_team_names_on_every_listing_route(team_router, monkeypatch):
|
||||
gate: _Gate = _Gate(hidden=frozenset({"team-gpt"}))
|
||||
_register(monkeypatch, gate)
|
||||
|
||||
assert await _v1_models(_team_member()) == ["team-chat"]
|
||||
assert await _v1_model_info_names(_team_member()) == ["team-chat"]
|
||||
assert await _model_groups(_team_member()) == ["model_name_team1_def"]
|
||||
assert await _model_by_id_status("team-gpt", _team_member()) == 404
|
||||
assert len(gate.seen) == 4
|
||||
assert all(sorted(seen) == ["team-chat", "team-gpt"] for seen in gate.seen)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_alias_follows_its_hidden_target(monkeypatch):
|
||||
_install_router(
|
||||
monkeypatch,
|
||||
_deployment("open-model"),
|
||||
_deployment("restricted-model"),
|
||||
model_group_alias={"mini": "restricted-model", "wide": "open-model"},
|
||||
)
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
|
||||
|
||||
assert sorted(await _v1_models(_non_admin())) == ["open-model", "wide"]
|
||||
assert sorted(await _v1_model_info_names(_non_admin())) == ["open-model", "wide"]
|
||||
assert sorted(await _model_groups(_non_admin())) == ["open-model", "wide"]
|
||||
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"mini"})))
|
||||
|
||||
assert sorted(await _v1_models(_non_admin())) == ["open-model", "restricted-model", "wide"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_alias_of_a_team_model_follows_its_hidden_public_name(monkeypatch):
|
||||
_install_router(
|
||||
monkeypatch,
|
||||
_deployment("gpt-4"),
|
||||
_deployment("model_name_team1_abc", team_id="team1", team_public_model_name="team-gpt"),
|
||||
model_group_alias={"team-alias": "model_name_team1_abc"},
|
||||
)
|
||||
caller: UserAPIKeyAuth = _non_admin(
|
||||
user_id="u",
|
||||
team_id="team1",
|
||||
team_models=["model_name_team1_abc", "team-alias"],
|
||||
models=["model_name_team1_abc", "team-alias"],
|
||||
)
|
||||
_register(monkeypatch, _Gate(hidden=frozenset()))
|
||||
assert sorted(await _v1_models(caller)) == ["team-alias", "team-gpt"]
|
||||
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"team-gpt"})))
|
||||
assert await _v1_models(caller) == []
|
||||
assert await _model_groups(caller) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v1_model_info_offers_only_the_rows_the_caller_would_see(monkeypatch):
|
||||
_install_router(monkeypatch, _deployment("open-model"), _deployment("hidden-model", discoverable=False))
|
||||
gate: _Gate = _Gate()
|
||||
_register(monkeypatch, gate)
|
||||
|
||||
assert await _v1_model_info_names(_non_admin()) == ["open-model"]
|
||||
assert await _v1_model_info_names(_admin()) == ["open-model", "hidden-model"]
|
||||
assert gate.seen == [("open-model",), ("open-model", "hidden-model")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_listing_keeps_its_order_whatever_order_the_callback_returns(monkeypatch):
|
||||
_install_router(monkeypatch, _deployment("a"), _deployment("b"), _deployment("c"))
|
||||
_register(monkeypatch, _ReversingGate())
|
||||
|
||||
assert await _v1_models(_non_admin()) == ["a", "b", "c"]
|
||||
assert await _v1_model_info_names(_non_admin()) == ["a", "b", "c"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_callback_returning_a_string_is_an_error_not_an_empty_listing(two_model_router, monkeypatch):
|
||||
_register(monkeypatch, _StringReturningGate())
|
||||
|
||||
with pytest.raises(ProxyException, match=r"_StringReturningGate\.async_filter_listed_models") as raised:
|
||||
await _v1_models(_non_admin())
|
||||
assert raised.value.code == "500"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_alias_of_a_hidden_model_is_not_listed(two_model_router, monkeypatch):
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
|
||||
caller = _non_admin(aliases={"mini": "restricted-model", "wide": "open-model"})
|
||||
|
||||
assert await _v1_models(caller) == ["open-model", "wide"]
|
||||
assert await _model_by_id_status("mini", caller) == 404
|
||||
assert await _model_by_id_status("wide", caller) == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_error_reaches_the_caller(two_model_router, monkeypatch):
|
||||
_register(monkeypatch, _RaisingGate())
|
||||
|
||||
with pytest.raises(HTTPException) as raised:
|
||||
await _v1_models(_non_admin())
|
||||
assert raised.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hidden_model_still_routes_for_direct_requests(two_model_router, monkeypatch):
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
|
||||
assert "restricted-model" not in await _v1_models(_non_admin())
|
||||
|
||||
deployment = two_model_router.get_available_deployment(
|
||||
model="restricted-model", messages=[{"role": "user", "content": "hi"}]
|
||||
)
|
||||
assert deployment["model_name"] == "restricted-model"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_group_info_offers_a2a_agent_groups_to_the_callback(two_model_router, monkeypatch):
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
monkeypatch.setattr(
|
||||
global_agent_registry,
|
||||
"agent_list",
|
||||
[AgentResponse(agent_id="agent-1", agent_name="helper", agent_card_params={})],
|
||||
)
|
||||
caller = _non_admin(object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="p1", agents=["agent-1"]))
|
||||
gate: _Gate = _Gate()
|
||||
_register(monkeypatch, gate)
|
||||
|
||||
assert await _model_groups(caller) == ["open-model", "restricted-model", "a2a/helper"]
|
||||
assert gate.seen == [("open-model", "restricted-model", "a2a/helper")]
|
||||
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"a2a/helper", "restricted-model"})))
|
||||
|
||||
assert await _model_groups(caller) == ["open-model"]
|
||||
|
||||
|
||||
async def _v1_model_info_by_deployment_id(deployment_id: str, user_api_key_dict: UserAPIKeyAuth) -> int | list[str]:
|
||||
try:
|
||||
response = await proxy_server.model_info_v1(user_api_key_dict=user_api_key_dict, litellm_model_id=deployment_id)
|
||||
except HTTPException as error:
|
||||
return error.status_code
|
||||
return [row["model_name"] for row in json.loads(response.body)["data"]]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v1_model_info_by_deployment_id_answers_like_an_unknown_id_for_a_hidden_model(
|
||||
two_model_router, monkeypatch
|
||||
):
|
||||
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
|
||||
|
||||
assert await _v1_model_info_by_deployment_id("restricted-model-id", _non_admin()) == 400
|
||||
assert await _v1_model_info_by_deployment_id("no-such-id", _non_admin()) == 400
|
||||
assert await _v1_model_info_by_deployment_id("open-model-id", _non_admin()) == ["open-model"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v1_model_info_by_deployment_id_offers_the_public_team_name(team_router, monkeypatch):
|
||||
gate: _Gate = _Gate(hidden=frozenset({"team-gpt"}))
|
||||
_register(monkeypatch, gate)
|
||||
|
||||
assert await _v1_model_info_by_deployment_id("model_name_team1_abc-id", _team_member()) == 400
|
||||
assert await _v1_model_info_by_deployment_id("model_name_team1_def-id", _team_member()) == ["team-chat"]
|
||||
assert gate.seen == [("team-gpt",), ("team-chat",)]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v1_model_info_by_deployment_id_offers_the_name_its_listing_shows_in_legacy_mode(
|
||||
team_router, monkeypatch
|
||||
):
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {"use_team_public_model_name": False})
|
||||
gate: _Gate = _Gate(hidden=frozenset({"team-gpt"}))
|
||||
_register(monkeypatch, gate)
|
||||
|
||||
assert await _v1_model_info_names(_team_member()) == ["team-chat"]
|
||||
assert await _v1_model_info_by_deployment_id("model_name_team1_abc-id", _team_member()) == 400
|
||||
assert gate.seen[-1] == ("team-gpt",)
|
||||
|
|
@ -3224,3 +3224,106 @@ class TestLibpqSslParamTranslation:
|
|||
assert query["sslmode"] == ["require"]
|
||||
assert query["sslcert"] == ["/certs/rds-bundle.pem"]
|
||||
assert query["sslaccept"] == ["strict"]
|
||||
|
||||
|
||||
@pytest.mark.xdist_group("proxy_cli")
|
||||
class TestValidateConfigFlag:
|
||||
def test_validate_config_valid_config_exits_zero(self, tmp_path, monkeypatch):
|
||||
from click.testing import CliRunner
|
||||
|
||||
from litellm.proxy.proxy_cli import run_server
|
||||
|
||||
monkeypatch.delenv("DATABASE_URL", raising=False)
|
||||
monkeypatch.delenv("DIRECT_URL", raising=False)
|
||||
config_path = tmp_path / "config.yaml"
|
||||
config_path.write_text(
|
||||
yaml.safe_dump(
|
||||
{
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-4o",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o",
|
||||
"api_key": "sk-fake",
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
result = CliRunner().invoke(run_server, ["--config", str(config_path), "--validate_config"])
|
||||
|
||||
assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}"
|
||||
assert "config OK" in result.output
|
||||
|
||||
def test_validate_config_invalid_mcp_server_exits_one(self, tmp_path, monkeypatch):
|
||||
from click.testing import CliRunner
|
||||
|
||||
from litellm.proxy.proxy_cli import run_server
|
||||
|
||||
monkeypatch.delenv("DATABASE_URL", raising=False)
|
||||
monkeypatch.delenv("DIRECT_URL", raising=False)
|
||||
config_path = tmp_path / "config.yaml"
|
||||
config_path.write_text(
|
||||
yaml.safe_dump(
|
||||
{
|
||||
"mcp_servers": {
|
||||
"zapier": {
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": "http",
|
||||
"per_server_oauth_discovery": "yes",
|
||||
}
|
||||
}
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
result = CliRunner().invoke(run_server, ["--config", str(config_path), "--validate_config"])
|
||||
|
||||
assert result.exit_code == 1, f"exit_code={result.exit_code}, output={result.output}"
|
||||
assert "per_server_oauth_discovery must be a boolean" in result.output
|
||||
|
||||
def test_validate_config_without_config_is_usage_error(self, monkeypatch):
|
||||
from click.testing import CliRunner
|
||||
|
||||
from litellm.proxy.proxy_cli import run_server
|
||||
|
||||
result = CliRunner().invoke(run_server, ["--validate_config"])
|
||||
|
||||
assert result.exit_code != 0
|
||||
assert "--validate_config requires --config" in result.output
|
||||
|
||||
@patch("subprocess.Popen")
|
||||
def test_validate_config_with_ollama_model_does_not_start_ollama(self, mock_popen, tmp_path, monkeypatch):
|
||||
from click.testing import CliRunner
|
||||
|
||||
from litellm.proxy.proxy_cli import run_server
|
||||
|
||||
monkeypatch.delenv("DATABASE_URL", raising=False)
|
||||
monkeypatch.delenv("DIRECT_URL", raising=False)
|
||||
config_path = tmp_path / "config.yaml"
|
||||
config_path.write_text(
|
||||
yaml.safe_dump(
|
||||
{
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-4o",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o",
|
||||
"api_key": "sk-fake",
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
result = CliRunner().invoke(
|
||||
run_server,
|
||||
["--config", str(config_path), "--model", "ollama/llama3", "--validate_config"],
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}"
|
||||
assert "config OK" in result.output
|
||||
mock_popen.assert_not_called()
|
||||
|
|
|
|||
|
|
@ -340,7 +340,7 @@ def test_a_dockerfile_directory_entry_is_stale_because_only_an_exact_path_exempt
|
|||
def test_a_workflow_that_names_a_file_clears_it_from_the_slice_check():
|
||||
named = coverage._workflow_named_tokens()
|
||||
assert named, "the workflows must name some test paths or the check proves nothing"
|
||||
assert any(coverage._token_covers(token, "tests/local_testing/test_caching_handler.py") for token in named)
|
||||
assert any(coverage._token_covers(token, "tests/proxy_unit_tests/test_proxy_custom_logger.py") for token in named)
|
||||
|
||||
|
||||
def test_the_slice_check_credits_only_workflows_never_the_circleci_config():
|
||||
|
|
|
|||
|
|
@ -107,7 +107,7 @@ CI = [".github/workflows/test-litellm-ui-unit.yml"]
|
|||
),
|
||||
(
|
||||
"cost-map-only",
|
||||
["model_prices_and_context_window.json", "tests/proxy_unit_tests/test_y.py"],
|
||||
["model_prices_and_context_window.json", "tests/unit/proxy/test_y.py"],
|
||||
"run",
|
||||
),
|
||||
(
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ import litellm # noqa: E402 # litellm reads LITELLM_LOCAL_MODEL_COST_MAP at im
|
|||
import litellm.router as litellm_router_module # noqa: E402 # same import-time dependency
|
||||
import litellm.utils as litellm_utils_module # noqa: E402 # same import-time dependency
|
||||
|
||||
LOOPBACK_HOSTS: Final = ["127.0.0.1", "::1"]
|
||||
LOOPBACK_HOSTS: Final = ["127.0.0.1", "::1", "localhost"]
|
||||
AMBIENT_AZURE_CREDENTIAL_ENV_VARS: Final = (
|
||||
"AZURE_AD_TOKEN",
|
||||
"AZURE_TENANT_ID",
|
||||
|
|
|
|||
0
tests/unit/enterprise/integrations/__init__.py
Normal file
0
tests/unit/enterprise/integrations/__init__.py
Normal file
|
|
@ -13,7 +13,6 @@ import asyncio
|
|||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
import os
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
|
@ -165,9 +164,9 @@ async def test_prometheus_metric_tracking():
|
|||
"model_name": "gpt-5-mini", # openai model name
|
||||
"litellm_params": { # params for litellm completion/embedding call
|
||||
"model": "azure/gpt-4.1-mini",
|
||||
"api_key": os.getenv("AZURE_AI_API_KEY"),
|
||||
"api_version": os.getenv("AZURE_AI_API_VERSION"),
|
||||
"api_base": os.getenv("AZURE_AI_API_BASE"),
|
||||
"api_key": "sk-azure-unit-test",
|
||||
"api_version": "2025-01-01-preview",
|
||||
"api_base": "https://unit-test.openai.azure.com",
|
||||
},
|
||||
"model_info": {"id": "azure-model-id"},
|
||||
},
|
||||
|
|
@ -180,9 +179,6 @@ async def test_prometheus_metric_tracking():
|
|||
},
|
||||
],
|
||||
provider_budget_config=provider_budget_config,
|
||||
redis_host=os.getenv("REDIS_HOST"),
|
||||
redis_port=int(os.getenv("REDIS_PORT", 6379)),
|
||||
redis_password=os.getenv("REDIS_PASSWORD"),
|
||||
)
|
||||
|
||||
try:
|
||||
0
tests/unit/enterprise/proxy/__init__.py
Normal file
0
tests/unit/enterprise/proxy/__init__.py
Normal file
0
tests/unit/enterprise/proxy/auth/__init__.py
Normal file
0
tests/unit/enterprise/proxy/auth/__init__.py
Normal file
0
tests/unit/enterprise/proxy/guardrails/__init__.py
Normal file
0
tests/unit/enterprise/proxy/guardrails/__init__.py
Normal file
0
tests/unit/enterprise/proxy/hooks/__init__.py
Normal file
0
tests/unit/enterprise/proxy/hooks/__init__.py
Normal file
0
tests/unit/gateway/__init__.py
Normal file
0
tests/unit/gateway/__init__.py
Normal file
0
tests/unit/litellm_proxy_extras/__init__.py
Normal file
0
tests/unit/litellm_proxy_extras/__init__.py
Normal file
|
|
@ -9,7 +9,7 @@ import pytest
|
|||
sys.path.insert(
|
||||
0,
|
||||
os.path.abspath(
|
||||
os.path.join(os.path.dirname(__file__), "../../litellm-proxy-extras")
|
||||
os.path.join(os.path.dirname(__file__), "../../../litellm-proxy-extras")
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -23,7 +23,7 @@ from litellm_proxy_extras.utils import (
|
|||
_MIGRATIONS_DIR = os.path.abspath(
|
||||
os.path.join(
|
||||
os.path.dirname(__file__),
|
||||
"../../litellm-proxy-extras/litellm_proxy_extras/migrations",
|
||||
"../../../litellm-proxy-extras/litellm_proxy_extras/migrations",
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -999,7 +999,7 @@ class TestJWTKeyMappingCascade:
|
|||
schema_paths = glob.glob(
|
||||
os.path.abspath(
|
||||
os.path.join(
|
||||
os.path.dirname(__file__), "../../**/schema.prisma"
|
||||
os.path.dirname(__file__), "../../../**/schema.prisma"
|
||||
)
|
||||
),
|
||||
recursive=True,
|
||||
0
tests/unit/proxy/__init__.py
Normal file
0
tests/unit/proxy/__init__.py
Normal file
0
tests/unit/proxy/auth/__init__.py
Normal file
0
tests/unit/proxy/auth/__init__.py
Normal file
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue